Skip to content
Closed
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
22 changes: 21 additions & 1 deletion components/src/dynamo/sglang/_compat.py
Original file line number Diff line number Diff line change
Expand Up @@ -216,14 +216,33 @@ async def start_profile_compat(tokenizer_manager: Any, body: dict[str, Any]) ->
await start_profile(**body)


def set_resolved_server_arg(server_args: Any, **fields: Any) -> None:
"""Mutate SGLang ``ServerArgs`` fields after resolution.

Newer SGLang freezes ``ServerArgs`` once ``__post_init__`` resolves it: a
bare ``server_args.field = value`` raises ``AttributeError`` ("assigned
after resolution") and must go through ``ServerArgs.override()``, the single
sanctioned post-resolution mutation point. Older SGLang -- and the
``SimpleNamespace`` diffusion stub -- have no ``override`` and accept direct
assignment. Remove the fallback when the minimum supported SGLang carries
``override``.
"""
override = getattr(server_args, "override", None)
if callable(override):
override("dynamo", **fields)
else:
for name, value in fields.items():
setattr(server_args, name, value)


def enable_disjoint_streaming_output(server_args: Any) -> None:
"""Enable SGLang's disjoint streaming output.

Diffusion workers pass a ``SimpleNamespace`` stub that does not carry the
field, so this is a no-op when the attribute is absent.
"""
if hasattr(server_args, "incremental_streaming_output"):
server_args.incremental_streaming_output = True
set_resolved_server_arg(server_args, incremental_streaming_output=True)


__all__ = [
Expand All @@ -232,5 +251,6 @@ def enable_disjoint_streaming_output(server_args: Any) -> None:
"ensure_sglang_top_level_exports",
"filter_supported_async_generate_kwargs",
"require_reasoning_kwargs",
"set_resolved_server_arg",
"start_profile_compat",
]
5 changes: 3 additions & 2 deletions components/src/dynamo/sglang/args.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@
from dynamo.sglang._compat import (
enable_disjoint_streaming_output,
ensure_sglang_tensor_image_size,
set_resolved_server_arg,
)
from dynamo.sglang.backend_args import DynamoSGLangArgGroup, DynamoSGLangConfig

Expand Down Expand Up @@ -600,7 +601,7 @@ async def parse_args(
fpm_trace_relay_supported=fpm_trace_relay_supported,
)
if fpm_source and not getattr(server_args, "enable_forward_pass_metrics", False):
server_args.enable_forward_pass_metrics = True
set_resolved_server_arg(server_args, enable_forward_pass_metrics=True)
logging.info("Enabled forward_pass_metrics from %s", fpm_source)

# Auto-detect diffusion worker mode if dllm_algorithm
Expand All @@ -616,7 +617,7 @@ async def parse_args(
server_args.dllm_algorithm
and getattr(server_args, "max_running_requests", None) is None
):
server_args.max_running_requests = 8
set_resolved_server_arg(server_args, max_running_requests=8)
logging.info("Defaulting max_running_requests to 8 for diffusion worker")

