Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
193 changes: 193 additions & 0 deletions common/sagemaker_args.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,193 @@
"""Translate SM_VLLM_* environment variables into vLLM server CLI arguments.

SageMaker passes configuration as environment variables, so the entrypoint has to
turn ``SM_VLLM_TENSOR_PARALLEL_SIZE=8`` into ``--tensor-parallel-size 8``.
The translation mirrors what vLLM already does for ``--config file.yaml`` in
``vllm/utils/argparse_utils.py``:

* JSON array -> the flag plus one argv token per element
* JSON object -> the flag plus a single argv token (the object)
* anything else -> the flag plus the raw value as a single token

Because the decision is made from the *value's* syntax rather than a per-flag table,
no list of flags has to be maintained as vLLM adds arguments. The exception is the
handful of flags where vLLM deliberately wants an array as one token; those are
listed in SINGLE_TOKEN_FLAGS below.

Tokens are written to stdout NUL-delimited so the shell can read them back into an
array without word-splitting values that contain spaces or newlines. Informational
messages go to stderr to keep stdout parseable.
"""

import json
import os
import sys
from typing import Any, Iterable, List, Mapping, Optional

PREFIX = "SM_VLLM_"
ARG_PREFIX = "--"
DEFAULT_PORT = "8080"
MODEL_DIR = "/opt/ml/model"

# Flags whose value stays a single argv token even when it is a JSON array.
# vLLM strips nargs from these on purpose (see FrontendArgs._customize_cli_kwargs in
# vllm/entrypoints/openai/cli_args.py): the first three are typed ``json.loads`` so the
# array *is* the value, and --middleware is ``action="append"``, taking one value per
# occurrence.
SINGLE_TOKEN_FLAGS = frozenset(
{
"--allowed-origins",
"--allowed-methods",
"--allowed-headers",
"--middleware",
}
)


def flag_for(env_key: str) -> str:
"""``SM_VLLM_TENSOR_PARALLEL_SIZE`` -> ``--tensor-parallel-size``."""
name = env_key[len(PREFIX) :].lower().replace("_", "-")
return f"{ARG_PREFIX}{name}"


def as_token(element: Any) -> str:
"""Render one element of a JSON array as a single argv token."""
if isinstance(element, (dict, list)):
return json.dumps(element, separators=(",", ":"))
if isinstance(element, bool):
return "true" if element else "false"
if element is None:
return ""
return str(element)


def json_object_sequence(value: str) -> Optional[List[dict]]:
"""Parse ``{...} {...}`` into its objects.

Returns None if the string is not a whitespace-separated run of JSON objects, so
callers can fall back to passing the value through untouched.
"""
decoder = json.JSONDecoder()
objects: List[dict] = []
index = 0
length = len(value)
while index < length:
try:
parsed, end = decoder.raw_decode(value, index)
except ValueError:
return None
if not isinstance(parsed, dict):
return None
objects.append(parsed)
index = end
while index < length and value[index].isspace():
index += 1
return objects or None


def tokens_for(flag: str, value: str) -> Optional[List[str]]:
"""Return the argv tokens that follow `flag`.

An empty list means "emit the flag with no values" (a boolean-style flag); None
means "omit the flag entirely", which is what an empty JSON array asks for since
argparse rejects a nargs='+' flag with zero values.
"""
if not value:
return []

stripped = value.strip()

if flag in SINGLE_TOKEN_FLAGS:
return [value]

if stripped.startswith("[") and stripped.endswith("]"):
try:
parsed = json.loads(stripped)
except ValueError:
return [value]
if isinstance(parsed, list):
if not parsed:
return None
return [as_token(element) for element in parsed]
return [value]

if stripped.startswith("{") and stripped.endswith("}"):
objects = json_object_sequence(stripped)
if objects is not None and len(objects) > 1:
return [json.dumps(obj, separators=(",", ":")) for obj in objects]
# A single JSON object is passed through verbatim
return [value]

return [value]


def resolve_model(env: Mapping[str, str], model_dir: str) -> List[str]:
"""Pick the model source when SM_VLLM_MODEL is not set.

Precedence: an explicit SM_VLLM_MODEL (handled by the generic loop, so nothing is
added here), then a populated model dir, then HF_MODEL_ID.
"""
if env.get(f"{PREFIX}MODEL"):
return []

