From 9da998a882c810cad5bb739e691a457fa61eb1f1 Mon Sep 17 00:00:00 2001 From: Yuhao Yang <47235274+yhyang201@users.noreply.github.com> Date: Thu, 16 Apr 2026 23:51:32 +0800 Subject: [PATCH] [diffusion] feat: disaggregated diffusion (#21701) --- docs/diffusion/disaggregation.md | 237 +++ .../runtime/disaggregation/__init__.py | 2 + .../runtime/disaggregation/disagg_args.py | 193 +++ .../runtime/disaggregation/dispatch_policy.py | 165 ++ .../runtime/disaggregation/metrics.py | 133 ++ .../runtime/disaggregation/orchestrator.py | 1076 ++++++++++++ .../runtime/disaggregation/request_state.py | 165 ++ .../runtime/disaggregation/roles.py | 78 + .../runtime/disaggregation/scheduler_mixin.py | 1508 +++++++++++++++++ .../disaggregation/transport/__init__.py | 2 + .../disaggregation/transport/allocator.py | 200 +++ .../disaggregation/transport/buffer.py | 272 +++ .../runtime/disaggregation/transport/codec.py | 198 +++ .../disaggregation/transport/engine.py | 126 ++ .../disaggregation/transport/manager.py | 387 +++++ .../disaggregation/transport/protocol.py | 145 ++ .../runtime/distributed/group_coordinator.py | 4 +- .../runtime/entrypoints/cli/serve.py | 7 +- .../runtime/entrypoints/http_server.py | 26 + .../runtime/entrypoints/utils.py | 7 + .../multimodal_gen/runtime/launch_server.py | 471 ++++- .../runtime/managers/gpu_worker.py | 18 +- .../runtime/managers/scheduler.py | 54 +- .../pipelines_core/composed_pipeline_base.py | 120 +- .../runtime/pipelines_core/stages/base.py | 6 + .../runtime/pipelines_core/stages/decoding.py | 6 + .../pipelines_core/stages/denoising.py | 5 + .../multimodal_gen/runtime/server_args.py | 59 +- .../multimodal_gen/runtime/utils/common.py | 31 +- .../sglang/multimodal_gen/test/run_suite.py | 10 +- .../test/server/test_disagg_server.py | 388 +++++ .../test/unit/test_server_args.py | 129 +- 32 files changed, 6182 insertions(+), 46 deletions(-) create mode 100644 docs/diffusion/disaggregation.md create mode 100644 python/sglang/multimodal_gen/runtime/disaggregation/__init__.py create mode 100644 python/sglang/multimodal_gen/runtime/disaggregation/disagg_args.py create mode 100644 python/sglang/multimodal_gen/runtime/disaggregation/dispatch_policy.py create mode 100644 python/sglang/multimodal_gen/runtime/disaggregation/metrics.py create mode 100644 python/sglang/multimodal_gen/runtime/disaggregation/orchestrator.py create mode 100644 python/sglang/multimodal_gen/runtime/disaggregation/request_state.py create mode 100644 python/sglang/multimodal_gen/runtime/disaggregation/roles.py create mode 100644 python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py create mode 100644 python/sglang/multimodal_gen/runtime/disaggregation/transport/__init__.py create mode 100644 python/sglang/multimodal_gen/runtime/disaggregation/transport/allocator.py create mode 100644 python/sglang/multimodal_gen/runtime/disaggregation/transport/buffer.py create mode 100644 python/sglang/multimodal_gen/runtime/disaggregation/transport/codec.py create mode 100644 python/sglang/multimodal_gen/runtime/disaggregation/transport/engine.py create mode 100644 python/sglang/multimodal_gen/runtime/disaggregation/transport/manager.py create mode 100644 python/sglang/multimodal_gen/runtime/disaggregation/transport/protocol.py create mode 100755 python/sglang/multimodal_gen/test/server/test_disagg_server.py diff --git a/docs/diffusion/disaggregation.md b/docs/diffusion/disaggregation.md new file mode 100644 index 000000000..57bc2c4f1 --- /dev/null +++ b/docs/diffusion/disaggregation.md @@ -0,0 +1,237 @@ +# Disaggregated Diffusion Pipeline + +Split a monolithic text-to-video/image pipeline into independent **Encoder**, **Denoiser**, and **Decoder** roles, each running on its own GPU(s). A central **DiffusionServer** routes requests through the pipeline. + +## Quick Start + +Disaggregation is controlled by a single flag: `--disagg-role`. Each component is launched independently, just like LLM PD disaggregation. + +| `--disagg-role` | What it runs | +|----------------|--------------| +| `monolithic` | (Default) Standard single-server mode | +| `encoder` | All stages with the default `RoleType.ENCODER` affinity: `InputValidationStage`, `TextEncodingStage` (plus `ImageEncodingStage` / `ImageVAEEncodingStage` for image-conditioned pipelines), `LatentPreparationStage`, `TimestepPreparationStage`, and any model-specific "before denoising" stage (e.g. `QwenImageLayeredBeforeDenoisingStage`, `GlmImageBeforeDenoisingStage`). | +| `denoiser` | `DenoisingStage` (and its subclasses: `CausalDMDDenoisingStage`, `DmdDenoisingStage`, `LTX2AVDenoisingStage`, `LTX2RefinementStage`, `Hunyuan3DShapeDenoisingStage`, ...) — the DiT forward loop plus the scheduler stepping it drives. | +| `decoder` | `DecodingStage` (VAE decode) and its subclasses (`LTX2AVDecodingStage`, `HeliosDecodingStage`, ...). | +| `server` | DiffusionServer head node + HTTP server (no GPU) | + +> Each stage declares its role via the `role_affinity` property on `PipelineStage` (default `ENCODER`). When `--disagg-role` is not `monolithic`, the pipeline only instantiates stages whose affinity matches, so the above table is the source of truth for what actually runs in each process. + +### Single-Machine Example (Verified) + +The following commands have been tested end-to-end on an 8×H200 machine with +`Wan-AI/Wan2.1-T2V-1.3B-Diffusers`. Each role runs on a separate GPU via +`--base-gpu-id`; the `server` head node requires no GPU. + +```bash +# Terminal 1: Encoder (GPU 0) +sglang serve --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \ + --disagg-role encoder \ + --disagg-server-addr tcp://127.0.0.1:19655 \ + --scheduler-port 19000 \ + --num-gpus 1 --base-gpu-id 0 + +# Terminal 2: Denoiser (GPU 1) +sglang serve --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \ + --disagg-role denoiser \ + --disagg-server-addr tcp://127.0.0.1:19655 \ + --scheduler-port 19001 \ + --num-gpus 1 --base-gpu-id 1 + +# Terminal 3: Decoder (GPU 2) +sglang serve --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \ + --disagg-role decoder \ + --disagg-server-addr tcp://127.0.0.1:19655 \ + --scheduler-port 19002 \ + --num-gpus 1 --base-gpu-id 2 + +# Terminal 4: DiffusionServer head (no GPU, receives HTTP requests) +sglang serve --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \ + --disagg-role server \ + --encoder-urls "tcp://127.0.0.1:19000" \ + --denoiser-urls "tcp://127.0.0.1:19001" \ + --decoder-urls "tcp://127.0.0.1:19002" \ + --host 0.0.0.0 --port 22000 \ + --scheduler-port 19655 + +# Send request (video generation) +curl http://127.0.0.1:22000/v1/videos \ + -H "Content-Type: application/json" \ + -d '{"model": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers", "prompt": "A curious raccoon exploring a garden, cinematic", "size": "832x480"}' +``` + +> **Tested result (8×H200):** +> Encoder 2.3 s (TextEncoding) → Denoiser 312.8 s (50 steps, layerwise offload) → Decoder 7.1 s (VAE decode). +> Total ~322 s for 81-frame 1024×1024 video. + +> **Tip:** `--base-gpu-id` controls which physical GPU the role uses. +> Encoder and Decoder can share a GPU (e.g. both `--base-gpu-id 0`) to save resources, +> but make sure the combined GPU memory is sufficient. + +### Multi-Machine Example + +The exact same CLI pattern — just replace `127.0.0.1` with actual IPs and add +RDMA flags for direct transfer: + +```bash +# Machine A (10.0.0.1): Encoder +sglang serve --model-path Wan-AI/Wan2.1-T2V-14B-Diffusers \ + --disagg-role encoder \ + --disagg-server-addr tcp://10.0.0.4:19655 \ + --scheduler-port 19000 \ + --num-gpus 1 \ + --disagg-p2p-hostname 10.0.0.1 --disagg-ib-device mlx5_0 + +# Machine B (10.0.0.2): Denoiser (4 GPUs with SP) +sglang serve --model-path Wan-AI/Wan2.1-T2V-14B-Diffusers \ + --disagg-role denoiser \ + --disagg-server-addr tcp://10.0.0.4:19655 \ + --scheduler-port 19001 \ + --num-gpus 4 --denoiser-sp 4 --denoiser-ulysses 2 --denoiser-ring 2 \ + --disagg-p2p-hostname 10.0.0.2 --disagg-ib-device mlx5_0 + +# Machine C (10.0.0.3): Decoder +sglang serve --model-path Wan-AI/Wan2.1-T2V-14B-Diffusers \ + --disagg-role decoder \ + --disagg-server-addr tcp://10.0.0.4:19655 \ + --scheduler-port 19002 \ + --num-gpus 1 \ + --disagg-p2p-hostname 10.0.0.3 --disagg-ib-device mlx5_0 + +# Machine D (10.0.0.4): DiffusionServer head +sglang serve --model-path Wan-AI/Wan2.1-T2V-14B-Diffusers \ + --disagg-role server \ + --encoder-urls "tcp://10.0.0.1:19000" \ + --denoiser-urls "tcp://10.0.0.2:19001" \ + --decoder-urls "tcp://10.0.0.3:19002" \ + --host 0.0.0.0 --port 30000 \ + --scheduler-port 19655 \ + --disagg-dispatch-policy max_free_slots +``` + +> ZMQ handles startup order gracefully — instances and head can start in any order. + +## Multiple Instances per Role + +Use semicolons in `--*-urls` to register multiple instances: + +```bash +# 2 encoders + 2 denoisers (4-GPU SP each) + 1 decoder +sglang serve --model-path ... --disagg-role server \ + --encoder-urls "tcp://10.0.0.1:35000;tcp://10.0.0.2:35000" \ + --denoiser-urls "tcp://10.0.0.3:35000;tcp://10.0.0.4:35000" \ + --decoder-urls "tcp://10.0.0.5:35000" +``` + +## Port Convention + +Result endpoints are derived deterministically from the head node's `--scheduler-port` (default: 5555): + +| Socket | Port | +|--------|------| +| DS frontend (ROUTER) | `scheduler_port` | +| Encoder result (PULL) | `scheduler_port + 1` | +| Denoiser result (PULL) | `scheduler_port + 2` | +| Decoder result (PULL) | `scheduler_port + 3` | + +Role instances derive their result endpoint automatically from `--disagg-server-addr`. No manual endpoint configuration needed. + +## Transfer Mechanism + +Tensor data between roles (encoder→denoiser, denoiser→decoder) is transferred via a P2P transfer engine. The DiffusionServer only routes lightweight control messages (alloc/push/ready); actual tensor data flows directly between instances. + +**mooncake-transfer-engine** is required for disaggregated diffusion. It provides RDMA for direct GPU-to-GPU data movement. + +```bash +pip install mooncake-transfer-engine +``` + +### Transfer Flow + +1. **Sender** (encoder/denoiser) stages tensors: async copy to transfer buffer (GPU or CPU pinned, depending on GPUDirect support), overlapped with metadata JSON serialization. +2. **Sender** sends `transfer_staged` control message to DiffusionServer (metadata only, no tensor data). +3. **DiffusionServer** sends `transfer_alloc` to receiver → receiver allocates buffer slot → replies `transfer_allocated`. +4. **DiffusionServer** sends `transfer_push` to receiver with sender's address info. +5. **Receiver** pulls data via transfer engine (Mooncake RDMA or mock), sends `transfer_ready`. +6. **Receiver** loads tensors async on a dedicated transfer stream, overlapped with the previous request's compute. + +Decoder results (final output) flow back through DiffusionServer as raw ZMQ frames to the HTTP client. + +### RDMA Flags + +| Flag | Default | Description | +|------|---------|-------------| +| `--disagg-p2p-hostname` | `127.0.0.1` | RDMA-reachable hostname/IP of this instance | +| `--disagg-ib-device` | `None` | InfiniBand device (e.g., `mlx5_0`, `mlx5_roce0`) | +| `--disagg-transfer-pool-size` | 256 MiB | Pinned memory pool per instance | + +Set `--disagg-p2p-hostname` to the actual IP on each machine. For multi-machine, `--disagg-ib-device` specifies the RDMA NIC. + +## Per-Role Parallelism + +| Flag | Description | +|------|-------------| +| `--encoder-tp` | Encoder tensor parallelism | +| `--denoiser-tp` / `--denoiser-sp` / `--denoiser-ulysses` / `--denoiser-ring` | Denoiser parallelism | +| `--decoder-tp` | Decoder tensor parallelism | + +If not specified, parallelism is auto-derived from `--num-gpus`. + +## Other Options + +| Flag | Default | Description | +|------|---------|-------------| +| `--disagg-timeout` | `600` | Timeout (seconds) for pending requests | +| `--disagg-dispatch-policy` | `round_robin` | `round_robin` or `max_free_slots` | + +## Python API + +For programmatic single-machine deployment, `launch_pool_disagg_server()` is available: + +```python +from sglang.multimodal_gen.runtime.server_args import ServerArgs +from sglang.multimodal_gen.runtime.launch_server import launch_pool_disagg_server + +server_args = ServerArgs.from_kwargs( + model_path="Wan-AI/Wan2.1-T2V-14B-Diffusers", + denoiser_sp=4, denoiser_ulysses=2, denoiser_ring=2, + disagg_ib_device="mlx5_0", +) + +launch_pool_disagg_server( + server_args, + encoder_gpus=[[0]], + denoiser_gpus=[[1, 2, 3, 4], [5, 6, 7, 8]], + decoder_gpus=[[0]], +) +``` + +## Architecture + +``` +Client ─── HTTP (port 30000) ──► FastAPI Server + │ + ▼ + DiffusionServer (ROUTER, scheduler_port) + ┌───────┼───────┐ + PUSH work │ │ │ PUSH work + ▼ │ ▼ + Encoder[0..N] │ Decoder[0..K] + │ │ ▲ + P2P tensor │ │ │ P2P tensor + transfer ▼ │ │ transfer + Denoiser[0..M] ─────┘ + │ + PULL results ◄────┘ (decoder → DS → client) +``` + +### Request State Machine + +``` +PENDING → ENCODER_WAITING → ENCODER_RUNNING → ENCODER_DONE + │ + DENOISING_WAITING → DENOISING_RUNNING → DENOISING_DONE + │ + DECODER_WAITING → DECODER_RUNNING → DONE +``` + +Any state can transition to `FAILED` or `TIMED_OUT`. diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/__init__.py b/python/sglang/multimodal_gen/runtime/disaggregation/__init__.py new file mode 100644 index 000000000..f39a88e11 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/disaggregation/__init__.py @@ -0,0 +1,2 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Disaggregation support for diffusion pipelines.""" diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/disagg_args.py b/python/sglang/multimodal_gen/runtime/disaggregation/disagg_args.py new file mode 100644 index 000000000..07fffbb6d --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/disaggregation/disagg_args.py @@ -0,0 +1,193 @@ +# 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. +""" + +from __future__ import annotations + +import argparse +from typing import TYPE_CHECKING + +from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType + +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 ───────────────────────────────────────────────── + + +def add_disagg_cli_args(parser: argparse.ArgumentParser) -> None: + """Register all disaggregated-diffusion CLI arguments as a group.""" + + 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.", + ) + + # 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).", + ) + + +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"]) diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/dispatch_policy.py b/python/sglang/multimodal_gen/runtime/disaggregation/dispatch_policy.py new file mode 100644 index 000000000..1382a5ef3 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/disaggregation/dispatch_policy.py @@ -0,0 +1,165 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Dispatch policies for multi-instance disaggregated diffusion pipelines.""" + +import abc +import logging +import threading + +logger = logging.getLogger(__name__) + + +class DispatchPolicy(abc.ABC): + def __init__(self, num_instances: int): + if num_instances < 1: + raise ValueError(f"num_instances must be >= 1, got {num_instances}") + self._num_instances = num_instances + + @property + def num_instances(self) -> int: + return self._num_instances + + @abc.abstractmethod + def select(self, active_counts: list[int] | None = None) -> int: ... + + def select_with_capacity(self, free_slots: list[int]) -> int | None: + """Select an instance that has free capacity, or None if all full.""" + if not any(s > 0 for s in free_slots): + return None + return self.select(active_counts=None) + + def record_completion(self, instance_id: int) -> None: + pass + + +class RoundRobin(DispatchPolicy): + def __init__(self, num_instances: int): + super().__init__(num_instances) + self._lock = threading.Lock() + self._next = 0 + + def select(self, active_counts: list[int] | None = None) -> int: + with self._lock: + chosen = self._next + self._next = (self._next + 1) % self._num_instances + return chosen + + def select_with_capacity(self, free_slots: list[int]) -> int | None: + with self._lock: + for _ in range(self._num_instances): + idx = self._next + self._next = (self._next + 1) % self._num_instances + if free_slots[idx] > 0: + return idx + return None + + +class MaxFreeSlotsFirst(DispatchPolicy): + """Dispatch to the instance with the most free slots.""" + + def __init__(self, num_instances: int, max_slots_per_instance: int = 1): + super().__init__(num_instances) + self._max_slots = max_slots_per_instance + self._lock = threading.Lock() + self._tiebreak = 0 + + def select(self, active_counts: list[int] | None = None) -> int: + with self._lock: + if active_counts is None or len(active_counts) != self._num_instances: + chosen = self._tiebreak % self._num_instances + self._tiebreak += 1 + return chosen + + best_id = 0 + best_free = self._max_slots - active_counts[0] + for i in range(1, self._num_instances): + free = self._max_slots - active_counts[i] + if free > best_free: + best_free = free + best_id = i + elif free == best_free: + if i == (self._tiebreak % self._num_instances): + best_id = i + + self._tiebreak += 1 + + if best_free <= 0: + logger.warning( + "All %d instances are at capacity (%d slots each), " + "dispatching to instance %d anyway", + self._num_instances, + self._max_slots, + best_id, + ) + + return best_id + + def select_with_capacity(self, free_slots: list[int]) -> int | None: + with self._lock: + best_id = -1 + best_free = 0 + for i in range(self._num_instances): + if free_slots[i] > best_free: + best_free = free_slots[i] + best_id = i + elif free_slots[i] == best_free and best_free > 0: + if i == (self._tiebreak % self._num_instances): + best_id = i + + self._tiebreak += 1 + + if best_id < 0: + return None + return best_id + + +class PoolDispatcher: + """Wraps three independent dispatch policies for encoder/denoiser/decoder pools.""" + + def __init__( + self, + num_encoders: int, + num_denoisers: int, + num_decoders: int, + policy_name: str = "round_robin", + **kwargs, + ): + self.encoder_policy = create_dispatch_policy( + policy_name, num_encoders, **kwargs + ) + self.denoiser_policy = create_dispatch_policy( + policy_name, num_denoisers, **kwargs + ) + self.decoder_policy = create_dispatch_policy( + policy_name, num_decoders, **kwargs + ) + + def select_encoder(self, active_counts: list[int] | None = None) -> int: + return self.encoder_policy.select(active_counts) + + def select_denoiser(self, active_counts: list[int] | None = None) -> int: + return self.denoiser_policy.select(active_counts) + + def select_decoder(self, active_counts: list[int] | None = None) -> int: + return self.decoder_policy.select(active_counts) + + def select_encoder_with_capacity(self, free_slots: list[int]) -> int | None: + return self.encoder_policy.select_with_capacity(free_slots) + + def select_denoiser_with_capacity(self, free_slots: list[int]) -> int | None: + return self.denoiser_policy.select_with_capacity(free_slots) + + def select_decoder_with_capacity(self, free_slots: list[int]) -> int | None: + return self.decoder_policy.select_with_capacity(free_slots) + + +def create_dispatch_policy(name: str, num_instances: int, **kwargs) -> DispatchPolicy: + policies = { + "round_robin": RoundRobin, + "max_free_slots": MaxFreeSlotsFirst, + } + cls = policies.get(name) + if cls is None: + raise ValueError( + f"Unknown dispatch policy '{name}'. Available: {list(policies.keys())}" + ) + return cls(num_instances=num_instances, **kwargs) diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/metrics.py b/python/sglang/multimodal_gen/runtime/disaggregation/metrics.py new file mode 100644 index 000000000..cfe0b9c70 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/disaggregation/metrics.py @@ -0,0 +1,133 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Observability metrics for disaggregated diffusion pipelines.""" + +import threading +import time +from dataclasses import dataclass + + +@dataclass +class _RequestTiming: + start_time: float + stage_start: float = 0.0 + + +@dataclass +class RoleStats: + role: str + requests_completed: int = 0 + requests_failed: int = 0 + requests_in_flight: int = 0 + requests_timed_out: int = 0 + queue_depth: int = 0 + last_latency_s: float = 0.0 + avg_latency_s: float = 0.0 + max_latency_s: float = 0.0 + throughput_rps: float = 0.0 + uptime_s: float = 0.0 + + def to_dict(self) -> dict: + return { + "role": self.role, + "requests_completed": self.requests_completed, + "requests_failed": self.requests_failed, + "requests_in_flight": self.requests_in_flight, + "requests_timed_out": self.requests_timed_out, + "queue_depth": self.queue_depth, + "last_latency_s": round(self.last_latency_s, 4), + "avg_latency_s": round(self.avg_latency_s, 4), + "max_latency_s": round(self.max_latency_s, 4), + "throughput_rps": round(self.throughput_rps, 4), + "uptime_s": round(self.uptime_s, 1), + } + + +class DisaggMetrics: + """Thread-safe metrics collector for a single disagg role.""" + + def __init__(self, role: str): + self._role = role + self._lock = threading.Lock() + self._start_time = time.monotonic() + + self._completed = 0 + self._failed = 0 + self._timed_out = 0 + + self._in_flight: dict[str, _RequestTiming] = {} + + self._last_latency = 0.0 + self._max_latency = 0.0 + self._total_latency = 0.0 + + self._completion_times: list[float] = [] + self._throughput_window_s = 60.0 + + self._queue_depth = 0 + + @property + def role(self) -> str: + return self._role + + def record_request_start(self, request_id: str) -> None: + with self._lock: + self._in_flight[request_id] = _RequestTiming(start_time=time.monotonic()) + + def record_request_complete(self, request_id: str) -> None: + now = time.monotonic() + with self._lock: + timing = self._in_flight.pop(request_id, None) + if timing is not None: + latency = now - timing.start_time + self._last_latency = latency + self._max_latency = max(self._max_latency, latency) + self._total_latency += latency + + self._completed += 1 + self._completion_times.append(now) + self._prune_completion_times(now) + + def record_request_failed(self, request_id: str) -> None: + with self._lock: + self._in_flight.pop(request_id, None) + self._failed += 1 + + def record_request_timeout(self, request_id: str) -> None: + with self._lock: + self._in_flight.pop(request_id, None) + self._timed_out += 1 + + def update_queue_depth(self, depth: int) -> None: + with self._lock: + self._queue_depth = depth + + def snapshot(self) -> RoleStats: + now = time.monotonic() + with self._lock: + self._prune_completion_times(now) + total = self._completed + self._failed + avg_latency = self._total_latency / total if total > 0 else 0.0 + rps = ( + len(self._completion_times) / self._throughput_window_s + if self._completion_times + else 0.0 + ) + + return RoleStats( + role=self._role, + requests_completed=self._completed, + requests_failed=self._failed, + requests_in_flight=len(self._in_flight), + requests_timed_out=self._timed_out, + queue_depth=self._queue_depth, + last_latency_s=self._last_latency, + avg_latency_s=avg_latency, + max_latency_s=self._max_latency, + throughput_rps=rps, + uptime_s=now - self._start_time, + ) + + def _prune_completion_times(self, now: float) -> None: + cutoff = now - self._throughput_window_s + while self._completion_times and self._completion_times[0] < cutoff: + self._completion_times.pop(0) diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/orchestrator.py b/python/sglang/multimodal_gen/runtime/disaggregation/orchestrator.py new file mode 100644 index 000000000..05a158def --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/disaggregation/orchestrator.py @@ -0,0 +1,1076 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Central request router for disaggregated diffusion pipelines.""" + +import json +import logging +import pickle +import threading +import time +from collections import deque +from dataclasses import dataclass + +import zmq + +from sglang.multimodal_gen.runtime.disaggregation.dispatch_policy import ( + PoolDispatcher, +) +from sglang.multimodal_gen.runtime.disaggregation.request_state import ( + RequestState, + RequestTracker, +) +from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType +from sglang.multimodal_gen.runtime.disaggregation.transport.codec import ( + unpack_tensors, +) +from sglang.multimodal_gen.runtime.disaggregation.transport.protocol import ( + TransferAllocMsg, + TransferMsgType, + TransferPushMsg, + TransferReadyMsg, + decode_transfer_msg, + encode_transfer_msg, + is_transfer_message, +) +from sglang.multimodal_gen.runtime.utils.common import get_zmq_socket + +logger = logging.getLogger(__name__) + + +@dataclass +class _EncoderTTAEntry: + request_id: str + client_identity: bytes + payload: bytes + + +@dataclass +class _TransferRequestState: + sender_session_id: str = "" + sender_pool_ptr: int = 0 + sender_slot_offset: int = 0 + data_size: int = 0 + manifest: dict = None + scalar_fields: dict = None + receiver_session_id: str = "" + receiver_pool_ptr: int = 0 + receiver_slot_offset: int = 0 + sender_instance: int = -1 + receiver_instance: int = -1 + prealloc_slot_id: int | None = None + + def __post_init__(self): + if self.manifest is None: + self.manifest = {} + if self.scalar_fields is None: + self.scalar_fields = {} + + +@dataclass +class _RoleTTAEntry: + request_id: str + transfer_state: _TransferRequestState | None = None + + +class DiffusionServer: + """Global pipeline orchestrator for N:M:K disaggregated diffusion. + + Capacity-aware dispatch with FreeBufferSlots per instance and TTA queues. + """ + + def __init__( + self, + frontend_endpoint: str, + encoder_work_endpoints: list[str], + denoiser_work_endpoints: list[str], + decoder_work_endpoints: list[str], + encoder_result_endpoint: str, + denoiser_result_endpoint: str, + decoder_result_endpoint: str, + dispatch_policy_name: str = "round_robin", + timeout_s: float = 600.0, + encoder_capacity: int = 4, + denoiser_capacity: int = 2, + decoder_capacity: int = 4, + p2p_mode: bool = True, + ): + self._frontend_endpoint = frontend_endpoint + self._encoder_work_endpoints = encoder_work_endpoints + self._denoiser_work_endpoints = denoiser_work_endpoints + self._decoder_work_endpoints = decoder_work_endpoints + self._encoder_result_endpoint = encoder_result_endpoint + self._denoiser_result_endpoint = denoiser_result_endpoint + self._decoder_result_endpoint = decoder_result_endpoint + + self._num_encoders = len(encoder_work_endpoints) + self._num_denoisers = len(denoiser_work_endpoints) + self._num_decoders = len(decoder_work_endpoints) + self._timeout_s = timeout_s + + self._tracker = RequestTracker() + self._dispatcher = PoolDispatcher( + num_encoders=self._num_encoders, + num_denoisers=self._num_denoisers, + num_decoders=self._num_decoders, + policy_name=dispatch_policy_name, + ) + + self._context = zmq.Context(io_threads=2) + self._running = False + self._ready = threading.Event() + self._thread: threading.Thread | None = None + + self._pending: dict[str, bytes] = {} # request_id -> client ZMQ identity + self._lock = threading.Lock() + + # FreeBufferSlots per instance + self._encoder_free_slots = [encoder_capacity] * self._num_encoders + self._denoiser_free_slots = [denoiser_capacity] * self._num_denoisers + self._decoder_free_slots = [decoder_capacity] * self._num_decoders + + # TTA queues per role type + self._encoder_tta: deque[_EncoderTTAEntry] = deque() + self._denoiser_tta: deque[_RoleTTAEntry] = deque() + self._decoder_tta: deque[_RoleTTAEntry] = deque() + + self._transfer_mode = p2p_mode + self._transfer_state: dict[str, _TransferRequestState] = {} + + # Per-instance registration: instance_idx -> {session_id, pool_ptr, pool_size} + # Keyed by the same index used to build the PUSH work-socket list + # (i.e. the index into --encoder/denoiser/decoder-urls). The index is + # resolved from the registering instance's work_endpoint so the control + # plane (work PUSH) and the data plane (RDMA session_id / pool_ptr / + # preallocated slots) stay consistent regardless of startup order. + self._encoder_peers: dict[int, dict] = {} + self._denoiser_peers: dict[int, dict] = {} + self._decoder_peers: dict[int, dict] = {} + + # work_endpoint -> index lookup tables, built from the --*-urls args + self._encoder_endpoint_to_idx = { + ep: i for i, ep in enumerate(encoder_work_endpoints) + } + self._denoiser_endpoint_to_idx = { + ep: i for i, ep in enumerate(denoiser_work_endpoints) + } + self._decoder_endpoint_to_idx = { + ep: i for i, ep in enumerate(decoder_work_endpoints) + } + + @property + def tracker(self) -> RequestTracker: + return self._tracker + + @property + def dispatcher(self) -> PoolDispatcher: + return self._dispatcher + + def start(self) -> None: + if self._running: + return + self._running = True + self._thread = threading.Thread( + target=self._event_loop, + name="DiffusionServer", + daemon=True, + ) + self._thread.start() + logger.info( + "DiffusionServer started: frontend=%s, " + "%d encoder(s), %d denoiser(s), %d decoder(s), policy=%s, " + "capacity=(%d/%d/%d)", + self._frontend_endpoint, + self._num_encoders, + self._num_denoisers, + self._num_decoders, + type(self._dispatcher.encoder_policy).__name__, + self._encoder_free_slots[0] if self._encoder_free_slots else 0, + self._denoiser_free_slots[0] if self._denoiser_free_slots else 0, + self._decoder_free_slots[0] if self._decoder_free_slots else 0, + ) + + def wait_ready(self, timeout: float = 30.0) -> bool: + """Block until the event loop has bound all sockets, or *timeout* elapses.""" + return self._ready.wait(timeout=timeout) + + def stop(self) -> None: + self._running = False + if self._thread is not None: + self._thread.join(timeout=5.0) + self._thread = None + + def _event_loop(self) -> None: + frontend, _ = get_zmq_socket( + self._context, zmq.ROUTER, self._frontend_endpoint, bind=True + ) + + encoder_pushes: list[zmq.Socket] = [] + for i, ep in enumerate(self._encoder_work_endpoints): + sock, _ = get_zmq_socket(self._context, zmq.PUSH, ep, bind=False) + encoder_pushes.append(sock) + + denoiser_pushes: list[zmq.Socket] = [] + for i, ep in enumerate(self._denoiser_work_endpoints): + sock, _ = get_zmq_socket(self._context, zmq.PUSH, ep, bind=False) + denoiser_pushes.append(sock) + + decoder_pushes: list[zmq.Socket] = [] + for i, ep in enumerate(self._decoder_work_endpoints): + sock, _ = get_zmq_socket(self._context, zmq.PUSH, ep, bind=False) + decoder_pushes.append(sock) + + encoder_result_pull, _ = get_zmq_socket( + self._context, zmq.PULL, self._encoder_result_endpoint, bind=True + ) + denoiser_result_pull, _ = get_zmq_socket( + self._context, zmq.PULL, self._denoiser_result_endpoint, bind=True + ) + decoder_result_pull, _ = get_zmq_socket( + self._context, zmq.PULL, self._decoder_result_endpoint, bind=True + ) + + poller = zmq.Poller() + poller.register(frontend, zmq.POLLIN) + poller.register(encoder_result_pull, zmq.POLLIN) + poller.register(denoiser_result_pull, zmq.POLLIN) + poller.register(decoder_result_pull, zmq.POLLIN) + + self._encoder_pushes = encoder_pushes + self._denoiser_pushes = denoiser_pushes + self._decoder_pushes = decoder_pushes + self._frontend = frontend + + self._ready.set() + + all_sockets = ( + [frontend, encoder_result_pull, denoiser_result_pull, decoder_result_pull] + + encoder_pushes + + denoiser_pushes + + decoder_pushes + ) + + try: + while self._running: + events = dict(poller.poll(timeout=10)) + + self._handle_timeouts() + + if frontend in events: + self._handle_client_request(frontend) + + if encoder_result_pull in events: + self._handle_role_result(encoder_result_pull, RoleType.ENCODER) + + if denoiser_result_pull in events: + self._handle_role_result(denoiser_result_pull, RoleType.DENOISER) + + if decoder_result_pull in events: + self._handle_role_result(decoder_result_pull, RoleType.DECODER) + + self._drain_all_queues() + + except Exception: + logger.exception("DiffusionServer event loop error") + finally: + for sock in all_sockets: + sock.close() + self._context.destroy(linger=0) + + def _handle_role_result(self, result_pull: zmq.Socket, role: RoleType) -> None: + try: + frames = result_pull.recv_multipart(zmq.NOBLOCK, copy=True) + except zmq.Again: + return + + if is_transfer_message(frames): + self._handle_transfer_result(frames, role) + return + + if role == RoleType.DECODER: + self._handle_decoder_result_frames(frames) + else: + # Non-transfer frames from encoder/denoiser are error results + # sent via send_tensors (e.g., _disagg_error). + self._handle_role_error_frames(frames, role) + + def _handle_role_error_frames(self, frames: list, role: RoleType) -> None: + """Handle non-transfer error results from encoder/denoiser roles.""" + try: + tensor_fields, scalar_fields = unpack_tensors(frames, device="cpu") + except Exception as e: + logger.warning( + "DiffusionServer: failed to unpack non-transfer frames from %s: %s", + role.value, + e, + ) + return + + request_id = scalar_fields.get("request_id") + disagg_error = scalar_fields.get("_disagg_error") + + if request_id and disagg_error: + logger.error( + "DiffusionServer: %s error for %s: %s", + role.value, + request_id, + disagg_error, + ) + self._complete_with_error(request_id, f"{role.value} error: {disagg_error}") + elif request_id: + logger.warning( + "DiffusionServer: non-transfer frames from %s for %s without error", + role.value, + request_id, + ) + else: + logger.warning( + "DiffusionServer: non-transfer frames from %s without request_id", + role.value, + ) + + def _handle_client_request(self, frontend: zmq.Socket) -> None: + try: + parts = frontend.recv_multipart(zmq.NOBLOCK) + except zmq.Again: + return + + if len(parts) < 3: + return + + client_identity = parts[0] + payload = parts[-1] + + try: + reqs = pickle.loads(payload) + except (pickle.UnpicklingError, EOFError): + logger.warning("DiffusionServer: failed to deserialize request") + return + + if not isinstance(reqs, list): + reqs = [reqs] + + req = reqs[0] + + if isinstance(req, dict) or not hasattr(req, "request_id"): + # Send empty reply so REQ socket doesn't hang + try: + frontend.send_multipart( + [client_identity, b"", pickle.dumps({"status": "ignored"})], + zmq.NOBLOCK, + ) + except zmq.Again: + pass + return + + request_id = getattr(req, "request_id", None) + if request_id is None: + request_id = f"ds-{time.monotonic()}" + + try: + self._tracker.submit(request_id) + except ValueError: + logger.warning("DiffusionServer: duplicate request_id %s", request_id) + return + + with self._lock: + self._pending[request_id] = client_identity + + try: + self._tracker.transition(request_id, RequestState.ENCODER_WAITING) + except ValueError: + pass + self._encoder_tta.append( + _EncoderTTAEntry( + request_id=request_id, + client_identity=client_identity, + payload=payload, + ) + ) + logger.debug( + "DiffusionServer: queued %s to encoder_tta", + request_id, + ) + + def _handle_decoder_result_frames(self, frames: list) -> None: + from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import ( + OutputBatch, + ) + + request_id = self._extract_request_id(frames) + if request_id is None: + logger.warning("DiffusionServer: decoder result missing request_id") + return + + logger.debug("DiffusionServer: decoder result %s", request_id) + record = self._tracker.get(request_id) + if record and record.decoder_instance is not None: + self._decoder_free_slots[record.decoder_instance] += 1 + + tensor_fields, scalar_fields = unpack_tensors(frames, device="cpu") + + output_batch = OutputBatch( + output=tensor_fields.get("output"), + audio=tensor_fields.get("audio"), + audio_sample_rate=scalar_fields.get("audio_sample_rate"), + error=scalar_fields.get("error"), + ) + + try: + if output_batch.error: + self._tracker.transition( + request_id, RequestState.FAILED, error=output_batch.error + ) + else: + self._tracker.transition(request_id, RequestState.DONE) + except ValueError: + pass + + with self._lock: + client_identity = self._pending.pop(request_id, None) + + if client_identity is None: + logger.warning( + "DiffusionServer: no pending client for decoder result %s", + request_id, + ) + self._tracker.remove(request_id) + return + + try: + self._frontend.send_multipart( + [client_identity, b"", pickle.dumps(output_batch)] + ) + except zmq.ZMQError as e: + logger.error( + "DiffusionServer: failed to send result for %s: %s", + request_id, + e, + ) + + logger.debug("DiffusionServer: returned result for %s", request_id) + self._transfer_state.pop(request_id, None) + self._tracker.remove(request_id) + + def _dispatch_to_encoder( + self, request_id: str, payload: bytes, encoder_idx: int + ) -> None: + self._encoder_free_slots[encoder_idx] -= 1 + + try: + self._tracker.transition( + request_id, + RequestState.ENCODER_RUNNING, + encoder_instance=encoder_idx, + ) + except ValueError: + pass + + self._encoder_pushes[encoder_idx].send_multipart( + [request_id.encode("utf-8"), payload] + ) + logger.debug( + "DiffusionServer: dispatched %s to encoder[%d] (free=%d)", + request_id, + encoder_idx, + self._encoder_free_slots[encoder_idx], + ) + + def _drain_all_queues(self) -> None: + self._drain_encoder_tta() + self._drain_denoiser_tta() + self._drain_decoder_tta() + + def _drain_encoder_tta(self) -> None: + while self._encoder_tta: + idx = self._dispatcher.select_encoder_with_capacity( + self._encoder_free_slots + ) + if idx is None: + break + entry = self._encoder_tta.popleft() + self._dispatch_to_encoder(entry.request_id, entry.payload, idx) + + def _drain_denoiser_tta(self) -> None: + while self._denoiser_tta: + idx = self._dispatcher.select_denoiser_with_capacity( + self._denoiser_free_slots + ) + if idx is None: + break + entry = self._denoiser_tta.popleft() + self._transfer_dispatch_to_denoiser( + entry.request_id, entry.transfer_state, idx + ) + + def _drain_decoder_tta(self) -> None: + while self._decoder_tta: + idx = self._dispatcher.select_decoder_with_capacity( + self._decoder_free_slots + ) + if idx is None: + break + entry = self._decoder_tta.popleft() + self._transfer_dispatch_to_decoder( + entry.request_id, entry.transfer_state, idx + ) + + def _extract_request_id(self, frames: list) -> str | None: + try: + metadata = json.loads(frames[0]) + return metadata.get("scalar_fields", {}).get("request_id") + except (json.JSONDecodeError, IndexError, TypeError): + return None + + def _complete_with_error(self, request_id: str, error_msg: str) -> None: + from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import ( + OutputBatch, + ) + + logger.error("DiffusionServer: %s — %s", request_id, error_msg) + + try: + self._tracker.transition(request_id, RequestState.FAILED, error=error_msg) + except ValueError: + pass + + with self._lock: + client_identity = self._pending.pop(request_id, None) + + if client_identity is None: + self._tracker.remove(request_id) + return + + error_batch = OutputBatch(error=error_msg) + try: + self._frontend.send_multipart( + [client_identity, b"", pickle.dumps(error_batch)] + ) + except zmq.ZMQError as e: + logger.error( + "DiffusionServer: failed to send error for %s: %s", + request_id, + e, + ) + + self._tracker.remove(request_id) + + def _handle_timeouts(self) -> None: + timed_out = self._tracker.find_timed_out(self._timeout_s) + for request_id in timed_out: + # Free the slot for the timed-out request + record = self._tracker.get(request_id) + if record: + self._free_slot_for_record(record) + + self._complete_with_error( + request_id, + f"DiffusionServer timeout: request {request_id} " + f"not completed within {self._timeout_s}s", + ) + + if timed_out: + timed_set = set(timed_out) + self._encoder_tta = deque( + e for e in self._encoder_tta if e.request_id not in timed_set + ) + self._denoiser_tta = deque( + e for e in self._denoiser_tta if e.request_id not in timed_set + ) + self._decoder_tta = deque( + e for e in self._decoder_tta if e.request_id not in timed_set + ) + + def _free_slot_for_record(self, record) -> None: + if ( + record.state in (RequestState.ENCODER_RUNNING, RequestState.ENCODER_DONE) + and record.encoder_instance is not None + ): + self._encoder_free_slots[record.encoder_instance] += 1 + if ( + record.state + in (RequestState.DENOISING_RUNNING, RequestState.DENOISING_DONE) + and record.denoiser_instance is not None + ): + self._denoiser_free_slots[record.denoiser_instance] += 1 + if ( + record.state == RequestState.DECODER_RUNNING + and record.decoder_instance is not None + ): + self._decoder_free_slots[record.decoder_instance] += 1 + + def _handle_transfer_result(self, frames: list, role: RoleType) -> None: + try: + msg = decode_transfer_msg(frames) + except (ValueError, Exception) as e: + logger.error("DiffusionServer: failed to decode transfer message: %s", e) + return + + msg_type = msg.get("msg_type") + + if msg_type == TransferMsgType.REGISTER: + self._handle_transfer_register(msg) + elif msg_type == TransferMsgType.STAGED: + self._handle_transfer_staged(msg) + elif msg_type == TransferMsgType.ALLOCATED: + self._handle_transfer_allocated(msg) + elif msg_type == TransferMsgType.PUSHED: + self._handle_transfer_pushed(msg) + elif msg_type == TransferMsgType.DONE: + self._handle_transfer_done(msg, role) + else: + logger.warning("DiffusionServer: unknown transfer msg_type=%s", msg_type) + + def _handle_transfer_register(self, msg: dict) -> None: + try: + role = RoleType.from_string(msg.get("role", "")) + except ValueError: + logger.warning( + "DiffusionServer transfer: unknown role in register: %s", + msg.get("role"), + ) + return + + work_endpoint = msg.get("work_endpoint", "") + if role == RoleType.ENCODER: + endpoint_to_idx = self._encoder_endpoint_to_idx + peers = self._encoder_peers + elif role == RoleType.DENOISER: + endpoint_to_idx = self._denoiser_endpoint_to_idx + peers = self._denoiser_peers + elif role == RoleType.DECODER: + endpoint_to_idx = self._decoder_endpoint_to_idx + peers = self._decoder_peers + else: + logger.warning( + "DiffusionServer transfer: unsupported role in register: %s", role + ) + return + + idx = endpoint_to_idx.get(work_endpoint) + if idx is None: + # Fail loudly: without a URL match, the control plane (work PUSH) + # and data plane (RDMA dest) would drift silently. + logger.error( + "DiffusionServer transfer: register for role=%s with unknown " + "work_endpoint=%r (known=%s); dropping registration", + role.value, + work_endpoint, + list(endpoint_to_idx.keys()), + ) + return + + info = { + "session_id": msg.get("session_id", ""), + "pool_ptr": msg.get("pool_ptr", 0), + "pool_size": msg.get("pool_size", 0), + "work_endpoint": work_endpoint, + } + prealloc = msg.get("preallocated_slots", []) + info["free_preallocated_slots"] = list(prealloc) + peers[idx] = info + + logger.info( + "DiffusionServer transfer: registered %s[%d] work_endpoint=%s " + "session=%s pool_ptr=%#x prealloc=%d", + role, + idx, + work_endpoint, + info["session_id"], + info["pool_ptr"], + len(prealloc), + ) + + def _handle_transfer_staged(self, msg: dict) -> None: + request_id = msg["request_id"] + logger.debug("DiffusionServer transfer: encoder staged %s", request_id) + record = self._tracker.get(request_id) + encoder_idx = record.encoder_instance if record else 0 + + p2p = _TransferRequestState( + sender_session_id=msg.get("session_id", ""), + sender_pool_ptr=msg.get("pool_ptr", 0), + sender_slot_offset=msg.get("slot_offset", 0), + data_size=msg.get("data_size", 0), + manifest=msg.get("manifest", {}), + scalar_fields=msg.get("scalar_fields", {}), + sender_instance=encoder_idx, + ) + self._transfer_state[request_id] = p2p + + # Encoder slot freed later in _handle_transfer_pushed after RDMA completes + try: + self._tracker.transition(request_id, RequestState.ENCODER_DONE) + except ValueError: + pass + + try: + self._tracker.transition(request_id, RequestState.DENOISING_WAITING) + except ValueError: + pass + self._denoiser_tta.append( + _RoleTTAEntry(request_id=request_id, transfer_state=p2p) + ) + + def _try_fast_path_push( + self, + request_id: str, + p2p: _TransferRequestState, + receiver_peer_info: dict, + sender_pushes: list, + receiver_role_label: str, + receiver_idx: int, + ) -> bool: + """Try to dispatch via a pre-allocated receive slot (fast path). + + If the receiver already registered a free prealloc slot large enough + for this transfer, claim it and send a ``TransferPushMsg`` directly + to the sender so RDMA can start immediately. Returns True when the + fast path is used; False when the caller must fall back to the + round-trip alloc path. + """ + free_slots = receiver_peer_info.get("free_preallocated_slots", []) + if not (free_slots and free_slots[0].get("size", 0) >= p2p.data_size): + return False + + slot_info = free_slots.pop(0) + p2p.receiver_session_id = receiver_peer_info.get("session_id", "") + p2p.receiver_pool_ptr = receiver_peer_info.get("pool_ptr", 0) + p2p.receiver_slot_offset = slot_info["offset"] + p2p.prealloc_slot_id = slot_info.get("slot_id") + + push_msg = TransferPushMsg( + request_id=request_id, + dest_session_id=p2p.receiver_session_id, + dest_addr=slot_info["addr"], + transfer_size=p2p.data_size, + ) + sender_pushes[p2p.sender_instance].send_multipart(encode_transfer_msg(push_msg)) + logger.debug( + "DiffusionServer transfer: fast-path push to %s[%d] for %s " + "(prealloc slot %s, %d bytes)", + receiver_role_label, + receiver_idx, + request_id, + slot_info.get("slot_id"), + p2p.data_size, + ) + return True + + def _send_slow_path_alloc( + self, + request_id: str, + p2p: _TransferRequestState, + receiver_pushes: list, + receiver_idx: int, + source_role: str, + ) -> None: + """Ask the receiver to allocate a slot (slow path). + + Used when the receiver has no free prealloc slot large enough. The + receiver will respond with ``transfer_allocated``; see + :meth:`_handle_transfer_allocated`. + """ + alloc_msg = TransferAllocMsg( + request_id=request_id, + data_size=p2p.data_size, + source_role=source_role, + ) + receiver_pushes[receiver_idx].send_multipart(encode_transfer_msg(alloc_msg)) + + def _transfer_dispatch_to_denoiser( + self, request_id: str, p2p: _TransferRequestState, denoiser_idx: int + ) -> None: + self._denoiser_free_slots[denoiser_idx] -= 1 + p2p.receiver_instance = denoiser_idx + + try: + self._tracker.transition( + request_id, + RequestState.DENOISING_RUNNING, + denoiser_instance=denoiser_idx, + ) + except ValueError: + pass + + peer_info = self._denoiser_peers.get(denoiser_idx, {}) + if not self._try_fast_path_push( + request_id=request_id, + p2p=p2p, + receiver_peer_info=peer_info, + sender_pushes=self._encoder_pushes, + receiver_role_label="denoiser", + receiver_idx=denoiser_idx, + ): + self._send_slow_path_alloc( + request_id=request_id, + p2p=p2p, + receiver_pushes=self._denoiser_pushes, + receiver_idx=denoiser_idx, + source_role="encoder", + ) + + def _handle_transfer_allocated(self, msg: dict) -> None: + request_id = msg["request_id"] + p2p = self._transfer_state.get(request_id) + if p2p is None: + logger.warning( + "DiffusionServer transfer: no state for allocated %s", request_id + ) + return + + p2p.receiver_session_id = msg.get("session_id", "") + p2p.receiver_pool_ptr = msg.get("pool_ptr", 0) + p2p.receiver_slot_offset = msg.get("slot_offset", 0) + + dest_addr = p2p.receiver_pool_ptr + p2p.receiver_slot_offset + push_msg = TransferPushMsg( + request_id=request_id, + dest_session_id=p2p.receiver_session_id, + dest_addr=dest_addr, + transfer_size=p2p.data_size, + ) + + sender_idx = p2p.sender_instance + record = self._tracker.get(request_id) + if record and record.state in ( + RequestState.DECODER_RUNNING, + RequestState.DECODER_WAITING, + ): + self._denoiser_pushes[sender_idx].send_multipart( + encode_transfer_msg(push_msg) + ) + else: + self._encoder_pushes[sender_idx].send_multipart( + encode_transfer_msg(push_msg) + ) + + def _handle_transfer_pushed(self, msg: dict) -> None: + request_id = msg["request_id"] + logger.debug("DiffusionServer transfer: pushed %s", request_id) + p2p = self._transfer_state.get(request_id) + if p2p is None: + logger.warning( + "DiffusionServer transfer: no state for pushed %s", request_id + ) + return + + # Use record state (not sender_idx) to determine sender role, + # because encoder and denoiser can share the same instance index. + record = self._tracker.get(request_id) + if record and record.state in ( + RequestState.DENOISING_RUNNING, + RequestState.DENOISING_WAITING, + RequestState.DENOISING_DONE, + ): + if record.encoder_instance is not None: + self._encoder_free_slots[record.encoder_instance] += 1 + elif record and record.state in ( + RequestState.DECODER_RUNNING, + RequestState.DECODER_WAITING, + ): + if record.denoiser_instance is not None: + self._denoiser_free_slots[record.denoiser_instance] += 1 + + scalar_fields = dict(p2p.scalar_fields) if p2p.scalar_fields else {} + if p2p.prealloc_slot_id is not None: + scalar_fields["_prealloc_slot_id"] = p2p.prealloc_slot_id + ready_msg = TransferReadyMsg( + request_id=request_id, + manifest=p2p.manifest, + slot_offset=p2p.receiver_slot_offset, + scalar_fields=scalar_fields, + ) + + receiver_idx = p2p.receiver_instance + record = self._tracker.get(request_id) + if record and record.state in ( + RequestState.DENOISING_RUNNING, + RequestState.DENOISING_WAITING, + ): + self._denoiser_pushes[receiver_idx].send_multipart( + encode_transfer_msg(ready_msg) + ) + elif record and record.state in ( + RequestState.DECODER_RUNNING, + RequestState.DECODER_WAITING, + ): + self._decoder_pushes[receiver_idx].send_multipart( + encode_transfer_msg(ready_msg) + ) + + logger.debug( + "DiffusionServer transfer: notified receiver for %s (data ready)", + request_id, + ) + + def _recycle_prealloc_slot( + self, p2p: _TransferRequestState, role: RoleType + ) -> None: + if p2p is None or p2p.prealloc_slot_id is None: + return + receiver_idx = p2p.receiver_instance + if role == RoleType.DENOISER: + peer_info = self._denoiser_peers.get(receiver_idx, {}) + elif role == RoleType.DECODER: + peer_info = self._decoder_peers.get(receiver_idx, {}) + else: + return + free_list = peer_info.get("free_preallocated_slots", []) + free_list.append( + { + "offset": p2p.receiver_slot_offset, + "size": p2p.data_size, + "slot_id": p2p.prealloc_slot_id, + "addr": p2p.receiver_pool_ptr + p2p.receiver_slot_offset, + } + ) + p2p.prealloc_slot_id = None + + def _handle_transfer_done(self, msg: dict, role: RoleType) -> None: + request_id = msg.get("request_id", "") + logger.debug( + "DiffusionServer transfer: done %s role=%s", + request_id, + role.value, + ) + error = msg.get("error") + p2p = self._transfer_state.get(request_id) + + if role == RoleType.DENOISER: + record = self._tracker.get(request_id) + + if p2p is not None: + self._recycle_prealloc_slot(p2p, RoleType.DENOISER) + + if error: + if record and record.denoiser_instance is not None: + self._denoiser_free_slots[record.denoiser_instance] += 1 + self._complete_with_error(request_id, f"Denoiser error: {error}") + return + + try: + self._tracker.transition(request_id, RequestState.DENOISING_DONE) + except ValueError: + pass + + if p2p is not None and msg.get("staged_for_decoder"): + # Denoiser slot freed later in _handle_transfer_pushed + p2p.sender_session_id = msg.get("session_id", "") + p2p.sender_pool_ptr = msg.get("pool_ptr", 0) + p2p.sender_slot_offset = msg.get("slot_offset", 0) + p2p.data_size = msg.get("data_size", 0) + p2p.manifest = msg.get("manifest", {}) + p2p.scalar_fields = msg.get("scalar_fields", {}) + p2p.sender_instance = record.denoiser_instance if record else 0 + + try: + self._tracker.transition(request_id, RequestState.DECODER_WAITING) + except ValueError: + pass + self._decoder_tta.append( + _RoleTTAEntry(request_id=request_id, transfer_state=p2p) + ) + else: + if record and record.denoiser_instance is not None: + self._denoiser_free_slots[record.denoiser_instance] += 1 + + elif role == RoleType.DECODER: + if p2p is not None: + self._recycle_prealloc_slot(p2p, RoleType.DECODER) + + record = self._tracker.get(request_id) + if record and record.decoder_instance is not None: + self._decoder_free_slots[record.decoder_instance] += 1 + + if error: + self._complete_with_error(request_id, f"Decoder error: {error}") + else: + try: + self._tracker.transition(request_id, RequestState.DONE) + except ValueError: + pass + + self._transfer_return_to_client_from_msg(request_id, msg) + + self._transfer_state.pop(request_id, None) + + def _transfer_dispatch_to_decoder( + self, request_id: str, p2p: _TransferRequestState, decoder_idx: int + ) -> None: + self._decoder_free_slots[decoder_idx] -= 1 + p2p.receiver_instance = decoder_idx + + try: + self._tracker.transition( + request_id, + RequestState.DECODER_RUNNING, + decoder_instance=decoder_idx, + ) + except ValueError: + pass + + peer_info = self._decoder_peers.get(decoder_idx, {}) + if not self._try_fast_path_push( + request_id=request_id, + p2p=p2p, + receiver_peer_info=peer_info, + sender_pushes=self._denoiser_pushes, + receiver_role_label="decoder", + receiver_idx=decoder_idx, + ): + self._send_slow_path_alloc( + request_id=request_id, + p2p=p2p, + receiver_pushes=self._decoder_pushes, + receiver_idx=decoder_idx, + source_role="denoiser", + ) + + def _transfer_return_to_client_from_msg(self, request_id: str, msg: dict) -> None: + from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import ( + OutputBatch, + ) + + with self._lock: + client_identity = self._pending.pop(request_id, None) + + if client_identity is None: + self._tracker.remove(request_id) + return + + output_batch = OutputBatch(error=msg.get("error")) + + try: + self._frontend.send_multipart( + [client_identity, b"", pickle.dumps(output_batch)] + ) + except zmq.ZMQError as e: + logger.error( + "DiffusionServer transfer: failed to send result for %s: %s", + request_id, + e, + ) + self._tracker.remove(request_id) + + def get_stats(self) -> dict: + with self._lock: + pending_count = len(self._pending) + return { + "role": "diffusion_server", + "transfer_mode": self._transfer_mode, + "num_encoders": self._num_encoders, + "num_denoisers": self._num_denoisers, + "num_decoders": self._num_decoders, + "pending_requests": pending_count, + "dispatch_policy": type(self._dispatcher.encoder_policy).__name__, + "encoder_free_slots": list(self._encoder_free_slots), + "denoiser_free_slots": list(self._denoiser_free_slots), + "decoder_free_slots": list(self._decoder_free_slots), + "encoder_tta_depth": len(self._encoder_tta), + "denoiser_tta_depth": len(self._denoiser_tta), + "decoder_tta_depth": len(self._decoder_tta), + "transfer_active_transfers": len(self._transfer_state), + "encoder_peers": len(self._encoder_peers), + "denoiser_peers": len(self._denoiser_peers), + "decoder_peers": len(self._decoder_peers), + "tracker": self._tracker.snapshot(), + } diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/request_state.py b/python/sglang/multimodal_gen/runtime/disaggregation/request_state.py new file mode 100644 index 000000000..7c80906bc --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/disaggregation/request_state.py @@ -0,0 +1,165 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Request state machine for disaggregated diffusion pipelines.""" + +import enum +import logging +import threading +import time +from dataclasses import dataclass, field + +logger = logging.getLogger(__name__) + + +class RequestState(enum.Enum): + """Lifecycle states for a disagg pipeline request. + + *_WAITING: request queued, awaiting a free buffer slot. + *_RUNNING: request dispatched to a specific instance. + """ + + PENDING = "pending" + ENCODER_WAITING = "encoder_waiting" + ENCODER_RUNNING = "encoder_running" + ENCODER_DONE = "encoder_done" + DENOISING_WAITING = "denoising_waiting" + DENOISING_RUNNING = "denoising_running" + DENOISING_DONE = "denoising_done" + DECODER_WAITING = "decoder_waiting" + DECODER_RUNNING = "decoder_running" + DONE = "done" + FAILED = "failed" + TIMED_OUT = "timed_out" + + +_TERMINAL_STATES = {RequestState.DONE, RequestState.FAILED, RequestState.TIMED_OUT} +_ACTIVE_STATES = set(RequestState) - _TERMINAL_STATES + +# Normal (non-failure) transitions. FAILED and TIMED_OUT are handled +# separately in transition() — any active state can reach them. +_VALID_TRANSITIONS: dict[RequestState, set[RequestState]] = { + RequestState.PENDING: {RequestState.ENCODER_WAITING, RequestState.ENCODER_RUNNING}, + RequestState.ENCODER_WAITING: {RequestState.ENCODER_RUNNING}, + RequestState.ENCODER_RUNNING: {RequestState.ENCODER_DONE}, + RequestState.ENCODER_DONE: { + RequestState.DENOISING_WAITING, + RequestState.DENOISING_RUNNING, + }, + RequestState.DENOISING_WAITING: {RequestState.DENOISING_RUNNING}, + RequestState.DENOISING_RUNNING: {RequestState.DENOISING_DONE}, + RequestState.DENOISING_DONE: { + RequestState.DECODER_WAITING, + RequestState.DECODER_RUNNING, + }, + RequestState.DECODER_WAITING: {RequestState.DECODER_RUNNING}, + RequestState.DECODER_RUNNING: {RequestState.DONE}, +} + + +@dataclass +class RequestRecord: + request_id: str + state: RequestState = RequestState.PENDING + submit_time: float = field(default_factory=time.monotonic) + last_transition_time: float = field(default_factory=time.monotonic) + encoder_instance: int | None = None + denoiser_instance: int | None = None + decoder_instance: int | None = None + error: str | None = None + + def elapsed_s(self) -> float: + return time.monotonic() - self.submit_time + + def is_terminal(self) -> bool: + return self.state in _TERMINAL_STATES + + +class RequestTracker: + """Thread-safe tracker for request state machines.""" + + def __init__(self): + self._lock = threading.Lock() + self._requests: dict[str, RequestRecord] = {} + + def submit(self, request_id: str) -> RequestRecord: + with self._lock: + if request_id in self._requests: + raise ValueError(f"Duplicate request_id: {request_id}") + record = RequestRecord(request_id=request_id) + self._requests[request_id] = record + return record + + def transition( + self, + request_id: str, + new_state: RequestState, + *, + error: str | None = None, + encoder_instance: int | None = None, + denoiser_instance: int | None = None, + decoder_instance: int | None = None, + ) -> RequestRecord: + with self._lock: + record = self._requests.get(request_id) + if record is None: + raise ValueError(f"Unknown request_id: {request_id}") + + old_state = record.state + + if new_state in _TERMINAL_STATES and new_state != RequestState.DONE: + # FAILED / TIMED_OUT: allowed from any active state + if old_state not in _ACTIVE_STATES: + raise ValueError( + f"Cannot transition {request_id} from terminal state " + f"{old_state.value} to {new_state.value}" + ) + elif new_state not in _VALID_TRANSITIONS.get(old_state, set()): + raise ValueError( + f"Invalid transition for {request_id}: " + f"{old_state.value} -> {new_state.value}" + ) + + record.state = new_state + record.last_transition_time = time.monotonic() + if error is not None: + record.error = error + if encoder_instance is not None: + record.encoder_instance = encoder_instance + if denoiser_instance is not None: + record.denoiser_instance = denoiser_instance + if decoder_instance is not None: + record.decoder_instance = decoder_instance + + logger.debug( + "Request %s: %s -> %s", request_id, old_state.value, new_state.value + ) + return record + + def get(self, request_id: str) -> RequestRecord | None: + with self._lock: + return self._requests.get(request_id) + + def remove(self, request_id: str) -> RequestRecord | None: + with self._lock: + return self._requests.pop(request_id, None) + + def find_timed_out(self, timeout_s: float) -> list[str]: + now = time.monotonic() + with self._lock: + return [ + r.request_id + for r in self._requests.values() + if r.state in _ACTIVE_STATES and (now - r.submit_time) > timeout_s + ] + + def snapshot(self) -> dict: + with self._lock: + state_counts = {} + for r in self._requests.values(): + state_counts[r.state.value] = state_counts.get(r.state.value, 0) + 1 + return { + "total": len(self._requests), + "active": sum( + 1 for r in self._requests.values() if not r.is_terminal() + ), + "by_state": state_counts, + } diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/roles.py b/python/sglang/multimodal_gen/runtime/disaggregation/roles.py new file mode 100644 index 000000000..b85b9244b --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/disaggregation/roles.py @@ -0,0 +1,78 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Role definitions for diffusion pipeline disaggregation.""" + +from enum import Enum + +_ROLE_ALIASES = {"denoising": "denoiser"} + + +class RoleType(str, Enum): + MONOLITHIC = "monolithic" + ENCODER = "encoder" + DENOISER = "denoiser" + DECODER = "decoder" + SERVER = "server" # Head node (no GPU, routes requests) + + @classmethod + def from_string(cls, value: str) -> "RoleType": + v = _ROLE_ALIASES.get(value.lower(), value.lower()) + try: + return cls(v) + except ValueError: + raise ValueError( + f"Invalid role: {value}. Must be one of: {', '.join([r.value for r in cls])}" + ) from None + + @classmethod + def choices(cls) -> list[str]: + return [role.value for role in cls] + + +def get_module_role(module_name: str) -> "RoleType | None": + """Classify a module name to its primary role. Returns None for shared modules.""" + encoder_prefixes = ( + "text_encoder", + "tokenizer", + "image_encoder", + "image_processor", + "processor", + "connectors", + ) + if any( + module_name == p or module_name.startswith(p + "_") for p in encoder_prefixes + ): + return RoleType.ENCODER + + denoising_prefixes = ("transformer",) + if any( + module_name == p or module_name.startswith(p + "_") for p in denoising_prefixes + ): + return RoleType.DENOISER + + decoder_prefixes = ("vae", "audio_vae", "video_vae", "vocoder") + if any( + module_name == p or module_name.startswith(p + "_") for p in decoder_prefixes + ): + return RoleType.DECODER + + return None + + +def filter_modules_for_role(module_names: list[str], role: "RoleType") -> list[str]: + """Filter module names to only those needed by the given role.""" + if role in (RoleType.MONOLITHIC, RoleType.SERVER): + return module_names + + filtered = [] + for name in module_names: + module_role = get_module_role(name) + + if module_role is None: + filtered.append(name) + elif module_role == role: + filtered.append(name) + elif role == RoleType.ENCODER and module_role == RoleType.DECODER: + # Encoder also needs VAE for ImageVAEEncoding stages + filtered.append(name) + + return filtered diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py b/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py new file mode 100644 index 000000000..670d0dedb --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py @@ -0,0 +1,1508 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Mixin that adds disaggregated diffusion scheduling to the Scheduler. + +Extracted from scheduler.py to keep the core scheduler lean. +All transfer, compute, and event-loop logic for disaggregated roles +(encoder / denoiser / decoder) lives here. +""" + +from __future__ import annotations + +import dataclasses +import json +import logging +import pickle +import queue +import threading +import time +from typing import TYPE_CHECKING, Any + +import torch +import zmq + +from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams +from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType +from sglang.multimodal_gen.runtime.disaggregation.transport.buffer import ( + TransferTensorBuffer, +) +from sglang.multimodal_gen.runtime.disaggregation.transport.codec import ( + send_tensors, +) +from sglang.multimodal_gen.runtime.disaggregation.transport.engine import ( + create_transfer_engine, +) +from sglang.multimodal_gen.runtime.disaggregation.transport.manager import ( + DiffusionTransferManager, +) +from sglang.multimodal_gen.runtime.disaggregation.transport.protocol import ( + TRANSFER_MAGIC, + TransferAllocatedMsg, + TransferDoneMsg, + TransferMsgType, + TransferPushedMsg, + TransferRegisterMsg, + decode_transfer_msg, + encode_transfer_msg, + is_transfer_message, +) +from sglang.multimodal_gen.runtime.pipelines_core import Req +from sglang.multimodal_gen.runtime.utils.common import get_zmq_socket +from sglang.multimodal_gen.runtime.utils.distributed import broadcast_pyobj +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger + +if TYPE_CHECKING: + from sglang.multimodal_gen.runtime.managers.scheduler import Scheduler + +logger = init_logger(__name__) + +# --------------------------------------------------------------------------- +# Field extraction: split Req into tensors (transfer buffer) and scalars (JSON) +# --------------------------------------------------------------------------- + +# Fields that should never be transferred (non-serializable, internal, or receiver rebuilds) +_EXCLUDE_FIELDS = frozenset( + { + "sampling_params", + "generator", + "modules", + "metrics", + "extra_step_kwargs", + "extra", + "condition_image", + "vae_image", + "pixel_values", + "preprocessed_image", + "image_embeds", + "original_condition_image_size", + "vae_image_sizes", + "output", + "audio", + "audio_sample_rate", + "trajectory_timesteps", + "trajectory_latents", + "trajectory_audio_latents", + "timestep", + "step_index", + "prompt_template", + "max_sequence_length", + } +) + +# Sampling-params fields that should never be transferred across roles: +# - data_type / supported_resolutions: enums / non-JSON classvars reconstructed on the receiver +# - teacache_params: model-specific object, not JSON-safe +# - output_* / save_output / return_*: output-side concerns owned by the decoder role +# +# Everything else on SamplingParams is forwarded automatically via a field-walk +# below; this keeps new request-level features (e.g. Qwen-Image's +# true_cfg_scale, guidance_rescale, cfg_normalization, ...) from silently +# getting dropped just because nobody remembered to add them to a whitelist. +_SAMPLING_PARAMS_EXCLUDE_FIELDS = frozenset( + { + "data_type", + "supported_resolutions", + "teacache_params", + } +) + +_BASE_SP_DEFAULTS: dict[str, Any] = {} +for _f in dataclasses.fields(SamplingParams): + if _f.default is not dataclasses.MISSING: + _BASE_SP_DEFAULTS[_f.name] = _f.default + + +def _is_tensor_like(value) -> bool: + if isinstance(value, torch.Tensor): + return True + if isinstance(value, list) and value and isinstance(value[0], torch.Tensor): + return True + return False + + +def _to_json_serializable(value): + if isinstance(value, torch.Tensor): + return value.tolist() + if isinstance(value, (list, tuple)): + converted = [] + for item in value: + if isinstance(item, torch.Tensor): + converted.append(item.tolist()) + else: + converted.append(item) + return converted + return value + + +def _is_default(value, field_info) -> bool: + if field_info.default is not dataclasses.MISSING: + return value == field_info.default + if field_info.default_factory is not dataclasses.MISSING: + if isinstance(value, (list, dict)) and len(value) == 0: + return True + return False + + +def _extract_extra_fields(extra: dict, scalar_fields: dict) -> None: + """Extract JSON-serializable entries from Req.extra into scalar_fields.""" + for key, value in extra.items(): + if key.startswith("_"): + continue + try: + json.dumps(value) + scalar_fields[f"_extra_{key}"] = value + except (TypeError, ValueError, OverflowError): + pass + + +def extract_transfer_fields(req) -> tuple[dict, dict]: + """Extract all transferable fields from a Req, split into tensors and scalars.""" + tensor_fields = {} + scalar_fields = {} + _debug_transfer = logger.isEnabledFor(logging.DEBUG) + + for f in dataclasses.fields(req): + if f.name in _EXCLUDE_FIELDS: + continue + + value = getattr(req, f.name, None) + if value is None: + continue + if _is_default(value, f): + continue + + if _is_tensor_like(value): + tensor_fields[f.name] = value + else: + try: + scalar_fields[f.name] = _to_json_serializable(value) + except (TypeError, ValueError): + pass + + extra = getattr(req, "extra", None) + if extra: + _extract_extra_fields(extra, scalar_fields) + + sp = getattr(req, "sampling_params", None) + if sp is not None: + # Forward every non-default, JSON-safe SamplingParams field, not a + # narrow whitelist. Previously only a handful of fields were carried + # across roles, which silently dropped per-request config like + # Qwen-Image's true_cfg_scale (and any future feature added to + # SamplingParams). Using a field-walk keeps the disagg boundary + # feature-complete without needing to edit this list. + for f in dataclasses.fields(sp): + name = f.name + if name in _SAMPLING_PARAMS_EXCLUDE_FIELDS: + continue + if name in scalar_fields: + # Req-level field already took precedence (or upstream Req + # explicitly set it). + continue + value = getattr(sp, name, None) + if value is None: + continue + base_default = _BASE_SP_DEFAULTS.get(name, dataclasses.MISSING) + if base_default is not dataclasses.MISSING and value == base_default: + continue + try: + scalar_fields[name] = _to_json_serializable(value) + except (TypeError, ValueError): + pass + + if _debug_transfer: + import torch as _torch + + for _n, _t in tensor_fields.items(): + if isinstance(_t, _torch.Tensor): + _sz = _t.nelement() * _t.element_size() + logger.debug( + "transfer_field %s shape=%s dtype=%s size=%d", + _n, + list(_t.shape), + _t.dtype, + _sz, + ) + elif isinstance(_t, list): + for _i, _ti in enumerate(_t): + if isinstance(_ti, _torch.Tensor): + _sz = _ti.nelement() * _ti.element_size() + logger.debug( + "transfer_field %s[%d] shape=%s dtype=%s size=%d", + _n, + _i, + list(_ti.shape), + _ti.dtype, + _sz, + ) + + return tensor_fields, scalar_fields + + +# --------------------------------------------------------------------------- +# Helpers for broadcasting Req contents across SP/CFG/TP ranks +# --------------------------------------------------------------------------- + +# Sentinel marker key used to distinguish "list of tensors" from a regular +# nested dict when round-tripping through GroupCoordinator.broadcast_tensor_dict +# (which only natively understands tensor / nested-dict values). +_LIST_MARKER_KEY = "__is_list__" + + +def _pack_tensor_fields_for_broadcast(tensor_fields: dict) -> dict: + """Pack ``tensor_fields`` into a structure ``broadcast_tensor_dict`` accepts. + + ``broadcast_tensor_dict`` understands dict-of-tensor values (recursively), + but not list-of-tensor values. Several Req fields (``prompt_embeds``, + ``image_embeds``, ...) are lists of tensors, so we encode each list as a + nested dict whose tensors are keyed by their stringified index, with a + sentinel ``__is_list__`` flag to disambiguate from real nested dicts. + """ + packed: dict = {} + for key, value in tensor_fields.items(): + if isinstance(value, torch.Tensor): + packed[key] = value + elif isinstance(value, list): + sub: dict = {_LIST_MARKER_KEY: True} + for i, item in enumerate(value): + if isinstance(item, torch.Tensor): + sub[str(i)] = item + packed[key] = sub + # Anything else (e.g. None, scalars) is intentionally dropped — the + # scalar_fields broadcast covers non-tensor metadata. + return packed + + +def _unpack_tensor_fields_from_broadcast(packed: dict) -> dict: + """Inverse of :func:`_pack_tensor_fields_for_broadcast`.""" + out: dict = {} + for key, value in packed.items(): + if isinstance(value, dict) and value.get(_LIST_MARKER_KEY) is True: + indexed = [(int(k), v) for k, v in value.items() if k != _LIST_MARKER_KEY] + indexed.sort(key=lambda kv: kv[0]) + out[key] = [v for _, v in indexed] + else: + out[key] = value + return out + + +class SchedulerDisaggMixin: + """Disaggregated diffusion scheduling: transfer, compute, event loops.""" + + # ------------------------------------------------------------------ + # Initialization + # ------------------------------------------------------------------ + + def _init_disagg_state(self: Scheduler, server_args, local_rank: int) -> None: + """Initialize all disaggregation state, sockets, and transfer infrastructure.""" + from sglang.multimodal_gen.runtime.disaggregation.metrics import DisaggMetrics + + self._disagg_role = server_args.disagg_role + self._disagg_timeout_s = float(getattr(server_args, "disagg_timeout", 600)) + self._disagg_metrics = None + self._disagg_mode = getattr(server_args, "disagg_mode", False) + self._pool_work_pull = None + self._pool_result_push = None + self._transfer_manager = None + self._transfer_stream = None + self._rdma_push_queue = None + self._rdma_push_thread = None + self._rdma_push_zmq = None + self._compute_ready_queue = None + self._recv_prefetch_thread = None + + if self._disagg_role != RoleType.MONOLITHIC: + self._disagg_metrics = DisaggMetrics(role=self._disagg_role.value) + device = torch.device(f"cuda:{local_rank}") + self._transfer_stream = torch.cuda.Stream(device=device) + self._init_disagg_sockets() + self._init_disagg_transfer_manager() + + def _init_disagg_sockets(self: Scheduler): + """Initialize ZMQ sockets for disaggregated mode (DiffusionServer-mediated). + + Only rank 0 creates ZMQ sockets. Non-rank-0 processes participate + via NCCL broadcast from rank 0 (see _disagg_recv_work). + """ + if self.gpu_id != 0: + logger.info( + "Pool mode %s rank %d: no ZMQ sockets (non-rank-0)", + self._disagg_role.value.upper(), + self.gpu_id, + ) + return + + sa = self.server_args + + # PULL: receive work from DiffusionServer + self._pool_work_pull, _ = get_zmq_socket( + self.context, + zmq.PULL, + sa.pool_work_endpoint, + bind=True, + max_bind_retries=5, + same_port=True, + ) + # PUSH: send results to DiffusionServer + self._pool_result_push, _ = get_zmq_socket( + self.context, zmq.PUSH, sa.pool_result_endpoint, bind=False + ) + logger.info( + "Disagg %s rank 0: work_pull=%s, result_push=%s", + self._disagg_role.value.upper(), + sa.pool_work_endpoint, + sa.pool_result_endpoint, + ) + + def _init_disagg_transfer_manager(self: Scheduler): + """Initialize TransferManager for transfer mode (rank 0 only). + + Creates a TransferTensorBuffer (pinned memory pool) and a + BaseTransferEngine, then wraps them in a DiffusionTransferManager. + Also sends a transfer_register message to DiffusionServer. + """ + if self.gpu_id != 0: + return + + sa = self.server_args + + # Pool size: configurable, default 256 MiB + pool_size = getattr(sa, "disagg_transfer_pool_size", 256 * 1024 * 1024) + + # Create transfer engine. + # NOTE: self.gpu_id is the role-internal rank (0..num_role_gpus-1), + # not the physical GPU index. In disagg mode with --base-gpu-id > 0, + # the physical device is self.worker.local_rank. Mooncake needs the + # physical index to pin the right NIC and register GPUDirect buffers. + hostname = getattr(sa, "disagg_p2p_hostname", "127.0.0.1") + ib_device = getattr(sa, "disagg_ib_device", None) + physical_gpu_id = self.worker.local_rank + engine = create_transfer_engine( + hostname=hostname, + gpu_id=physical_gpu_id, + ib_device=ib_device, + ) + + # Use GPU buffer when engine supports GPUDirect RDMA, CPU pinned otherwise + device = f"cuda:{physical_gpu_id}" if engine.supports_gpu_direct else "cpu" + buffer = TransferTensorBuffer( + pool_size=pool_size, device=device, role_name=self._disagg_role.value + ) + + # Create transfer manager + self._transfer_manager = DiffusionTransferManager(engine=engine, buffer=buffer) + + # Pre-allocate receive slots for receivers (denoiser/decoder) + self._preallocated_slots: dict[int, object] = {} + preallocated_slot_info = [] + if self._disagg_role in (RoleType.DENOISER, RoleType.DECODER): + capacity = getattr(sa, "disagg_prealloc_slots", 2) + typical_size = 64 * 1024 * 1024 # 64 MiB per slot + for i in range(capacity): + slot = buffer.allocate(typical_size, f"prealloc_{i}") + if slot is not None: + self._preallocated_slots[i] = slot + preallocated_slot_info.append( + { + "offset": slot.offset, + "size": slot.size, + "slot_id": i, + "addr": self._transfer_manager.pool_data_ptr + slot.offset, + } + ) + if preallocated_slot_info: + logger.info( + "Transfer %s: pre-allocated %d receive slots", + self._disagg_role.value.upper(), + len(preallocated_slot_info), + ) + + # Register with DiffusionServer. + # Include our own work_endpoint so DS can key the peer by URL index, + # not by registration order (startup order is not guaranteed to match + # --encoder/denoiser/decoder-urls ordering). + register_msg = TransferRegisterMsg( + role=self._disagg_role.value, + session_id=self._transfer_manager.session_id, + pool_ptr=self._transfer_manager.pool_data_ptr, + pool_size=self._transfer_manager.pool_size, + work_endpoint=sa.pool_work_endpoint, + preallocated_slots=preallocated_slot_info, + ) + self._pool_result_push.send_multipart(encode_transfer_msg(register_msg)) + logger.info( + "Transfer %s: registered with DS (session=%s, pool=%d bytes, prealloc=%d)", + self._disagg_role.value.upper(), + self._transfer_manager.session_id, + pool_size, + len(preallocated_slot_info), + ) + + # RDMA push thread for sender roles (encoder/denoiser) + if self._disagg_role in (RoleType.ENCODER, RoleType.DENOISER): + self._rdma_push_queue = queue.Queue(maxsize=4) + self._rdma_push_zmq, _ = get_zmq_socket( + self.context, + zmq.PUSH, + sa.pool_result_endpoint, + bind=False, + ) + self._rdma_push_thread = threading.Thread( + target=self._rdma_push_loop, + daemon=True, + name=f"rdma-push-{self._disagg_role.value}", + ) + self._rdma_push_thread.start() + logger.info( + "Transfer %s: RDMA push thread started", + self._disagg_role.value.upper(), + ) + + # Recv prefetch thread for receiver roles (denoiser/decoder) + # Rank 0 only (bg thread does ZMQ recv + load; multi-rank gets + # scalar fields via broadcast_pyobj from the main thread). + if self._disagg_role in (RoleType.DENOISER, RoleType.DECODER): + self._compute_ready_queue = queue.Queue(maxsize=4) + self._recv_prefetch_thread = threading.Thread( + target=self._recv_prefetch_loop, + daemon=True, + name=f"recv-prefetch-{self._disagg_role.value}", + ) + self._recv_prefetch_thread.start() + logger.info( + "Transfer %s: recv prefetch thread started", + self._disagg_role.value.upper(), + ) + + # ------------------------------------------------------------------ + # Background threads + # ------------------------------------------------------------------ + + def _rdma_push_loop(self: Scheduler): + """Background thread: execute RDMA push + notify DS. + + Runs push_to_peer (blocking RDMA) on a dedicated thread so the + main event loop can immediately start processing the next request. + """ + role_name = self._disagg_role.value.upper() + while True: + item = self._rdma_push_queue.get() + if item is None: + break # Shutdown signal + request_id, dest_session_id, dest_addr, transfer_size = item + try: + success = self._transfer_manager.push_to_peer( + request_id=request_id, + dest_session_id=dest_session_id, + dest_addr=dest_addr, + transfer_size=transfer_size, + ) + if success: + self._transfer_manager.free_staged(request_id) + + pushed_msg = TransferPushedMsg(request_id=request_id) + self._rdma_push_zmq.send_multipart(encode_transfer_msg(pushed_msg)) + + if not success: + logger.error( + "Transfer %s: RDMA push failed for %s", role_name, request_id + ) + except Exception: + logger.exception( + "Transfer %s: RDMA push thread error for %s", role_name, request_id + ) + + def _recv_prefetch_loop(self: Scheduler): + """Background thread: recv transfer messages and prefetch tensor loads. + + For transfer_ready: loads tensors + builds Req in this thread, then + enqueues the ready-to-compute item. This allows loading of request N+1 + to overlap with compute of request N on the main thread. + + For transfer_alloc/push: passes them through to the main thread for handling. + """ + role_name = self._disagg_role.value.upper() + while self._running: + try: + raw_frames = self._pool_work_pull.recv_multipart() + frames = [bytes(f) for f in raw_frames] + + msg = decode_transfer_msg(frames) + msg_type = msg.get("msg_type", "") + + if msg_type == TransferMsgType.READY: + # Prefetch: load tensors + build Req in this thread + item = self._prefetch_transfer_ready(msg) + self._compute_ready_queue.put(("transfer_compute", item)) + elif msg_type == TransferMsgType.PUSH: + # Handle push directly in prefetch thread — it only + # enqueues to the RDMA bg thread (thread-safe queue). + # Critical for pipeline parallelism: if deferred to the + # main thread, this gets blocked behind the next request's + # GPU compute, preventing the previous request's output + # from reaching the decoder. + self._handle_transfer_push(msg) + else: + # alloc and other messages: pass to main thread + # (alloc sends on _pool_result_push which isn't thread-safe) + self._compute_ready_queue.put(("transfer_control", frames)) + + except zmq.ZMQError as e: + if not self._running: + break + logger.error("Transfer %s recv prefetch: ZMQ error: %s", role_name, e) + except Exception: + logger.exception("Transfer %s recv prefetch: error", role_name) + + def _prefetch_transfer_ready(self: Scheduler, msg: dict) -> tuple: + """Load tensors from transfer buffer and build Req for a transfer_ready message. + + Called from the recv prefetch thread. Loads on _transfer_stream + and builds the Req, so the main thread can start compute immediately. + + Returns (req, load_event, request_id, role_name, prealloc_slot_id). + """ + request_id = msg["request_id"] + manifest = msg.get("manifest", {}) + scalar_fields = msg.get("scalar_fields", {}) + role_name = self._disagg_role.value.upper() + + if self._disagg_metrics: + self._disagg_metrics.record_request_start(request_id) + + # Pre-allocated slot handling + prealloc_slot_id = scalar_fields.pop("_prealloc_slot_id", None) + if ( + prealloc_slot_id is not None + and prealloc_slot_id in self._preallocated_slots + ): + slot = self._preallocated_slots[prealloc_slot_id] + self._transfer_manager.register_prealloc_as_receive(request_id, slot) + + # Load tensors on transfer_stream (non-blocking) + local_device = f"cuda:{self.worker.local_rank}" + tensors, load_event = self._transfer_manager.load_tensors_async( + request_id, + manifest, + device=local_device, + stream=self._transfer_stream, + ) + + # NOTE: Do NOT free the receive slot here. The async load is still + # in progress. The slot must remain valid until the main thread waits + # on load_event. Freeing is done in _disagg_prefetch_event_loop. + + # Build Req (CPU work, overlapped with load) + req = self._build_disagg_req(scalar_fields, tensors) + + # NOTE: Do NOT call scheduler_mod.set_timesteps() here! + # This runs on the prefetch thread. set_timesteps mutates shared + # scheduler state (self.sigmas), which would corrupt the currently + # running denoising loop on the main thread. Deferred to main thread + # in _disagg_prefetch_event_loop, right before compute. + + return (req, load_event, request_id, role_name, prealloc_slot_id, scalar_fields) + + # ------------------------------------------------------------------ + # Broadcast + # ------------------------------------------------------------------ + + def _broadcast_to_all_ranks(self: Scheduler, data): + """Broadcast *data* from rank 0 to all other ranks. + + Rank 0 passes the real payload; non-rank-0 passes ``None``. + Broadcasts through all applicable groups (SP, CFG, TP). + """ + sa = self.server_args + + if sa.sp_degree != 1: + data = broadcast_pyobj( + data, + self.worker.sp_group.rank, + self.worker.sp_cpu_group, + src=self.worker.sp_group.ranks[0], + ) + + if sa.enable_cfg_parallel: + data = broadcast_pyobj( + data, + self.worker.cfg_group.rank, + self.worker.cfg_cpu_group, + src=self.worker.cfg_group.ranks[0], + ) + + if sa.tp_size > 1: + data = broadcast_pyobj( + data, + self.worker.tp_group.rank, + self.worker.tp_cpu_group, + src=self.worker.tp_group.ranks[0], + ) + + return data + + def _is_multi_rank(self: Scheduler) -> bool: + sa = self.server_args + return sa.sp_degree != 1 or sa.tp_size > 1 or sa.enable_cfg_parallel + + def _broadcast_tensor_dict_to_all_ranks( + self: Scheduler, tensor_dict: dict | None + ) -> dict | None: + """Broadcast a tensor dict from rank 0 to non-rank-0 via NCCL. + + Uses ``GroupCoordinator.broadcast_tensor_dict`` which ships tensor + metadata over the CPU group and the tensor payload over the device + (NCCL) group, so large GPU buffers never bounce through CPU. + """ + sa = self.server_args + + if sa.sp_degree != 1: + tensor_dict = self.worker.sp_group.broadcast_tensor_dict(tensor_dict, src=0) + if sa.enable_cfg_parallel: + tensor_dict = self.worker.cfg_group.broadcast_tensor_dict( + tensor_dict, src=0 + ) + if sa.tp_size > 1: + tensor_dict = self.worker.tp_group.broadcast_tensor_dict(tensor_dict, src=0) + return tensor_dict + + def _broadcast_req_to_all_ranks(self: Scheduler, req: Req | None) -> Req | None: + """Broadcast a fully-loaded Req (scalars + GPU tensors) from rank 0. + + Required for multi-rank denoiser/decoder in disagg mode: only rank 0 + owns the TransferManager and RDMA-loads tensors into GPU memory. All + other ranks must see the same Req before entering ``execute_forward``, + otherwise REPLICATED stages (e.g. denoising) blow up on empty tensor + fields because ``ParallelExecutor`` never broadcasts the batch for + that paradigm. + + Tensor fields travel over NCCL (stays on GPU); scalar fields travel + as a small pickled object over the CPU group. + """ + if not self._is_multi_rank(): + return req + + is_rank0 = self.gpu_id == 0 + + if is_rank0: + assert req is not None, "rank 0 must pass a loaded Req" + tensor_fields, scalar_fields = extract_transfer_fields(req) + packed_tensors = _pack_tensor_fields_for_broadcast(tensor_fields) + else: + scalar_fields = None + packed_tensors = None + + # 1. Scalars via CPU pyobj broadcast. + scalar_fields = self._broadcast_to_all_ranks(scalar_fields) + + # 2. Tensors via NCCL broadcast — keeps GPU buffers on device. + packed_tensors = self._broadcast_tensor_dict_to_all_ranks(packed_tensors) + + if is_rank0: + return req + + tensor_fields = _unpack_tensor_fields_from_broadcast(packed_tensors or {}) + # Move tensors onto this rank's physical device. The broadcast + # allocates receive tensors on the receiver's default CUDA device + # (set via torch.cuda.set_device(local_rank) during init), which is + # already the right physical GPU — the .to() is effectively a no-op + # but makes the invariant explicit for future readers. + local_device = torch.device(f"cuda:{self.worker.local_rank}") + for key, value in list(tensor_fields.items()): + if isinstance(value, torch.Tensor): + tensor_fields[key] = value.to(local_device, non_blocking=True) + elif isinstance(value, list): + tensor_fields[key] = [ + ( + t.to(local_device, non_blocking=True) + if isinstance(t, torch.Tensor) + else t + ) + for t in value + ] + return self._build_disagg_req(scalar_fields or {}, tensor_fields) + + # ------------------------------------------------------------------ + # Event loops + # ------------------------------------------------------------------ + + def _disagg_recv_work(self: Scheduler) -> list[bytes] | None: + """Receive work frames in pool mode, with multi-rank broadcast. + + Rank 0: recv from ZMQ PULL socket, broadcast to other ranks. + Non-rank-0: receive via NCCL broadcast from rank 0. + + Returns list of bytes frames, or None on shutdown. + """ + if self.gpu_id == 0: + raw_frames = self._pool_work_pull.recv_multipart() + frames = [bytes(f) for f in raw_frames] + else: + frames = None + + return self._broadcast_to_all_ranks(frames) + + def _disagg_prefetch_event_loop(self: Scheduler, role_name: str) -> None: + """Event loop for transfer receiver roles with recv prefetch thread (rank 0). + + The recv thread reads from ZMQ and prefetches tensor loads. + This loop reads from _compute_ready_queue: + - "transfer_compute": load already done, wait_event + free slot + → broadcast scalar_fields to non-rank-0 → compute + - "transfer_control": alloc/push messages, handle on main thread + → broadcast "skip" so non-rank-0 doesn't hang + - queue timeout: broadcast "skip" + - shutdown: broadcast None + """ + is_multi_rank = ( + self.server_args.sp_degree != 1 + or self.server_args.tp_size > 1 + or self.server_args.enable_cfg_parallel + ) + + while self._running: + try: + try: + msg_type, data = self._compute_ready_queue.get(timeout=1.0) + except queue.Empty: + if is_multi_rank: + self._broadcast_to_all_ranks(("skip",)) + continue + + if msg_type == "transfer_compute": + # Load already done by recv thread + req, load_event, request_id, rn, prealloc_slot_id, scalar_fields = ( + data + ) + # Wait for load to complete on compute stream + if load_event is not None: + torch.cuda.current_stream().wait_event(load_event) + # Now safe to free the receive slot + if prealloc_slot_id is not None: + with self._transfer_manager._lock: + self._transfer_manager._pending_receives.pop( + request_id, None + ) + else: + self._transfer_manager.free_receive_slot(request_id) + # Broadcast the full Req (scalar + tensor fields) to + # non-rank-0 ranks. Tensors ride NCCL on the SP/CFG/TP + # groups so downstream REPLICATED stages (e.g. denoising) + # see identical inputs on every rank — without this, the + # non-rank-0 ranks would enter execute_forward with empty + # prompt_embeds and fail verify_input. + if is_multi_rank: + self._broadcast_to_all_ranks(("compute",)) + self._broadcast_req_to_all_ranks(req) + # Init scheduler timesteps on main thread (safe — no + # concurrent denoising loop can be running here). + if self._disagg_role == RoleType.DENOISER: + scheduler_mod = self.worker.pipeline.get_module("scheduler") + num_steps = getattr(req, "num_inference_steps", None) + if scheduler_mod is not None and num_steps is not None: + device = torch.device(f"cuda:{self.worker.local_rank}") + extra_kwargs = {} + mu = req.extra.get("mu") if hasattr(req, "extra") else None + if mu is not None: + extra_kwargs["mu"] = mu + scheduler_mod.set_timesteps( + num_steps, device=device, **extra_kwargs + ) + # Run compute + if self._disagg_role == RoleType.DENOISER: + self._disagg_denoiser_compute(req, request_id, rn) + elif self._disagg_role == RoleType.DECODER: + self._disagg_decoder_compute(req, request_id, rn) + + elif msg_type == "transfer_control": + # alloc, push messages — handle on main thread (rank 0 only) + if is_multi_rank: + self._broadcast_to_all_ranks(("skip",)) + self._handle_transfer_msg(data) + + self._consecutive_error_count = 0 + + except Exception as e: + self._consecutive_error_count += 1 + logger.error( + "Pool %s rank %d prefetch loop: error (attempt %d/%d): %s", + role_name, + self.gpu_id, + self._consecutive_error_count, + self._max_consecutive_errors, + e, + exc_info=True, + ) + if self._consecutive_error_count >= self._max_consecutive_errors: + raise RuntimeError( + f"Pool {role_name} rank {self.gpu_id} terminated after " + f"{self._max_consecutive_errors} consecutive errors: {e}" + ) from e + + # Shutdown: notify non-rank-0 to exit + if is_multi_rank: + self._broadcast_to_all_ranks(None) + self._cleanup_disagg() + + def _disagg_non_rank0_event_loop(self: Scheduler) -> None: + """Event loop for non-rank-0 receivers in multi-rank prefetch mode. + + Blocks on broadcast from rank 0: + - ("compute", scalar_fields): build minimal Req → execute_forward + - ("skip",): continue (rank 0 handled a control msg or timed out) + - None: shutdown, exit loop + """ + role_name = self._disagg_role.value.upper() + logger.info( + "Pool %s rank %d: entering non-rank-0 prefetch loop", + role_name, + self.gpu_id, + ) + + while True: + try: + msg = self._broadcast_to_all_ranks(None) + + if msg is None: + # Shutdown signal + break + + if isinstance(msg, tuple) and len(msg) >= 1 and msg[0] == "compute": + # Participate in the companion tensor broadcast so this + # rank sees the full Req (scalars + GPU tensors). Without + # the tensor half, REPLICATED stages would see empty + # prompt_embeds on non-rank-0 and fail verify_input. + req = self._broadcast_req_to_all_ranks(None) + self._disagg_compute_non_rank0(req) + # else: ("skip",) — continue + + except Exception as e: + self._consecutive_error_count += 1 + logger.error( + "Pool %s rank %d non-rank-0 loop: error (attempt %d/%d): %s", + role_name, + self.gpu_id, + self._consecutive_error_count, + self._max_consecutive_errors, + e, + exc_info=True, + ) + if self._consecutive_error_count >= self._max_consecutive_errors: + raise RuntimeError( + f"Pool {role_name} rank {self.gpu_id} terminated after " + f"{self._max_consecutive_errors} consecutive errors: {e}" + ) from e + + self._cleanup_disagg() + + def _disagg_event_loop(self: Scheduler) -> None: + """Event loop for all roles in pool mode (DiffusionServer-mediated). + + Multi-rank support: + - Rank 0 receives from ZMQ, broadcasts to other ranks via NCCL + - All ranks process work (execute_forward with SP/TP sharding) + - Only rank 0 sends results back to DiffusionServer + + Transfer: + - Transfer control messages (transfer_alloc, transfer_push) are rank-0-only. + - transfer_ready is broadcast to all ranks for compute. + - Encoder receives pickled Req, runs compute, stages output for transfer. + - Denoiser/decoder only receive transfer control messages. + + Receiver prefetch paths: + - Rank 0: _disagg_prefetch_event_loop (reads from compute_ready_queue) + - Non-rank-0 in multi-rank: _disagg_non_rank0_event_loop (broadcast) + - Encoder (any rank): existing _disagg_recv_work while loop below + """ + role_name = self._disagg_role.value.upper() + is_rank0 = self.gpu_id == 0 + is_multi_rank = ( + self.server_args.sp_degree != 1 + or self.server_args.tp_size > 1 + or self.server_args.enable_cfg_parallel + ) + use_prefetch = self._compute_ready_queue is not None + logger.info( + "Pool mode %s rank %d event loop started " "(multi_rank=%s, prefetch=%s)", + role_name, + self.gpu_id, + is_multi_rank, + use_prefetch, + ) + + # Rank 0 receiver with prefetch queue → prefetch event loop + if use_prefetch: + self._disagg_prefetch_event_loop(role_name) + return + + # Non-rank-0 receiver in multi-rank → broadcast-based loop + if ( + not is_rank0 + and is_multi_rank + and self._disagg_role in (RoleType.DENOISER, RoleType.DECODER) + ): + self._disagg_non_rank0_event_loop() + return + + while self._running: + try: + # All ranks receive work (rank 0 via ZMQ, others via broadcast) + frames = self._disagg_recv_work() + + # Transfer dispatch: check on ALL ranks (frames are broadcast) + if self._is_transfer_frames(frames): + if is_rank0: + # Rank 0: handle all transfer messages + self._handle_transfer_msg(frames) + else: + # Non-rank-0: only participate in transfer_ready compute + self._handle_transfer_non_rank0(frames) + elif self._disagg_role == RoleType.ENCODER: + self._disagg_encoder_step( + send_tensors, + frames=frames, + ) + + self._consecutive_error_count = 0 + + except Exception as e: + self._consecutive_error_count += 1 + logger.error( + "Pool %s rank %d: error (attempt %d/%d): %s", + role_name, + self.gpu_id, + self._consecutive_error_count, + self._max_consecutive_errors, + e, + exc_info=True, + ) + if self._consecutive_error_count >= self._max_consecutive_errors: + raise RuntimeError( + f"Pool {role_name} rank {self.gpu_id} terminated after " + f"{self._max_consecutive_errors} consecutive errors: {e}" + ) from e + + self._cleanup_disagg() + + def _cleanup_disagg(self: Scheduler): + """Clean up all pool mode resources (sockets, threads, transfer manager).""" + # Shutdown RDMA push thread + if self._rdma_push_queue is not None: + self._rdma_push_queue.put(None) + if self._rdma_push_thread is not None: + self._rdma_push_thread.join(timeout=5) + if self._rdma_push_zmq is not None: + self._rdma_push_zmq.close() + # Recv prefetch thread stops when self._running = False + if self._recv_prefetch_thread is not None: + self._recv_prefetch_thread.join(timeout=5) + if self._transfer_manager is not None: + self._transfer_manager.cleanup() + if self._pool_work_pull is not None: + self._pool_work_pull.close() + if self._pool_result_push is not None: + self._pool_result_push.close() + + # ------------------------------------------------------------------ + # Transfer message handling + # ------------------------------------------------------------------ + + @staticmethod + def _is_transfer_frames(frames: list) -> bool: + """Check if ZMQ multipart frames carry a transfer control message.""" + return is_transfer_message(frames) + + def _handle_transfer_msg(self: Scheduler, frames: list) -> None: + """Dispatch a transfer control message to the appropriate handler (rank 0).""" + msg = decode_transfer_msg(frames) + msg_type = msg.get("msg_type", "") + request_id = msg.get("request_id", "") + + logger.debug( + "Transfer %s: received %s for %s", + self._disagg_role.value.upper(), + msg_type, + request_id, + ) + + if msg_type == TransferMsgType.ALLOC: + self._handle_transfer_alloc(msg) + elif msg_type == TransferMsgType.PUSH: + self._handle_transfer_push(msg) + elif msg_type == TransferMsgType.READY: + self._handle_transfer_ready(msg) + else: + logger.warning( + "Transfer %s: unknown message type %s", + self._disagg_role.value.upper(), + msg_type, + ) + + def _handle_transfer_non_rank0(self: Scheduler, frames: list) -> None: + """Handle transfer messages on non-rank-0 workers. + + Only transfer_ready requires non-rank-0 participation (for compute). + transfer_alloc and transfer_push are rank-0-only operations — skip them. + """ + msg = decode_transfer_msg(frames) + msg_type = msg.get("msg_type", "") + + if msg_type == TransferMsgType.READY: + # Non-rank-0 has no TransferManager, so rank 0 loads tensors from + # the RDMA buffer and broadcasts the full Req (scalars + tensors) + # over NCCL. Participate in the matching broadcast here. + req = self._broadcast_req_to_all_ranks(None) + self._disagg_compute_non_rank0(req) + # else: transfer_alloc, transfer_push — skip (rank-0-only operations) + + def _handle_transfer_alloc(self: Scheduler, msg: dict) -> None: + """Handle transfer_alloc: allocate a receive slot and reply with transfer_allocated.""" + request_id = msg["request_id"] + data_size = msg.get("data_size", 0) + + pending = self._transfer_manager.allocate_receive_slot(request_id, data_size) + if pending is None: + logger.error( + "Transfer %s: failed to allocate receive slot for %s (%d bytes)", + self._disagg_role.value.upper(), + request_id, + data_size, + ) + return + + allocated_msg = TransferAllocatedMsg( + request_id=request_id, + session_id=self._transfer_manager.session_id, + pool_ptr=self._transfer_manager.pool_data_ptr, + slot_offset=pending.slot.offset, + slot_size=pending.slot.size, + ) + self._pool_result_push.send_multipart(encode_transfer_msg(allocated_msg)) + + logger.debug( + "Transfer %s: allocated receive slot for %s (offset=%d, size=%d)", + self._disagg_role.value.upper(), + request_id, + pending.slot.offset, + pending.slot.size, + ) + + def _handle_transfer_push(self: Scheduler, msg: dict) -> None: + """Handle transfer_push: RDMA push staged data to peer, reply with transfer_pushed. + + If RDMA push thread is active, enqueue non-blocking. + Otherwise fall back to blocking push (e.g., during shutdown). + """ + request_id = msg["request_id"] + dest_session_id = msg.get("dest_session_id", "") + dest_addr = msg.get("dest_addr", 0) + transfer_size = msg.get("transfer_size", 0) + + if self._rdma_push_queue is not None: + # Non-blocking: enqueue to RDMA push thread + self._rdma_push_queue.put( + ( + request_id, + dest_session_id, + dest_addr, + transfer_size, + ) + ) + return + + # Fallback: blocking push on main thread + success = self._transfer_manager.push_to_peer( + request_id=request_id, + dest_session_id=dest_session_id, + dest_addr=dest_addr, + transfer_size=transfer_size, + ) + + if success: + self._transfer_manager.free_staged(request_id) + + pushed_msg = TransferPushedMsg(request_id=request_id) + self._pool_result_push.send_multipart(encode_transfer_msg(pushed_msg)) + + if not success: + logger.error( + "Transfer %s: RDMA push failed for %s", + self._disagg_role.value.upper(), + request_id, + ) + + def _handle_transfer_ready(self: Scheduler, msg: dict) -> None: + """Handle transfer_ready: load tensors from buffer, run compute, send result. + + Overlap tensor load with Req construction and scheduler init. + After the RDMA data arrives: + 1. Start load on transfer_stream (non-blocking) + 2. Build Req from scalar fields + tensors (CPU, overlapped) + 3. Init scheduler timesteps if denoiser (CPU, overlapped) + 4. Wait for load before compute + 5. Run the role's compute + """ + + request_id = msg["request_id"] + manifest = msg.get("manifest", {}) + scalar_fields = msg.get("scalar_fields", {}) + role_name = self._disagg_role.value.upper() + + if self._disagg_metrics: + self._disagg_metrics.record_request_start(request_id) + + # If using a pre-allocated slot, register it as pending receive + prealloc_slot_id = scalar_fields.pop("_prealloc_slot_id", None) + if ( + prealloc_slot_id is not None + and prealloc_slot_id in self._preallocated_slots + ): + slot = self._preallocated_slots[prealloc_slot_id] + self._transfer_manager.register_prealloc_as_receive(request_id, slot) + + # 1. Start load on transfer_stream (non-blocking) + local_device = f"cuda:{self.worker.local_rank}" + tensors, load_event = self._transfer_manager.load_tensors_async( + request_id, + manifest, + device=local_device, + stream=self._transfer_stream, + ) + + # 2. Build Req from scalar fields + tensors (CPU work, overlapped) + req = self._build_disagg_req(scalar_fields, tensors) + + # 3. Init scheduler timesteps if denoiser (CPU work, overlapped) + if self._disagg_role == RoleType.DENOISER: + scheduler_mod = self.worker.pipeline.get_module("scheduler") + num_steps = getattr(req, "num_inference_steps", None) + if scheduler_mod is not None and num_steps is not None: + device = torch.device(local_device) + extra_kwargs = {} + mu = req.extra.get("mu") if hasattr(req, "extra") else None + if mu is not None: + extra_kwargs["mu"] = mu + scheduler_mod.set_timesteps(num_steps, device=device, **extra_kwargs) + + # 4. Wait for load before compute (GPU must see the data) + if load_event is not None: + torch.cuda.current_stream().wait_event(load_event) + + # 5. Free receive slot after load completes (data is on compute GPU) + if prealloc_slot_id is not None: + # Pre-allocated slot: just remove from pending receives, don't free buffer + with self._transfer_manager._lock: + self._transfer_manager._pending_receives.pop(request_id, None) + else: + self._transfer_manager.free_receive_slot(request_id) + + # 6. In multi-rank mode, broadcast the fully-loaded Req to the other + # ranks so REPLICATED stages see identical inputs everywhere. See + # the prefetch-loop variant for the matching receiver broadcast. + if self._is_multi_rank(): + self._broadcast_req_to_all_ranks(req) + + # 7. Run compute + if self._disagg_role == RoleType.DENOISER: + self._disagg_denoiser_compute(req, request_id, role_name) + elif self._disagg_role == RoleType.DECODER: + self._disagg_decoder_compute(req, request_id, role_name) + + # ------------------------------------------------------------------ + # Compute + # ------------------------------------------------------------------ + + def _disagg_compute_non_rank0(self: Scheduler, req: Req) -> None: + """Non-rank-0 compute: enter execute_forward with a Req received via + NCCL broadcast from rank 0. + + The Req already contains tensor fields materialized on this rank's + GPU (see ``_broadcast_req_to_all_ranks``), so REPLICATED stages such + as denoising have non-empty prompt_embeds and verify_input passes. + + Used by both the non-prefetch path (:meth:`_handle_transfer_non_rank0`) + and the prefetch non-rank-0 loop + (:meth:`_disagg_non_rank0_event_loop`). + """ + if self._disagg_role == RoleType.DENOISER: + # Initialize scheduler timesteps (same as rank 0) + scheduler_mod = self.worker.pipeline.get_module("scheduler") + num_steps = getattr(req, "num_inference_steps", None) + if scheduler_mod is not None and num_steps is not None: + device = torch.device(f"cuda:{self.worker.local_rank}") + extra_kwargs = {} + mu = req.extra.get("mu") if hasattr(req, "extra") else None + if mu is not None: + extra_kwargs["mu"] = mu + scheduler_mod.set_timesteps(num_steps, device=device, **extra_kwargs) + + self.worker.execute_forward([req], return_req=True) + + elif self._disagg_role == RoleType.DECODER: + req.save_output = False + req.return_file_paths_only = False + self.worker.execute_forward([req]) + + def _build_disagg_req(self: Scheduler, scalar_fields: dict, tensors: dict) -> Req: + """Reconstruct a Req from transfer scalar fields and loaded GPU tensors. + + Initializes all dataclass field defaults first, then overlays + scalar and tensor fields from the transfer message. + """ + req = object.__new__(Req) + # Initialize all dataclass fields with their defaults + for f in dataclasses.fields(Req): + if f.default is not dataclasses.MISSING: + object.__setattr__(req, f.name, f.default) + elif f.default_factory is not dataclasses.MISSING: + object.__setattr__(req, f.name, f.default_factory()) + # Ensure sampling_params is not None so __getattr__ delegation works + object.__setattr__(req, "sampling_params", SamplingParams()) + # Restore _extra_* prefixed fields into req.extra dict + extra_keys = [k for k in scalar_fields if k.startswith("_extra_")] + for key in extra_keys: + req.extra[key[len("_extra_") :]] = scalar_fields.pop(key) + for key, value in scalar_fields.items(): + setattr(req, key, value) + # Set tensor fields + for key, value in tensors.items(): + setattr(req, key, value) + # Recreate torch.Generator from seed (not serializable over transfer) + seed = scalar_fields.get("seed") + if seed is not None: + gen = torch.Generator(device="cpu") + gen.manual_seed(int(seed)) + req.generator = gen + req.validate() + return req + + def _disagg_denoiser_compute( + self: Scheduler, req: Req, request_id: str, role_name: str + ) -> None: + """Run denoiser compute in transfer mode, then stage output for decoder. + + Note: Scheduler timestep init is done in _handle_transfer_ready + to overlap with tensor loading. + """ + # Run denoising + start_time = time.monotonic() + result = self.worker.execute_forward([req], return_req=True) + duration_s = time.monotonic() - start_time + + if not isinstance(result, Req): + error_msg = getattr(result, "error", "denoiser error") + done_msg = TransferDoneMsg(request_id=request_id, error=str(error_msg)) + self._pool_result_push.send_multipart(encode_transfer_msg(done_msg)) + if self._disagg_metrics: + self._disagg_metrics.record_request_failed(request_id) + return + + # Stage denoiser output for decoder transfer (async staging) + tensor_fields, scalar_fields = extract_transfer_fields(result) + + # 1. Stage tensors on transfer_stream (non-blocking) + staged, stage_event = self._transfer_manager.stage_tensors_async( + request_id=request_id, + tensor_fields=tensor_fields, + scalar_fields=scalar_fields, + stream=self._transfer_stream, + ) + + if staged is None: + done_msg = TransferDoneMsg( + request_id=request_id, + error="Failed to stage denoiser output for decoder", + ) + self._pool_result_push.send_multipart(encode_transfer_msg(done_msg)) + if self._disagg_metrics: + self._disagg_metrics.record_request_failed(request_id) + return + + # 2. Build done_data dict while staging runs (CPU work, overlapped) + done_data = { + "msg_type": "transfer_done", + "request_id": request_id, + "staged_for_decoder": True, + "session_id": self._transfer_manager.session_id, + "pool_ptr": self._transfer_manager.pool_data_ptr, + "slot_offset": staged.slot.offset if staged.slot else 0, + "data_size": staged.slot.size if staged.slot else 0, + "manifest": staged.manifest, + "scalar_fields": staged.scalar_fields, + } + msg_bytes = json.dumps(done_data, separators=(",", ":")).encode("utf-8") + + # 3. Wait for staging to complete before sending + if stage_event is not None: + stage_event.synchronize() + + # 4. Send transfer_done with staged info + self._pool_result_push.send_multipart([TRANSFER_MAGIC, msg_bytes]) + + if self._disagg_metrics: + self._disagg_metrics.record_request_complete(request_id) + + logger.debug( + "Transfer DENOISER: processed %s in %.2f s, staged for decoder", + request_id, + duration_s, + ) + + def _disagg_decoder_compute( + self: Scheduler, req: Req, request_id: str, role_name: str + ) -> None: + """Run decoder compute in transfer mode, send result to DS. + + Decoder result is sent as raw ZMQ multipart frames (same format as + relay mode) so DiffusionServer handles it via _handle_decoder_result_frames + without hex/JSON overhead. + """ + + # Check for upstream error + disagg_error = getattr(req, "_disagg_error", None) + if disagg_error: + if self._pool_result_push is not None: + send_tensors( + self._pool_result_push, + {}, + { + "request_id": request_id, + "error": f"Upstream error: {disagg_error}", + }, + ) + return + + req.save_output = False + req.return_file_paths_only = False + + start_time = time.monotonic() + output_batch = self.worker.execute_forward([req]) + duration_s = time.monotonic() - start_time + + # Send result as raw ZMQ frames (no TRANSFER_MAGIC prefix). + # DiffusionServer will route it through _handle_decoder_result_frames, + # the same path as relay mode. + tensor_fields = {} + scalar_fields = {"request_id": request_id} + if output_batch.output is not None: + tensor_fields["output"] = output_batch.output + if output_batch.audio is not None: + tensor_fields["audio"] = output_batch.audio + if output_batch.audio_sample_rate is not None: + scalar_fields["audio_sample_rate"] = output_batch.audio_sample_rate + if output_batch.error is not None: + scalar_fields["error"] = output_batch.error + + if self._pool_result_push is not None: + send_tensors(self._pool_result_push, tensor_fields, scalar_fields) + + if self._disagg_metrics: + if output_batch.error: + self._disagg_metrics.record_request_failed(request_id) + else: + self._disagg_metrics.record_request_complete(request_id) + + logger.debug("Transfer DECODER: processed %s in %.2f s", request_id, duration_s) + + def _disagg_encoder_step( + self: Scheduler, + send_tensors_fn, + frames=None, + ): + """Single encoder step in pool mode.""" + # Receive: [request_id_bytes, pickled_req_bytes] + if frames is None: + frames = self._pool_work_pull.recv_multipart() + pickled_req = frames[-1] + reqs = pickle.loads(pickled_req) + if not isinstance(reqs, list): + reqs = [reqs] + + req = reqs[0] + request_id = getattr(req, "request_id", "unknown") + + if self._disagg_metrics: + self._disagg_metrics.record_request_start(request_id) + + # Run encoder stages + req_result = self.worker.execute_forward(reqs, return_req=True) + + if not isinstance(req_result, Req): + # Error — send error via scalar fields (rank 0 only) + if self._pool_result_push is not None: + error_msg = getattr(req_result, "error", "encoder error") + send_tensors_fn( + self._pool_result_push, + {}, + {"request_id": request_id, "_disagg_error": str(error_msg)}, + ) + if self._disagg_metrics: + self._disagg_metrics.record_request_failed(request_id) + return + + # Pack and send encoder output (rank 0 only sends) + tensor_fields, scalar_fields = extract_transfer_fields(req_result) + + if self._pool_result_push is not None: + if self._transfer_manager is not None: + # Transfer mode: stage tensors to TransferBuffer, send transfer_staged + self._disagg_encoder_transfer_stage( + request_id, tensor_fields, scalar_fields + ) + else: + # Fallback: send error (transfer manager not initialized) + send_tensors_fn( + self._pool_result_push, + {}, + {"request_id": request_id, "_disagg_error": "No transfer manager"}, + ) + + if self._disagg_metrics: + self._disagg_metrics.record_request_complete(request_id) + + logger.debug("Pool ENCODER: processed %s", request_id) + + def _disagg_encoder_transfer_stage( + self: Scheduler, request_id: str, tensor_fields: dict, scalar_fields: dict + ) -> None: + """Stage encoder output and send transfer_staged to DS. + + Overlap staging with metadata JSON serialization. + """ + # 1. Stage tensors on transfer_stream (non-blocking) + staged, stage_event = self._transfer_manager.stage_tensors_async( + request_id=request_id, + tensor_fields=tensor_fields, + scalar_fields=scalar_fields, + stream=self._transfer_stream, + ) + + if staged is None: + # Staging failed — send error via relay as fallback + send_tensors( + self._pool_result_push, + {}, + {"request_id": request_id, "_disagg_error": "Transfer staging failed"}, + ) + if self._disagg_metrics: + self._disagg_metrics.record_request_failed(request_id) + return + + # 2. Build transfer metadata dict while staging runs (CPU work, overlapped) + staged_data = { + "msg_type": "transfer_staged", + "request_id": request_id, + "data_size": staged.slot.size if staged.slot else 0, + "manifest": staged.manifest, + "session_id": self._transfer_manager.session_id, + "pool_ptr": self._transfer_manager.pool_data_ptr, + "slot_offset": staged.slot.offset if staged.slot else 0, + "scalar_fields": staged.scalar_fields, + } + msg_bytes = json.dumps(staged_data, separators=(",", ":")).encode("utf-8") + + # 3. Wait for staging to complete before sending (buffer must be ready) + if stage_event is not None: + stage_event.synchronize() + + # 4. Send transfer staged message + self._pool_result_push.send_multipart([TRANSFER_MAGIC, msg_bytes]) diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/transport/__init__.py b/python/sglang/multimodal_gen/runtime/disaggregation/transport/__init__.py new file mode 100644 index 000000000..c3a063e13 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/disaggregation/transport/__init__.py @@ -0,0 +1,2 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Transport layer for disaggregated diffusion pipelines.""" diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/transport/allocator.py b/python/sglang/multimodal_gen/runtime/disaggregation/transport/allocator.py new file mode 100644 index 000000000..49a5399a2 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/disaggregation/transport/allocator.py @@ -0,0 +1,200 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Buddy-system memory allocator for TransferTensorBuffer.""" + +from __future__ import annotations + +import logging +import threading +from dataclasses import dataclass +from functools import lru_cache + +logger = logging.getLogger(__name__) + + +@dataclass +class Block: + offset: int # byte offset from pool start + size: int + allocated: bool = False + request_id: str | None = None + + +class BuddyAllocator: + """Power-of-2 buddy-system allocator for pinned memory.""" + + def __init__(self, pool_size: int, min_block_size: int = 1 << 20): + if min_block_size <= 0 or (min_block_size & (min_block_size - 1)) != 0: + raise ValueError( + f"min_block_size must be a power of 2, got {min_block_size}" + ) + + self._min_block_size = min_block_size + self._pool_size = self._next_power_of_2(max(pool_size, min_block_size)) + self._lock = threading.Lock() + + # Free lists indexed by order: order 0 = min_block_size, order 1 = 2*min_block_size, ... + self._max_order = self._size_to_order(self._pool_size) + self._free_lists: list[list[int]] = [[] for _ in range(self._max_order + 1)] + + self._blocks: dict[int, Block] = {} + + root = Block(offset=0, size=self._pool_size) + self._blocks[0] = root + self._free_lists[self._max_order].append(0) + + self._allocated_bytes = 0 + self._num_allocations = 0 + + @property + def pool_size(self) -> int: + return self._pool_size + + def allocate(self, size: int, request_id: str | None = None) -> int | None: + """Allocate a block of at least `size` bytes. Returns offset or None.""" + if size <= 0: + raise ValueError(f"Allocation size must be positive, got {size}") + + alloc_size = max(self._next_power_of_2(size), self._min_block_size) + target_order = self._size_to_order(alloc_size) + + if target_order > self._max_order: + logger.warning( + "Requested size %d exceeds pool size %d", size, self._pool_size + ) + return None + + with self._lock: + return self._allocate_locked(target_order, request_id) + + def free(self, offset: int) -> bool: + """Free the block at the given offset and coalesce with buddy if possible.""" + with self._lock: + return self._free_locked(offset) + + def get_block_info(self, offset: int) -> Block | None: + with self._lock: + return self._blocks.get(offset) + + def get_stats(self) -> dict: + with self._lock: + free_blocks_by_order = {} + for order, offsets in enumerate(self._free_lists): + if offsets: + block_size = self._min_block_size << order + free_blocks_by_order[block_size] = len(offsets) + + return { + "pool_size": self._pool_size, + "min_block_size": self._min_block_size, + "allocated_bytes": self._allocated_bytes, + "free_bytes": self._pool_size - self._allocated_bytes, + "num_allocations": self._num_allocations, + "num_blocks": len(self._blocks), + "free_blocks_by_size": free_blocks_by_order, + } + + def count_free_slots(self, slot_size: int) -> int: + """Count how many allocations of the given size can fit.""" + if slot_size <= 0: + return 0 + alloc_size = max(self._next_power_of_2(slot_size), self._min_block_size) + + with self._lock: + count = 0 + for order in range(self._size_to_order(alloc_size), self._max_order + 1): + for _ in self._free_lists[order]: + block_size = self._min_block_size << order + count += block_size // alloc_size + return count + + # --- Internal (caller must hold self._lock) --- + + def _allocate_locked(self, target_order: int, request_id: str | None) -> int | None: + found_order = -1 + for order in range(target_order, self._max_order + 1): + if self._free_lists[order]: + found_order = order + break + + if found_order < 0: + return None + + offset = self._free_lists[found_order].pop(0) + block = self._blocks[offset] + + # Split down to target_order + while found_order > target_order: + found_order -= 1 + buddy_size = self._min_block_size << found_order + buddy_offset = offset + buddy_size + + buddy = Block(offset=buddy_offset, size=buddy_size) + self._blocks[buddy_offset] = buddy + self._free_lists[found_order].append(buddy_offset) + + block.size = buddy_size + + block.allocated = True + block.request_id = request_id + self._allocated_bytes += block.size + self._num_allocations += 1 + + return offset + + def _free_locked(self, offset: int) -> bool: + block = self._blocks.get(offset) + if block is None or not block.allocated: + return False + + block.allocated = False + block.request_id = None + self._allocated_bytes -= block.size + self._num_allocations -= 1 + + self._coalesce(block) + return True + + def _coalesce(self, block: Block) -> None: + """Recursively merge with buddy if both are free.""" + while block.size < self._pool_size: + buddy_offset = block.offset ^ block.size + buddy = self._blocks.get(buddy_offset) + + if buddy is None or buddy.allocated or buddy.size != block.size: + break + + order = self._size_to_order(buddy.size) + self._free_lists[order].remove(buddy_offset) + + if buddy_offset < block.offset: + del self._blocks[block.offset] + buddy.size *= 2 + block = buddy + else: + del self._blocks[buddy_offset] + block.size *= 2 + + order = self._size_to_order(block.size) + self._free_lists[order].append(block.offset) + + def _size_to_order(self, size: int) -> int: + order = 0 + s = self._min_block_size + while s < size: + s <<= 1 + order += 1 + return order + + @staticmethod + @lru_cache(maxsize=256) + def _next_power_of_2(n: int) -> int: + if n <= 0: + return 1 + n -= 1 + n |= n >> 1 + n |= n >> 2 + n |= n >> 4 + n |= n >> 8 + n |= n >> 16 + n |= n >> 32 + return n + 1 diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/transport/buffer.py b/python/sglang/multimodal_gen/runtime/disaggregation/transport/buffer.py new file mode 100644 index 000000000..e7a8fae8a --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/disaggregation/transport/buffer.py @@ -0,0 +1,272 @@ +# SPDX-License-Identifier: Apache-2.0 +"""TransferTensorBuffer: memory staging area for disaggregated tensor transfer.""" + +from __future__ import annotations + +import logging +from dataclasses import dataclass, field + +import torch + +from sglang.multimodal_gen.runtime.disaggregation.transport.allocator import ( + BuddyAllocator, +) +from sglang.multimodal_gen.runtime.disaggregation.transport.codec import ( + str_to_dtype, +) + +logger = logging.getLogger(__name__) + + +@dataclass +class SlotHandle: + request_id: str + offset: int # byte offset in the pool + size: int # allocated size in bytes + tensor_views: dict[str, torch.Tensor | list[torch.Tensor]] = field( + default_factory=dict + ) + + +class TransferTensorBuffer: + """Memory pool for staging tensor payloads between roles. + + Wraps a contiguous block of memory (CPU pinned or GPU) with a BuddyAllocator. + """ + + def __init__( + self, + pool_size: int, + min_block_size: int = 1 << 20, + role_name: str = "unknown", + device: str = "cpu", + ): + self._role_name = role_name + self._device = device + self._allocator = BuddyAllocator(pool_size, min_block_size) + actual_size = self._allocator.pool_size + + if device == "cpu": + self._pool = torch.empty(actual_size, dtype=torch.uint8, pin_memory=True) + else: + self._pool = torch.empty(actual_size, dtype=torch.uint8, device=device) + self._pool_ptr = self._pool.data_ptr() + + pool_location = "pinned CPU" if device == "cpu" else f"GPU ({device})" + logger.info( + "TransferTensorBuffer[%s]: allocated %d MiB %s memory " + "(min_block=%d KiB)", + role_name, + actual_size >> 20, + pool_location, + min_block_size >> 10, + ) + + @property + def pool_size(self) -> int: + return self._allocator.pool_size + + @property + def device(self) -> str: + return self._device + + @property + def pool_data_ptr(self) -> int: + return self._pool_ptr + + def allocate(self, size: int, request_id: str) -> SlotHandle | None: + """Allocate a slot. Returns None if pool is full.""" + offset = self._allocator.allocate(size, request_id=request_id) + if offset is None: + logger.warning( + "TransferTensorBuffer[%s]: allocation failed for %s (%d bytes). " + "Pool stats: %s", + self._role_name, + request_id, + size, + self._allocator.get_stats(), + ) + return None + + block = self._allocator.get_block_info(offset) + return SlotHandle( + request_id=request_id, + offset=offset, + size=block.size if block else size, + ) + + def free(self, handle: SlotHandle) -> bool: + return self._allocator.free(handle.offset) + + def write_tensor( + self, + handle: SlotHandle, + name: str, + tensor: torch.Tensor, + byte_offset: int = 0, + stream: torch.cuda.Stream | None = None, + ) -> int: + """Copy a tensor into the pool slot. Returns bytes written.""" + src_tensor = tensor.contiguous() + nbytes = src_tensor.numel() * src_tensor.element_size() + + if byte_offset + nbytes > handle.size: + raise ValueError( + f"Write exceeds slot: offset={byte_offset}, nbytes={nbytes}, " + f"slot_size={handle.size}" + ) + + dst = self._pool[ + handle.offset + byte_offset : handle.offset + byte_offset + nbytes + ] + src_bytes = src_tensor.view(torch.uint8).reshape(-1) + + if stream is not None: + with torch.cuda.stream(stream): + dst.copy_(src_bytes, non_blocking=True) + else: + dst.copy_(src_bytes, non_blocking=True) + + return nbytes + + def read_tensor( + self, + handle: SlotHandle, + shape: list[int], + dtype: torch.dtype, + byte_offset: int = 0, + device: torch.device | str = "cpu", + stream: torch.cuda.Stream | None = None, + ) -> torch.Tensor: + """Read a tensor from the pool slot. Returns a clone on target device.""" + nbytes = 1 + for s in shape: + nbytes *= s + nbytes *= torch.tensor([], dtype=dtype).element_size() + + raw = self._pool[ + handle.offset + byte_offset : handle.offset + byte_offset + nbytes + ] + src = raw.view(dtype).reshape(shape) + + pool_dev = str(self._pool.device) + target_dev = str(device) + + same_device = pool_dev == target_dev + + if same_device: + # Clone to decouple tensor lifetime from pool slot + if stream is not None: + with torch.cuda.stream(stream): + return src.clone() + return src.clone() + + if stream is not None: + with torch.cuda.stream(stream): + return src.to(device, non_blocking=True) + return src.to(device, non_blocking=True) + + def write_tensors_from_gpu( + self, + handle: SlotHandle, + tensors: dict[str, torch.Tensor | list[torch.Tensor] | None], + stream: torch.cuda.Stream | None = None, + ) -> dict[str, list[dict]]: + """Batch-write GPU tensors into a slot. Returns a manifest for later reads.""" + manifest: dict[str, list[dict]] = {} + byte_offset = 0 + + # Ensure copy stream sees all prior compute kernels + if stream is not None: + stream.wait_stream(torch.cuda.current_stream()) + + for name, value in tensors.items(): + if value is None: + continue + + entries = [] + if isinstance(value, torch.Tensor): + nbytes = self.write_tensor(handle, name, value, byte_offset, stream) + entries.append( + { + "offset": byte_offset, + "shape": list(value.shape), + "dtype": str(value.dtype).replace("torch.", ""), + } + ) + byte_offset += nbytes + byte_offset = (byte_offset + 511) & ~511 # align to 512B + + elif isinstance(value, list): + for i, t in enumerate(value): + if t is None: + continue + nbytes = self.write_tensor( + handle, f"{name}[{i}]", t, byte_offset, stream + ) + entries.append( + { + "offset": byte_offset, + "shape": list(t.shape), + "dtype": str(t.dtype).replace("torch.", ""), + "list_index": i, + } + ) + byte_offset += nbytes + byte_offset = (byte_offset + 511) & ~511 + + if entries: + manifest[name] = entries + + return manifest + + def read_tensors_from_manifest( + self, + handle: SlotHandle, + manifest: dict[str, list[dict]], + device: torch.device | str = "cpu", + stream: torch.cuda.Stream | None = None, + ) -> dict[str, torch.Tensor | list[torch.Tensor]]: + """Batch-read tensors from a slot using a manifest.""" + result: dict[str, torch.Tensor | list[torch.Tensor]] = {} + + for name, entries in manifest.items(): + if not entries: + continue + has_list_index = any("list_index" in e for e in entries) + + if has_list_index: + max_idx = max(e.get("list_index", 0) for e in entries) + 1 + tensors = [None] * max_idx + for entry in entries: + t = self.read_tensor( + handle, + entry["shape"], + str_to_dtype(entry["dtype"]), + entry["offset"], + device, + stream, + ) + tensors[entry["list_index"]] = t + result[name] = tensors + else: + entry = entries[0] + result[name] = self.read_tensor( + handle, + entry["shape"], + str_to_dtype(entry["dtype"]), + entry["offset"], + device, + stream, + ) + + return result + + def free_slots_count(self, typical_request_size: int) -> int: + """Estimate how many requests of typical size can still be buffered.""" + return self._allocator.count_free_slots(typical_request_size) + + def get_stats(self) -> dict: + alloc_stats = self._allocator.get_stats() + alloc_stats["role"] = self._role_name + return alloc_stats diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/transport/codec.py b/python/sglang/multimodal_gen/runtime/disaggregation/transport/codec.py new file mode 100644 index 000000000..51664b6b4 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/disaggregation/transport/codec.py @@ -0,0 +1,198 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Zero-copy tensor codec for ZMQ multipart messages. + +Frame 0: JSON metadata (tensor descriptors + scalar fields) +Frame 1-N: Raw tensor data buffers (one per tensor) +""" + +import ctypes +import json +import logging +from dataclasses import dataclass + +import torch +import zmq + +logger = logging.getLogger(__name__) + +_DTYPE_TO_STR = { + torch.float16: "float16", + torch.float32: "float32", + torch.float64: "float64", + torch.bfloat16: "bfloat16", + torch.int8: "int8", + torch.int16: "int16", + torch.int32: "int32", + torch.int64: "int64", + torch.uint8: "uint8", + torch.bool: "bool", +} +_STR_TO_DTYPE = {v: k for k, v in _DTYPE_TO_STR.items()} + + +def dtype_to_str(dtype: torch.dtype) -> str: + s = _DTYPE_TO_STR.get(dtype) + if s is None: + raise ValueError(f"Unsupported dtype: {dtype}") + return s + + +def str_to_dtype(s: str) -> torch.dtype: + d = _STR_TO_DTYPE.get(s) + if d is None: + raise ValueError(f"Unknown dtype string: {s}") + return d + + +class TensorWrapper: + """Expose a CPU-contiguous tensor's data buffer for zero-copy ZMQ send.""" + + def __init__(self, tensor: torch.Tensor): + if tensor.is_cuda: + tensor = tensor.cpu() + if not tensor.is_contiguous(): + tensor = tensor.contiguous() + self.tensor = tensor + data_ptr = tensor.data_ptr() + total_bytes = tensor.numel() * tensor.element_size() + self._c_buf = (ctypes.c_char * total_bytes).from_address(data_ptr) + self._view = memoryview(self._c_buf) + + +@dataclass +class TensorDescriptor: + field_name: str + shape: list[int] + dtype: str + list_index: int = -1 # -1 means not part of a list + + def to_dict(self) -> dict: + return { + "field_name": self.field_name, + "shape": self.shape, + "dtype": self.dtype, + "list_index": self.list_index, + } + + @classmethod + def from_dict(cls, d: dict) -> "TensorDescriptor": + return cls( + field_name=d["field_name"], + shape=d["shape"], + dtype=d["dtype"], + list_index=d.get("list_index", -1), + ) + + +def pack_tensors( + tensor_fields: dict[str, torch.Tensor | list[torch.Tensor] | None], + scalar_fields: dict | None = None, +) -> tuple[bytes, list[TensorWrapper]]: + """Pack tensor fields into metadata + buffer list for send_multipart.""" + descriptors = [] + buffers = [] + + for field_name, value in tensor_fields.items(): + if value is None: + continue + + if isinstance(value, torch.Tensor): + wrapper = TensorWrapper(value) + descriptors.append( + TensorDescriptor( + field_name=field_name, + shape=list(value.shape), + dtype=dtype_to_str(value.dtype), + ) + ) + buffers.append(wrapper) + + elif isinstance(value, list): + for i, t in enumerate(value): + if t is None: + continue + if not isinstance(t, torch.Tensor): + raise TypeError( + f"Expected Tensor in list for field '{field_name}', " + f"got {type(t)}" + ) + wrapper = TensorWrapper(t) + descriptors.append( + TensorDescriptor( + field_name=field_name, + shape=list(t.shape), + dtype=dtype_to_str(t.dtype), + list_index=i, + ) + ) + buffers.append(wrapper) + + metadata = { + "tensor_descriptors": [d.to_dict() for d in descriptors], + "scalar_fields": scalar_fields or {}, + } + metadata_bytes = json.dumps(metadata, separators=(",", ":")).encode("utf-8") + return metadata_bytes, buffers + + +def send_tensors( + socket: zmq.Socket, + tensor_fields: dict[str, torch.Tensor | list[torch.Tensor] | None], + scalar_fields: dict | None = None, + flags: int = 0, +) -> None: + """Send tensors over ZMQ using multipart with zero-copy.""" + metadata_bytes, buffers = pack_tensors(tensor_fields, scalar_fields) + parts: list = [metadata_bytes] + parts.extend(w._view if isinstance(w, TensorWrapper) else w for w in buffers) + socket.send_multipart(parts, flags=flags, copy=True) + + +def unpack_tensors( + parts: list, + device: str | torch.device = "cpu", +) -> tuple[dict[str, torch.Tensor | list[torch.Tensor]], dict]: + """Unpack multipart message frames into tensor fields and scalar fields.""" + metadata_frame = parts[0] + metadata_bytes = ( + bytes(metadata_frame.buffer) + if hasattr(metadata_frame, "buffer") + else bytes(metadata_frame) + ) + metadata = json.loads(metadata_bytes) + + descriptors = [ + TensorDescriptor.from_dict(d) for d in metadata["tensor_descriptors"] + ] + scalar_fields = metadata.get("scalar_fields", {}) + + if len(parts) - 1 != len(descriptors): + raise ValueError( + f"Expected {len(descriptors)} tensor frames, got {len(parts) - 1}" + ) + + tensor_fields: dict[str, torch.Tensor | list[torch.Tensor]] = {} + list_sizes: dict[str, int] = {} + for desc in descriptors: + if desc.list_index >= 0: + current_max = list_sizes.get(desc.field_name, 0) + list_sizes[desc.field_name] = max(current_max, desc.list_index + 1) + + for field_name, size in list_sizes.items(): + tensor_fields[field_name] = [None] * size + + for i, desc in enumerate(descriptors): + frame = parts[i + 1] + buf = frame.buffer if hasattr(frame, "buffer") else bytes(frame) + dtype = str_to_dtype(desc.dtype) + # clone() to own the memory (decouple from ZMQ buffer lifetime) + tensor = torch.frombuffer(buf, dtype=dtype).reshape(desc.shape).clone() + if device != "cpu" and device != torch.device("cpu"): + tensor = tensor.to(device) + + if desc.list_index >= 0: + tensor_fields[desc.field_name][desc.list_index] = tensor + else: + tensor_fields[desc.field_name] = tensor + + return tensor_fields, scalar_fields diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/transport/engine.py b/python/sglang/multimodal_gen/runtime/disaggregation/transport/engine.py new file mode 100644 index 000000000..90fdd31af --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/disaggregation/transport/engine.py @@ -0,0 +1,126 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Transfer engine abstraction for tensor transfer between role instances.""" + +import logging +from abc import ABC, abstractmethod + +logger = logging.getLogger(__name__) + +_MOONCAKE_AVAILABLE = None + + +def _check_mooncake() -> bool: + global _MOONCAKE_AVAILABLE + if _MOONCAKE_AVAILABLE is None: + try: + from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import ( # noqa: F401 + MooncakeTransferEngine as _MTE, + ) + + _MOONCAKE_AVAILABLE = True + except ImportError: + _MOONCAKE_AVAILABLE = False + return _MOONCAKE_AVAILABLE + + +class BaseTransferEngine(ABC): + """Abstract transfer engine for data movement between roles.""" + + @property + def supports_gpu_direct(self) -> bool: + return False + + @property + @abstractmethod + def session_id(self) -> str: ... + + @abstractmethod + def register_buffer(self, ptr: int, length: int) -> None: ... + + @abstractmethod + def deregister_buffer(self, ptr: int) -> None: ... + + @abstractmethod + def transfer_sync( + self, dst_session_id: str, src_addr: int, dst_addr: int, length: int + ) -> int: + """Returns 0 on success, negative on failure.""" + + @abstractmethod + def batch_transfer_sync( + self, + dst_session_id: str, + src_addrs: list[int], + dst_addrs: list[int], + lengths: list[int], + ) -> int: ... + + +class MooncakeDiffusionEngine(BaseTransferEngine): + """Production engine backed by MooncakeTransferEngine (RDMA).""" + + @property + def supports_gpu_direct(self) -> bool: + return True + + def __init__( + self, + hostname: str, + gpu_id: int = 0, + ib_device: str | None = None, + ): + from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import ( + MooncakeTransferEngine, + ) + + self._engine = MooncakeTransferEngine( + hostname=hostname, + gpu_id=gpu_id, + ib_device=ib_device, + ) + logger.info( + "MooncakeDiffusionEngine initialized: session_id=%s", + self._engine.session_id, + ) + + @property + def session_id(self) -> str: + return self._engine.session_id + + def register_buffer(self, ptr: int, length: int) -> None: + self._engine.register(ptr, length) + + def deregister_buffer(self, ptr: int) -> None: + self._engine.deregister(ptr) + + def transfer_sync( + self, dst_session_id: str, src_addr: int, dst_addr: int, length: int + ) -> int: + return self._engine.transfer_sync(dst_session_id, src_addr, dst_addr, length) + + def batch_transfer_sync( + self, + dst_session_id: str, + src_addrs: list[int], + dst_addrs: list[int], + lengths: list[int], + ) -> int: + return self._engine.batch_transfer_sync( + dst_session_id, src_addrs, dst_addrs, lengths + ) + + +def create_transfer_engine( + hostname: str = "127.0.0.1", + gpu_id: int = 0, + ib_device: str | None = None, +) -> BaseTransferEngine: + """Factory: returns MooncakeDiffusionEngine if mooncake is available.""" + if not _check_mooncake(): + raise RuntimeError( + "Mooncake transfer engine is required for disaggregated diffusion " + "but is not installed. Please install mooncake first." + ) + return MooncakeDiffusionEngine( + hostname=hostname, gpu_id=gpu_id, ib_device=ib_device + ) diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/transport/manager.py b/python/sglang/multimodal_gen/runtime/disaggregation/transport/manager.py new file mode 100644 index 000000000..9647c0bb4 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/disaggregation/transport/manager.py @@ -0,0 +1,387 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Per-instance transfer manager for disaggregated diffusion roles.""" + +import logging +import threading +from dataclasses import dataclass, field + +import torch + +from sglang.multimodal_gen.runtime.disaggregation.transport.buffer import ( + SlotHandle, + TransferTensorBuffer, +) +from sglang.multimodal_gen.runtime.disaggregation.transport.engine import ( + BaseTransferEngine, +) + +logger = logging.getLogger(__name__) + + +@dataclass +class StagedTransfer: + request_id: str + slot: SlotHandle + manifest: dict + scalar_fields: dict = field(default_factory=dict) + + +@dataclass +class PendingReceive: + request_id: str + slot: SlotHandle + + +class DiffusionTransferManager: + """Manages tensor transfers for a single role instance. + + Owns a TransferTensorBuffer (memory pool) and a BaseTransferEngine (RDMA or mock). + """ + + def __init__( + self, + engine: BaseTransferEngine, + buffer: TransferTensorBuffer, + ): + self._engine = engine + self._buffer = buffer + self._lock = threading.Lock() + + self._engine.register_buffer(self._buffer.pool_data_ptr, self._buffer.pool_size) + + self._staged: dict[str, StagedTransfer] = {} + self._pending_receives: dict[str, PendingReceive] = {} + + logger.info( + "DiffusionTransferManager initialized: session=%s, pool=%d bytes", + self._engine.session_id, + self._buffer.pool_size, + ) + + @property + def session_id(self) -> str: + return self._engine.session_id + + @property + def pool_data_ptr(self) -> int: + return self._buffer.pool_data_ptr + + @property + def pool_size(self) -> int: + return self._buffer.pool_size + + def stage_tensors( + self, + request_id: str, + tensor_fields: dict[str, torch.Tensor | list[torch.Tensor] | None], + scalar_fields: dict | None = None, + stream: torch.cuda.Stream | None = None, + ) -> StagedTransfer | None: + """Stage GPU tensors into the local TransferBuffer. Returns None on allocation failure.""" + total_size = 0 + for name, t in tensor_fields.items(): + if t is None: + continue + if isinstance(t, list): + for ti in t: + total_size += ti.nelement() * ti.element_size() + else: + total_size += t.nelement() * t.element_size() + + if total_size == 0: + staged = StagedTransfer( + request_id=request_id, + slot=None, + manifest={}, + scalar_fields=scalar_fields or {}, + ) + with self._lock: + self._staged[request_id] = staged + return staged + + slot = self._buffer.allocate(total_size, request_id) + if slot is None: + logger.warning( + "TransferManager: failed to allocate %d bytes for %s", + total_size, + request_id, + ) + return None + + manifest = self._buffer.write_tensors_from_gpu(slot, tensor_fields, stream) + + if stream is not None: + stream.synchronize() + elif torch.cuda.is_available(): + torch.cuda.synchronize() + + staged = StagedTransfer( + request_id=request_id, + slot=slot, + manifest=manifest, + scalar_fields=scalar_fields or {}, + ) + with self._lock: + self._staged[request_id] = staged + + logger.debug( + "TransferManager: staged %s (%d bytes, offset=%d)", + request_id, + total_size, + slot.offset, + ) + return staged + + def stage_tensors_async( + self, + request_id: str, + tensor_fields: dict[str, torch.Tensor | list[torch.Tensor] | None], + scalar_fields: dict | None = None, + stream: torch.cuda.Stream | None = None, + ) -> tuple[StagedTransfer | None, torch.cuda.Event | None]: + """Stage GPU tensors, returning a CUDA event instead of blocking. + + Caller MUST wait on the event before reading buffer data. + """ + total_size = 0 + for name, t in tensor_fields.items(): + if t is None: + continue + if isinstance(t, list): + for ti in t: + total_size += ti.nelement() * ti.element_size() + else: + total_size += t.nelement() * t.element_size() + + if total_size == 0: + staged = StagedTransfer( + request_id=request_id, + slot=None, + manifest={}, + scalar_fields=scalar_fields or {}, + ) + with self._lock: + self._staged[request_id] = staged + return staged, None + + slot = self._buffer.allocate(total_size, request_id) + if slot is None: + logger.warning( + "TransferManager: failed to allocate %d bytes for %s", + total_size, + request_id, + ) + return None, None + + manifest = self._buffer.write_tensors_from_gpu(slot, tensor_fields, stream) + + d2h_event = None + if stream is not None: + d2h_event = torch.cuda.Event() + d2h_event.record(stream) + elif torch.cuda.is_available(): + d2h_event = torch.cuda.Event() + d2h_event.record(torch.cuda.current_stream()) + + staged = StagedTransfer( + request_id=request_id, + slot=slot, + manifest=manifest, + scalar_fields=scalar_fields or {}, + ) + with self._lock: + self._staged[request_id] = staged + + logger.debug( + "TransferManager: staged_async %s (%d bytes, offset=%d)", + request_id, + total_size, + slot.offset, + ) + return staged, d2h_event + + def load_tensors_async( + self, + request_id: str, + manifest: dict, + device: torch.device | str = "cuda", + stream: torch.cuda.Stream | None = None, + ) -> tuple[dict[str, torch.Tensor | list[torch.Tensor]], torch.cuda.Event | None]: + """Load tensors from receive slot to GPU, returning a CUDA event. + + Caller MUST wait on the event before using the returned tensors. + """ + with self._lock: + pending = self._pending_receives.get(request_id) + + if pending is None: + raise ValueError( + f"TransferManager: no pending receive slot for {request_id}" + ) + + tensors = self._buffer.read_tensors_from_manifest( + pending.slot, manifest, device=device, stream=stream + ) + + load_event = None + if stream is not None: + load_event = torch.cuda.Event() + load_event.record(stream) + elif torch.cuda.is_available(): + load_event = torch.cuda.Event() + load_event.record(torch.cuda.current_stream()) + + logger.debug( + "TransferManager: loaded_async %d tensor fields for %s to %s", + len(tensors), + request_id, + device, + ) + return tensors, load_event + + def push_to_peer( + self, + request_id: str, + dest_session_id: str, + dest_addr: int, + transfer_size: int, + ) -> bool: + """Push staged data to a remote peer's buffer via RDMA. Returns True on success.""" + with self._lock: + staged = self._staged.get(request_id) + + if staged is None: + logger.error("TransferManager: no staged transfer for %s", request_id) + return False + + if staged.slot is None: + return True + + src_addr = self._buffer.pool_data_ptr + staged.slot.offset + ret = self._engine.transfer_sync( + dest_session_id, src_addr, dest_addr, transfer_size + ) + + if ret == 0: + logger.debug( + "TransferManager: pushed %s (%d bytes) to %s", + request_id, + transfer_size, + dest_session_id, + ) + else: + logger.error( + "TransferManager: RDMA push failed for %s (ret=%d)", + request_id, + ret, + ) + + return ret == 0 + + def free_staged(self, request_id: str) -> None: + with self._lock: + staged = self._staged.pop(request_id, None) + + if staged and staged.slot is not None: + self._buffer.free(staged.slot) + logger.debug("TransferManager: freed staged slot for %s", request_id) + + def allocate_receive_slot( + self, request_id: str, size: int + ) -> PendingReceive | None: + """Allocate a local buffer slot to receive incoming data.""" + slot = self._buffer.allocate(size, request_id) + if slot is None: + logger.warning( + "TransferManager: failed to allocate receive slot (%d bytes) for %s", + size, + request_id, + ) + return None + + pending = PendingReceive(request_id=request_id, slot=slot) + with self._lock: + self._pending_receives[request_id] = pending + + logger.debug( + "TransferManager: allocated receive slot for %s (offset=%d, size=%d)", + request_id, + slot.offset, + slot.size, + ) + return pending + + def load_tensors( + self, + request_id: str, + manifest: dict, + device: torch.device | str = "cuda", + stream: torch.cuda.Stream | None = None, + ) -> dict[str, torch.Tensor | list[torch.Tensor]]: + """Load tensors from a receive slot into GPU memory.""" + with self._lock: + pending = self._pending_receives.get(request_id) + + if pending is None: + raise ValueError( + f"TransferManager: no pending receive slot for {request_id}" + ) + + tensors = self._buffer.read_tensors_from_manifest( + pending.slot, manifest, device=device, stream=stream + ) + + if stream is not None: + stream.synchronize() + elif torch.cuda.is_available(): + torch.cuda.synchronize() + + logger.debug( + "TransferManager: loaded %d tensor fields for %s to %s", + len(tensors), + request_id, + device, + ) + return tensors + + def register_prealloc_as_receive( + self, request_id: str, slot: "SlotHandle" + ) -> "PendingReceive": + """Register a pre-allocated slot as a pending receive (fast path).""" + pending = PendingReceive(request_id=request_id, slot=slot) + with self._lock: + self._pending_receives[request_id] = pending + return pending + + def free_receive_slot(self, request_id: str) -> None: + with self._lock: + pending = self._pending_receives.pop(request_id, None) + + if pending: + self._buffer.free(pending.slot) + logger.debug("TransferManager: freed receive slot for %s", request_id) + + def get_receive_slot_addr(self, request_id: str) -> int | None: + with self._lock: + pending = self._pending_receives.get(request_id) + if pending is None: + return None + return self._buffer.pool_data_ptr + pending.slot.offset + + def get_receive_slot_offset(self, request_id: str) -> int | None: + with self._lock: + pending = self._pending_receives.get(request_id) + if pending is None: + return None + return pending.slot.offset + + def get_staged_info(self, request_id: str) -> StagedTransfer | None: + with self._lock: + return self._staged.get(request_id) + + def free_slots_count(self, typical_size: int = 64 * 1024 * 1024) -> int: + return self._buffer.free_slots_count(typical_size) + + def cleanup(self) -> None: + self._engine.deregister_buffer(self._buffer.pool_data_ptr) + logger.info("DiffusionTransferManager cleaned up") diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/transport/protocol.py b/python/sglang/multimodal_gen/runtime/disaggregation/transport/protocol.py new file mode 100644 index 000000000..347bf0be5 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/disaggregation/transport/protocol.py @@ -0,0 +1,145 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Transfer protocol messages for disaggregated diffusion. + +All messages are sent as ZMQ multipart with a b"__transfer__" discriminator +in frame[0] and JSON payload in frame[1]. +""" + +import json +import logging +from dataclasses import asdict, dataclass, field +from typing import Any + +logger = logging.getLogger(__name__) + +TRANSFER_MAGIC = b"__transfer__" + + +class TransferMsgType: + # Instance → DiffusionServer + STAGED = "transfer_staged" + ALLOCATED = "transfer_allocated" + PUSHED = "transfer_pushed" + DONE = "transfer_done" + + # DiffusionServer → Instance + ALLOC = "transfer_alloc" + PUSH = "transfer_push" + READY = "transfer_ready" + + # Registration + REGISTER = "transfer_register" + REGISTER_ACK = "transfer_register_ack" + + +@dataclass +class TransferStagedMsg: + msg_type: str = TransferMsgType.STAGED + request_id: str = "" + data_size: int = 0 + manifest: dict = None + session_id: str = "" + pool_ptr: int = 0 + slot_offset: int = 0 + + def __post_init__(self): + if self.manifest is None: + self.manifest = {} + + +@dataclass +class TransferAllocMsg: + msg_type: str = TransferMsgType.ALLOC + request_id: str = "" + data_size: int = 0 + source_role: str = "" + + +@dataclass +class TransferAllocatedMsg: + msg_type: str = TransferMsgType.ALLOCATED + request_id: str = "" + session_id: str = "" + pool_ptr: int = 0 + slot_offset: int = 0 + slot_size: int = 0 + + +@dataclass +class TransferPushMsg: + msg_type: str = TransferMsgType.PUSH + request_id: str = "" + dest_session_id: str = "" + dest_addr: int = 0 + transfer_size: int = 0 + + +@dataclass +class TransferPushedMsg: + msg_type: str = TransferMsgType.PUSHED + request_id: str = "" + + +@dataclass +class TransferReadyMsg: + msg_type: str = TransferMsgType.READY + request_id: str = "" + manifest: dict = None + slot_offset: int = 0 + scalar_fields: dict = None + + def __post_init__(self): + if self.manifest is None: + self.manifest = {} + if self.scalar_fields is None: + self.scalar_fields = {} + + +@dataclass +class TransferDoneMsg: + msg_type: str = TransferMsgType.DONE + request_id: str = "" + error: str | None = None + + +@dataclass +class TransferRegisterMsg: + msg_type: str = TransferMsgType.REGISTER + role: str = "" + session_id: str = "" + pool_ptr: int = 0 + pool_size: int = 0 + # The instance's own work endpoint (e.g. tcp://host:port). Used by the + # DiffusionServer to key peer info by URL index (i.e. the same index used + # to build the PUSH work-socket list), so the control plane and the RDMA + # data plane cannot drift when instances register in a different order + # than --*-urls. + work_endpoint: str = "" + # Pre-allocated receive slots: [{"offset": int, "size": int, "slot_id": int, "addr": int}] + preallocated_slots: list = field(default_factory=list) + + +def encode_transfer_msg(msg: Any) -> list[bytes]: + """Encode as [TRANSFER_MAGIC, json_payload_bytes].""" + if hasattr(msg, "__dataclass_fields__"): + d = asdict(msg) + elif isinstance(msg, dict): + d = msg + else: + raise TypeError(f"Cannot encode transfer message: {type(msg)}") + + return [TRANSFER_MAGIC, json.dumps(d, separators=(",", ":")).encode("utf-8")] + + +def decode_transfer_msg(frames: list[bytes]) -> dict: + if len(frames) < 2 or frames[0] != TRANSFER_MAGIC: + raise ValueError(f"Not a transfer message: frame[0]={frames[0]!r}") + return json.loads(frames[1]) + + +def is_transfer_message(frames: list) -> bool: + return len(frames) >= 2 and ( + frames[0] == TRANSFER_MAGIC + or (isinstance(frames[0], memoryview) and bytes(frames[0]) == TRANSFER_MAGIC) + or (hasattr(frames[0], "bytes") and frames[0].bytes == TRANSFER_MAGIC) + ) diff --git a/python/sglang/multimodal_gen/runtime/distributed/group_coordinator.py b/python/sglang/multimodal_gen/runtime/distributed/group_coordinator.py index cabd056ff..cda96325d 100644 --- a/python/sglang/multimodal_gen/runtime/distributed/group_coordinator.py +++ b/python/sglang/multimodal_gen/runtime/distributed/group_coordinator.py @@ -435,9 +435,9 @@ class GroupCoordinator: # Bypass the function if we are using only 1 GPU. if self.world_size == 1: return obj - if self.shm_broadcaster is not None: + if self.mq_broadcaster is not None: assert src == 0, "Shared memory broadcaster only supports src=0" - return self.shm_broadcaster.broadcast_object(obj) + return self.mq_broadcaster.broadcast_object(obj) if self.rank_in_group == src: torch.distributed.broadcast_object_list( [obj], src=self.ranks[src], group=self.cpu_group diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/cli/serve.py b/python/sglang/multimodal_gen/runtime/entrypoints/cli/serve.py index baa908c77..a5171d8e5 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/cli/serve.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/cli/serve.py @@ -8,7 +8,9 @@ from typing import cast from sglang.multimodal_gen.apps.webui import run_sgl_diffusion_webui from sglang.multimodal_gen.runtime.entrypoints.cli.cli_types import CLISubcommand -from sglang.multimodal_gen.runtime.launch_server import launch_server +from sglang.multimodal_gen.runtime.launch_server import ( + dispatch_launch, +) from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.utils import FlexibleArgumentParser @@ -31,7 +33,8 @@ def add_multimodal_gen_serve_args(parser: argparse.ArgumentParser): def execute_serve_cmd(args: argparse.Namespace, unknown_args: list[str] | None = None): """The entry point for the serve command.""" server_args = ServerArgs.from_cli_args(args, unknown_args) - launch_server(server_args) + + dispatch_launch(server_args) if server_args.webui: run_sgl_diffusion_webui(server_args) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py b/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py index cc6af08fc..63162cb66 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py @@ -162,6 +162,32 @@ async def health_generate(): return {"status": "ok"} +@health_router.get("/stats") +async def stats_endpoint(request: Request): + """Get runtime statistics including disagg pipeline metrics. + + Returns queue depth, request counts, latency, throughput, etc. + Sends a GetDisaggStatsReq to the scheduler via ZMQ and returns the result. + """ + from sglang.multimodal_gen.runtime.entrypoints.utils import GetDisaggStatsReq + + server_args: ServerArgs = request.app.state.server_args + response: dict = { + "status": "ok", + "model_path": server_args.model_path, + } + + # Query the scheduler for disagg metrics + try: + stats_response = await async_scheduler_client.forward(GetDisaggStatsReq()) + if hasattr(stats_response, "output") and stats_response.output is not None: + response["disagg"] = stats_response.output + except Exception as e: + response["disagg"] = {"error": str(e)} + + return response + + def make_serializable(obj): """Recursively converts Tensors to None for JSON serialization.""" if isinstance(obj, torch.Tensor): diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/utils.py b/python/sglang/multimodal_gen/runtime/entrypoints/utils.py index 765a38f7f..a8b210bbd 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/utils.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/utils.py @@ -69,6 +69,13 @@ class ShutdownReq: pass +@dataclass +class GetDisaggStatsReq: + """Request to get disagg pipeline metrics from the scheduler.""" + + pass + + def format_lora_message( lora_nickname: Union[str, List[str]], target: Union[str, List[str]], diff --git a/python/sglang/multimodal_gen/runtime/launch_server.py b/python/sglang/multimodal_gen/runtime/launch_server.py index 5b60c844a..3a1bafa69 100644 --- a/python/sglang/multimodal_gen/runtime/launch_server.py +++ b/python/sglang/multimodal_gen/runtime/launch_server.py @@ -1,5 +1,6 @@ # Copied and adapted from: https://github.com/hao-ai-lab/FastVideo +import dataclasses import multiprocessing as mp import os import signal @@ -9,6 +10,10 @@ import threading import psutil import uvicorn +from sglang.multimodal_gen.runtime.disaggregation.orchestrator import ( + DiffusionServer, +) +from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType from sglang.multimodal_gen.runtime.entrypoints.http_server import create_app from sglang.multimodal_gen.runtime.managers.gpu_worker import run_scheduler_process from sglang.multimodal_gen.runtime.server_args import ( @@ -16,9 +21,28 @@ from sglang.multimodal_gen.runtime.server_args import ( prepare_server_args, set_global_server_args, ) +from sglang.multimodal_gen.runtime.utils.common import is_port_available from sglang.multimodal_gen.runtime.utils.logging_utils import configure_logger, logger +def _find_available_port( + start: int = 10000, avoid: set[int] | None = None, max_attempts: int = 100 +) -> int: + """Find an available port starting from *start*, skipping ports in *avoid*.""" + if avoid is None: + avoid = set() + port = max(1024, min(start, 65535)) + for _ in range(max_attempts): + if port not in avoid and is_port_available(port): + return port + port += 1 + if port > 65535: + port = 1024 + raise RuntimeError( + f"No available port found after {max_attempts} attempts (start={start})" + ) + + def kill_process_tree(parent_pid, include_parent: bool = True, skip_pid: int = None): """Kill the process and all its child processes.""" # Remove sigchld handler to avoid spammy logs. @@ -185,6 +209,237 @@ def launch_server(server_args: ServerArgs, launch_http_server: bool = True): return processes +def launch_pool_disagg_server( + server_args: ServerArgs, + encoder_gpus: list[list[int]], + denoiser_gpus: list[list[int]], + decoder_gpus: list[list[int]], + launch_http_server: bool = True, +): + """Launch a pool-based disaggregated server with N:M:K independent role instances. + + DiffusionServer orchestrates the full pipeline, dispatching at every + role transition (Encoder → Denoiser → Decoder). + + Args: + server_args: Base server configuration + encoder_gpus: List of GPU ID lists, one per encoder instance. + e.g., [[0], [2]] for 2 encoder instances on GPUs 0 and 2. + denoiser_gpus: List of GPU ID lists, one per denoiser instance. + e.g., [[1], [3]] for 2 denoiser instances. + decoder_gpus: List of GPU ID lists, one per decoder instance. + e.g., [[0], [2]] for 2 decoder instances (can share with encoder). + launch_http_server: Whether to launch the HTTP server. + + Example: + launch_pool_disagg_server(server_args, + encoder_gpus=[[0], [2]], + denoiser_gpus=[[1], [3]], + decoder_gpus=[[0], [2]], + ) + """ + configure_logger(server_args) + + num_encoders = len(encoder_gpus) + num_denoisers = len(denoiser_gpus) + num_decoders = len(decoder_gpus) + logger.info( + "Starting pool disagg server: %d encoder(s), %d denoiser(s), %d decoder(s)...", + num_encoders, + num_denoisers, + num_decoders, + ) + + host = server_args.host or "127.0.0.1" + + def find_port(start): + return _find_available_port(start) + + # Allocate endpoints + port_cursor = server_args.scheduler_port + 3000 + + # Per-instance work endpoints (instance binds PULL, DS connects PUSH) + encoder_work_endpoints = [] + for i in range(num_encoders): + p = find_port(port_cursor) + encoder_work_endpoints.append(f"tcp://{host}:{p}") + port_cursor = p + 1 + + denoiser_work_endpoints = [] + for i in range(num_denoisers): + p = find_port(port_cursor) + denoiser_work_endpoints.append(f"tcp://{host}:{p}") + port_cursor = p + 1 + + decoder_work_endpoints = [] + for i in range(num_decoders): + p = find_port(port_cursor) + decoder_work_endpoints.append(f"tcp://{host}:{p}") + port_cursor = p + 1 + + # Per-role-type result endpoints (DS binds PULL, instances connect PUSH) + # Use deterministic convention: scheduler_port + {1,2,3} + base_port = server_args.scheduler_port + encoder_result_ep = f"tcp://{host}:{base_port + 1}" + denoiser_result_ep = f"tcp://{host}:{base_port + 2}" + decoder_result_ep = f"tcp://{host}:{base_port + 3}" + + logger.info( + "Pool endpoints allocated: %d work + 3 result endpoints", + num_encoders + num_denoisers + num_decoders, + ) + + # Launch all role instances + all_processes = [] + + role_configs = [ + (RoleType.ENCODER, encoder_gpus, encoder_work_endpoints, encoder_result_ep), + ( + RoleType.DENOISER, + denoiser_gpus, + denoiser_work_endpoints, + denoiser_result_ep, + ), + (RoleType.DECODER, decoder_gpus, decoder_work_endpoints, decoder_result_ep), + ] + + for role_type, gpu_lists, work_eps, result_ep in role_configs: + for inst_idx, gpu_ids in enumerate(gpu_lists): + num_role_gpus = len(gpu_ids) + + # Per-role parallelism: use explicit overrides if set, else None (auto-derive) + role_par = server_args.get_role_parallelism(role_type) + + role_overrides = { + "disagg_role": role_type, + "disagg_mode": True, + "pool_work_endpoint": work_eps[inst_idx], + "pool_result_endpoint": result_ep, + "num_gpus": num_role_gpus, + "warmup": role_type == RoleType.ENCODER, + "scheduler_port": find_port(port_cursor), + "master_port": find_port(port_cursor + 100), + # Per-role parallelism (None = auto-derive from num_gpus) + "tp_size": role_par["tp_size"], + "sp_degree": role_par["sp_degree"], + "ulysses_degree": role_par["ulysses_degree"], + "ring_degree": role_par["ring_degree"], + } + port_cursor = role_overrides["master_port"] + 100 + + base_dict = { + f.name: getattr(server_args, f.name) + for f in dataclasses.fields(server_args) + } + base_dict.update(role_overrides) + base_dict.pop("pipeline_config", None) + role_args = ServerArgs.from_kwargs(**base_dict) + + pool_ctx = mp.get_context("spawn") + inst_readers = [] + + # Spawn all ranks first — NCCL init blocks until all ranks connect + for rank_idx in range(num_role_gpus): + reader, writer = pool_ctx.Pipe(duplex=False) + gpu_id = gpu_ids[rank_idx] + + process = pool_ctx.Process( + target=_run_disagg_role_process, + args=(gpu_id, rank_idx, rank_idx, role_args, writer, [], []), + name=f"sglang-pool-{role_type.value}-{inst_idx}-r{rank_idx}", + daemon=True, + ) + process.start() + all_processes.append(process) + inst_readers.append(reader) + + # Wait for all ranks to be ready (after all are spawned) + for rank_idx, reader in enumerate(inst_readers): + try: + data = reader.recv() + except EOFError: + logger.error( + "Pool %s[%d] rank %d is dead.", + role_type.value, + inst_idx, + rank_idx, + ) + raise + if data.get("status") != "ready": + raise RuntimeError( + f"Pool {role_type.value}[{inst_idx}] rank {rank_idx} " + "failed to initialize." + ) + reader.close() + + logger.info( + "Pool %s[%d] ready on GPU(s) %s (work=%s)", + role_type.value.upper(), + inst_idx, + gpu_ids, + work_eps[inst_idx], + ) + + logger.info("All pool role instances ready") + + # Start DiffusionServer + frontend_endpoint = f"tcp://{host}:{server_args.scheduler_port}" + + diffusion_server = DiffusionServer( + frontend_endpoint=frontend_endpoint, + encoder_work_endpoints=encoder_work_endpoints, + denoiser_work_endpoints=denoiser_work_endpoints, + decoder_work_endpoints=decoder_work_endpoints, + encoder_result_endpoint=encoder_result_ep, + denoiser_result_endpoint=denoiser_result_ep, + decoder_result_endpoint=decoder_result_ep, + dispatch_policy_name=server_args.disagg_dispatch_policy, + timeout_s=float(server_args.disagg_timeout), + ) + diffusion_server.start() + + if not diffusion_server.wait_ready(timeout=30.0): + raise RuntimeError("DiffusionServer failed to bind sockets within 30 seconds") + + if launch_http_server: + logger.info( + "Starting FastAPI server (connected to DiffusionServer at port %d).", + server_args.scheduler_port, + ) + launch_http_server_only(server_args) + + return all_processes + + +def _run_disagg_role_process( + gpu_id: int, + _local_rank: int, + rank: int, + server_args: ServerArgs, + pipe_writer: mp.connection.Connection, + task_pipes: list, + result_pipes: list, +): + """Entry point for a disagg role process. + + Uses the physical GPU index (gpu_id) as local_rank so that + torch.cuda.set_device(local_rank) selects the correct GPU. + This avoids relying on CUDA_VISIBLE_DEVICES remapping, which + may not work if CUDA was pre-initialized in the parent process. + """ + run_scheduler_process( + local_rank=gpu_id, + rank=rank, + master_port=server_args.master_port, + server_args=server_args, + pipe_writer=pipe_writer, + task_pipe_r=None, + result_pipe_w=None, + task_pipes_to_slaves=task_pipes, + result_pipes_from_slaves=result_pipes, + ) + + def launch_http_server_only(server_args): # set for endpoints to access global_server_args set_global_server_args(server_args) @@ -199,10 +454,224 @@ def launch_http_server_only(server_args): ) +def parse_url_string(url_str: str) -> list[str]: + """Parse a semicolon-separated URL string into a list. + + Example: "tcp://10.0.0.1:35000;tcp://10.0.0.2:35000" -> ["tcp://...", "tcp://..."] + """ + return [u.strip() for u in url_str.split(";") if u.strip()] + + +def launch_disagg_server(server_args: ServerArgs): + """Launch DiffusionServer head node + HTTP server (--disagg-role server). + + No GPU workers are spawned. Connects to remote role instances + specified by --encoder-urls, --denoiser-urls, --decoder-urls. + + Result endpoints use deterministic convention: + encoder result: scheduler_port + 1 + denoiser result: scheduler_port + 2 + decoder result: scheduler_port + 3 + """ + configure_logger(server_args) + + for name, val in [ + ("--encoder-urls", server_args.encoder_urls), + ("--denoiser-urls", server_args.denoiser_urls), + ("--decoder-urls", server_args.decoder_urls), + ]: + if val is None: + raise ValueError(f"{name} is required for --disagg-role server") + + host = server_args.host or "127.0.0.1" + base_port = server_args.scheduler_port + + encoder_work_endpoints = parse_url_string(server_args.encoder_urls) + denoiser_work_endpoints = parse_url_string(server_args.denoiser_urls) + decoder_work_endpoints = parse_url_string(server_args.decoder_urls) + + encoder_result_ep = f"tcp://{host}:{base_port + 1}" + denoiser_result_ep = f"tcp://{host}:{base_port + 2}" + decoder_result_ep = f"tcp://{host}:{base_port + 3}" + + frontend_endpoint = f"tcp://{host}:{base_port}" + + logger.info( + "Starting DiffusionServer: %d encoder(s), %d denoiser(s), %d decoder(s)", + len(encoder_work_endpoints), + len(denoiser_work_endpoints), + len(decoder_work_endpoints), + ) + logger.info(" Frontend: %s", frontend_endpoint) + logger.info(" Encoder work endpoints: %s", encoder_work_endpoints) + logger.info(" Denoiser work endpoints: %s", denoiser_work_endpoints) + logger.info(" Decoder work endpoints: %s", decoder_work_endpoints) + logger.info( + " Result endpoints: encoder=%s, denoiser=%s, decoder=%s", + encoder_result_ep, + denoiser_result_ep, + decoder_result_ep, + ) + + diffusion_server = DiffusionServer( + frontend_endpoint=frontend_endpoint, + encoder_work_endpoints=encoder_work_endpoints, + denoiser_work_endpoints=denoiser_work_endpoints, + decoder_work_endpoints=decoder_work_endpoints, + encoder_result_endpoint=encoder_result_ep, + denoiser_result_endpoint=denoiser_result_ep, + decoder_result_endpoint=decoder_result_ep, + dispatch_policy_name=server_args.disagg_dispatch_policy, + timeout_s=float(server_args.disagg_timeout), + ) + diffusion_server.start() + + if not diffusion_server.wait_ready(timeout=30.0): + raise RuntimeError("DiffusionServer failed to bind sockets within 30 seconds") + + logger.info( + "Starting HTTP server (connected to DiffusionServer at port %d).", + base_port, + ) + launch_http_server_only(server_args) + + +def launch_disagg_role(server_args: ServerArgs): + """Launch a standalone disaggregated role instance (--disagg-role encoder/denoising/decoder). + + The instance: + 1. Binds its work PULL socket on tcp://0.0.0.0:{scheduler_port} + 2. Connects its result PUSH socket to the DiffusionServer head node + (derived from --disagg-server-addr + role offset) + 3. Spawns GPU worker processes for the assigned role. + """ + configure_logger(server_args) + + role_type = server_args.disagg_role + if server_args.disagg_server_addr is None: + raise ValueError( + "--disagg-server-addr is required for --disagg-role " f"{role_type.value}" + ) + + # Derive endpoints + work_endpoint = server_args.derive_pool_work_endpoint() + result_endpoint = server_args.derive_pool_result_endpoint() + + logger.info( + "Starting disagg role: %s, num_gpus=%d", + role_type.value, + server_args.num_gpus, + ) + logger.info(" Work endpoint (bind): %s", work_endpoint) + logger.info(" Result endpoint (connect): %s", result_endpoint) + logger.info( + " P2P: hostname=%s, ib_device=%s, pool_size=%d", + server_args.disagg_p2p_hostname, + server_args.disagg_ib_device, + server_args.disagg_transfer_pool_size, + ) + + # Build role-specific ServerArgs + # Use a different port for the scheduler's internal ROUTER socket to avoid + # conflicting with the pool work PULL socket (both bind on scheduler_port). + internal_scheduler_port = _find_available_port( + start=server_args.scheduler_port + 100, avoid={server_args.scheduler_port} + ) + + role_par = server_args.get_role_parallelism(role_type) + role_overrides = { + "disagg_role": role_type, + "disagg_mode": True, + "pool_work_endpoint": work_endpoint, + "pool_result_endpoint": result_endpoint, + "warmup": role_type == RoleType.ENCODER, + "scheduler_port": internal_scheduler_port, + # Per-role parallelism (None = auto-derive from num_gpus) + "tp_size": role_par["tp_size"], + "sp_degree": role_par["sp_degree"], + "ulysses_degree": role_par["ulysses_degree"], + "ring_degree": role_par["ring_degree"], + } + + base_dict = { + f.name: getattr(server_args, f.name) for f in dataclasses.fields(server_args) + } + base_dict.update(role_overrides) + base_dict.pop("pipeline_config", None) + role_args = ServerArgs.from_kwargs(**base_dict) + + # Spawn GPU worker processes + # NOTE: All ranks must be spawned before waiting for ready signals, + # because NCCL init_process_group blocks until all ranks connect. + num_gpus = server_args.num_gpus + base_gpu_id = server_args.base_gpu_id + pool_ctx = mp.get_context("spawn") + processes = [] + readers = [] + + for rank_idx in range(num_gpus): + reader, writer = pool_ctx.Pipe(duplex=False) + gpu_id = base_gpu_id + rank_idx + + process = pool_ctx.Process( + target=_run_disagg_role_process, + args=(gpu_id, rank_idx, rank_idx, role_args, writer, [], []), + name=f"sglang-{role_type.value}-r{rank_idx}", + daemon=True, + ) + process.start() + processes.append(process) + readers.append(reader) + + # Wait for all ranks to be ready (after all are spawned) + for rank_idx, reader in enumerate(readers): + try: + data = reader.recv() + except EOFError: + logger.error( + "Role %s rank %d is dead.", + role_type.value, + rank_idx, + ) + raise + if data.get("status") != "ready": + raise RuntimeError( + f"Role {role_type.value} rank {rank_idx} failed to initialize." + ) + reader.close() + + logger.info( + "Role %s ready (%d GPU(s), work=%s)", + role_type.value.upper(), + num_gpus, + work_endpoint, + ) + + # Block until interrupted + try: + for p in processes: + p.join() + except KeyboardInterrupt: + logger.info("Role %s shutting down.", role_type.value) + + +def dispatch_launch(server_args: ServerArgs): + """Route to the correct launch function based on --disagg-role.""" + role = server_args.disagg_role + if role == RoleType.MONOLITHIC: + launch_server(server_args) + elif role == RoleType.SERVER: + launch_disagg_server(server_args) + elif role in (RoleType.ENCODER, RoleType.DENOISER, RoleType.DECODER): + launch_disagg_role(server_args) + else: + raise ValueError(f"Unknown disagg_role: {role}") + + if __name__ == "__main__": server_args = prepare_server_args(sys.argv[1:]) try: - launch_server(server_args) + dispatch_launch(server_args) finally: kill_process_tree(os.getpid(), include_parent=False) diff --git a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py index 5ea209060..daa2b6e31 100644 --- a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py +++ b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py @@ -208,9 +208,16 @@ class GPUWorker: f"Related offload server args to disable: {suggested_args_str}" ) - def execute_forward(self, batch: List[Req]) -> OutputBatch: + def execute_forward( + self, batch: List[Req], return_req: bool = False + ) -> OutputBatch | Req: """ Execute a forward pass. + + Args: + batch: List of requests to process. + return_req: If True, return the raw Req instead of OutputBatch. + Used by disaggregated pipelines to access intermediate tensors. """ assert self.pipeline is not None req = batch[0] @@ -229,6 +236,11 @@ class GPUWorker: req.log(server_args=self.server_args) result = self.pipeline.forward(req, self.server_args) + # For disagg roles, return raw Req to let the caller handle + # the role-to-role tensor transfer before OutputBatch conversion. + if return_req and isinstance(result, Req): + return result + if isinstance(result, Req): output_batch = OutputBatch( output=result.output, @@ -261,7 +273,8 @@ class GPUWorker: self.do_mem_analysis(output_batch) duration_ms = (time.monotonic() - start_time) * 1000 - output_batch.metrics.total_duration_ms = duration_ms + if output_batch.metrics is not None: + output_batch.metrics.total_duration_ms = duration_ms # Save output to file and return file path only if requested. Avoid the serialization # and deserialization overhead between scheduler_client and gpu_worker. @@ -526,6 +539,7 @@ def run_scheduler_process( port_args=port_args, task_pipes_to_slaves=task_pipes_to_slaves, result_pipes_from_slaves=result_pipes_from_slaves, + local_rank=local_rank, ) logger.info(f"Worker {rank}: Scheduler loop started.") pipe_writer.send( diff --git a/python/sglang/multimodal_gen/runtime/managers/scheduler.py b/python/sglang/multimodal_gen/runtime/managers/scheduler.py index 1c26655f9..a3f1d379c 100644 --- a/python/sglang/multimodal_gen/runtime/managers/scheduler.py +++ b/python/sglang/multimodal_gen/runtime/managers/scheduler.py @@ -10,6 +10,10 @@ from typing import Any, List import zmq +from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType +from sglang.multimodal_gen.runtime.disaggregation.scheduler_mixin import ( + SchedulerDisaggMixin, +) from sglang.multimodal_gen.runtime.distributed import get_world_group from sglang.multimodal_gen.runtime.entrypoints.openai.utils import ( _parse_size, @@ -20,6 +24,7 @@ from sglang.multimodal_gen.runtime.entrypoints.post_training.io_struct import ( UpdateWeightFromDiskReqInput, ) from sglang.multimodal_gen.runtime.entrypoints.utils import ( + GetDisaggStatsReq, ListLorasReq, MergeLoraWeightsReq, SetLoraReq, @@ -43,7 +48,7 @@ logger = init_logger(__name__) MINIMUM_PICTURE_BASE64_FOR_WARMUP = "data:image/jpg;base64,iVBORw0KGgoAAAANSUhEUgAAACAAAAAgCAYAAABzenr0AAAACXBIWXMAAA7EAAAOxAGVKw4bAAAAbUlEQVRYhe3VsQ2AMAxE0Y/lIgNQULD/OqyCMgCihCKSG4yRuKuiNH6JLsoEbMACOGBcua9HOR7Y6w6swBwMy0qLTpkeI77qdEBpBFAHBBDAGH8WrwJKI4AAegUCfAKgEgpQDvh3CR3oQCuav58qlAw73kKCSgAAAABJRU5ErkJggg==" -class Scheduler: +class Scheduler(SchedulerDisaggMixin): """ Runs the main event loop for the rank 0 worker. It listens for external requests via ZMQ and coordinates with other workers. @@ -57,10 +62,17 @@ class Scheduler: port_args: PortArgs, task_pipes_to_slaves: list = None, result_pipes_from_slaves: list = None, + local_rank: int | None = None, ): self.server_args = server_args self.port_args = port_args + # local_rank is the physical GPU index for torch.cuda.set_device. + # In non-disagg mode, it equals gpu_id. In disagg mode, it may differ + # (e.g., denoiser rank 0 on physical GPU 1). + if local_rank is None: + local_rank = gpu_id + set_global_server_args(server_args=server_args) # Inter-process Communication @@ -76,7 +88,7 @@ class Scheduler: self.receiver = None worker = GPUWorker( - local_rank=gpu_id, + local_rank=local_rank, master_port=port_args.master_port, rank=gpu_id, server_args=server_args, @@ -95,6 +107,7 @@ class Scheduler: List[Req]: self._handle_generation, ListLorasReq: self._handle_list_loras, ShutdownReq: self._handle_shutdown, + GetDisaggStatsReq: self._handle_get_disagg_stats, UpdateWeightFromDiskReqInput: self._handle_update_weights_from_disk, GetWeightsChecksumReqInput: self._handle_get_weights_checksum, } @@ -114,6 +127,21 @@ class Scheduler: self._max_consecutive_errors = 3 self._consecutive_error_count = 0 + self._init_disagg_state(server_args, local_rank) + + def get_disagg_metrics(self) -> dict | None: + """Return disagg role metrics snapshot, or None if not in disagg mode.""" + if self._disagg_metrics is None: + return None + return self._disagg_metrics.snapshot().to_dict() + + def _handle_get_disagg_stats(self, _reqs: List[Any]) -> OutputBatch: + """Handle stats request — return disagg metrics via OutputBatch.output.""" + stats = self.get_disagg_metrics() + return OutputBatch( + output=stats or {"role": "monolithic", "message": "not in disagg mode"} + ) + def _handle_set_lora(self, reqs: List[Any]) -> OutputBatch: # TODO: return set status # TODO: return with SetLoRAResponse or something more appropriate @@ -166,6 +194,7 @@ class Scheduler: ) else: logger.info("Processing warmup req...") + return self.worker.execute_forward(reqs) def return_result( @@ -362,12 +391,20 @@ class Scheduler: The main event loop that listens for ZMQ requests. Handles abortion """ + # Pool mode: all roles use the pool event loop + if self._disagg_role != RoleType.MONOLITHIC: + self._disagg_event_loop() + return logger.debug( f"Rank 0 scheduler listening on tcp://*:{self.server_args.scheduler_port}" ) while self._running: + # Update queue depth for metrics + if self._disagg_metrics: + self._disagg_metrics.update_queue_depth(len(self.waiting_queue)) + # 1: receive requests try: new_reqs = self.recv_reqs() @@ -403,6 +440,10 @@ class Scheduler: try: processed_req = reqs[0] + is_warmup = ( + processed_req.is_warmup if isinstance(processed_req, Req) else False + ) + handler = self.request_handlers.get(type(processed_req)) if handler: output_batch = handler(reqs) @@ -415,16 +456,10 @@ class Scheduler: f"Error executing request in scheduler event loop: {e}", exc_info=True, ) - # Determine appropriate error response format - output_batch = ( - OutputBatch(error=str(e)) - if reqs and isinstance(reqs[0], Req) - else OutputBatch(error=str(e)) - ) + output_batch = OutputBatch(error=str(e)) # 3. return results try: - # log warmup info is_warmup = ( processed_req.is_warmup if isinstance(processed_req, Req) else False ) @@ -457,6 +492,7 @@ class Scheduler: if self.receiver is not None: self.receiver.close() + self._cleanup_disagg() self.context.destroy(linger=0) def _broadcast_task(self, payload: dict[str, Any]) -> None: diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py b/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py index 732233c48..260aec33e 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py @@ -14,6 +14,10 @@ from typing import Any, Callable, Literal, cast import torch from tqdm import tqdm +from sglang.multimodal_gen.runtime.disaggregation.roles import ( + RoleType, + filter_modules_for_role, +) from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import ( PipelineComponentLoader, ) @@ -82,6 +86,7 @@ class ComposedPipelineBase(ABC): use. The pipeline should be stateless and not hold any batch state. """ self.server_args = server_args + self._disagg_role = server_args.disagg_role self.model_path: str = model_path self._stages: list[PipelineStage] = [] @@ -94,6 +99,20 @@ class ComposedPipelineBase(ABC): if self._required_config_modules is None: raise NotImplementedError("Subclass must set _required_config_modules") + # Filter modules based on disaggregation role + if self._disagg_role != RoleType.MONOLITHIC: + original_modules = list(self._required_config_modules) + self._required_config_modules = filter_modules_for_role( + self._required_config_modules, self._disagg_role + ) + skipped = set(original_modules) - set(self._required_config_modules) + if skipped: + logger.info( + "Disagg role=%s: skipping modules %s", + self._disagg_role.value, + sorted(skipped), + ) + # [module_name, gpu memory usage] self.memory_usages: dict[str, float] = {} # Load modules directly in initialization @@ -169,6 +188,70 @@ class ComposedPipelineBase(ABC): """ return + # --- Config-name → pipeline_config attribute mapping --- + _CONFIG_ATTR_MAP: dict[str, str] = { + "vae": "vae_config", + "video_vae": "vae_config", + "audio_vae": "audio_vae_config", + } + + def _init_skipped_component_configs( + self, + full_model_index: dict[str, Any], + server_args: ServerArgs, + ) -> None: + """Read HF JSON configs for skipped components and run + update_model_arch + post_init so pipeline_config is fully + initialized without loading weights. + """ + from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import ( + get_diffusers_component_config, + ) + + required = set(self.required_config_modules) + for module_name in full_model_index: + if module_name in required: + continue # will be loaded normally + cfg_attr = self._CONFIG_ATTR_MAP.get(module_name) + if cfg_attr is None: + continue # not a config we need to patch + + pipeline_cfg = getattr(server_args.pipeline_config, cfg_attr, None) + if pipeline_cfg is None: + continue + + try: + component_path = self._resolve_component_path( + server_args, module_name, module_name + ) + hf_config = get_diffusers_component_config( + component_path=component_path + ) + hf_config.pop("_class_name", None) + hf_config.pop("_diffusers_version", None) + pipeline_cfg.update_model_arch(hf_config) + if hasattr(pipeline_cfg, "post_init"): + pipeline_cfg.post_init() + logger.info( + "Disagg role=%s: initialized %s config from HF JSON " + "(spatial_compression_ratio=%s)", + self._disagg_role.value, + module_name, + getattr( + getattr(pipeline_cfg, "arch_config", None), + "spatial_compression_ratio", + "N/A", + ), + ) + except Exception as e: + logger.warning( + "Disagg role=%s: failed to read HF config for skipped " + "component %s: %s", + self._disagg_role.value, + module_name, + e, + ) + def _resolve_component_path( self, server_args: ServerArgs, module_name: str, load_module_name: str ) -> str: @@ -214,7 +297,24 @@ class ComposedPipelineBase(ABC): "MoE pipeline detected. Adding transformer_2 to self.required_config_modules..." ) if "transformer_2" not in self.required_config_modules: - self.required_config_modules.append("transformer_2") + # Re-apply disagg role filter: only add transformer_2 if the + # role actually needs denoising modules. + from sglang.multimodal_gen.runtime.disaggregation.roles import ( + get_module_role, + ) + + module_role = get_module_role("transformer_2") + if ( + self._disagg_role == RoleType.MONOLITHIC + or module_role is None + or module_role == self._disagg_role + ): + self.required_config_modules.append("transformer_2") + else: + logger.info( + "Disagg role=%s: skipping dynamically added module transformer_2", + self._disagg_role.value, + ) else: logger.info( "Boundary ratio found in model_index.json without transformers; " @@ -237,6 +337,11 @@ class ComposedPipelineBase(ABC): len(model_index) > 1 ), "model_index.json must contain at least one pipeline module" + # In disagg mode, read HF config for skipped components (e.g., VAE) + # so that update_model_arch + post_init can derive pipeline_config. + if self._disagg_role != RoleType.MONOLITHIC: + self._init_skipped_component_configs(model_index, server_args) + model_index = { required_module: model_index[required_module] for required_module in self.required_config_modules @@ -341,6 +446,19 @@ class ComposedPipelineBase(ABC): assert self.modules is not None, "No modules are registered" + # Filter stages based on disaggregation role + if self._disagg_role != RoleType.MONOLITHIC: + if stage.role_affinity != self._disagg_role: + if stage_name is None: + stage_name = self._infer_stage_name(stage) + logger.info( + "Disagg role=%s: skipping stage %s (affinity=%s)", + self._disagg_role.value, + stage_name, + stage.role_affinity.value, + ) + return self + if stage_name is None: stage_name = self._infer_stage_name(stage) if stage_name in self._stage_name_mapping: diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/base.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/base.py index e55243294..34778ee98 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/base.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/base.py @@ -13,6 +13,7 @@ from enum import Enum, auto import torch +from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import ( VerificationResult, @@ -103,6 +104,11 @@ class PipelineStage(ABC): """ pass + # Default role affinity: ENCODER. Override in subclasses for DENOISING/DECODER. + @property + def role_affinity(self) -> RoleType: + return RoleType.ENCODER + # execute on all ranks by default @property def parallelism_type(self) -> StageParallelismType: diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py index 29ef27fe4..b8e25636b 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py @@ -56,6 +56,12 @@ class DecodingStage(PipelineStage): output format (e.g., pixel values). """ + @property + def role_affinity(self): + from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType + + return RoleType.DECODER + def __init__(self, vae, pipeline=None, component_name: str = "vae") -> None: super().__init__() self.vae: ParallelTiledVAE = vae diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py index afd2bcc44..bfe4bbea2 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py @@ -35,6 +35,7 @@ from sglang.multimodal_gen.runtime.cache.cache_dit_integration import ( refresh_context_on_dual_transformer, refresh_context_on_transformer, ) +from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType from sglang.multimodal_gen.runtime.distributed import ( cfg_model_parallel_all_reduce, get_local_torch_device, @@ -152,6 +153,10 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): the initial noise into the final output. """ + @property + def role_affinity(self): + return RoleType.DENOISER + def __init__( self, transformer, scheduler, pipeline=None, transformer_2=None, vae=None ) -> None: diff --git a/python/sglang/multimodal_gen/runtime/server_args.py b/python/sglang/multimodal_gen/runtime/server_args.py index 159fcea46..49e9f3e34 100644 --- a/python/sglang/multimodal_gen/runtime/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args.py @@ -23,6 +23,12 @@ from sglang.multimodal_gen import envs from sglang.multimodal_gen.configs.models.encoders import T5Config from sglang.multimodal_gen.configs.pipeline_configs.base import PipelineConfig 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, ) @@ -94,7 +100,7 @@ class Backend(str, Enum): @dataclasses.dataclass -class ServerArgs: +class ServerArgs(DisaggArgsMixin): # Model and path configuration (for convenience) model_path: str @@ -230,10 +236,35 @@ class ServerArgs: # 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 + disagg_role: RoleType = RoleType.MONOLITHIC + disagg_timeout: int = 600 + disagg_dispatch_policy: str = "round_robin" + disagg_mode: bool = False + disagg_server_addr: str | None = None + encoder_urls: str | None = None + denoiser_urls: str | None = None + decoder_urls: str | None = None + encoder_tp: int | None = None + denoiser_tp: int | None = None + denoiser_sp: int | None = None + denoiser_ulysses: int | None = None + denoiser_ring: 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 + # Logging log_level: str = "info" uvicorn_access_log_exclude_prefixes: list[str] = field(default_factory=list) + # get_role_parallelism, derive_pool_*_endpoint — from DisaggArgsMixin + @property def broker_port(self) -> int: return self.port + 1 @@ -411,12 +442,22 @@ class ServerArgs: ) def _adjust_network_ports(self): + # Disagg role instances (encoder/denoiser/decoder) don't serve HTTP, + # so skip settling the HTTP port to avoid unnecessary port collisions. + needs_http = self.disagg_role in ( + RoleType.MONOLITHIC, + RoleType.SERVER, + ) + if self.strict_ports: - self._require_port(self.port, "HTTP") + if needs_http: + self._require_port(self.port, "HTTP") self._require_port(self.scheduler_port, "Scheduler") - self._require_port(self.master_port, "Master") + if self.master_port is not None: + self._require_port(self.master_port, "Master") else: - self.port = self.settle_port(self.port) + if needs_http: + self.port = self.settle_port(self.port) initial_scheduler_port = self.scheduler_port + ( random.randint(0, 100) if self.scheduler_port == 5555 else 0 ) @@ -611,6 +652,9 @@ class ServerArgs: # configure logger before use configure_logger(server_args=self) + # Convert string disagg_role to enum (from CLI/config) + convert_disagg_role_string(self.__dict__) + # 1. adjust parameters self._adjust_parameters() @@ -699,6 +743,7 @@ class ServerArgs: default=ServerArgs.num_gpus, help="The number of GPUs to use.", ) + parser.add_argument( "--tp-size", type=int, @@ -758,6 +803,9 @@ class ServerArgs: "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) + # Prompt text file for batch processing parser.add_argument( "--prompt-file-path", @@ -1139,6 +1187,9 @@ class ServerArgs: if "backend" in kwargs and isinstance(kwargs["backend"], str): kwargs["backend"] = Backend.from_string(kwargs["backend"]) + # Convert disagg_role string to enum if necessary + convert_disagg_role_string(kwargs) + kwargs["pipeline_config"] = PipelineConfig.from_kwargs(kwargs) return cls(**kwargs) diff --git a/python/sglang/multimodal_gen/runtime/utils/common.py b/python/sglang/multimodal_gen/runtime/utils/common.py index e2fb8a88d..410f2b890 100644 --- a/python/sglang/multimodal_gen/runtime/utils/common.py +++ b/python/sglang/multimodal_gen/runtime/utils/common.py @@ -130,6 +130,7 @@ def get_zmq_socket( endpoint: str, bind: bool, max_bind_retries: int = 10, + same_port: bool = False, ) -> tuple[zmq.Socket, str]: """ Create and configure a ZMQ socket. @@ -140,10 +141,13 @@ def get_zmq_socket( endpoint: Endpoint string (e.g., "tcp://localhost:5555") bind: Whether to bind (True) or connect (False) max_bind_retries: Maximum number of retries if bind fails due to address already in use + same_port: If True, retry on the same port instead of incrementing. + Useful when the port must be fixed (e.g., disagg sockets where + DiffusionServer connects to a pre-determined port). Returns: A tuple of (socket, actual_endpoint). The actual_endpoint may differ from the - requested endpoint if bind retry was needed. + requested endpoint if bind retry was needed (and same_port is False). """ mem = psutil.virtual_memory() total_mem = mem.total / 1024**3 @@ -182,13 +186,15 @@ def get_zmq_socket( port_match = re.search(r":(\d+)$", endpoint) if port_match and max_bind_retries > 1: + import time as _time + original_port = int(port_match.group(1)) last_exception = None for attempt in range(max_bind_retries): try: current_endpoint = endpoint - if attempt > 0: + if attempt > 0 and not same_port: # Try next port (increment by 42 to match settle_port logic) current_port = original_port + attempt * 42 current_endpoint = re.sub( @@ -198,6 +204,11 @@ def get_zmq_socket( f"ZMQ bind failed for port {original_port + (attempt - 1) * 42}, " f"retrying with port {current_port} (attempt {attempt + 1}/{max_bind_retries})" ) + elif attempt > 0: + logger.info( + f"ZMQ bind attempt {attempt + 1}/{max_bind_retries} " + f"on same port {original_port}..." + ) socket.bind(current_endpoint) @@ -212,7 +223,21 @@ def get_zmq_socket( except zmq.ZMQError as e: last_exception = e if e.errno == zmq.EADDRINUSE and attempt < max_bind_retries - 1: - # Address already in use, try next port + # Address already in use, retry + # Longer sleep for same_port (waiting for TIME_WAIT release) + _time.sleep(1.0 if same_port else 0.5) + # Re-create socket since ZMQ socket state may be invalid after failed bind + socket.close() + socket = context.socket(socket_type) + if endpoint.find("[") != -1: + socket.setsockopt(zmq.IPV6, 1) + if socket_type == zmq.PUSH: + set_send_opt() + elif socket_type == zmq.PULL: + set_recv_opt() + elif socket_type in [zmq.DEALER, zmq.REQ, zmq.REP, zmq.ROUTER]: + set_send_opt() + set_recv_opt() continue elif attempt == max_bind_retries - 1: # Last attempt failed diff --git a/python/sglang/multimodal_gen/test/run_suite.py b/python/sglang/multimodal_gen/test/run_suite.py index 9c1437b20..9cccf3f24 100644 --- a/python/sglang/multimodal_gen/test/run_suite.py +++ b/python/sglang/multimodal_gen/test/run_suite.py @@ -88,7 +88,9 @@ STANDALONE_FILES = { "../cli/test_generate_t2i_perf.py", "test_update_weights_from_disk.py", ], - "2-gpu": [], + "2-gpu": [ + "test_disagg_server.py", + ], } # New standalone files may omit an estimate once to learn the real CI runtime. @@ -99,7 +101,11 @@ STANDALONE_FILE_EST_TIMES = { "../cli/test_generate_t2i_perf.py": 240.0, "test_update_weights_from_disk.py": 480.0, }, - "2-gpu": {}, + "2-gpu": { + # Two disagg clusters × (~3 min startup + ~1 min generate) ≈ 8 min. + # Raise if CI reports a higher measured time. + "test_disagg_server.py": 600.0, + }, } # Backward-compatible suite view for scripts that still operate on file lists. diff --git a/python/sglang/multimodal_gen/test/server/test_disagg_server.py b/python/sglang/multimodal_gen/test/server/test_disagg_server.py new file mode 100755 index 000000000..e696c8f59 --- /dev/null +++ b/python/sglang/multimodal_gen/test/server/test_disagg_server.py @@ -0,0 +1,388 @@ +"""End-to-end tests for disaggregated diffusion. + +Launches encoder / denoiser / decoder role instances plus a DiffusionServer +head, sends a generation request through the HTTP front-end, and verifies +that a non-empty output comes back. + +Two configurations are covered: + +1. :class:`TestDisaggZImage1Rank` — 1 rank per role (baseline disagg path). +2. :class:`TestDisaggZImage2RankDenoiser` — denoiser with + ``--denoiser-sp 2`` across 2 GPUs. Exercises the multi-rank receive path + where only rank 0 owns the RDMA TransferManager and must broadcast + prompt/image tensors to non-rank-0 ranks before + ``execute_forward`` — without that broadcast the denoising stage fails + ``verify_input`` on an empty ``prompt_embeds``. + +Run directly: + + pytest -v python/sglang/multimodal_gen/test/server/test_disagg_server.py + pytest -v ... -k ZImage1Rank # one class + pytest -v ... -k test_generates_image # one test +""" + +from __future__ import annotations + +import base64 +import os +import signal +import subprocess +import time +import unittest +from pathlib import Path + +import requests +import torch + +from sglang.multimodal_gen.test.test_utils import ( + DEFAULT_SMALL_MODEL_NAME_FOR_TEST, + find_free_port, + wait_for_server_health, +) +from sglang.test.test_utils import CustomTestCase + +HOST = "127.0.0.1" +_LOG_DIR = Path(os.environ.get("SGLANG_TEST_LOG_DIR", "/tmp")) + +# Env knob: bump if a cold HF download is needed on a fresh CI runner. +_STARTUP_TIMEOUT_S = float(os.environ.get("SGLANG_DISAGG_STARTUP_TIMEOUT", "600")) + + +# --------------------------------------------------------------------------- +# Process management +# --------------------------------------------------------------------------- + + +def _kill_tree(pid: int) -> None: + try: + os.killpg(os.getpgid(pid), signal.SIGKILL) + except (ProcessLookupError, PermissionError): + pass + + +def _wait_for_log(path: Path, message: str, timeout: float) -> bool: + deadline = time.time() + timeout + while time.time() < deadline: + if path.exists(): + try: + if message in path.read_text(errors="ignore"): + return True + except OSError: + pass + time.sleep(2) + return False + + +def _tail_log(path: Path, n: int = 50) -> str: + if not path.exists(): + return f"" + try: + lines = path.read_text(errors="ignore").splitlines() + except OSError as e: + return f"" + return "\n".join(lines[-n:]) + + +# --------------------------------------------------------------------------- +# Disagg cluster helper +# --------------------------------------------------------------------------- + + +class DisaggCluster: + """Launch encoder / denoiser / decoder / server as separate processes. + + ``gpu_layout`` is a mapping role → list of physical GPU ids. The length + of each list determines ``--num-gpus`` for that role, and the first id is + passed as ``--base-gpu-id``. For a multi-rank role the GPUs must be + contiguous starting from ``base-gpu-id`` (sglang derives local_rank from + ``base-gpu-id + rank``). + """ + + def __init__( + self, + model: str, + name: str, + gpu_layout: dict[str, list[int]], + extra_role_args: dict[str, list[str]] | None = None, + startup_timeout: float = _STARTUP_TIMEOUT_S, + ) -> None: + self.model = model + self.name = name + self.gpu_layout = gpu_layout + self.extra_role_args = extra_role_args or {} + self.startup_timeout = startup_timeout + self._procs: list[subprocess.Popen] = [] + self._fhs: list = [] + self._logs: dict[str, Path] = {} + self._alloc_ports() + + def _alloc_ports(self) -> None: + self.base_port = find_free_port(HOST) + self.api_port = find_free_port(HOST) + self._role_ports = { + "encoder": find_free_port(HOST), + "denoiser": find_free_port(HOST), + "decoder": find_free_port(HOST), + } + + # -- context manager ----------------------------------------------------- + + def __enter__(self) -> "DisaggCluster": + for attempt in range(3): + try: + self._launch_roles() + self._launch_server_head() + self._warmup() + return self + except Exception as e: + print( + f"[disagg-test] Cluster {self.name} attempt {attempt + 1} " + f"failed: {e}", + flush=True, + ) + self.stop() + self._alloc_ports() + if attempt == 2: + raise + return self # unreachable + + def __exit__(self, *exc) -> None: + self.stop() + + # -- internals ----------------------------------------------------------- + + def _start_proc(self, cmd: list[str], log_path: Path) -> subprocess.Popen: + fh = open(log_path, "w") + proc = subprocess.Popen( + cmd, + stdout=fh, + stderr=subprocess.STDOUT, + preexec_fn=os.setsid, + env=os.environ.copy(), + ) + self._procs.append(proc) + self._fhs.append(fh) + return proc + + def _launch_roles(self) -> None: + for role in ("encoder", "denoiser", "decoder"): + port = self._role_ports[role] + gpus = self.gpu_layout[role] + log = _LOG_DIR / f"disagg_{self.name}_{role}.log" + self._logs[role] = log + + cmd = [ + "sglang", + "serve", + "--model-path", + self.model, + "--disagg-role", + role, + "--disagg-server-addr", + f"tcp://{HOST}:{self.base_port}", + "--scheduler-port", + str(port), + "--num-gpus", + str(len(gpus)), + "--base-gpu-id", + str(gpus[0]), + "--log-level", + "info", + *self.extra_role_args.get(role, []), + ] + self._start_proc(cmd, log) + + ready_msg = f"Role {role.upper()} ready" + if not _wait_for_log(log, ready_msg, self.startup_timeout): + raise RuntimeError( + f"{role} failed to start for {self.name}. Log tail:\n" + f"{_tail_log(log)}" + ) + + def _launch_server_head(self) -> None: + log = _LOG_DIR / f"disagg_{self.name}_server.log" + self._logs["server"] = log + cmd = [ + "sglang", + "serve", + "--model-path", + self.model, + "--disagg-role", + "server", + "--encoder-urls", + f"tcp://{HOST}:{self._role_ports['encoder']}", + "--denoiser-urls", + f"tcp://{HOST}:{self._role_ports['denoiser']}", + "--decoder-urls", + f"tcp://{HOST}:{self._role_ports['decoder']}", + "--scheduler-port", + str(self.base_port), + "--port", + str(self.api_port), + "--host", + HOST, + "--disagg-timeout", + "120", + "--log-level", + "info", + ] + self._start_proc(cmd, log) + try: + wait_for_server_health( + f"http://{HOST}:{self.api_port}", + path="/v1/models", + timeout=self.startup_timeout, + ) + except Exception as e: + raise RuntimeError( + f"server head failed to become healthy for {self.name}: {e}\n" + f"Server log tail:\n{_tail_log(log)}" + ) from e + + def _warmup(self) -> None: + """Send a warmup request to establish RDMA connections.""" + try: + _generate_image(self.api_port, self.model) + except Exception as e: + raise RuntimeError( + f"Warmup request failed for {self.name}: {e}\n" + f"Server log tail:\n{_tail_log(self._logs.get('server', Path('/dev/null')))}" + ) from e + + def stop(self) -> None: + for proc in self._procs: + _kill_tree(proc.pid) + for fh in self._fhs: + try: + fh.close() + except OSError: + pass + # Give OS a moment to release ports before the next test. + time.sleep(3) + self._procs.clear() + self._fhs.clear() + + +# --------------------------------------------------------------------------- +# Request helpers +# --------------------------------------------------------------------------- + + +def _generate_image(api_port: int, model: str) -> bytes: + # Use raw requests (openai SDK pulls in a lot and complicates CI deps). + resp = requests.post( + f"http://{HOST}:{api_port}/v1/images/generations", + json={ + "model": model, + "prompt": "A sunset over mountains", + "n": 1, + "size": "1024x1024", + "response_format": "b64_json", + }, + timeout=600, + ) + if resp.status_code != 200: + print( + f"[disagg-test] Server returned {resp.status_code}: {resp.text[:2000]}", + flush=True, + ) + resp.raise_for_status() + data = resp.json() + return base64.b64decode(data["data"][0]["b64_json"]) + + +# --------------------------------------------------------------------------- +# Test classes +# --------------------------------------------------------------------------- + + +def _require_gpus(n: int) -> None: + available = torch.cuda.device_count() if torch.cuda.is_available() else 0 + if available < n: + raise unittest.SkipTest(f"need {n} GPUs, have {available}") + + +class _DisaggTestBase(CustomTestCase): + """Shared setup: launch cluster once per class, tear down at the end.""" + + model: str = DEFAULT_SMALL_MODEL_NAME_FOR_TEST + required_gpus: int = 2 + cluster_name: str = "" + gpu_layout: dict[str, list[int]] = {} + extra_role_args: dict[str, list[str]] = {} + + cluster: DisaggCluster | None = None + + @classmethod + def setUpClass(cls) -> None: + super().setUpClass() + _require_gpus(cls.required_gpus) + cls.cluster = DisaggCluster( + model=cls.model, + name=cls.cluster_name, + gpu_layout=cls.gpu_layout, + extra_role_args=cls.extra_role_args, + ) + cls.cluster.__enter__() + + @classmethod + def tearDownClass(cls) -> None: + if cls.cluster is not None: + # Dump log tails for debugging CI failures + for role_name, log_path in cls.cluster._logs.items(): + print( + f"\n=== [{cls.cluster_name}] {role_name} log tail ===", + flush=True, + ) + print(_tail_log(log_path, n=80), flush=True) + cls.cluster.stop() + cls.cluster = None + super().tearDownClass() + + +class TestDisaggZImage1Rank(_DisaggTestBase): + """Baseline: 1 rank per role, 2 physical GPUs.""" + + cluster_name = "zimage_1rank" + required_gpus = 2 + gpu_layout = { + "encoder": [0], + "denoiser": [1], + "decoder": [0], + } + + def test_generates_image(self) -> None: + assert self.cluster is not None + img = _generate_image(self.cluster.api_port, self.model) + # A real PNG is well above 1 KB; catches empty / error responses. + self.assertGreater(len(img), 1_000, f"image too small: {len(img)} bytes") + + +class TestDisaggZImage2RankDenoiser(_DisaggTestBase): + """Multi-rank denoiser (``--denoiser-sp 2``) on 2 GPUs. + + Regression guard for the bug where non-rank-0 denoiser ranks entered + ``execute_forward`` with an empty Req because ``ParallelExecutor``'s + REPLICATED stage does not broadcast the batch. With the fix, rank 0 + broadcasts both scalar and tensor fields over NCCL before compute. + """ + + cluster_name = "zimage_sp2" + required_gpus = 2 + gpu_layout = { + "encoder": [0], + "denoiser": [0, 1], + "decoder": [0], + } + extra_role_args = { + "denoiser": ["--denoiser-sp", "2"], + } + + def test_generates_image_with_sp2_denoiser(self) -> None: + assert self.cluster is not None + img = _generate_image(self.cluster.api_port, self.model) + self.assertGreater(len(img), 1_000, f"image too small: {len(img)} bytes") + + +if __name__ == "__main__": + unittest.main() diff --git a/python/sglang/multimodal_gen/test/unit/test_server_args.py b/python/sglang/multimodal_gen/test/unit/test_server_args.py index 5c6e519e6..3151287f6 100644 --- a/python/sglang/multimodal_gen/test/unit/test_server_args.py +++ b/python/sglang/multimodal_gen/test/unit/test_server_args.py @@ -81,6 +81,115 @@ class TestModelIdResolution(unittest.TestCase): _get_config_info("/data/no-such-model", model_id="NonExistentModelXYZ") +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) + + def test_defaults_are_none(self): + args = self._from_dict({"model_path": "/fake"}) + from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType + + for role in [RoleType.ENCODER, RoleType.DENOISER, RoleType.DECODER]: + par = args.get_role_parallelism(role) + self.assertIsNone(par["tp_size"]) + self.assertIsNone(par["sp_degree"]) + self.assertIsNone(par["ulysses_degree"]) + self.assertIsNone(par["ring_degree"]) + + def test_encoder_overrides(self): + args = self._from_dict({"model_path": "/fake", "encoder_tp": 2}) + from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType + + par = args.get_role_parallelism(RoleType.ENCODER) + self.assertEqual(par["tp_size"], 2) + self.assertIsNone(par["sp_degree"]) + self.assertIsNone(par["ulysses_degree"]) + self.assertIsNone(par["ring_degree"]) + + def test_denoiser_overrides(self): + args = self._from_dict( + { + "model_path": "/fake", + "denoiser_tp": 1, + "denoiser_sp": 8, + "denoiser_ulysses": 4, + "denoiser_ring": 2, + } + ) + from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType + + par = args.get_role_parallelism(RoleType.DENOISER) + self.assertEqual(par["tp_size"], 1) + self.assertEqual(par["sp_degree"], 8) + self.assertEqual(par["ulysses_degree"], 4) + self.assertEqual(par["ring_degree"], 2) + + def test_decoder_overrides(self): + args = self._from_dict({"model_path": "/fake", "decoder_tp": 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["ulysses_degree"]) + self.assertIsNone(par["ring_degree"]) + + 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 + + par = args.get_role_parallelism(RoleType.MONOLITHIC) + self.assertIsNone(par["tp_size"]) + self.assertIsNone(par["sp_degree"]) + + def test_mixed_roles_independent(self): + """Per-role args don't interfere with each other.""" + args = self._from_dict( + { + "model_path": "/fake", + "encoder_tp": 1, + "denoiser_tp": 2, + "decoder_tp": 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) + + def test_cli_args_parsed(self): + """Per-role parallelism args are parsed from CLI.""" + parser = FlexibleArgumentParser() + ServerArgs.add_cli_args(parser) + argv = [ + "--model-path", + "/fake", + "--denoiser-tp", + "2", + "--denoiser-sp", + "4", + "--denoiser-ulysses", + "2", + "--denoiser-ring", + "2", + "--encoder-tp", + "1", + ] + args, unknown = parser.parse_known_args(argv) + self.assertEqual(args.denoiser_tp, 2) + self.assertEqual(args.denoiser_sp, 4) + self.assertEqual(args.denoiser_ulysses, 2) + self.assertEqual(args.denoiser_ring, 2) + self.assertEqual(args.encoder_tp, 1) + self.assertIsNone(args.decoder_tp) + + class TestPipelineResolutionCliOverride(unittest.TestCase): def setUp(self): _get_config_info.cache_clear() @@ -102,25 +211,5 @@ class TestPipelineResolutionCliOverride(unittest.TestCase): self.assertEqual(server_args.pipeline_config.resolution, 768) -class TestComponentPathParsing(unittest.TestCase): - def test_extract_component_paths_accepts_config_expanded_keys(self): - component_paths, remaining = ServerArgs._extract_component_paths( - [ - "--component-paths.spatial-upsampler", - "/tmp/latent_upsampler", - "--component_paths.distilled-lora=/tmp/distilled.safetensors", - ] - ) - - self.assertEqual( - component_paths, - { - "spatial_upsampler": "/tmp/latent_upsampler", - "distilled_lora": "/tmp/distilled.safetensors", - }, - ) - self.assertEqual(remaining, []) - - if __name__ == "__main__": unittest.main()