[diffusion] feat: disaggregated diffusion (#21701)

This commit is contained in:
Yuhao Yang
2026-04-16 23:51:32 +08:00
committed by GitHub
parent 14bcdfca21
commit 9da998a882
32 changed files with 6182 additions and 46 deletions
+237
View File
@@ -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`.
@@ -0,0 +1,2 @@
# SPDX-License-Identifier: Apache-2.0
"""Disaggregation support for diffusion pipelines."""
@@ -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"])
@@ -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)
@@ -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)
File diff suppressed because it is too large Load Diff
@@ -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,
}
@@ -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
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,2 @@
# SPDX-License-Identifier: Apache-2.0
"""Transport layer for disaggregated diffusion pipelines."""
@@ -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
@@ -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
@@ -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
@@ -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
)
@@ -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")
@@ -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)
)
@@ -435,9 +435,9 @@ class GroupCoordinator:
# Bypass the function if we are using only 1 GPU. # Bypass the function if we are using only 1 GPU.
if self.world_size == 1: if self.world_size == 1:
return obj 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" 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: if self.rank_in_group == src:
torch.distributed.broadcast_object_list( torch.distributed.broadcast_object_list(
[obj], src=self.ranks[src], group=self.cpu_group [obj], src=self.ranks[src], group=self.cpu_group
@@ -8,7 +8,9 @@ from typing import cast
from sglang.multimodal_gen.apps.webui import run_sgl_diffusion_webui 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.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.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.utils import FlexibleArgumentParser 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): def execute_serve_cmd(args: argparse.Namespace, unknown_args: list[str] | None = None):
"""The entry point for the serve command.""" """The entry point for the serve command."""
server_args = ServerArgs.from_cli_args(args, unknown_args) server_args = ServerArgs.from_cli_args(args, unknown_args)
launch_server(server_args)
dispatch_launch(server_args)
if server_args.webui: if server_args.webui:
run_sgl_diffusion_webui(server_args) run_sgl_diffusion_webui(server_args)
@@ -162,6 +162,32 @@ async def health_generate():
return {"status": "ok"} 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): def make_serializable(obj):
"""Recursively converts Tensors to None for JSON serialization.""" """Recursively converts Tensors to None for JSON serialization."""
if isinstance(obj, torch.Tensor): if isinstance(obj, torch.Tensor):
@@ -69,6 +69,13 @@ class ShutdownReq:
pass pass
@dataclass
class GetDisaggStatsReq:
"""Request to get disagg pipeline metrics from the scheduler."""
pass
def format_lora_message( def format_lora_message(
lora_nickname: Union[str, List[str]], lora_nickname: Union[str, List[str]],
target: Union[str, List[str]], target: Union[str, List[str]],
@@ -1,5 +1,6 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo # Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
import dataclasses
import multiprocessing as mp import multiprocessing as mp
import os import os
import signal import signal
@@ -9,6 +10,10 @@ import threading
import psutil import psutil
import uvicorn 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.entrypoints.http_server import create_app
from sglang.multimodal_gen.runtime.managers.gpu_worker import run_scheduler_process from sglang.multimodal_gen.runtime.managers.gpu_worker import run_scheduler_process
from sglang.multimodal_gen.runtime.server_args import ( from sglang.multimodal_gen.runtime.server_args import (
@@ -16,9 +21,28 @@ from sglang.multimodal_gen.runtime.server_args import (
prepare_server_args, prepare_server_args,
set_global_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 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): def kill_process_tree(parent_pid, include_parent: bool = True, skip_pid: int = None):
"""Kill the process and all its child processes.""" """Kill the process and all its child processes."""
# Remove sigchld handler to avoid spammy logs. # Remove sigchld handler to avoid spammy logs.
@@ -185,6 +209,237 @@ def launch_server(server_args: ServerArgs, launch_http_server: bool = True):
return processes 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): def launch_http_server_only(server_args):
# set for endpoints to access global_server_args # set for endpoints to access global_server_args
set_global_server_args(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__": if __name__ == "__main__":
server_args = prepare_server_args(sys.argv[1:]) server_args = prepare_server_args(sys.argv[1:])
try: try:
launch_server(server_args) dispatch_launch(server_args)
finally: finally:
kill_process_tree(os.getpid(), include_parent=False) kill_process_tree(os.getpid(), include_parent=False)
@@ -208,9 +208,16 @@ class GPUWorker:
f"Related offload server args to disable: {suggested_args_str}" 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. 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 assert self.pipeline is not None
req = batch[0] req = batch[0]
@@ -229,6 +236,11 @@ class GPUWorker:
req.log(server_args=self.server_args) req.log(server_args=self.server_args)
result = self.pipeline.forward(req, 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): if isinstance(result, Req):
output_batch = OutputBatch( output_batch = OutputBatch(
output=result.output, output=result.output,
@@ -261,7 +273,8 @@ class GPUWorker:
self.do_mem_analysis(output_batch) self.do_mem_analysis(output_batch)
duration_ms = (time.monotonic() - start_time) * 1000 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 # Save output to file and return file path only if requested. Avoid the serialization
# and deserialization overhead between scheduler_client and gpu_worker. # and deserialization overhead between scheduler_client and gpu_worker.
@@ -526,6 +539,7 @@ def run_scheduler_process(
port_args=port_args, port_args=port_args,
task_pipes_to_slaves=task_pipes_to_slaves, task_pipes_to_slaves=task_pipes_to_slaves,
result_pipes_from_slaves=result_pipes_from_slaves, result_pipes_from_slaves=result_pipes_from_slaves,
local_rank=local_rank,
) )
logger.info(f"Worker {rank}: Scheduler loop started.") logger.info(f"Worker {rank}: Scheduler loop started.")
pipe_writer.send( pipe_writer.send(
@@ -10,6 +10,10 @@ from typing import Any, List
import zmq 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.distributed import get_world_group
from sglang.multimodal_gen.runtime.entrypoints.openai.utils import ( from sglang.multimodal_gen.runtime.entrypoints.openai.utils import (
_parse_size, _parse_size,
@@ -20,6 +24,7 @@ from sglang.multimodal_gen.runtime.entrypoints.post_training.io_struct import (
UpdateWeightFromDiskReqInput, UpdateWeightFromDiskReqInput,
) )
from sglang.multimodal_gen.runtime.entrypoints.utils import ( from sglang.multimodal_gen.runtime.entrypoints.utils import (
GetDisaggStatsReq,
ListLorasReq, ListLorasReq,
MergeLoraWeightsReq, MergeLoraWeightsReq,
SetLoraReq, SetLoraReq,
@@ -43,7 +48,7 @@ logger = init_logger(__name__)
MINIMUM_PICTURE_BASE64_FOR_WARMUP = "data:image/jpg;base64,iVBORw0KGgoAAAANSUhEUgAAACAAAAAgCAYAAABzenr0AAAACXBIWXMAAA7EAAAOxAGVKw4bAAAAbUlEQVRYhe3VsQ2AMAxE0Y/lIgNQULD/OqyCMgCihCKSG4yRuKuiNH6JLsoEbMACOGBcua9HOR7Y6w6swBwMy0qLTpkeI77qdEBpBFAHBBDAGH8WrwJKI4AAegUCfAKgEgpQDvh3CR3oQCuav58qlAw73kKCSgAAAABJRU5ErkJggg==" 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. Runs the main event loop for the rank 0 worker.
It listens for external requests via ZMQ and coordinates with other workers. It listens for external requests via ZMQ and coordinates with other workers.
@@ -57,10 +62,17 @@ class Scheduler:
port_args: PortArgs, port_args: PortArgs,
task_pipes_to_slaves: list = None, task_pipes_to_slaves: list = None,
result_pipes_from_slaves: list = None, result_pipes_from_slaves: list = None,
local_rank: int | None = None,
): ):
self.server_args = server_args self.server_args = server_args
self.port_args = port_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) set_global_server_args(server_args=server_args)
# Inter-process Communication # Inter-process Communication
@@ -76,7 +88,7 @@ class Scheduler:
self.receiver = None self.receiver = None
worker = GPUWorker( worker = GPUWorker(
local_rank=gpu_id, local_rank=local_rank,
master_port=port_args.master_port, master_port=port_args.master_port,
rank=gpu_id, rank=gpu_id,
server_args=server_args, server_args=server_args,
@@ -95,6 +107,7 @@ class Scheduler:
List[Req]: self._handle_generation, List[Req]: self._handle_generation,
ListLorasReq: self._handle_list_loras, ListLorasReq: self._handle_list_loras,
ShutdownReq: self._handle_shutdown, ShutdownReq: self._handle_shutdown,
GetDisaggStatsReq: self._handle_get_disagg_stats,
UpdateWeightFromDiskReqInput: self._handle_update_weights_from_disk, UpdateWeightFromDiskReqInput: self._handle_update_weights_from_disk,
GetWeightsChecksumReqInput: self._handle_get_weights_checksum, GetWeightsChecksumReqInput: self._handle_get_weights_checksum,
} }
@@ -114,6 +127,21 @@ class Scheduler:
self._max_consecutive_errors = 3 self._max_consecutive_errors = 3
self._consecutive_error_count = 0 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: def _handle_set_lora(self, reqs: List[Any]) -> OutputBatch:
# TODO: return set status # TODO: return set status
# TODO: return with SetLoRAResponse or something more appropriate # TODO: return with SetLoRAResponse or something more appropriate
@@ -166,6 +194,7 @@ class Scheduler:
) )
else: else:
logger.info("Processing warmup req...") logger.info("Processing warmup req...")
return self.worker.execute_forward(reqs) return self.worker.execute_forward(reqs)
def return_result( def return_result(
@@ -362,12 +391,20 @@ class Scheduler:
The main event loop that listens for ZMQ requests. The main event loop that listens for ZMQ requests.
Handles abortion Handles abortion
""" """
# Pool mode: all roles use the pool event loop
if self._disagg_role != RoleType.MONOLITHIC:
self._disagg_event_loop()
return
logger.debug( logger.debug(
f"Rank 0 scheduler listening on tcp://*:{self.server_args.scheduler_port}" f"Rank 0 scheduler listening on tcp://*:{self.server_args.scheduler_port}"
) )
while self._running: 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 # 1: receive requests
try: try:
new_reqs = self.recv_reqs() new_reqs = self.recv_reqs()
@@ -403,6 +440,10 @@ class Scheduler:
try: try:
processed_req = reqs[0] 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)) handler = self.request_handlers.get(type(processed_req))
if handler: if handler:
output_batch = handler(reqs) output_batch = handler(reqs)
@@ -415,16 +456,10 @@ class Scheduler:
f"Error executing request in scheduler event loop: {e}", f"Error executing request in scheduler event loop: {e}",
exc_info=True, exc_info=True,
) )
# Determine appropriate error response format output_batch = OutputBatch(error=str(e))
output_batch = (
OutputBatch(error=str(e))
if reqs and isinstance(reqs[0], Req)
else OutputBatch(error=str(e))
)
# 3. return results # 3. return results
try: try:
# log warmup info
is_warmup = ( is_warmup = (
processed_req.is_warmup if isinstance(processed_req, Req) else False processed_req.is_warmup if isinstance(processed_req, Req) else False
) )
@@ -457,6 +492,7 @@ class Scheduler:
if self.receiver is not None: if self.receiver is not None:
self.receiver.close() self.receiver.close()
self._cleanup_disagg()
self.context.destroy(linger=0) self.context.destroy(linger=0)
def _broadcast_task(self, payload: dict[str, Any]) -> None: def _broadcast_task(self, payload: dict[str, Any]) -> None:
@@ -14,6 +14,10 @@ from typing import Any, Callable, Literal, cast
import torch import torch
from tqdm import tqdm 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 ( from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
PipelineComponentLoader, PipelineComponentLoader,
) )
@@ -82,6 +86,7 @@ class ComposedPipelineBase(ABC):
use. The pipeline should be stateless and not hold any batch state. use. The pipeline should be stateless and not hold any batch state.
""" """
self.server_args = server_args self.server_args = server_args
self._disagg_role = server_args.disagg_role
self.model_path: str = model_path self.model_path: str = model_path
self._stages: list[PipelineStage] = [] self._stages: list[PipelineStage] = []
@@ -94,6 +99,20 @@ class ComposedPipelineBase(ABC):
if self._required_config_modules is None: if self._required_config_modules is None:
raise NotImplementedError("Subclass must set _required_config_modules") 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] # [module_name, gpu memory usage]
self.memory_usages: dict[str, float] = {} self.memory_usages: dict[str, float] = {}
# Load modules directly in initialization # Load modules directly in initialization
@@ -169,6 +188,70 @@ class ComposedPipelineBase(ABC):
""" """
return 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( def _resolve_component_path(
self, server_args: ServerArgs, module_name: str, load_module_name: str self, server_args: ServerArgs, module_name: str, load_module_name: str
) -> str: ) -> str:
@@ -214,7 +297,24 @@ class ComposedPipelineBase(ABC):
"MoE pipeline detected. Adding transformer_2 to self.required_config_modules..." "MoE pipeline detected. Adding transformer_2 to self.required_config_modules..."
) )
if "transformer_2" not in 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: else:
logger.info( logger.info(
"Boundary ratio found in model_index.json without transformers; " "Boundary ratio found in model_index.json without transformers; "
@@ -237,6 +337,11 @@ class ComposedPipelineBase(ABC):
len(model_index) > 1 len(model_index) > 1
), "model_index.json must contain at least one pipeline module" ), "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 = { model_index = {
required_module: model_index[required_module] required_module: model_index[required_module]
for required_module in self.required_config_modules 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" 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: if stage_name is None:
stage_name = self._infer_stage_name(stage) stage_name = self._infer_stage_name(stage)
if stage_name in self._stage_name_mapping: if stage_name in self._stage_name_mapping:
@@ -13,6 +13,7 @@ from enum import Enum, auto
import torch 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.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
VerificationResult, VerificationResult,
@@ -103,6 +104,11 @@ class PipelineStage(ABC):
""" """
pass 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 # execute on all ranks by default
@property @property
def parallelism_type(self) -> StageParallelismType: def parallelism_type(self) -> StageParallelismType:
@@ -56,6 +56,12 @@ class DecodingStage(PipelineStage):
output format (e.g., pixel values). 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: def __init__(self, vae, pipeline=None, component_name: str = "vae") -> None:
super().__init__() super().__init__()
self.vae: ParallelTiledVAE = vae self.vae: ParallelTiledVAE = vae
@@ -35,6 +35,7 @@ from sglang.multimodal_gen.runtime.cache.cache_dit_integration import (
refresh_context_on_dual_transformer, refresh_context_on_dual_transformer,
refresh_context_on_transformer, refresh_context_on_transformer,
) )
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
from sglang.multimodal_gen.runtime.distributed import ( from sglang.multimodal_gen.runtime.distributed import (
cfg_model_parallel_all_reduce, cfg_model_parallel_all_reduce,
get_local_torch_device, get_local_torch_device,
@@ -152,6 +153,10 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
the initial noise into the final output. the initial noise into the final output.
""" """
@property
def role_affinity(self):
return RoleType.DENOISER
def __init__( def __init__(
self, transformer, scheduler, pipeline=None, transformer_2=None, vae=None self, transformer, scheduler, pipeline=None, transformer_2=None, vae=None
) -> None: ) -> None:
@@ -23,6 +23,12 @@ from sglang.multimodal_gen import envs
from sglang.multimodal_gen.configs.models.encoders import T5Config from sglang.multimodal_gen.configs.models.encoders import T5Config
from sglang.multimodal_gen.configs.pipeline_configs.base import PipelineConfig from sglang.multimodal_gen.configs.pipeline_configs.base import PipelineConfig
from sglang.multimodal_gen.configs.quantization.nunchaku import NunchakuSVDQuantArgs 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 ( from sglang.multimodal_gen.runtime.layers.quantization.configs.nunchaku_config import (
NunchakuConfig, NunchakuConfig,
) )
@@ -94,7 +100,7 @@ class Backend(str, Enum):
@dataclasses.dataclass @dataclasses.dataclass
class ServerArgs: class ServerArgs(DisaggArgsMixin):
# Model and path configuration (for convenience) # Model and path configuration (for convenience)
model_path: str model_path: str
@@ -230,10 +236,35 @@ class ServerArgs:
# MoE parameters used by Wan2.2 # MoE parameters used by Wan2.2
boundary_ratio: float | None = None 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 # Logging
log_level: str = "info" log_level: str = "info"
uvicorn_access_log_exclude_prefixes: list[str] = field(default_factory=list) uvicorn_access_log_exclude_prefixes: list[str] = field(default_factory=list)
# get_role_parallelism, derive_pool_*_endpoint — from DisaggArgsMixin
@property @property
def broker_port(self) -> int: def broker_port(self) -> int:
return self.port + 1 return self.port + 1
@@ -411,12 +442,22 @@ class ServerArgs:
) )
def _adjust_network_ports(self): 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: 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.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: else:
self.port = self.settle_port(self.port) if needs_http:
self.port = self.settle_port(self.port)
initial_scheduler_port = self.scheduler_port + ( initial_scheduler_port = self.scheduler_port + (
random.randint(0, 100) if self.scheduler_port == 5555 else 0 random.randint(0, 100) if self.scheduler_port == 5555 else 0
) )
@@ -611,6 +652,9 @@ class ServerArgs:
# configure logger before use # configure logger before use
configure_logger(server_args=self) configure_logger(server_args=self)
# Convert string disagg_role to enum (from CLI/config)
convert_disagg_role_string(self.__dict__)
# 1. adjust parameters # 1. adjust parameters
self._adjust_parameters() self._adjust_parameters()
@@ -699,6 +743,7 @@ class ServerArgs:
default=ServerArgs.num_gpus, default=ServerArgs.num_gpus,
help="The number of GPUs to use.", help="The number of GPUs to use.",
) )
parser.add_argument( parser.add_argument(
"--tp-size", "--tp-size",
type=int, type=int,
@@ -758,6 +803,9 @@ class ServerArgs:
"Increase this value if you encounter 'Connection closed by peer' errors after the service is idle. ", "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 # Prompt text file for batch processing
parser.add_argument( parser.add_argument(
"--prompt-file-path", "--prompt-file-path",
@@ -1139,6 +1187,9 @@ class ServerArgs:
if "backend" in kwargs and isinstance(kwargs["backend"], str): if "backend" in kwargs and isinstance(kwargs["backend"], str):
kwargs["backend"] = Backend.from_string(kwargs["backend"]) 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) kwargs["pipeline_config"] = PipelineConfig.from_kwargs(kwargs)
return cls(**kwargs) return cls(**kwargs)
@@ -130,6 +130,7 @@ def get_zmq_socket(
endpoint: str, endpoint: str,
bind: bool, bind: bool,
max_bind_retries: int = 10, max_bind_retries: int = 10,
same_port: bool = False,
) -> tuple[zmq.Socket, str]: ) -> tuple[zmq.Socket, str]:
""" """
Create and configure a ZMQ socket. Create and configure a ZMQ socket.
@@ -140,10 +141,13 @@ def get_zmq_socket(
endpoint: Endpoint string (e.g., "tcp://localhost:5555") endpoint: Endpoint string (e.g., "tcp://localhost:5555")
bind: Whether to bind (True) or connect (False) bind: Whether to bind (True) or connect (False)
max_bind_retries: Maximum number of retries if bind fails due to address already in use 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: Returns:
A tuple of (socket, actual_endpoint). The actual_endpoint may differ from the 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() mem = psutil.virtual_memory()
total_mem = mem.total / 1024**3 total_mem = mem.total / 1024**3
@@ -182,13 +186,15 @@ def get_zmq_socket(
port_match = re.search(r":(\d+)$", endpoint) port_match = re.search(r":(\d+)$", endpoint)
if port_match and max_bind_retries > 1: if port_match and max_bind_retries > 1:
import time as _time
original_port = int(port_match.group(1)) original_port = int(port_match.group(1))
last_exception = None last_exception = None
for attempt in range(max_bind_retries): for attempt in range(max_bind_retries):
try: try:
current_endpoint = endpoint current_endpoint = endpoint
if attempt > 0: if attempt > 0 and not same_port:
# Try next port (increment by 42 to match settle_port logic) # Try next port (increment by 42 to match settle_port logic)
current_port = original_port + attempt * 42 current_port = original_port + attempt * 42
current_endpoint = re.sub( current_endpoint = re.sub(
@@ -198,6 +204,11 @@ def get_zmq_socket(
f"ZMQ bind failed for port {original_port + (attempt - 1) * 42}, " f"ZMQ bind failed for port {original_port + (attempt - 1) * 42}, "
f"retrying with port {current_port} (attempt {attempt + 1}/{max_bind_retries})" 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) socket.bind(current_endpoint)
@@ -212,7 +223,21 @@ def get_zmq_socket(
except zmq.ZMQError as e: except zmq.ZMQError as e:
last_exception = e last_exception = e
if e.errno == zmq.EADDRINUSE and attempt < max_bind_retries - 1: 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 continue
elif attempt == max_bind_retries - 1: elif attempt == max_bind_retries - 1:
# Last attempt failed # Last attempt failed
@@ -88,7 +88,9 @@ STANDALONE_FILES = {
"../cli/test_generate_t2i_perf.py", "../cli/test_generate_t2i_perf.py",
"test_update_weights_from_disk.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. # 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, "../cli/test_generate_t2i_perf.py": 240.0,
"test_update_weights_from_disk.py": 480.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. # Backward-compatible suite view for scripts that still operate on file lists.
@@ -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"<no log at {path}>"
try:
lines = path.read_text(errors="ignore").splitlines()
except OSError as e:
return f"<log read failed: {e}>"
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()
@@ -81,6 +81,115 @@ class TestModelIdResolution(unittest.TestCase):
_get_config_info("/data/no-such-model", model_id="NonExistentModelXYZ") _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): class TestPipelineResolutionCliOverride(unittest.TestCase):
def setUp(self): def setUp(self):
_get_config_info.cache_clear() _get_config_info.cache_clear()
@@ -102,25 +211,5 @@ class TestPipelineResolutionCliOverride(unittest.TestCase):
self.assertEqual(server_args.pipeline_config.resolution, 768) 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__": if __name__ == "__main__":
unittest.main() unittest.main()