if os.path.isdir(model_dir) and os.listdir(model_dir):
log(f"INFO: {PREFIX}MODEL not set, auto-detected model at {model_dir}")
return ["--model", model_dir]

if env.get("HF_MODEL_ID"):
log(f"INFO: {PREFIX}MODEL not set, using HF_MODEL_ID={env['HF_MODEL_ID']}")
return ["--model", env["HF_MODEL_ID"]]

log(
f"WARNING: No model specified. Set {PREFIX}MODEL, HF_MODEL_ID, "
f"or mount a model to {model_dir}."
)
return []


def log(message: str) -> None:
print(message, file=sys.stderr)


def build_args(env: Mapping[str, str], model_dir: str = MODEL_DIR) -> List[str]:
"""Build the full argv list (minus the program itself) for the vLLM server."""
args: List[str] = ["--port", DEFAULT_PORT]
args += resolve_model(env, model_dir)

for key in sorted(env):
if not key.startswith(PREFIX):
continue
flag = flag_for(key)
value = env[key]
lowered = value.strip().lower()

# Boolean flags: true -> bare flag, false -> omitted entirely.
if lowered == "true":
args.append(flag)
continue
if lowered == "false":
continue

tokens = tokens_for(flag, value)
if tokens is None:
log(f"WARNING: {key} is an empty list; skipping {flag}.")
continue
args.append(flag)
args += tokens

return args


def emit(tokens: Iterable[str]) -> None:
"""Write tokens NUL-delimited for `mapfile -t -d ''` on the shell side."""
sys.stdout.write("".join(f"{token}\0" for token in tokens))


def main() -> None:
args = build_args(os.environ)
log(f"INFO: vLLM server arguments: {args}")
emit(args)


if __name__ == "__main__":
main()
26 changes: 26 additions & 0 deletions common/serve
Original file line number Diff line number Diff line change
@@ -0,0 +1,26 @@
#!/bin/bash
# SageMaker hosting entrypoint for the vLLM Neuron DLC.
#
# SageMaker launches an inference container as `docker run <image> serve`. This image's
# ENTRYPOINT (common/vllm_entrypoint.py) execs argv as-is, so exposing an executable
# named `serve` on PATH satisfies SageMaker without changing the ENTRYPOINT, keeping the
# existing EC2/k8s usage (`docker run <image> vllm serve ...`) working unchanged.
#
# Adapted from aws/deep-learning-containers scripts/docker/vllm/sagemaker_entrypoint.sh
# (Apache-2.0). vLLM >= 0.24 already registers /ping and /invocations via
# vllm.entrypoints.serve.sagemaker.api_router, so no routing middleware is needed here.
set -euo pipefail

ARGS_FILE=$(mktemp)
trap 'rm -f "${ARGS_FILE}"' EXIT
if ! python3 /usr/local/bin/sagemaker_args.py >"${ARGS_FILE}"; then
echo "ERROR: failed to build vLLM arguments from SM_VLLM_* environment variables" >&2
exit 1
fi

ARGS=()
while IFS= read -r -d '' token; do
ARGS+=("${token}")
done <"${ARGS_FILE}"

exec standard-supervisor python3 -m vllm.entrypoints.openai.api_server "${ARGS[@]}"
6 changes: 6 additions & 0 deletions vllm/inference/0.24.0.1.1.0/Dockerfile.neuronx
Original file line number Diff line number Diff line change
Expand Up @@ -119,6 +119,12 @@ RUN printf '%s\n' \

COPY --chmod=755 vllm_entrypoint.py neuron-monitor.sh deep_learning_container.py /usr/local/bin/

# SageMaker hosting support. SageMaker launches an inference container as
# `docker run <image> serve`, so exposing an executable named `serve` on PATH lets the
# existing ENTRYPOINT (vllm_entrypoint.py, which execs its argv) start the vLLM server.
COPY --chmod=755 serve /usr/local/bin/serve
COPY --chmod=755 sagemaker_args.py /usr/local/bin/sagemaker_args.py

### Mount Point ###
# When launching the container, mount the code directory to /workspace
ARG APP_MOUNT=/workspace
Expand Down