feat(cli): add extensible serve backend plugins (#34753)

This commit is contained in:
Mick
2026-08-14 13:57:59 +08:00
committed by GitHub
parent 463981922c
commit 46d84f4b48
7 changed files with 796 additions and 85 deletions
+147 -83
View File
@@ -4,7 +4,14 @@ import argparse
import logging
import os
from sglang.cli.utils import get_is_diffusion_model, get_model_path
from sglang.cli.serve_backends import (
SERVE_BACKEND_API_VERSION,
ServeBackend,
ServeBackendDetection,
ServeBackendRegistry,
ServeRequest,
)
from sglang.cli.utils import get_is_diffusion_model, get_model_path, try_get_model_path
from sglang.srt.utils import kill_process_tree
from sglang.srt.utils.common import suppress_noisy_warnings
@@ -14,36 +21,35 @@ logger = logging.getLogger(__name__)
def _extract_model_type_override(extra_argv):
"""Extract and remove --model-type override from argv."""
model_type = "auto"
"""Extract and remove the backend selector from argv.
Validation is deferred to :class:`ServeBackendRegistry` because installed
out-of-tree entry points dynamically extend the accepted values.
"""
backend_name = "auto"
filtered_argv = []
i = 0
while i < len(extra_argv):
arg = extra_argv[i]
if arg == "--model-type":
if i + 1 >= len(extra_argv):
raise Exception(
"Error: --model-type requires a value. "
"Valid values are: auto, llm, diffusion."
)
model_type = extra_argv[i + 1]
raise ValueError("Error: --model-type requires a value.")
backend_name = extra_argv[i + 1]
i += 2
continue
if arg.startswith("--model-type="):
model_type = arg.split("=", 1)[1]
backend_name = arg.split("=", 1)[1]
i += 1
continue
filtered_argv.append(arg)
i += 1
if model_type not in ("auto", "llm", "diffusion"):
raise Exception(
f"Error: invalid --model-type '{model_type}'. "
"Valid values are: auto, llm, diffusion."
)
return model_type, filtered_argv
if not backend_name:
raise ValueError("Error: --model-type requires a non-empty value.")
return backend_name, filtered_argv
def _normalize_positional_model_path(extra_argv):
@@ -53,91 +59,149 @@ def _normalize_positional_model_path(extra_argv):
return extra_argv, False
def serve(args, extra_argv):
if any(h in extra_argv for h in ("-h", "--help")):
# Since the server type is determined by the model, and we don't have a model path,
# we can't show the exact help. Instead, we show a general help message and then
# the help for both possible server types.
print(
"Usage: sglang serve <model-name-or-path> [additional-arguments]\n"
" or: sglang serve --model-path <model-name-or-path> [additional-arguments]\n\n"
"This command can launch either a standard language model server or a diffusion model server.\n"
"The server type is determined by the --model-path.\n"
"Optional override: --model-type {auto,llm,diffusion} "
"(default: auto, fallback to LLM on detection failure)."
def _print_llm_help(_request: ServeRequest) -> None:
from sglang.srt.server_args import prepare_server_args
try:
prepare_server_args(["--help"])
except SystemExit:
pass # argparse --help calls sys.exit
def _print_diffusion_help(_request: ServeRequest) -> None:
try:
from sglang.multimodal_gen.runtime.entrypoints.cli.serve import (
add_multimodal_gen_serve_args,
)
print("\n--- Help for Standard Language Model Server ---")
from sglang.srt.server_args import prepare_server_args
parser = argparse.ArgumentParser(
prog="sglang serve",
description="SGLang Diffusion Model Serving",
)
add_multimodal_gen_serve_args(parser)
parser.print_help()
except ImportError:
print(
"Diffusion model support is not available. "
'Install with: pip install "sglang[diffusion]"'
)
try:
prepare_server_args(["--help"])
except SystemExit:
pass # argparse --help calls sys.exit
print("\n--- Help for Diffusion Model Server ---")
try:
from sglang.multimodal_gen.runtime.entrypoints.cli.serve import (
add_multimodal_gen_serve_args,
)
def _run_llm(request: ServeRequest) -> None:
if any(arg in request.argv for arg in ("-h", "--help")):
_print_llm_help(request)
return
parser = argparse.ArgumentParser(
prog="sglang serve",
description="SGLang Diffusion Model Serving",
)
add_multimodal_gen_serve_args(parser)
parser.print_help()
except ImportError:
print(
"Diffusion model support is not available. "
'Install with: pip install "sglang[diffusion]"'
)
from sglang.launch_server import run_server
from sglang.srt.server_args import prepare_server_args
server_args = prepare_server_args(list(request.argv))
run_server(server_args)
def _detect_diffusion(request: ServeRequest) -> ServeBackendDetection:
if request.model_path is None:
return ServeBackendDetection.UNKNOWN
if get_is_diffusion_model(request.model_path):
return ServeBackendDetection.MATCH
return ServeBackendDetection.NO_MATCH
def _run_diffusion(request: ServeRequest) -> None:
if any(arg in request.argv for arg in ("-h", "--help")):
_print_diffusion_help(request)
return
from sglang.multimodal_gen.runtime.entrypoints.cli.serve import (
add_multimodal_gen_serve_args,
execute_serve_cmd,
)
parser = argparse.ArgumentParser(description="SGLang Diffusion Model Serving")
add_multimodal_gen_serve_args(parser)
parsed_args, remaining_argv = parser.parse_known_args(list(request.argv))
if request.model_path_is_positional:
parsed_args._sglang_explicit_arg_names = {"model_path"}
execute_serve_cmd(parsed_args, remaining_argv)
def _create_backend_registry() -> ServeBackendRegistry:
return ServeBackendRegistry(
{
"llm": ServeBackend(
api_version=SERVE_BACKEND_API_VERSION,
run=_run_llm,
),
"diffusion": ServeBackend(
api_version=SERVE_BACKEND_API_VERSION,
run=_run_diffusion,
detect=_detect_diffusion,
),
}
)
def _print_general_help(registry: ServeBackendRegistry) -> None:
available = ",".join(("auto", *registry.available_names))
print(
"Usage: sglang serve <model-name-or-path> [additional-arguments]\n"
" or: sglang serve --model-path <model-name-or-path> "
"[additional-arguments]\n\n"
"The serving backend is detected from the model by default. Installed "
"out-of-tree projects can add more backends.\n"
f"Optional override: --model-type {{{available}}} "
"(default: auto, fallback to LLM when no backend matches).\n"
"Use `sglang serve --model-type BACKEND --help` for backend-specific "
"help."
)
print("\n--- Help for Standard Language Model Server ---")
_print_llm_help(ServeRequest(argv=("--help",), model_path=None))
print("\n--- Help for Diffusion Model Server ---")
_print_diffusion_help(ServeRequest(argv=("--help",), model_path=None))
def serve(args, extra_argv):
del args # The top-level parser currently has no serve-owned fields.
backend_name, dispatch_argv = _extract_model_type_override(extra_argv)
dispatch_argv, positional_model_path = _normalize_positional_model_path(
dispatch_argv
)
request = ServeRequest(
argv=tuple(dispatch_argv),
model_path=try_get_model_path(dispatch_argv),
model_path_is_positional=positional_model_path,
)
registry = _create_backend_registry()
if any(h in request.argv for h in ("-h", "--help")):
if backend_name == "auto":
_print_general_help(registry)
else:
registry.get(backend_name).backend.run(request)
return
from sglang.srt.plugins import load_plugins
load_plugins()
model_type, dispatch_argv = _extract_model_type_override(extra_argv)
dispatch_argv, positional_model_path = _normalize_positional_model_path(
dispatch_argv
)
model_path = get_model_path(dispatch_argv)
try:
if model_type == "auto":
is_diffusion_model = get_is_diffusion_model(model_path)
if is_diffusion_model:
logger.info("Diffusion model detected")
if backend_name == "auto":
registered = registry.auto_detect(request)
logger.info("Selected serve backend %r", registered.name)
else:
is_diffusion_model = model_type == "diffusion"
registered = registry.get(backend_name)
logger.info(
"Dispatch override enabled: --model-type=%s " "(skip auto detection)",
model_type,
backend_name,
)
if is_diffusion_model:
# Logic for Diffusion Models
from sglang.multimodal_gen.runtime.entrypoints.cli.serve import (
add_multimodal_gen_serve_args,
execute_serve_cmd,
)
if registered.backend.requires_model_path and request.model_path is None:
get_model_path(request.argv) # Raise the existing actionable CLI error.
parser = argparse.ArgumentParser(
description="SGLang Diffusion Model Serving"
)
add_multimodal_gen_serve_args(parser)
parsed_args, remaining_argv = parser.parse_known_args(dispatch_argv)
if positional_model_path:
parsed_args._sglang_explicit_arg_names = {"model_path"}
execute_serve_cmd(parsed_args, remaining_argv)
else:
# Logic for Standard Language Models
from sglang.launch_server import run_server
from sglang.srt.server_args import prepare_server_args
server_args = prepare_server_args(dispatch_argv)
run_server(server_args)
registered.backend.run(request)
finally:
kill_process_tree(os.getpid(), include_parent=False)
+231
View File
@@ -0,0 +1,231 @@
# SPDX-License-Identifier: Apache-2.0
"""Extension API and discovery for ``sglang serve`` backends.
Out-of-tree projects register a zero-argument factory in the
``sglang.serve_backends`` entry point group. The factory returns a
:class:`ServeBackend`; its entry point name becomes a valid ``--model-type``.
The module is intentionally lightweight. Importing it must not initialize an
inference runtime or import an out-of-tree backend implementation.
"""
from __future__ import annotations
import logging
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from enum import Enum
from importlib.metadata import EntryPoint, entry_points
logger = logging.getLogger(__name__)
SERVE_BACKENDS_GROUP = "sglang.serve_backends"
SERVE_BACKEND_API_VERSION = 1
RESERVED_SERVE_BACKEND_NAMES = frozenset({"auto"})
class ServeBackendDetection(str, Enum):
"""Result returned by a serve backend's optional detector."""
MATCH = "match"
NO_MATCH = "no_match"
UNKNOWN = "unknown"
@dataclass(frozen=True)
class ServeRequest:
"""Normalized command line forwarded from ``sglang serve`` to a backend."""
argv: tuple[str, ...]
model_path: str | None
model_path_is_positional: bool = False
ServeBackendRunner = Callable[[ServeRequest], None]
ServeBackendDetector = Callable[[ServeRequest], ServeBackendDetection]
@dataclass(frozen=True)
class ServeBackend:
"""Implementation contract for built-in and out-of-tree serve backends.
``run`` must parse the backend-owned arguments in ``request.argv``. It must
also honor ``-h`` and ``--help`` without launching a server. For a real
launch, it should block for the server lifetime; SGLang applies its common
child-process cleanup after ``run`` returns or raises.
``detect`` is optional. Backends without a detector remain available by
explicit ``--model-type`` but do not participate in automatic routing.
Detector implementations should be lightweight and return ``UNKNOWN`` on
inconclusive I/O or metadata errors.
"""
api_version: int
run: ServeBackendRunner
detect: ServeBackendDetector | None = None
requires_model_path: bool = True
@dataclass(frozen=True)
class RegisteredServeBackend:
"""A loaded backend together with its discovery metadata."""
name: str
backend: ServeBackend
distribution: str | None = None
class ServeBackendRegistry:
"""Registry of built-in and installed out-of-tree serve backends."""
def __init__(self, builtins: Mapping[str, ServeBackend]) -> None:
invalid_builtin_names = set(builtins) & RESERVED_SERVE_BACKEND_NAMES
if invalid_builtin_names:
names = ", ".join(sorted(invalid_builtin_names))
raise ValueError(f"Reserved serve backend names cannot be used: {names}")
self._builtins = dict(builtins)
self._entry_points = self._discover_entry_points()
self._loaded: dict[str, RegisteredServeBackend] = {
name: RegisteredServeBackend(name=name, backend=backend)
for name, backend in self._builtins.items()
}
reserved = (set(self._builtins) | RESERVED_SERVE_BACKEND_NAMES) & set(
self._entry_points
)
if reserved:
names = ", ".join(sorted(reserved))
raise RuntimeError(
"Out-of-tree serve backends cannot replace reserved or built-in "
f"backends: {names}"
)
@staticmethod
def _discover_entry_points() -> dict[str, list[EntryPoint]]:
discovered: dict[str, list[EntryPoint]] = {}
for entry_point in entry_points(group=SERVE_BACKENDS_GROUP):
discovered.setdefault(entry_point.name, []).append(entry_point)
return discovered
@property
def available_names(self) -> tuple[str, ...]:
"""Return backend names without importing out-of-tree packages."""
external_names = sorted(set(self._entry_points) - set(self._builtins))
return (*self._builtins, *external_names)
def get(self, name: str) -> RegisteredServeBackend:
"""Return one backend, importing only the explicitly requested plugin."""
if name in self._loaded:
return self._loaded[name]
candidates = self._entry_points.get(name, [])
if not candidates:
available = ", ".join(("auto", *self.available_names))
raise ValueError(
f"Unknown serve backend {name!r}. Available values: {available}."
)
if len(candidates) > 1:
providers = ", ".join(
sorted(
self._entry_point_provider(candidate) for candidate in candidates
)
)
raise RuntimeError(
f"Multiple distributions register serve backend {name!r}: "
f"{providers}. Uninstall one provider or choose another backend name."
)
entry_point = candidates[0]
try:
factory = entry_point.load()
if not callable(factory):
raise TypeError("the entry point must resolve to a callable factory")
backend = factory()
except Exception as exc:
raise RuntimeError(
f"Failed to load serve backend {name!r} from "
f"{self._entry_point_provider(entry_point)}: {exc}"
) from exc
if not isinstance(backend, ServeBackend):
raise TypeError(
f"Serve backend {name!r} factory returned {type(backend).__name__}; "
"expected sglang.cli.serve_backends.ServeBackend."
)
if backend.api_version != SERVE_BACKEND_API_VERSION:
raise RuntimeError(
f"Serve backend {name!r} uses API version {backend.api_version}; "
f"this SGLang release requires version {SERVE_BACKEND_API_VERSION}."
)
registered = RegisteredServeBackend(
name=name,
backend=backend,
distribution=self._entry_point_distribution(entry_point),
)
self._loaded[name] = registered
return registered
def auto_detect(self, request: ServeRequest) -> RegisteredServeBackend:
"""Resolve a unique detector match, falling back to the ``llm`` backend."""
matches: list[RegisteredServeBackend] = []
for name in self.available_names:
if name == "llm":
# LLM preserves the historical fallback role instead of matching
# every Hugging Face repository.
continue
try:
registered = self.get(name)
detector = registered.backend.detect
if detector is None:
continue
result = detector(request)
if not isinstance(result, ServeBackendDetection):
raise TypeError(
"detector must return ServeBackendDetection, got "
f"{type(result).__name__}"
)
if result is ServeBackendDetection.MATCH:
matches.append(registered)
except Exception as exc:
# An unrelated optional extension must not make the default LLM
# path unusable. Explicit selection remains strict via get().
logger.warning(
"Skipping automatic detection for serve backend %r: %s",
name,
exc,
)
if len(matches) > 1:
names = ", ".join(match.name for match in matches)
raise RuntimeError(
f"Multiple serve backends matched this request: {names}. "
"Select one explicitly with --model-type BACKEND."
)
if matches:
return matches[0]
return self.get("llm")
@classmethod
def _entry_point_distribution(cls, entry_point: EntryPoint) -> str | None:
distribution = getattr(entry_point, "dist", None)
return getattr(distribution, "name", None)
@classmethod
def _entry_point_provider(cls, entry_point: EntryPoint) -> str:
return cls._entry_point_distribution(entry_point) or entry_point.value
__all__ = [
"SERVE_BACKEND_API_VERSION",
"SERVE_BACKENDS_GROUP",
"ServeBackend",
"ServeBackendDetection",
"ServeBackendRegistry",
"ServeRequest",
]
+10 -2
View File
@@ -96,8 +96,9 @@ def get_is_diffusion_model(model_path: str) -> bool:
return False
def get_model_path(extra_argv):
# Find the model_path argument
def try_get_model_path(extra_argv) -> str | None:
"""Return a model path from command-line arguments when one is present."""
model_path = None
for i, arg in enumerate(extra_argv):
if arg in ("--model-path", "--model"):
@@ -108,6 +109,13 @@ def get_model_path(extra_argv):
model_path = arg.split("=", 1)[1]
break
return model_path
def get_model_path(extra_argv):
# Find the model_path argument
model_path = try_get_model_path(extra_argv)
if model_path is None:
# Fallback for --help or other cases where model-path is not provided
if any(h in extra_argv for h in ["-h", "--help"]):