[diffusion] chore: disagg server args, launch helpers, and warmup utils (#26119)

This commit is contained in:
Heyang Huang
2026-06-04 13:40:39 +08:00
committed by GitHub
parent 0e0ecc11ff
commit 858e5a5109
14 changed files with 657 additions and 234 deletions
+6
View File
@@ -26,6 +26,7 @@ if TYPE_CHECKING:
SGLANG_DIFFUSION_TRACE_FUNCTION: int = 0
SGLANG_DIFFUSION_WORKER_MULTIPROC_METHOD: str = "fork"
SGLANG_DIFFUSION_TARGET_DEVICE: str = "cuda"
SGLANG_DIFFUSION_PLATFORM_OVERRIDE: str = ""
MAX_JOBS: str | None = None
NVCC_THREADS: str | None = None
CMAKE_BUILD_TYPE: str | None = None
@@ -239,6 +240,11 @@ environment_variables: dict[str, Callable[[], Any]] = {
"SGLANG_DIFFUSION_WORKER_MULTIPROC_METHOD": _lazy_str(
"SGLANG_DIFFUSION_WORKER_MULTIPROC_METHOD", "fork"
),
# Internal per-worker platform override used by disaggregated role launch.
# Empty means normal platform auto-detection.
"SGLANG_DIFFUSION_PLATFORM_OVERRIDE": _lazy_str(
"SGLANG_DIFFUSION_PLATFORM_OVERRIDE", ""
),
# Enables torch profiler if set. Path to the directory where torch profiler
# traces are saved. Note that it must be an absolute path.
"SGLANG_DIFFUSION_TORCH_PROFILER_DIR": _lazy_path(
@@ -1,193 +1,28 @@
# SPDX-License-Identifier: Apache-2.0
"""Disaggregated diffusion CLI arguments and helper methods.
All disagg-related dataclass fields, argparse registration, and endpoint
derivation logic live here. ``ServerArgs`` inherits from
``DisaggArgsMixin`` so the fields appear on the top-level config object.
"""
"""Compatibility shim for disaggregated diffusion argument helpers."""
from __future__ import annotations
import argparse
from typing import TYPE_CHECKING
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
from sglang.multimodal_gen.runtime.server_args_disagg import DisaggServerArgsMixin
if TYPE_CHECKING:
pass
# ── Port offsets for disagg result endpoints (deterministic convention) ──
DISAGG_RESULT_PORT_OFFSETS: dict[RoleType, int] = {
RoleType.ENCODER: 1,
RoleType.DENOISER: 2,
RoleType.DECODER: 3,
}
class DisaggArgsMixin:
"""Methods for disaggregated diffusion, mixed into ``ServerArgs``.
The dataclass **fields** remain in ``ServerArgs`` (to avoid MRO
ordering issues with ``@dataclass`` inheritance). This mixin only
provides the methods that operate on those fields.
"""
def get_role_parallelism(self, role_type: RoleType) -> dict[str, int | None]:
"""Return per-role parallelism overrides for the given role.
Returns a dict with keys tp_size, sp_degree, ulysses_degree,
ring_degree. Values are ``None`` when not explicitly set
(auto-derive from ``num_gpus``).
"""
_none: dict[str, int | None] = {
"tp_size": None,
"sp_degree": None,
"ulysses_degree": None,
"ring_degree": None,
}
if role_type == RoleType.ENCODER:
return {**_none, "tp_size": self.encoder_tp}
elif role_type == RoleType.DENOISER:
return {
"tp_size": self.denoiser_tp,
"sp_degree": self.denoiser_sp,
"ulysses_degree": self.denoiser_ulysses,
"ring_degree": self.denoiser_ring,
}
elif role_type == RoleType.DECODER:
return {**_none, "tp_size": self.decoder_tp}
return _none
def derive_pool_result_endpoint(self) -> str:
"""Derive the result PUSH endpoint from ``disagg_server_addr`` + role.
Convention: DS binds result PULL on ``scheduler_port + {1,2,3}``
for encoder / denoiser / decoder.
"""
if self.disagg_server_addr is None:
raise ValueError("disagg_server_addr is required for per-role launch")
addr = self.disagg_server_addr
if addr.startswith("tcp://"):
addr = addr[len("tcp://") :]
host, port_str = addr.rsplit(":", 1)
base_port = int(port_str)
offset = DISAGG_RESULT_PORT_OFFSETS[self.disagg_role]
return f"tcp://{host}:{base_port + offset}"
def derive_pool_work_endpoint(self) -> str:
"""Derive the work PULL bind endpoint for a standalone role instance."""
return f"tcp://0.0.0.0:{self.scheduler_port}"
# ── CLI registration ─────────────────────────────────────────────────
# Keep the historical disagg_args import path working.
DISAGG_RESULT_PORT_OFFSETS = DisaggServerArgsMixin.DISAGG_RESULT_PORT_OFFSETS
DisaggArgsMixin = DisaggServerArgsMixin
def add_disagg_cli_args(parser: argparse.ArgumentParser) -> None:
"""Register all disaggregated-diffusion CLI arguments as a group."""
"""Register disaggregated-diffusion CLI args through ServerArgs."""
g = parser.add_argument_group(
"Disaggregated diffusion",
"Split the pipeline into independent Encoder / Denoiser / Decoder "
"roles, each on its own GPU(s). A DiffusionServer head node routes "
"requests. See docs/disaggregation.md for details.",
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
# Core
g.add_argument(
"--base-gpu-id",
type=int,
default=0,
help="Starting GPU ID for this instance. Used with --disagg-role "
"to place role instances on specific GPUs without CUDA_VISIBLE_DEVICES.",
)
g.add_argument(
"--disagg-role",
type=str,
default=RoleType.MONOLITHIC.value,
choices=RoleType.choices(),
help="Role for disaggregated pipeline. "
"'monolithic' (default): single server. "
"'encoder' / 'denoiser' / 'decoder': role instance. "
"'server': DiffusionServer head node (no GPU). "
"Role instances require --disagg-server-addr. "
"Server requires --encoder-urls, --denoiser-urls, --decoder-urls.",
)
g.add_argument(
"--disagg-server-addr",
type=str,
default=None,
help="DiffusionServer head node address (tcp://HOST:PORT). "
"Required for role instances.",
)
g.add_argument(
"--disagg-timeout",
type=int,
default=600,
help="Timeout in seconds for pending disagg requests (default: 600).",
)
g.add_argument(
"--disagg-dispatch-policy",
type=str,
default="round_robin",
choices=["round_robin", "max_free_slots"],
help="Dispatch policy: 'round_robin' or 'max_free_slots' (default: round_robin).",
)
# Server head: remote instance URLs
g.add_argument(
"--encoder-urls",
type=str,
default=None,
help="Encoder work endpoints (semicolon-separated). "
"Example: 'tcp://10.0.0.1:35000;tcp://10.0.0.2:35000'.",
)
g.add_argument(
"--denoiser-urls",
type=str,
default=None,
help="Denoiser work endpoints (semicolon-separated).",
)
g.add_argument(
"--decoder-urls",
type=str,
default=None,
help="Decoder work endpoints (semicolon-separated).",
)
# Per-role parallelism
g.add_argument("--encoder-tp", type=int, default=None, help="Encoder TP degree.")
g.add_argument("--denoiser-tp", type=int, default=None, help="Denoiser TP degree.")
g.add_argument("--denoiser-sp", type=int, default=None, help="Denoiser SP degree.")
g.add_argument(
"--denoiser-ulysses", type=int, default=None, help="Denoiser Ulysses degree."
)
g.add_argument(
"--denoiser-ring", type=int, default=None, help="Denoiser Ring degree."
)
g.add_argument("--decoder-tp", type=int, default=None, help="Decoder TP degree.")
# P2P transfer engine
g.add_argument(
"--disagg-transfer-pool-size",
type=int,
default=256 * 1024 * 1024,
help="P2P transfer buffer pool size in bytes (default: 256 MiB).",
)
g.add_argument(
"--disagg-p2p-hostname",
type=str,
default="127.0.0.1",
help="RDMA-reachable hostname/IP of this instance (default: 127.0.0.1).",
)
g.add_argument(
"--disagg-ib-device",
type=str,
default=None,
help="InfiniBand device for RDMA transfers (e.g., mlx5_0).",
)
ServerArgs.add_disagg_cli_args(parser)
def convert_disagg_role_string(kwargs: dict) -> None:
"""Convert ``disagg_role`` from string to ``RoleType`` enum in-place."""
if "disagg_role" in kwargs and isinstance(kwargs["disagg_role"], str):
kwargs["disagg_role"] = RoleType.from_string(kwargs["disagg_role"])
@@ -3,8 +3,8 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/platforms/__init__.py
import os
import traceback
from typing import TYPE_CHECKING
# imported by other files, do not remove
from sglang.multimodal_gen.runtime.platforms.interface import ( # noqa: F401
@@ -171,6 +171,23 @@ builtin_platform_plugins = {
def resolve_current_platform_cls_qualname() -> str:
forced_platform = os.environ.get("SGLANG_DIFFUSION_PLATFORM_OVERRIDE", "").strip()
if forced_platform:
forced_map = {
"cpu": "sglang.multimodal_gen.runtime.platforms.cpu.CpuPlatform",
"cuda": "sglang.multimodal_gen.runtime.platforms.cuda.CudaPlatform",
"rocm": "sglang.multimodal_gen.runtime.platforms.rocm.RocmPlatform",
"mps": "sglang.multimodal_gen.runtime.platforms.mps.MpsPlatform",
"npu": "sglang.multimodal_gen.runtime.platforms.npu.NPUPlatformBase",
"musa": "sglang.multimodal_gen.runtime.platforms.musa.MusaPlatform",
}
qualname = forced_map.get(forced_platform.lower())
if qualname is None:
raise ValueError(
f"Unsupported SGLANG_DIFFUSION_PLATFORM_OVERRIDE={forced_platform!r}"
)
return qualname
# TODO(will): if we need to support other platforms, we should consider if
# vLLM's plugin architecture is suitable for our needs.
@@ -27,6 +27,14 @@ class CpuPlatform(Platform):
device_type = "cpu"
dispatch_key = "CPU"
@classmethod
def get_local_torch_device(cls) -> torch.device:
return torch.device("cpu")
@classmethod
def get_torch_distributed_backend_str(cls) -> str:
return "gloo"
@classmethod
def get_cpu_architecture(cls) -> CpuArchEnum:
"""Get the CPU architecture."""
@@ -38,10 +46,6 @@ class CpuPlatform(Platform):
else:
return CpuArchEnum.UNSPECIFIED
@classmethod
def get_local_torch_device(cls) -> torch.device:
return torch.device("cpu")
@classmethod
def get_device_name(cls, device_id: int = 0) -> str:
return platform.processor()
@@ -70,7 +74,7 @@ class CpuPlatform(Platform):
@classmethod
def get_available_gpu_memory(
cls,
device_id: int = 0,
device_id: int | None = None,
distributed: bool = False,
empty_cache: bool = True,
cpu_group: Any = None,
@@ -91,10 +95,6 @@ class CpuPlatform(Platform):
return free_memory / (1 << 30)
@classmethod
def get_device_communicator_cls(cls) -> str:
return "sglang.multimodal_gen.runtime.distributed.device_communicators.cpu_communicator.CpuCommunicator"
@classmethod
def get_attn_backend_cls_str(
cls,
@@ -102,12 +102,21 @@ class CpuPlatform(Platform):
head_size: int,
dtype: torch.dtype,
) -> str:
if selected_backend not in (None, AttentionBackendEnum.TORCH_SDPA):
logger.warning(
"%s is not supported on CPU; falling back to Torch SDPA.",
selected_backend,
)
logger.info("Using Torch SDPA backend")
logger.info("Using Torch SDPA backend for CPU.")
return (
"sglang.multimodal_gen.runtime.layers.attention.backends.sdpa.SDPABackend"
)
@classmethod
def get_device_communicator_cls(cls) -> str:
return "sglang.multimodal_gen.runtime.distributed.device_communicators.cpu_communicator.CpuCommunicator"
@classmethod
def enable_dit_layerwise_offload_for_wan_by_default(cls) -> bool:
"""Whether to enable DIT layerwise offload by default on the current platform."""
@@ -189,7 +189,7 @@ class CudaPlatformBase(Platform):
@classmethod
def get_available_gpu_memory(
cls,
device_id: int = 0,
device_id: int | None = None,
distributed: bool = False,
empty_cache: bool = True,
cpu_group: Any = None,
@@ -197,8 +197,8 @@ class CudaPlatformBase(Platform):
if empty_cache:
torch.cuda.empty_cache()
if torch.distributed.is_initialized():
device_id = torch.distributed.get_rank()
if device_id is None:
device_id = torch.cuda.current_device()
device_props = torch.cuda.get_device_properties(device_id)
if device_props.is_integrated:
@@ -384,7 +384,7 @@ class Platform:
@classmethod
def get_available_gpu_memory(
cls,
device_id: int = 0,
device_id: int | None = None,
distributed: bool = False,
empty_cache: bool = True,
cpu_group: Any = None,
@@ -78,7 +78,7 @@ class MpsPlatform(Platform):
@classmethod
def get_available_gpu_memory(
cls,
device_id: int = 0,
device_id: int | None = None,
distributed: bool = False,
empty_cache: bool = True,
cpu_group: Any = None,
@@ -122,7 +122,7 @@ class MusaPlatformBase(Platform):
@classmethod
def get_available_gpu_memory(
cls,
device_id: int = 0,
device_id: int | None = None,
distributed: bool = False,
empty_cache: bool = True,
cpu_group: Any = None,
@@ -130,8 +130,8 @@ class MusaPlatformBase(Platform):
if empty_cache:
torch.cuda.empty_cache()
if torch.distributed.is_initialized():
device_id = torch.distributed.get_rank()
if device_id is None:
device_id = torch.cuda.current_device()
device_props = torch.cuda.get_device_properties(device_id)
if device_props.is_integrated:
@@ -84,7 +84,7 @@ class NPUPlatformBase(Platform):
@classmethod
def get_available_gpu_memory(
cls,
device_id: int = 0,
device_id: int | None = None,
distributed: bool = False,
empty_cache: bool = True,
cpu_group: Any = None,
@@ -92,6 +92,9 @@ class NPUPlatformBase(Platform):
if empty_cache:
torch.npu.empty_cache()
if device_id is None:
device_id = torch.npu.current_device()
free_gpu_memory, _ = torch.npu.mem_get_info(device_id)
if distributed:
@@ -75,7 +75,7 @@ class RocmPlatform(Platform):
@classmethod
def get_available_gpu_memory(
cls,
device_id: int = 0,
device_id: int | None = None,
distributed: bool = False,
empty_cache: bool = True,
cpu_group: Any = None,
@@ -83,6 +83,9 @@ class RocmPlatform(Platform):
if empty_cache:
torch.cuda.empty_cache()
if device_id is None:
device_id = torch.cuda.current_device()
free_gpu_memory, _ = torch.cuda.mem_get_info(device_id)
if distributed:
@@ -14,7 +14,7 @@ import sys
import tempfile
from dataclasses import field
from enum import Enum
from typing import Any, Optional
from typing import Any, Literal, Optional
import addict
import yaml
@@ -27,11 +27,6 @@ from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import (
is_ltx23_native_variant,
)
from sglang.multimodal_gen.configs.quantization.nunchaku import NunchakuSVDQuantArgs
from sglang.multimodal_gen.runtime.disaggregation.disagg_args import (
DisaggArgsMixin,
add_disagg_cli_args,
convert_disagg_role_string,
)
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
from sglang.multimodal_gen.runtime.layers.quantization.configs.nunchaku_config import (
NunchakuConfig,
@@ -52,9 +47,11 @@ from sglang.multimodal_gen.runtime.server_args_auto_tune import (
PERFORMANCE_MODES,
ServerArgsAutoTuner,
)
from sglang.multimodal_gen.runtime.server_args_disagg import DisaggServerArgsMixin
from sglang.multimodal_gen.runtime.utils.common import (
is_port_available,
is_valid_ipv6_address,
normalize_gpu_ids,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import (
_sanitize_for_logging,
@@ -116,7 +113,7 @@ class Backend(str, Enum):
@dataclasses.dataclass
class ServerArgs(DisaggArgsMixin):
class ServerArgs(DisaggServerArgsMixin):
# Model and path configuration (for convenience)
model_path: str
@@ -146,6 +143,8 @@ class ServerArgs(DisaggArgsMixin):
# Parallelism
num_gpus: int = 1
performance_mode: str = "auto"
base_gpu_id: int = 0
gpu_ids: list[int] | None = None
tp_size: Optional[int] = None
sp_degree: Optional[int] = None
# sequence parallelism
@@ -282,13 +281,21 @@ class ServerArgs(DisaggArgsMixin):
# MoE parameters used by Wan2.2
boundary_ratio: float | None = None
# Disaggregation — fields defined here, methods in DisaggArgsMixin,
# CLI registration in disagg_args.add_disagg_cli_args().
base_gpu_id: int = 0
# Disaggregation (pool mode only — launched via launch_pool_disagg_server())
disagg_role: RoleType = RoleType.MONOLITHIC
disagg_timeout: int = 600
disagg_timeout: int = 3600
disagg_downstream_wait_timeout: int = 1800
disagg_dispatch_policy: str = "round_robin"
disagg_mode: bool = False
disagg_instance_id: int = 0
disagg_max_slots_per_instance: int = 8
disagg_transfer_redundancy: float = 1.25
disagg_role_device: Literal["auto", "cpu", "cuda"] = "auto"
disagg_transfer_backend: Literal["auto", "mock", "mooncake"] = "auto"
disagg_transfer_pool_size: int = 256 * 1024 * 1024
disagg_transfer_pin_memory: Literal["auto", "off", "required"] = "auto"
disagg_p2p_hostname: str = "127.0.0.1"
disagg_ib_device: str | None = None
disagg_server_addr: str | None = None
encoder_urls: str | None = None
denoiser_urls: str | None = None
@@ -298,12 +305,12 @@ class ServerArgs(DisaggArgsMixin):
denoiser_sp: int | None = None
denoiser_ulysses: int | None = None
denoiser_ring: int | None = None
decoder_sp: int | None = None
decoder_tp: int | None = None
disagg_transfer_pool_size: int = 256 * 1024 * 1024
disagg_p2p_hostname: str = "127.0.0.1"
disagg_ib_device: str | None = None
pool_work_endpoint: str | None = None
pool_result_endpoint: str | None = None
pool_control_endpoint: str | None = None
pool_control_advertised_endpoint: str | None = None
# Logging
log_level: str = "info"
@@ -313,8 +320,6 @@ class ServerArgs(DisaggArgsMixin):
enable_trace: bool = False
otlp_traces_endpoint: str = "localhost:4317"
# get_role_parallelism, derive_pool_*_endpoint — from DisaggArgsMixin
@property
def broker_port(self) -> int:
return self.port + 1
@@ -334,6 +339,7 @@ class ServerArgs(DisaggArgsMixin):
"""set defaults and normalize values."""
auto_tuner = ServerArgsAutoTuner(self)
auto_tuner.adjust_based_on_performance_mode()
self._adjust_disagg_parallelism_aliases()
if auto_tuner.could_override_server_args():
self._adjust_offload()
auto_tuner.maybe_adjust_auto_default_layerwise_offload()
@@ -355,6 +361,21 @@ class ServerArgs(DisaggArgsMixin):
auto_tuner.finalize_auto_flags()
self.adjust_pipeline_config()
def _adjust_disagg_parallelism_aliases(self):
if self.decoder_tp is None:
return
if self.decoder_sp is not None and self.decoder_sp != self.decoder_tp:
raise ValueError(
"decoder_tp is deprecated in favor of decoder_sp; "
"please set only one of them or keep the same value."
)
if self.decoder_sp is None:
logger.warning(
"decoder_tp is deprecated and is treated as decoder_sp for "
"decoder/VAE parallel decode. Please use decoder_sp instead."
)
self.decoder_sp = self.decoder_tp
def _validate_parameters(self):
"""check consistency and raise errors for invalid configs"""
self._validate_pipeline()
@@ -729,16 +750,16 @@ class ServerArgs(DisaggArgsMixin):
ring_unspecified = self.ring_degree is None
cfg_unspecified = self.enable_cfg_parallel is None
if current_platform.is_cpu() and (self.tp_size or 1) > 1:
if self.tp_size is None:
self.tp_size = 1
if current_platform.is_cpu() and self.tp_size > 1:
# CPU platform reuse num_gpus to represent num cpu numa nodes as devices
self.num_gpus = self.tp_size
if self.hsdp_shard_dim is None:
self.hsdp_shard_dim = self.num_gpus
if self.tp_size is None:
self.tp_size = 1
# --cfg-parallel-size takes precedence over --enable-cfg-parallel bool.
if self.cfg_parallel_degree is not None:
if self.cfg_parallel_degree == 1:
@@ -983,7 +1004,9 @@ class ServerArgs(DisaggArgsMixin):
configure_logger(server_args=self)
# Convert string disagg_role to enum (from CLI/config)
convert_disagg_role_string(self.__dict__)
if isinstance(self.disagg_role, str):
self.disagg_role = RoleType.from_string(self.disagg_role)
self.gpu_ids = normalize_gpu_ids(self.gpu_ids)
# 1. adjust parameters
self._adjust_parameters()
@@ -1099,6 +1122,23 @@ class ServerArgs(DisaggArgsMixin):
default=ServerArgs.num_gpus,
help="The number of GPUs to use.",
)
parser.add_argument(
"--base-gpu-id",
type=int,
default=ServerArgs.base_gpu_id,
help="The starting GPU ID for this instance. Used with --disagg-role "
"to place role instances on specific GPUs without CUDA_VISIBLE_DEVICES.",
)
parser.add_argument(
"--gpu-ids",
nargs="+",
default=None,
help=(
"Physical GPU IDs for this instance, e.g. --gpu-ids 0 1 6 7 "
"or --gpu-ids 0,1,6,7. Overrides --base-gpu-id for standalone "
"disagg roles."
),
)
parser.add_argument(
"--tp-size",
type=int,
@@ -1170,8 +1210,7 @@ class ServerArgs(DisaggArgsMixin):
"Increase this value if you encounter 'Connection closed by peer' errors after the service is idle. ",
)
# Disaggregated diffusion args (defined in disagg_args.py)
add_disagg_cli_args(parser)
ServerArgs.add_disagg_cli_args(parser)
# Prompt text file for batch processing
parser.add_argument(
@@ -1771,7 +1810,8 @@ class ServerArgs(DisaggArgsMixin):
kwargs["backend"] = Backend.from_string(kwargs["backend"])
# Convert disagg_role string to enum if necessary
convert_disagg_role_string(kwargs)
if "disagg_role" in kwargs and isinstance(kwargs["disagg_role"], str):
kwargs["disagg_role"] = RoleType.from_string(kwargs["disagg_role"])
kwargs["pipeline_config"] = PipelineConfig.from_kwargs(kwargs)
kwargs["_explicit_arg_names"] = explicit_arg_names
@@ -0,0 +1,242 @@
# SPDX-License-Identifier: Apache-2.0
"""Disaggregated diffusion server argument helpers."""
from __future__ import annotations
from typing import ClassVar, Literal
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
from sglang.multimodal_gen.runtime.utils.common import (
format_tcp_endpoint,
parse_tcp_host_port,
)
from sglang.multimodal_gen.utils import FlexibleArgumentParser
class DisaggServerArgsMixin:
DISAGG_RESULT_PORT_OFFSETS: ClassVar[dict[RoleType, int]] = {
RoleType.ENCODER: 1,
RoleType.DENOISER: 2,
RoleType.DECODER: 3,
}
def get_role_parallelism(self, role_type: RoleType) -> dict[str, int | None]:
_none = {
"tp_size": None,
"sp_degree": None,
"ulysses_degree": None,
"ring_degree": None,
}
if role_type == RoleType.ENCODER:
return {**_none, "tp_size": self.encoder_tp}
if role_type == RoleType.DENOISER:
return {
"tp_size": self.denoiser_tp,
"sp_degree": self.denoiser_sp,
"ulysses_degree": self.denoiser_ulysses,
"ring_degree": self.denoiser_ring,
}
if role_type == RoleType.DECODER:
return {**_none, "sp_degree": self.decoder_sp}
return _none
def derive_pool_result_endpoint(self) -> str:
host, base_port = parse_tcp_host_port(
self.disagg_server_addr, "disagg_server_addr"
)
role = (
self.disagg_role
if isinstance(self.disagg_role, RoleType)
else RoleType.from_string(self.disagg_role)
)
try:
offset = self.DISAGG_RESULT_PORT_OFFSETS[role]
except KeyError as exc:
raise ValueError(
"pool result endpoints are only defined for encoder, denoiser, "
f"and decoder roles, got {role.value!r}"
) from exc
return format_tcp_endpoint(host, base_port + offset, "pool_result_endpoint")
def derive_pool_work_endpoint(self) -> str:
return format_tcp_endpoint("0.0.0.0", self.scheduler_port, "pool_work_endpoint")
def derive_pool_control_endpoint(self) -> str:
return format_tcp_endpoint(
"0.0.0.0", self.scheduler_port + 1, "pool_control_endpoint"
)
def derive_pool_control_advertised_endpoint(self) -> str:
host = self.host or self.disagg_p2p_hostname or "127.0.0.1"
if host == "0.0.0.0":
host = self.disagg_p2p_hostname or "127.0.0.1"
return format_tcp_endpoint(
host, self.scheduler_port + 1, "pool_control_advertised_endpoint"
)
def resolved_role_device(self) -> Literal["cpu", "cuda"]:
if self.disagg_role_device == "auto":
return "cpu" if self.num_gpus <= 0 else "cuda"
return self.disagg_role_device
@classmethod
def add_disagg_cli_args(cls, parser: FlexibleArgumentParser) -> None:
role_default = (
cls.disagg_role.value
if isinstance(cls.disagg_role, RoleType)
else cls.disagg_role
)
parser.add_argument(
"--disagg-role",
type=str,
default=role_default,
choices=RoleType.choices(),
help="Role for disaggregated pipeline.",
)
parser.add_argument(
"--disagg-timeout",
type=int,
default=cls.disagg_timeout,
help="Timeout in seconds for pending disagg requests. "
f"Default: {cls.disagg_timeout}.",
)
parser.add_argument(
"--disagg-downstream-wait-timeout",
type=int,
default=cls.disagg_downstream_wait_timeout,
help="Timeout in seconds while waiting for a downstream role slot. "
f"Default: {cls.disagg_downstream_wait_timeout}.",
)
parser.add_argument(
"--disagg-dispatch-policy",
type=str,
default=cls.disagg_dispatch_policy,
choices=["round_robin", "max_free_slots"],
help="Dispatch policy for pool mode disagg routing.",
)
parser.add_argument(
"--disagg-instance-id",
type=int,
default=cls.disagg_instance_id,
help="Stable per-role instance ID used by DiffusionServer registration.",
)
parser.add_argument(
"--disagg-max-slots-per-instance",
type=int,
default=cls.disagg_max_slots_per_instance,
help="Maximum concurrent transfer/computation slots tracked per instance.",
)
parser.add_argument(
"--disagg-transfer-redundancy",
type=float,
default=cls.disagg_transfer_redundancy,
help="Redundancy factor used when sizing transfer buffers from warmup.",
)
parser.add_argument(
"--disagg-role-device",
type=str,
default=cls.disagg_role_device,
choices=["auto", "cpu", "cuda"],
help=(
"Per-role device override. 'cpu' is intended for same-machine "
"encoder roles."
),
)
parser.add_argument(
"--disagg-transfer-backend",
type=str,
default=cls.disagg_transfer_backend,
choices=["auto", "mock", "mooncake"],
help="Transfer backend for multimodal diffusion disaggregation.",
)
parser.add_argument(
"--disagg-transfer-pool-size",
type=int,
default=cls.disagg_transfer_pool_size,
help="Size of the P2P transfer buffer pool in bytes.",
)
parser.add_argument(
"--disagg-transfer-pin-memory",
type=str,
default=cls.disagg_transfer_pin_memory,
choices=["auto", "off", "required"],
help="CUDA host-register same-host shared-memory transfer buffers.",
)
parser.add_argument(
"--disagg-p2p-hostname",
type=str,
default=cls.disagg_p2p_hostname,
help="Hostname for P2P transfer engine.",
)
parser.add_argument(
"--disagg-ib-device",
type=str,
default=cls.disagg_ib_device,
help="InfiniBand device for P2P RDMA transfers.",
)
parser.add_argument(
"--disagg-server-addr",
type=str,
default=cls.disagg_server_addr,
help="DiffusionServer head node address for per-role launch mode.",
)
parser.add_argument(
"--encoder-urls",
type=str,
default=cls.encoder_urls,
help="Encoder instance work endpoints for DiffusionServer head mode.",
)
parser.add_argument(
"--denoiser-urls",
type=str,
default=cls.denoiser_urls,
help="Denoiser instance work endpoints for DiffusionServer head mode.",
)
parser.add_argument(
"--decoder-urls",
type=str,
default=cls.decoder_urls,
help="Decoder instance work endpoints for DiffusionServer head mode.",
)
parser.add_argument(
"--encoder-tp",
type=int,
default=cls.encoder_tp,
help="Tensor parallelism for encoder role.",
)
parser.add_argument(
"--denoiser-tp",
type=int,
default=cls.denoiser_tp,
help="Tensor parallelism for denoiser role.",
)
parser.add_argument(
"--denoiser-sp",
type=int,
default=cls.denoiser_sp,
help="Sequence parallelism for denoiser role.",
)
parser.add_argument(
"--denoiser-ulysses",
type=int,
default=cls.denoiser_ulysses,
help="Ulysses SP degree for denoiser role.",
)
parser.add_argument(
"--denoiser-ring",
type=int,
default=cls.denoiser_ring,
help="Ring SP degree for denoiser role.",
)
parser.add_argument(
"--decoder-sp",
type=int,
default=cls.decoder_sp,
help="Sequence parallelism for decoder role.",
)
parser.add_argument(
"--decoder-tp",
type=int,
default=cls.decoder_tp,
help="Deprecated alias for --decoder-sp.",
)
@@ -9,6 +9,7 @@ import socket
import sys
import threading
from functools import lru_cache
from typing import Any
import psutil
import torch
@@ -78,6 +79,73 @@ def is_valid_ipv6_address(address: str) -> bool:
return False
def normalize_gpu_ids(gpu_ids: Any) -> list[int] | None:
if gpu_ids is None:
return None
if isinstance(gpu_ids, str):
values = [gpu_ids]
else:
values = list(gpu_ids)
tokens: list[str] = []
for value in values:
tokens.extend(part for part in str(value).replace(",", " ").split() if part)
if not tokens:
return []
parsed: list[int] = []
for token in tokens:
try:
gpu_id = int(token)
except ValueError as exc:
raise ValueError(
f"--gpu-ids contains a non-integer GPU id: {token}"
) from exc
if gpu_id < 0:
raise ValueError(f"--gpu-ids GPU ids must be non-negative: {gpu_id}")
parsed.append(gpu_id)
if len(set(parsed)) != len(parsed):
raise ValueError(f"--gpu-ids contains duplicate GPU ids: {parsed}")
return parsed
def parse_tcp_host_port(value: str | None, field_name: str) -> tuple[str, int]:
if value is None or not str(value).strip():
raise ValueError(f"{field_name} is required")
addr = str(value).strip()
if addr.startswith("tcp://"):
addr = addr[len("tcp://") :]
try:
host, port_str = addr.rsplit(":", 1)
except ValueError as exc:
raise ValueError(
f"{field_name} must be formatted as tcp://host:port or host:port"
) from exc
host = host.strip()
port_str = port_str.strip()
if not host or not port_str:
raise ValueError(f"{field_name} must include both host and port: {value!r}")
try:
port = int(port_str)
except ValueError as exc:
raise ValueError(f"{field_name} port must be an integer: {port_str}") from exc
if port < 0 or port > 65535:
raise ValueError(f"{field_name} port must be between 0 and 65535: {port}")
return host, port
def format_tcp_endpoint(host: str, port: int, field_name: str) -> str:
if port < 0 or port > 65535:
raise ValueError(f"{field_name} port must be between 0 and 65535: {port}")
return f"tcp://{host}:{port}"
def configure_ipv6(dist_init_addr):
addr = dist_init_addr
end = addr.find("]")
@@ -3,6 +3,7 @@ import os
import sys
import tempfile
import unittest
from contextlib import contextmanager
from unittest.mock import patch
from sglang.multimodal_gen.configs.models.fsdp import (
@@ -38,12 +39,62 @@ from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.utils import FlexibleArgumentParser
@contextmanager
def _mock_cuda_platform(
*,
memory_gb: int = 80,
available_memory_gb: int | dict[int, int] | None = None,
):
def get_available_gpu_memory(device_id=0, **_kwargs):
if isinstance(available_memory_gb, dict):
return available_memory_gb[device_id]
if available_memory_gb is not None:
return available_memory_gb
return memory_gb
with (
patch(
"sglang.multimodal_gen.runtime.platforms.current_platform.is_cpu",
return_value=False,
),
patch(
"sglang.multimodal_gen.runtime.platforms.current_platform.is_mps",
return_value=False,
),
patch(
"sglang.multimodal_gen.runtime.platforms.current_platform.is_cuda",
return_value=True,
),
patch(
"sglang.multimodal_gen.runtime.platforms.current_platform.get_device_total_memory",
return_value=memory_gb * 1024**3,
),
patch(
"sglang.multimodal_gen.runtime.platforms.current_platform.get_available_gpu_memory",
side_effect=get_available_gpu_memory,
),
patch(
"sglang.multimodal_gen.runtime.platforms.current_platform.enable_dit_layerwise_offload_for_wan_by_default",
return_value=True,
),
):
yield
def _from_dict_without_model_resolution(
kwargs, pipeline_config: PipelineConfig | None = None
):
pipeline_config = pipeline_config or QwenImagePipelineConfig()
with (
patch.object(PipelineConfig, "from_kwargs", return_value=pipeline_config),
_mock_cuda_platform(),
):
return ServerArgs.from_dict(kwargs)
class TestServerArgsPathExpansion(unittest.TestCase):
def _from_dict_without_model_resolution(self, kwargs):
with patch.object(
PipelineConfig, "from_kwargs", return_value=QwenImagePipelineConfig()
):
return ServerArgs.from_dict(kwargs)
return _from_dict_without_model_resolution(kwargs)
def test_tilde_model_path_is_expanded(self):
args = self._from_dict_without_model_resolution(
@@ -1306,10 +1357,7 @@ class TestPerRoleParallelism(unittest.TestCase):
"""Test per-role parallelism args and get_role_parallelism helper."""
def _from_dict(self, kwargs):
with patch.object(
PipelineConfig, "from_kwargs", return_value=QwenImagePipelineConfig()
):
return ServerArgs.from_dict(kwargs)
return _from_dict_without_model_resolution(kwargs)
def test_defaults_are_none(self):
args = self._from_dict({"model_path": "/fake"})
@@ -1351,15 +1399,34 @@ class TestPerRoleParallelism(unittest.TestCase):
self.assertEqual(par["ring_degree"], 2)
def test_decoder_overrides(self):
args = self._from_dict({"model_path": "/fake", "decoder_tp": 2})
args = self._from_dict({"model_path": "/fake", "decoder_sp": 2})
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
par = args.get_role_parallelism(RoleType.DECODER)
self.assertEqual(par["tp_size"], 2)
self.assertIsNone(par["sp_degree"])
self.assertIsNone(par["tp_size"])
self.assertEqual(par["sp_degree"], 2)
self.assertIsNone(par["ulysses_degree"])
self.assertIsNone(par["ring_degree"])
def test_decoder_tp_is_alias_of_decoder_sp(self):
args = self._from_dict({"model_path": "/fake", "decoder_tp": 2})
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
self.assertEqual(args.decoder_sp, 2)
par = args.get_role_parallelism(RoleType.DECODER)
self.assertIsNone(par["tp_size"])
self.assertEqual(par["sp_degree"], 2)
def test_conflicting_decoder_tp_and_decoder_sp_raise(self):
with self.assertRaisesRegex(ValueError, "decoder_tp is deprecated"):
self._from_dict(
{
"model_path": "/fake",
"decoder_tp": 2,
"decoder_sp": 4,
}
)
def test_monolithic_returns_all_none(self):
args = self._from_dict({"model_path": "/fake", "encoder_tp": 2})
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
@@ -1375,14 +1442,72 @@ class TestPerRoleParallelism(unittest.TestCase):
"model_path": "/fake",
"encoder_tp": 1,
"denoiser_tp": 2,
"decoder_tp": 4,
"decoder_sp": 4,
}
)
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
self.assertEqual(args.get_role_parallelism(RoleType.ENCODER)["tp_size"], 1)
self.assertEqual(args.get_role_parallelism(RoleType.DENOISER)["tp_size"], 2)
self.assertEqual(args.get_role_parallelism(RoleType.DECODER)["tp_size"], 4)
self.assertEqual(args.get_role_parallelism(RoleType.DECODER)["sp_degree"], 4)
def test_disagg_args_import_path_stays_compatible(self):
from sglang.multimodal_gen.runtime.disaggregation import disagg_args
from sglang.multimodal_gen.runtime.server_args_disagg import (
DisaggServerArgsMixin,
)
self.assertIs(disagg_args.DisaggArgsMixin, DisaggServerArgsMixin)
self.assertIs(
disagg_args.DISAGG_RESULT_PORT_OFFSETS,
DisaggServerArgsMixin.DISAGG_RESULT_PORT_OFFSETS,
)
def test_gpu_ids_normalize_lists_and_commas(self):
args = self._from_dict({"model_path": "/fake", "gpu_ids": ["0,1", "6", "7 8"]})
self.assertEqual(args.gpu_ids, [0, 1, 6, 7, 8])
def test_gpu_ids_reject_duplicates(self):
with self.assertRaisesRegex(ValueError, "duplicate GPU ids"):
self._from_dict({"model_path": "/fake", "gpu_ids": ["0,1", "1"]})
def test_pool_endpoints_use_role_and_scheduler_ports(self):
args = self._from_dict(
{
"model_path": "/fake",
"disagg_role": "denoiser",
"disagg_server_addr": "tcp://127.0.0.1:30000",
"scheduler_port": 5600,
"host": "0.0.0.0",
"disagg_p2p_hostname": "10.0.0.7",
}
)
self.assertEqual(args.derive_pool_result_endpoint(), "tcp://127.0.0.1:30002")
self.assertEqual(
args.derive_pool_work_endpoint(),
f"tcp://0.0.0.0:{args.scheduler_port}",
)
self.assertEqual(
args.derive_pool_control_endpoint(),
f"tcp://0.0.0.0:{args.scheduler_port + 1}",
)
self.assertEqual(
args.derive_pool_control_advertised_endpoint(),
f"tcp://10.0.0.7:{args.scheduler_port + 1}",
)
def test_pool_result_endpoint_validates_addr_and_role(self):
args = self._from_dict({"model_path": "/fake", "disagg_server_addr": "bad"})
with self.assertRaisesRegex(ValueError, "disagg_server_addr must be"):
args.derive_pool_result_endpoint()
args = self._from_dict(
{"model_path": "/fake", "disagg_server_addr": "127.0.0.1:30000"}
)
with self.assertRaisesRegex(ValueError, "only defined for encoder"):
args.derive_pool_result_endpoint()
def test_cli_args_parsed(self):
"""Per-role parallelism args are parsed from CLI."""
@@ -1401,6 +1526,8 @@ class TestPerRoleParallelism(unittest.TestCase):
"2",
"--encoder-tp",
"1",
"--decoder-sp",
"8",
]
args, unknown = parser.parse_known_args(argv)
self.assertEqual(args.denoiser_tp, 2)
@@ -1408,6 +1535,7 @@ class TestPerRoleParallelism(unittest.TestCase):
self.assertEqual(args.denoiser_ulysses, 2)
self.assertEqual(args.denoiser_ring, 2)
self.assertEqual(args.encoder_tp, 1)
self.assertEqual(args.decoder_sp, 8)
self.assertIsNone(args.decoder_tp)
@@ -1425,7 +1553,10 @@ class TestPipelineResolutionCliOverride(unittest.TestCase):
"768",
]
with patch.object(sys, "argv", ["sglang"] + argv):
with (
patch.object(sys, "argv", ["sglang"] + argv),
_mock_cuda_platform(),
):
args, unknown_args = parser.parse_known_args(argv)
server_args = ServerArgs.from_cli_args(args, unknown_args)
@@ -1441,7 +1572,10 @@ class TestPipelineResolutionCliOverride(unittest.TestCase):
"true",
]
with patch.object(sys, "argv", ["sglang"] + argv):
with (
patch.object(sys, "argv", ["sglang"] + argv),
_mock_cuda_platform(),
):
args, unknown_args = parser.parse_known_args(argv)
server_args = ServerArgs.from_cli_args(args, unknown_args)
@@ -1449,5 +1583,71 @@ class TestPipelineResolutionCliOverride(unittest.TestCase):
self.assertTrue(server_args.disable_autocast)
class TestDisaggTimeoutArgs(unittest.TestCase):
def test_disagg_defaults_match_reviewed_values(self):
args = _from_dict_without_model_resolution({"model_path": "/fake"})
self.assertEqual(args.disagg_max_slots_per_instance, 8)
self.assertEqual(args.disagg_downstream_wait_timeout, 1800)
self.assertEqual(args.disagg_timeout, 3600)
def test_downstream_wait_timeout_cli_arg_is_parsed(self):
parser = FlexibleArgumentParser()
ServerArgs.add_cli_args(parser)
argv = [
"--model-path",
"/fake",
"--disagg-downstream-wait-timeout",
"45",
]
args, _unknown = parser.parse_known_args(argv)
self.assertEqual(args.disagg_downstream_wait_timeout, 45)
def test_disagg_timeout_help_uses_current_defaults(self):
parser = FlexibleArgumentParser()
ServerArgs.add_cli_args(parser)
help_text = parser.format_help()
self.assertIn("Default: 3600.", help_text)
self.assertIn("Default: 1800.", help_text)
def test_disagg_role_alias_cli_arg_is_accepted(self):
parser = FlexibleArgumentParser()
ServerArgs.add_cli_args(parser)
args, _unknown = parser.parse_known_args(
["--model-path", "/fake", "--disagg-role", "denoising"]
)
self.assertEqual(args.disagg_role, "denoising")
def test_disagg_role_alias_normalizes_to_denoiser(self):
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
args = _from_dict_without_model_resolution(
{"model_path": "/fake", "disagg_role": "denoising"}
)
self.assertEqual(args.disagg_role, RoleType.DENOISER)
class TestDisaggTransferBackendArgs(unittest.TestCase):
def test_transfer_backend_defaults_to_auto(self):
args = _from_dict_without_model_resolution({"model_path": "/fake"})
self.assertEqual(args.disagg_transfer_backend, "auto")
def test_transfer_backend_cli_arg_is_parsed(self):
parser = FlexibleArgumentParser()
ServerArgs.add_cli_args(parser)
argv = [
"--model-path",
"/fake",
"--disagg-transfer-backend",
"mock",
]
args, _unknown = parser.parse_known_args(argv)
self.assertEqual(args.disagg_transfer_backend, "mock")
if __name__ == "__main__":
unittest.main()