feat: unify Docker serving configuration in YAML

- Add server.py --config serve.yaml; explicit CLI flags override YAML
- Add scripts/serve.sh and serve_runtime.py for the Compose lifecycle
- Template server/cpu ports and param mounts in docker-compose.yml
- Document schema in docs/developer/docker-serving.md and params guide
- Add tests for runtime parsing and server CLI merge logic
This commit is contained in:
2026-08-21 23:16:45 +08:00
parent dcc96de12a
commit cb21af38ba
10 changed files with 817 additions and 6 deletions
+189
View File
@@ -0,0 +1,189 @@
#!/usr/bin/env bash
set -euo pipefail
ROOT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")/.." && pwd)"
source "${ROOT_DIR}/scripts/docker/lib/train-common.sh"
COMPOSE_BASE=(
docker compose
--project-directory "${ROOT_DIR}"
--file "${ROOT_DIR}/docker-compose.yml"
)
usage() {
cat <<'EOF'
Usage: scripts/serve.sh <command> [CONFIG] [options]
CONFIG defaults to ./serve.yaml. The same file declares host runtime settings
under `runtime:` and server settings under `server:`.
Commands:
init [CONFIG] Create the model directory
preflight [CONFIG] Validate Docker, paths, GPU, and Compose
build [CONFIG] Build the serving image
up [CONFIG] Start the server container (detached)
run [CONFIG] Start the server container (foreground)
down [CONFIG] Stop and remove the server container
restart [CONFIG] Down, then up
logs [CONFIG] Follow server logs
status [CONFIG] Show container status
EOF
}
resolve_path() {
if [[ "$1" = /* ]]; then
printf '%s\n' "$1"
else
printf '%s/%s\n' "${ROOT_DIR}" "${1#./}"
fi
}
load_config() {
CONFIG_FILE="$(resolve_path "$1")"
[[ -f "${CONFIG_FILE}" ]] || die "Serving config not found: ${CONFIG_FILE}"
require_command python3
python3 -c 'import yaml' >/dev/null 2>&1 ||
die "PyYAML is required on the host (install python3-yaml)"
local exports
exports="$(python3 "${ROOT_DIR}/scripts/tools/serve_runtime.py" exports "${CONFIG_FILE}")" ||
die "Failed to load runtime configuration"
eval "${exports}"
if [[ -n "${SERVE_JOB_NAME}" ]]; then
validate_job_name "${SERVE_JOB_NAME}"
fi
}
compose() {
ASTRAI_UID="$(id -u)" ASTRAI_GID="$(id -g)" "${COMPOSE_BASE[@]}" "$@"
}
container_name() {
if [[ -n "${SERVE_JOB_NAME}" ]]; then
printf 'astrai-server-%s\n' "${SERVE_JOB_NAME}"
else
printf 'astrai-server\n'
fi
}
service_name() {
if [[ "${SERVE_GPU_ENABLED:-true}" == "false" ]]; then
printf 'server-cpu\n'
else
printf 'server\n'
fi
}
set_profile_args() {
PROFILE_ARGS=()
if [[ "${SERVE_GPU_ENABLED:-true}" == "false" ]]; then
PROFILE_ARGS=(--profile cpu)
fi
}
init_environment() {
mkdir -p "${SERVE_PARAM_DIR}"
log_info "Model: ${SERVE_PARAM_DIR}"
}
preflight() {
require_command docker
docker info >/dev/null 2>&1 || die "Docker daemon is unavailable"
[[ -d "${SERVE_PARAM_DIR}" ]] || die "Model directory not found: ${SERVE_PARAM_DIR}"
[[ -s "${SERVE_PARAM_DIR}/config.json" ]] ||
die "Model config not found: ${SERVE_PARAM_DIR}/config.json"
[[ -s "${SERVE_PARAM_DIR}/model.safetensors" ]] ||
die "Model weights not found: ${SERVE_PARAM_DIR}/model.safetensors"
if [[ "${SERVE_GPU_ENABLED}" == "false" ]] && [[ "${SERVE_DEVICE}" != "cpu" ]]; then
die "runtime.gpu.enabled is false but server.device is '${SERVE_DEVICE}'; use server.device: cpu"
fi
compose config --quiet
log_info "Preflight passed (service: $(service_name), device: ${SERVE_DEVICE})"
}
runtime_environment_args() {
RUNTIME_ENV_ARGS=()
local pair
while IFS= read -r -d '' pair; do
RUNTIME_ENV_ARGS+=(--env "${pair}")
done < <(python3 "${ROOT_DIR}/scripts/tools/serve_runtime.py" environment "${CONFIG_FILE}")
}
start_server() {
local foreground="$1"
shift
local container running
local -a run_options
preflight
runtime_environment_args
set_profile_args
container="$(container_name)"
running="$(docker inspect --format '{{.State.Running}}' "${container}" 2>/dev/null || true)"
[[ "${running}" != "true" ]] || die "Server is already running: ${container}"
docker rm "${container}" >/dev/null 2>&1 || true
run_options=(
--volume "${CONFIG_FILE}:/run/astrai/serve.yaml:ro"
"${RUNTIME_ENV_ARGS[@]}"
)
if [[ "${foreground}" == "true" ]]; then
compose "${PROFILE_ARGS[@]}" run --build --rm --service-ports \
"${run_options[@]}" "$(service_name)" \
python -m scripts.tools.server --config /run/astrai/serve.yaml "$@"
else
compose "${PROFILE_ARGS[@]}" run -d --build --service-ports \
--name "${container}" "${run_options[@]}" "$(service_name)" \
python -m scripts.tools.server --config /run/astrai/serve.yaml "$@"
log_info "Server started; run scripts/serve.sh logs ${CONFIG_FILE} to follow it"
fi
}
stop_server() {
local container
container="$(container_name)"
docker stop --timeout 30 "${container}" >/dev/null 2>&1 ||
log_warn "Server container is not running"
docker rm "${container}" >/dev/null 2>&1 || true
}
show_status() {
docker ps -a --filter "name=^/$(container_name)$"
}
main() {
local command="${1:-}" config="${SERVE_CONFIG_FILE:-${ROOT_DIR}/serve.yaml}"
[[ -n "${command}" ]] || { usage; exit 1; }
shift || true
if [[ "${command}" =~ ^(help|-h|--help)$ ]]; then
usage
return
fi
if [[ $# -gt 0 && "$1" != --* ]]; then
config="$1"
shift
fi
load_config "${config}"
case "${command}" in
init) init_environment ;;
preflight) preflight ;;
build)
set_profile_args
preflight
compose "${PROFILE_ARGS[@]}" build "$(service_name)"
;;
up) start_server false "$@" ;;
run) start_server true "$@" ;;
down) stop_server ;;
restart) stop_server; start_server false ;;
logs) docker logs -f --tail "${SERVE_LOG_TAIL:-200}" "$(container_name)" ;;
status) show_status ;;
*) die "Unknown command: ${command}" ;;
esac
}
main "$@"
+154
View File
@@ -0,0 +1,154 @@
"""Parse the host-side runtime section of a serving configuration.
The Compose wrapper needs a few container-side values on the host as well:
``server.port`` (the port the container listens on) and ``server.device``
(used by the preflight GPU consistency check). Everything else under
``server:`` is owned by ``scripts/tools/server.py --config`` inside the
container.
"""
import argparse
import re
import shlex
from pathlib import Path
import yaml
ENV_NAME = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$")
def _mapping(value, name: str) -> dict:
if value is None:
return {}
if not isinstance(value, dict):
raise ValueError(f"{name} must be a mapping")
return value
def _path(value, name: str, config_dir: Path) -> str:
if not isinstance(value, str) or not value.strip():
raise ValueError(f"runtime.paths.{name} is required")
path = Path(value).expanduser()
if not path.is_absolute():
path = config_dir / path
return str(path.resolve())
def _port(value, name: str) -> int:
if isinstance(value, bool) or not isinstance(value, int):
raise ValueError(f"{name} must be an integer")
if not 1 <= value <= 65535:
raise ValueError(f"{name} must be between 1 and 65535")
return value
def load_runtime(config_path: str) -> dict[str, str]:
path = Path(config_path).resolve()
with path.open(encoding="utf-8") as file:
config = yaml.safe_load(file) or {}
if not isinstance(config, dict):
raise ValueError("serving configuration must be a mapping")
runtime = _mapping(config.get("runtime"), "runtime")
if not runtime:
raise ValueError("top-level runtime section is required")
paths = _mapping(runtime.get("paths"), "paths")
gpu = _mapping(runtime.get("gpu"), "gpu")
container = _mapping(runtime.get("container"), "container")
environment = _mapping(runtime.get("environment"), "environment")
server = _mapping(config.get("server"), "server")
job_name = runtime.get("job_name", "")
if job_name and not isinstance(job_name, str):
raise ValueError("runtime.job_name must be a string")
if job_name and not re.fullmatch(r"[A-Za-z0-9][A-Za-z0-9._-]*", job_name):
raise ValueError(
"runtime.job_name must use letters, numbers, dot, underscore, or dash"
)
port = _port(runtime.get("port", 8000), "runtime.port")
container_port = _port(server.get("port", 8000), "server.port")
device = server.get("device", "cuda")
if not isinstance(device, str) or not device.strip():
raise ValueError("server.device must be a string")
gpu_enabled = gpu.get("enabled", True)
if not isinstance(gpu_enabled, bool):
raise ValueError("runtime.gpu.enabled must be a boolean")
devices = gpu.get("devices", "all")
if gpu_enabled:
if devices == "all":
visible_devices = ""
elif isinstance(devices, list) and len(devices) == 1:
text = str(devices[0])
if not text.isdigit():
raise ValueError(
"runtime.gpu.devices entries must be non-negative integers"
)
visible_devices = text
else:
raise ValueError(
"runtime.gpu.devices must be 'all' or a single-device list such as [0]"
)
else:
visible_devices = ""
if devices != "all":
raise ValueError(
"runtime.gpu.devices is ignored when runtime.gpu.enabled is false"
)
if device != "cpu":
raise ValueError(
"server.device must be 'cpu' when runtime.gpu.enabled is false"
)
values = {
"SERVE_JOB_NAME": job_name,
"SERVE_PORT": str(port),
"SERVE_CONTAINER_PORT": str(container_port),
"SERVE_PARAM_DIR": _path(paths.get("param", "./params"), "param", path.parent),
"SERVE_GPU_ENABLED": "true" if gpu_enabled else "false",
"SERVE_DEVICE": device,
"CUDA_VISIBLE_DEVICES": visible_devices,
"CUDA_TAG": str(container.get("cuda_tag", "cu128")),
}
for name, value in environment.items():
if not isinstance(name, str) or not ENV_NAME.fullmatch(name):
raise ValueError(f"invalid runtime.environment name: {name!r}")
if value is not None and not isinstance(value, (str, int, float, bool)):
raise ValueError(f"runtime.environment.{name} must be a scalar")
values["environment"] = environment
return values
def shell_exports(runtime: dict[str, str]) -> str:
return "\n".join(
f"export {name}={shlex.quote(value)}"
for name, value in runtime.items()
if name != "environment"
)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("command", choices=("exports", "environment"))
parser.add_argument("config")
args = parser.parse_args()
try:
runtime = load_runtime(args.config)
except (OSError, ValueError, yaml.YAMLError) as exc:
parser.error(str(exc))
if args.command == "exports":
print(shell_exports(runtime))
return
for name, value in runtime["environment"].items():
rendered = "" if value is None else str(value)
print(f"{name}={rendered}", end="\0")
if __name__ == "__main__":
main()
+128 -2
View File
@@ -2,16 +2,105 @@ from pathlib import Path
import click
import torch
import yaml
from click.core import ParameterSource
from astrai.inference import run_server
_DTYPES = ["bfloat16", "float16", "float32"]
_SERVER_KEYS = (
"host",
"port",
"reload",
"param_path",
"device",
"dtype",
"max_batch_size",
"max_seq_len",
)
def _merge_yaml_into_kwargs(
config_path: str,
passed_kwargs: dict,
explicit_keys: set[str] | None = None,
) -> dict:
"""Merge Click defaults, YAML server values, then explicit CLI values."""
with open(config_path, encoding="utf-8") as file:
config = yaml.safe_load(file) or {}
if not isinstance(config, dict):
raise click.UsageError(f"Serving config must be a mapping: {config_path}")
server = config.get("server") or {}
if not isinstance(server, dict):
raise click.UsageError("top-level server section must be a mapping")
unknown = sorted(set(server) - set(_SERVER_KEYS))
if unknown:
click.echo(
f"Warning: ignoring unknown server config keys: {', '.join(unknown)}",
err=True,
)
merged = dict(passed_kwargs)
merged.update({key: server[key] for key in _SERVER_KEYS if key in server})
if explicit_keys is None:
explicit_keys = set(passed_kwargs)
for key in explicit_keys:
if key in passed_kwargs:
merged[key] = passed_kwargs[key]
return merged
def _as_int(value, name: str) -> int | None:
if value is None:
return None
if isinstance(value, bool):
raise click.UsageError(f"{name} must be an integer")
try:
return int(value)
except (TypeError, ValueError):
raise click.UsageError(f"{name} must be an integer, got {value!r}") from None
def _resolve_server_config(
config_path: str,
passed_kwargs: dict,
explicit_keys: set[str] | None = None,
) -> dict:
"""Merge YAML values, then coerce and validate the resolved settings.
``explicit_keys`` are CLI flags that win over YAML; when None, YAML values
win over Click defaults.
"""
merged = _merge_yaml_into_kwargs(config_path, passed_kwargs, explicit_keys or set())
resolved = dict(merged)
resolved["port"] = _as_int(resolved["port"], "server.port") or 8000
resolved["max_batch_size"] = (
_as_int(resolved["max_batch_size"], "server.max_batch_size") or 16
)
resolved["max_seq_len"] = _as_int(resolved["max_seq_len"], "server.max_seq_len")
resolved["reload"] = bool(resolved["reload"])
if resolved["dtype"] not in _DTYPES:
raise click.UsageError(
f"server.dtype must be one of {', '.join(_DTYPES)}, got {resolved['dtype']!r}"
)
return resolved
@click.command(name="serve", help="Launch inference server (OpenAI-compatible API).")
@click.option(
"--config",
"-c",
"config_path",
type=click.Path(exists=True, dir_okay=False),
default=None,
help="Serving YAML config. CLI flags override YAML values.",
)
@click.option("--host", default="0.0.0.0", help="Host address.")
@click.option("--port", type=int, default=8000, help="Port number.")
@click.option("--reload", is_flag=True, help="Enable auto-reload for development.")
@click.option(
"--reload", is_flag=True, default=False, help="Enable auto-reload for development."
)
@click.option(
"--param_path",
type=click.Path(exists=True),
@@ -37,10 +126,47 @@ _DTYPES = ["bfloat16", "float16", "float32"]
default=None,
help="Maximum sequence length (KV cache size + prompt truncation). Uses model config if not set.",
)
@click.pass_context
def server_command(
host, port, reload, param_path, device, dtype, max_batch_size, max_seq_len
ctx,
config_path,
host,
port,
reload,
param_path,
device,
dtype,
max_batch_size,
max_seq_len,
):
"""Launch inference server (OpenAI-compatible API)."""
if config_path:
passed_kwargs = {
"host": host,
"port": port,
"reload": reload,
"param_path": param_path,
"device": device,
"dtype": dtype,
"max_batch_size": max_batch_size,
"max_seq_len": max_seq_len,
}
explicit_keys = {
key
for key in passed_kwargs
if ctx.get_parameter_source(key) is ParameterSource.COMMANDLINE
}
resolved = _resolve_server_config(config_path, passed_kwargs, explicit_keys)
host = resolved["host"]
port = resolved["port"]
reload = resolved["reload"]
param_path = resolved["param_path"]
device = resolved["device"]
dtype = resolved["dtype"]
max_batch_size = resolved["max_batch_size"]
max_seq_len = resolved["max_seq_len"]
click.echo(f"Config: {config_path}")
dtype_map = {
"bfloat16": torch.bfloat16,
"float16": torch.float16,