[diffusion] chore: disagg server args, launch helpers, and warmup utils (#26119)
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user