dynamo_config.namespace = parsed_namespace
Expand Down
5 changes: 3 additions & 2 deletions components/src/dynamo/sglang/init_llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
from dynamo.common.utils.endpoint_types import parse_endpoint_types
from dynamo.llm import ModelInput, ModelType, WorkerType
from dynamo.runtime import DistributedRuntime
from dynamo.sglang._compat import set_resolved_server_arg
from dynamo.sglang.args import Config
from dynamo.sglang.health_check import (
SglangDisaggHealthCheckPayload,
Expand Down Expand Up @@ -71,7 +72,7 @@ async def init_decode(
"created before the endpoint existed, so its FPM publisher bound "
"a different IPC path than the relay would subscribe to."
)
server_args.enable_forward_pass_metrics = False
set_resolved_server_arg(server_args, enable_forward_pass_metrics=False)
else:
set_forward_pass_metrics_worker_id(server_args, generate_endpoint)
start_time = time.time()
Expand Down Expand Up @@ -227,7 +228,7 @@ async def init_prefill(
"created before the endpoint existed, so its FPM publisher bound "
"a different IPC path than the relay would subscribe to."
)
server_args.enable_forward_pass_metrics = False
set_resolved_server_arg(server_args, enable_forward_pass_metrics=False)
else:
set_forward_pass_metrics_worker_id(server_args, generate_endpoint)
start_time = time.time()
Expand Down
5 changes: 4 additions & 1 deletion components/src/dynamo/sglang/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
)
from dynamo.common.utils.runtime import create_runtime
from dynamo.runtime.logging import configure_dynamo_logging
from dynamo.sglang._compat import set_resolved_server_arg
from dynamo.sglang.args import parse_args
from dynamo.sglang.init_diffusion import (
init_image_diffusion,
Expand Down Expand Up @@ -44,7 +45,9 @@ async def worker(argv: list[str] | None = None):
if config.server_args.load_format == "gms":
from gpu_memory_service.integrations.sglang import setup_gms

config.server_args.load_format = setup_gms(config.server_args)
set_resolved_server_arg(
config.server_args, load_format=setup_gms(config.server_args)
)

# Snapshot mode: engine must be created before runtime so CRIU captures no
# NATS/etcd connections.
Expand Down
8 changes: 6 additions & 2 deletions components/src/dynamo/sglang/publisher.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@
)
from dynamo.llm import KvEventPublisher, WorkerMetricsPublisher
from dynamo.runtime import Endpoint
from dynamo.sglang._compat import set_resolved_server_arg
from dynamo.sglang._disagg import SGLANG_WORKER_GROUP_ID_KEY, get_sglang_worker_group_id
from dynamo.sglang.args import Config
from dynamo.sglang.capacity import (
Expand All @@ -48,9 +49,12 @@ def set_forward_pass_metrics_worker_id(

import tempfile

server_args.forward_pass_metrics_worker_id = str(generate_endpoint.connection_id())
ipc_path = tempfile.NamedTemporaryFile(delete=False).name
server_args.forward_pass_metrics_ipc_name = f"ipc://{ipc_path}"
set_resolved_server_arg(
server_args,
forward_pass_metrics_worker_id=str(generate_endpoint.connection_id()),
forward_pass_metrics_ipc_name=f"ipc://{ipc_path}",
)


async def _resolve_multinode_leader_worker_id(
Expand Down
6 changes: 4 additions & 2 deletions components/src/dynamo/sglang/request_handlers/handler_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,7 +43,7 @@
)
from dynamo.llm.exceptions import EngineShutdown
from dynamo.runtime import DistributedRuntime
from dynamo.sglang._compat import start_profile_compat
from dynamo.sglang._compat import set_resolved_server_arg, start_profile_compat
from dynamo.sglang.args import Config
from dynamo.sglang.pause import SGLangEnginePauseController
from dynamo.sglang.publisher import DynamoSglangPublisher
Expand Down Expand Up @@ -948,7 +948,9 @@ async def update_weight_version(self, body: dict) -> dict:
if req.abort_all_requests:
self.engine.tokenizer_manager.abort_request(abort_all=True)

self.engine.tokenizer_manager.server_args.weight_version = req.new_version
set_resolved_server_arg(
self.engine.tokenizer_manager.server_args, weight_version=req.new_version
)
return {
"success": True,
"message": f"Weight version updated to {req.new_version}",
Expand Down
5 changes: 3 additions & 2 deletions components/src/dynamo/sglang/snapshot.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
SnapshotConfig,
configure_snapshot_capture_env,
)
from dynamo.sglang._compat import set_resolved_server_arg

from .pause import SGLangEnginePauseController

Expand Down Expand Up @@ -138,15 +139,15 @@ async def prepare_snapshot_engine(
# Enable memory_saver so GPU memory can be released for CRIU.
# When using GMS, weights use VA-stable unmap/remap (no CPU backup); GMS
# forbids enable_weights_cpu_backup. Otherwise use CPU backup for weights.
server_args.enable_memory_saver = True
set_resolved_server_arg(server_args, enable_memory_saver=True)
try:
from gpu_memory_service.integrations.sglang import is_gms_active

_using_gms = is_gms_active()
except ImportError:
_using_gms = False
if not _using_gms:
server_args.enable_weights_cpu_backup = True
set_resolved_server_arg(server_args, enable_weights_cpu_backup=True)

start_time = time.time()
engine = sgl.Engine(server_args=server_args)
Expand Down
40 changes: 40 additions & 0 deletions components/src/dynamo/sglang/tests/test_sglang_unit.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@
ensure_sglang_top_level_exports,
filter_supported_async_generate_kwargs,
require_reasoning_kwargs,
set_resolved_server_arg,
start_profile_compat,
)
from dynamo.sglang.args import (
Expand Down Expand Up @@ -428,6 +429,45 @@ async def start_profile(self, req=None):
assert manager.received is request


def test_set_resolved_server_arg_uses_sglang_override():
calls = []

class ResolvedServerArgs:
def override(self, source, **fields):
calls.append((source, fields))

server_args = ResolvedServerArgs()
set_resolved_server_arg(
server_args,
weight_version="v2",
enable_forward_pass_metrics=True,
)

assert calls == [
(
"dynamo",
{
"weight_version": "v2",
"enable_forward_pass_metrics": True,
},
)
]
assert not hasattr(server_args, "weight_version")


def test_set_resolved_server_arg_falls_back_for_legacy_objects():
server_args = SimpleNamespace()

set_resolved_server_arg(
server_args,
weight_version="v2",
enable_forward_pass_metrics=True,
)

assert server_args.weight_version == "v2"
assert server_args.enable_forward_pass_metrics is True


@pytest.mark.asyncio
async def test_custom_jinja_template_invalid_path(mock_sglang_cli):
"""Test that invalid file path raises FileNotFoundError."""
Expand Down
Loading