[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.
if self.world_size == 1:
return obj
if self.shm_broadcaster is not None:
if self.mq_broadcaster is not None:
assert src == 0, "Shared memory broadcaster only supports src=0"
return self.shm_broadcaster.broadcast_object(obj)
return self.mq_broadcaster.broadcast_object(obj)
if self.rank_in_group == src:
torch.distributed.broadcast_object_list(
[obj], src=self.ranks[src], group=self.cpu_group
@@ -8,7 +8,9 @@ from typing import cast
from sglang.multimodal_gen.apps.webui import run_sgl_diffusion_webui
from sglang.multimodal_gen.runtime.entrypoints.cli.cli_types import CLISubcommand
from sglang.multimodal_gen.runtime.launch_server import launch_server
from sglang.multimodal_gen.runtime.launch_server import (
dispatch_launch,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.utils import FlexibleArgumentParser
@@ -31,7 +33,8 @@ def add_multimodal_gen_serve_args(parser: argparse.ArgumentParser):
def execute_serve_cmd(args: argparse.Namespace, unknown_args: list[str] | None = None):
"""The entry point for the serve command."""
server_args = ServerArgs.from_cli_args(args, unknown_args)
launch_server(server_args)
dispatch_launch(server_args)
if server_args.webui:
run_sgl_diffusion_webui(server_args)
@@ -162,6 +162,32 @@ async def health_generate():
return {"status": "ok"}
@health_router.get("/stats")
async def stats_endpoint(request: Request):
"""Get runtime statistics including disagg pipeline metrics.
Returns queue depth, request counts, latency, throughput, etc.
Sends a GetDisaggStatsReq to the scheduler via ZMQ and returns the result.
"""
from sglang.multimodal_gen.runtime.entrypoints.utils import GetDisaggStatsReq
server_args: ServerArgs = request.app.state.server_args
response: dict = {
"status": "ok",
"model_path": server_args.model_path,
}
# Query the scheduler for disagg metrics
try:
stats_response = await async_scheduler_client.forward(GetDisaggStatsReq())
if hasattr(stats_response, "output") and stats_response.output is not None:
response["disagg"] = stats_response.output
except Exception as e:
response["disagg"] = {"error": str(e)}
return response
def make_serializable(obj):
"""Recursively converts Tensors to None for JSON serialization."""
if isinstance(obj, torch.Tensor):
@@ -69,6 +69,13 @@ class ShutdownReq:
pass
@dataclass
class GetDisaggStatsReq:
"""Request to get disagg pipeline metrics from the scheduler."""
pass
def format_lora_message(
lora_nickname: Union[str, List[str]],
target: Union[str, List[str]],
@@ -1,5 +1,6 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
import dataclasses
import multiprocessing as mp
import os
import signal
@@ -9,6 +10,10 @@ import threading
import psutil
import uvicorn
from sglang.multimodal_gen.runtime.disaggregation.orchestrator import (
DiffusionServer,
)
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
from sglang.multimodal_gen.runtime.entrypoints.http_server import create_app
from sglang.multimodal_gen.runtime.managers.gpu_worker import run_scheduler_process
from sglang.multimodal_gen.runtime.server_args import (
@@ -16,9 +21,28 @@ from sglang.multimodal_gen.runtime.server_args import (
prepare_server_args,
set_global_server_args,
)
from sglang.multimodal_gen.runtime.utils.common import is_port_available
from sglang.multimodal_gen.runtime.utils.logging_utils import configure_logger, logger
def _find_available_port(
start: int = 10000, avoid: set[int] | None = None, max_attempts: int = 100
) -> int:
"""Find an available port starting from *start*, skipping ports in *avoid*."""
if avoid is None:
avoid = set()
port = max(1024, min(start, 65535))
for _ in range(max_attempts):
if port not in avoid and is_port_available(port):
return port
port += 1
if port > 65535:
port = 1024
raise RuntimeError(
f"No available port found after {max_attempts} attempts (start={start})"
)
def kill_process_tree(parent_pid, include_parent: bool = True, skip_pid: int = None):
"""Kill the process and all its child processes."""
# Remove sigchld handler to avoid spammy logs.
@@ -185,6 +209,237 @@ def launch_server(server_args: ServerArgs, launch_http_server: bool = True):
return processes
def launch_pool_disagg_server(
server_args: ServerArgs,
encoder_gpus: list[list[int]],
denoiser_gpus: list[list[int]],
decoder_gpus: list[list[int]],
launch_http_server: bool = True,
):
"""Launch a pool-based disaggregated server with N:M:K independent role instances.
DiffusionServer orchestrates the full pipeline, dispatching at every
role transition (Encoder → Denoiser → Decoder).
Args:
server_args: Base server configuration
encoder_gpus: List of GPU ID lists, one per encoder instance.
e.g., [[0], [2]] for 2 encoder instances on GPUs 0 and 2.
denoiser_gpus: List of GPU ID lists, one per denoiser instance.
e.g., [[1], [3]] for 2 denoiser instances.
decoder_gpus: List of GPU ID lists, one per decoder instance.
e.g., [[0], [2]] for 2 decoder instances (can share with encoder).
launch_http_server: Whether to launch the HTTP server.
Example:
launch_pool_disagg_server(server_args,
encoder_gpus=[[0], [2]],
denoiser_gpus=[[1], [3]],
decoder_gpus=[[0], [2]],
)
"""
configure_logger(server_args)
num_encoders = len(encoder_gpus)
num_denoisers = len(denoiser_gpus)
num_decoders = len(decoder_gpus)
logger.info(
"Starting pool disagg server: %d encoder(s), %d denoiser(s), %d decoder(s)...",
num_encoders,
num_denoisers,
num_decoders,
)
host = server_args.host or "127.0.0.1"
def find_port(start):
return _find_available_port(start)
# Allocate endpoints
port_cursor = server_args.scheduler_port + 3000
# Per-instance work endpoints (instance binds PULL, DS connects PUSH)
encoder_work_endpoints = []
for i in range(num_encoders):
p = find_port(port_cursor)
encoder_work_endpoints.append(f"tcp://{host}:{p}")
port_cursor = p + 1
denoiser_work_endpoints = []
for i in range(num_denoisers):
p = find_port(port_cursor)
denoiser_work_endpoints.append(f"tcp://{host}:{p}")
port_cursor = p + 1
decoder_work_endpoints = []
for i in range(num_decoders):
p = find_port(port_cursor)
decoder_work_endpoints.append(f"tcp://{host}:{p}")
port_cursor = p + 1
# Per-role-type result endpoints (DS binds PULL, instances connect PUSH)
# Use deterministic convention: scheduler_port + {1,2,3}
base_port = server_args.scheduler_port
encoder_result_ep = f"tcp://{host}:{base_port + 1}"
denoiser_result_ep = f"tcp://{host}:{base_port + 2}"
decoder_result_ep = f"tcp://{host}:{base_port + 3}"
logger.info(
"Pool endpoints allocated: %d work + 3 result endpoints",
num_encoders + num_denoisers + num_decoders,
)
# Launch all role instances
all_processes = []
role_configs = [
(RoleType.ENCODER, encoder_gpus, encoder_work_endpoints, encoder_result_ep),
(
RoleType.DENOISER,
denoiser_gpus,
denoiser_work_endpoints,
denoiser_result_ep,
),
(RoleType.DECODER, decoder_gpus, decoder_work_endpoints, decoder_result_ep),
]
for role_type, gpu_lists, work_eps, result_ep in role_configs:
for inst_idx, gpu_ids in enumerate(gpu_lists):
num_role_gpus = len(gpu_ids)
# Per-role parallelism: use explicit overrides if set, else None (auto-derive)
role_par = server_args.get_role_parallelism(role_type)
role_overrides = {
"disagg_role": role_type,
"disagg_mode": True,
"pool_work_endpoint": work_eps[inst_idx],
"pool_result_endpoint": result_ep,
"num_gpus": num_role_gpus,
"warmup": role_type == RoleType.ENCODER,
"scheduler_port": find_port(port_cursor),
"master_port": find_port(port_cursor + 100),
# Per-role parallelism (None = auto-derive from num_gpus)
"tp_size": role_par["tp_size"],
"sp_degree": role_par["sp_degree"],
"ulysses_degree": role_par["ulysses_degree"],
"ring_degree": role_par["ring_degree"],
}
port_cursor = role_overrides["master_port"] + 100
base_dict = {
f.name: getattr(server_args, f.name)
for f in dataclasses.fields(server_args)
}
base_dict.update(role_overrides)
base_dict.pop("pipeline_config", None)
role_args = ServerArgs.from_kwargs(**base_dict)
pool_ctx = mp.get_context("spawn")
inst_readers = []
# Spawn all ranks first — NCCL init blocks until all ranks connect
for rank_idx in range(num_role_gpus):
reader, writer = pool_ctx.Pipe(duplex=False)
gpu_id = gpu_ids[rank_idx]
process = pool_ctx.Process(
target=_run_disagg_role_process,
args=(gpu_id, rank_idx, rank_idx, role_args, writer, [], []),
name=f"sglang-pool-{role_type.value}-{inst_idx}-r{rank_idx}",
daemon=True,
)
process.start()
all_processes.append(process)
inst_readers.append(reader)
# Wait for all ranks to be ready (after all are spawned)
for rank_idx, reader in enumerate(inst_readers):
try:
data = reader.recv()
except EOFError:
logger.error(
"Pool %s[%d] rank %d is dead.",
role_type.value,
inst_idx,
rank_idx,
)
raise
if data.get("status") != "ready":
raise RuntimeError(
f"Pool {role_type.value}[{inst_idx}] rank {rank_idx} "
"failed to initialize."
)
reader.close()
logger.info(
"Pool %s[%d] ready on GPU(s) %s (work=%s)",
role_type.value.upper(),
inst_idx,
gpu_ids,
work_eps[inst_idx],
)
logger.info("All pool role instances ready")
# Start DiffusionServer
frontend_endpoint = f"tcp://{host}:{server_args.scheduler_port}"
diffusion_server = DiffusionServer(
frontend_endpoint=frontend_endpoint,
encoder_work_endpoints=encoder_work_endpoints,
denoiser_work_endpoints=denoiser_work_endpoints,
decoder_work_endpoints=decoder_work_endpoints,
encoder_result_endpoint=encoder_result_ep,
denoiser_result_endpoint=denoiser_result_ep,
decoder_result_endpoint=decoder_result_ep,
dispatch_policy_name=server_args.disagg_dispatch_policy,
timeout_s=float(server_args.disagg_timeout),
)
diffusion_server.start()
if not diffusion_server.wait_ready(timeout=30.0):
raise RuntimeError("DiffusionServer failed to bind sockets within 30 seconds")
if launch_http_server:
logger.info(
"Starting FastAPI server (connected to DiffusionServer at port %d).",
server_args.scheduler_port,
)
launch_http_server_only(server_args)
return all_processes
def _run_disagg_role_process(
gpu_id: int,
_local_rank: int,
rank: int,
server_args: ServerArgs,
pipe_writer: mp.connection.Connection,
task_pipes: list,
result_pipes: list,
):
"""Entry point for a disagg role process.
Uses the physical GPU index (gpu_id) as local_rank so that
torch.cuda.set_device(local_rank) selects the correct GPU.
This avoids relying on CUDA_VISIBLE_DEVICES remapping, which
may not work if CUDA was pre-initialized in the parent process.
"""
run_scheduler_process(
local_rank=gpu_id,
rank=rank,
master_port=server_args.master_port,
server_args=server_args,
pipe_writer=pipe_writer,
task_pipe_r=None,
result_pipe_w=None,
task_pipes_to_slaves=task_pipes,
result_pipes_from_slaves=result_pipes,
)
def launch_http_server_only(server_args):
# set for endpoints to access global_server_args
set_global_server_args(server_args)
@@ -199,10 +454,224 @@ def launch_http_server_only(server_args):
)
def parse_url_string(url_str: str) -> list[str]:
"""Parse a semicolon-separated URL string into a list.
Example: "tcp://10.0.0.1:35000;tcp://10.0.0.2:35000" -> ["tcp://...", "tcp://..."]
"""
return [u.strip() for u in url_str.split(";") if u.strip()]
def launch_disagg_server(server_args: ServerArgs):
"""Launch DiffusionServer head node + HTTP server (--disagg-role server).
No GPU workers are spawned. Connects to remote role instances
specified by --encoder-urls, --denoiser-urls, --decoder-urls.
Result endpoints use deterministic convention:
encoder result: scheduler_port + 1
denoiser result: scheduler_port + 2
decoder result: scheduler_port + 3
"""
configure_logger(server_args)
for name, val in [
("--encoder-urls", server_args.encoder_urls),
("--denoiser-urls", server_args.denoiser_urls),
("--decoder-urls", server_args.decoder_urls),
]:
if val is None:
raise ValueError(f"{name} is required for --disagg-role server")
host = server_args.host or "127.0.0.1"
base_port = server_args.scheduler_port
encoder_work_endpoints = parse_url_string(server_args.encoder_urls)
denoiser_work_endpoints = parse_url_string(server_args.denoiser_urls)
decoder_work_endpoints = parse_url_string(server_args.decoder_urls)
encoder_result_ep = f"tcp://{host}:{base_port + 1}"
denoiser_result_ep = f"tcp://{host}:{base_port + 2}"
decoder_result_ep = f"tcp://{host}:{base_port + 3}"
frontend_endpoint = f"tcp://{host}:{base_port}"
logger.info(
"Starting DiffusionServer: %d encoder(s), %d denoiser(s), %d decoder(s)",
len(encoder_work_endpoints),
len(denoiser_work_endpoints),
len(decoder_work_endpoints),
)
logger.info(" Frontend: %s", frontend_endpoint)
logger.info(" Encoder work endpoints: %s", encoder_work_endpoints)
logger.info(" Denoiser work endpoints: %s", denoiser_work_endpoints)
logger.info(" Decoder work endpoints: %s", decoder_work_endpoints)
logger.info(
" Result endpoints: encoder=%s, denoiser=%s, decoder=%s",
encoder_result_ep,
denoiser_result_ep,
decoder_result_ep,
)
diffusion_server = DiffusionServer(
frontend_endpoint=frontend_endpoint,
encoder_work_endpoints=encoder_work_endpoints,
denoiser_work_endpoints=denoiser_work_endpoints,
decoder_work_endpoints=decoder_work_endpoints,
encoder_result_endpoint=encoder_result_ep,
denoiser_result_endpoint=denoiser_result_ep,
decoder_result_endpoint=decoder_result_ep,
dispatch_policy_name=server_args.disagg_dispatch_policy,
timeout_s=float(server_args.disagg_timeout),
)
diffusion_server.start()
if not diffusion_server.wait_ready(timeout=30.0):
raise RuntimeError("DiffusionServer failed to bind sockets within 30 seconds")
logger.info(
"Starting HTTP server (connected to DiffusionServer at port %d).",
base_port,
)
launch_http_server_only(server_args)
def launch_disagg_role(server_args: ServerArgs):
"""Launch a standalone disaggregated role instance (--disagg-role encoder/denoising/decoder).
The instance:
1. Binds its work PULL socket on tcp://0.0.0.0:{scheduler_port}
2. Connects its result PUSH socket to the DiffusionServer head node
(derived from --disagg-server-addr + role offset)
3. Spawns GPU worker processes for the assigned role.
"""
configure_logger(server_args)
role_type = server_args.disagg_role
if server_args.disagg_server_addr is None:
raise ValueError(
"--disagg-server-addr is required for --disagg-role " f"{role_type.value}"
)
# Derive endpoints
work_endpoint = server_args.derive_pool_work_endpoint()
result_endpoint = server_args.derive_pool_result_endpoint()
logger.info(
"Starting disagg role: %s, num_gpus=%d",
role_type.value,
server_args.num_gpus,
)
logger.info(" Work endpoint (bind): %s", work_endpoint)
logger.info(" Result endpoint (connect): %s", result_endpoint)
logger.info(
" P2P: hostname=%s, ib_device=%s, pool_size=%d",
server_args.disagg_p2p_hostname,
server_args.disagg_ib_device,
server_args.disagg_transfer_pool_size,
)
# Build role-specific ServerArgs
# Use a different port for the scheduler's internal ROUTER socket to avoid
# conflicting with the pool work PULL socket (both bind on scheduler_port).
internal_scheduler_port = _find_available_port(
start=server_args.scheduler_port + 100, avoid={server_args.scheduler_port}
)
role_par = server_args.get_role_parallelism(role_type)
role_overrides = {
"disagg_role": role_type,
"disagg_mode": True,
"pool_work_endpoint": work_endpoint,
"pool_result_endpoint": result_endpoint,
"warmup": role_type == RoleType.ENCODER,
"scheduler_port": internal_scheduler_port,
# Per-role parallelism (None = auto-derive from num_gpus)
"tp_size": role_par["tp_size"],
"sp_degree": role_par["sp_degree"],
"ulysses_degree": role_par["ulysses_degree"],
"ring_degree": role_par["ring_degree"],
}
base_dict = {
f.name: getattr(server_args, f.name) for f in dataclasses.fields(server_args)
}
base_dict.update(role_overrides)
base_dict.pop("pipeline_config", None)
role_args = ServerArgs.from_kwargs(**base_dict)
# Spawn GPU worker processes
# NOTE: All ranks must be spawned before waiting for ready signals,
# because NCCL init_process_group blocks until all ranks connect.
num_gpus = server_args.num_gpus
base_gpu_id = server_args.base_gpu_id
pool_ctx = mp.get_context("spawn")
processes = []
readers = []
for rank_idx in range(num_gpus):
reader, writer = pool_ctx.Pipe(duplex=False)
gpu_id = base_gpu_id + rank_idx
process = pool_ctx.Process(
target=_run_disagg_role_process,
args=(gpu_id, rank_idx, rank_idx, role_args, writer, [], []),
name=f"sglang-{role_type.value}-r{rank_idx}",
daemon=True,
)
process.start()
processes.append(process)
readers.append(reader)
# Wait for all ranks to be ready (after all are spawned)
for rank_idx, reader in enumerate(readers):
try:
data = reader.recv()
except EOFError:
logger.error(
"Role %s rank %d is dead.",
role_type.value,
rank_idx,
)
raise
if data.get("status") != "ready":
raise RuntimeError(
f"Role {role_type.value} rank {rank_idx} failed to initialize."
)
reader.close()
logger.info(
"Role %s ready (%d GPU(s), work=%s)",
role_type.value.upper(),
num_gpus,
work_endpoint,
)
# Block until interrupted
try:
for p in processes:
p.join()
except KeyboardInterrupt:
logger.info("Role %s shutting down.", role_type.value)
def dispatch_launch(server_args: ServerArgs):
"""Route to the correct launch function based on --disagg-role."""
role = server_args.disagg_role
if role == RoleType.MONOLITHIC:
launch_server(server_args)
elif role == RoleType.SERVER:
launch_disagg_server(server_args)
elif role in (RoleType.ENCODER, RoleType.DENOISER, RoleType.DECODER):
launch_disagg_role(server_args)
else:
raise ValueError(f"Unknown disagg_role: {role}")
if __name__ == "__main__":
server_args = prepare_server_args(sys.argv[1:])
try:
launch_server(server_args)
dispatch_launch(server_args)
finally:
kill_process_tree(os.getpid(), include_parent=False)
@@ -208,9 +208,16 @@ class GPUWorker:
f"Related offload server args to disable: {suggested_args_str}"
)
def execute_forward(self, batch: List[Req]) -> OutputBatch:
def execute_forward(
self, batch: List[Req], return_req: bool = False
) -> OutputBatch | Req:
"""
Execute a forward pass.
Args:
batch: List of requests to process.
return_req: If True, return the raw Req instead of OutputBatch.
Used by disaggregated pipelines to access intermediate tensors.
"""
assert self.pipeline is not None
req = batch[0]
@@ -229,6 +236,11 @@ class GPUWorker:
req.log(server_args=self.server_args)
result = self.pipeline.forward(req, self.server_args)
# For disagg roles, return raw Req to let the caller handle
# the role-to-role tensor transfer before OutputBatch conversion.
if return_req and isinstance(result, Req):
return result
if isinstance(result, Req):
output_batch = OutputBatch(
output=result.output,
@@ -261,7 +273,8 @@ class GPUWorker:
self.do_mem_analysis(output_batch)
duration_ms = (time.monotonic() - start_time) * 1000
output_batch.metrics.total_duration_ms = duration_ms
if output_batch.metrics is not None:
output_batch.metrics.total_duration_ms = duration_ms
# Save output to file and return file path only if requested. Avoid the serialization
# and deserialization overhead between scheduler_client and gpu_worker.
@@ -526,6 +539,7 @@ def run_scheduler_process(
port_args=port_args,
task_pipes_to_slaves=task_pipes_to_slaves,
result_pipes_from_slaves=result_pipes_from_slaves,
local_rank=local_rank,
)
logger.info(f"Worker {rank}: Scheduler loop started.")
pipe_writer.send(
@@ -10,6 +10,10 @@ from typing import Any, List
import zmq
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
from sglang.multimodal_gen.runtime.disaggregation.scheduler_mixin import (
SchedulerDisaggMixin,
)
from sglang.multimodal_gen.runtime.distributed import get_world_group
from sglang.multimodal_gen.runtime.entrypoints.openai.utils import (
_parse_size,
@@ -20,6 +24,7 @@ from sglang.multimodal_gen.runtime.entrypoints.post_training.io_struct import (
UpdateWeightFromDiskReqInput,
)
from sglang.multimodal_gen.runtime.entrypoints.utils import (
GetDisaggStatsReq,
ListLorasReq,
MergeLoraWeightsReq,
SetLoraReq,
@@ -43,7 +48,7 @@ logger = init_logger(__name__)
MINIMUM_PICTURE_BASE64_FOR_WARMUP = "data:image/jpg;base64,iVBORw0KGgoAAAANSUhEUgAAACAAAAAgCAYAAABzenr0AAAACXBIWXMAAA7EAAAOxAGVKw4bAAAAbUlEQVRYhe3VsQ2AMAxE0Y/lIgNQULD/OqyCMgCihCKSG4yRuKuiNH6JLsoEbMACOGBcua9HOR7Y6w6swBwMy0qLTpkeI77qdEBpBFAHBBDAGH8WrwJKI4AAegUCfAKgEgpQDvh3CR3oQCuav58qlAw73kKCSgAAAABJRU5ErkJggg=="
class Scheduler:
class Scheduler(SchedulerDisaggMixin):
"""
Runs the main event loop for the rank 0 worker.
It listens for external requests via ZMQ and coordinates with other workers.
@@ -57,10 +62,17 @@ class Scheduler:
port_args: PortArgs,
task_pipes_to_slaves: list = None,
result_pipes_from_slaves: list = None,
local_rank: int | None = None,
):
self.server_args = server_args
self.port_args = port_args
# local_rank is the physical GPU index for torch.cuda.set_device.
# In non-disagg mode, it equals gpu_id. In disagg mode, it may differ
# (e.g., denoiser rank 0 on physical GPU 1).
if local_rank is None:
local_rank = gpu_id
set_global_server_args(server_args=server_args)
# Inter-process Communication
@@ -76,7 +88,7 @@ class Scheduler:
self.receiver = None
worker = GPUWorker(
local_rank=gpu_id,
local_rank=local_rank,
master_port=port_args.master_port,
rank=gpu_id,
server_args=server_args,
@@ -95,6 +107,7 @@ class Scheduler:
List[Req]: self._handle_generation,
ListLorasReq: self._handle_list_loras,
ShutdownReq: self._handle_shutdown,
GetDisaggStatsReq: self._handle_get_disagg_stats,
UpdateWeightFromDiskReqInput: self._handle_update_weights_from_disk,
GetWeightsChecksumReqInput: self._handle_get_weights_checksum,
}
@@ -114,6 +127,21 @@ class Scheduler:
self._max_consecutive_errors = 3
self._consecutive_error_count = 0
self._init_disagg_state(server_args, local_rank)
def get_disagg_metrics(self) -> dict | None:
"""Return disagg role metrics snapshot, or None if not in disagg mode."""
if self._disagg_metrics is None:
return None
return self._disagg_metrics.snapshot().to_dict()
def _handle_get_disagg_stats(self, _reqs: List[Any]) -> OutputBatch:
"""Handle stats request — return disagg metrics via OutputBatch.output."""
stats = self.get_disagg_metrics()
return OutputBatch(
output=stats or {"role": "monolithic", "message": "not in disagg mode"}
)
def _handle_set_lora(self, reqs: List[Any]) -> OutputBatch:
# TODO: return set status
# TODO: return with SetLoRAResponse or something more appropriate
@@ -166,6 +194,7 @@ class Scheduler:
)
else:
logger.info("Processing warmup req...")
return self.worker.execute_forward(reqs)
def return_result(
@@ -362,12 +391,20 @@ class Scheduler:
The main event loop that listens for ZMQ requests.
Handles abortion
"""
# Pool mode: all roles use the pool event loop
if self._disagg_role != RoleType.MONOLITHIC:
self._disagg_event_loop()
return
logger.debug(
f"Rank 0 scheduler listening on tcp://*:{self.server_args.scheduler_port}"
)
while self._running:
# Update queue depth for metrics
if self._disagg_metrics:
self._disagg_metrics.update_queue_depth(len(self.waiting_queue))
# 1: receive requests
try:
new_reqs = self.recv_reqs()
@@ -403,6 +440,10 @@ class Scheduler:
try:
processed_req = reqs[0]
is_warmup = (
processed_req.is_warmup if isinstance(processed_req, Req) else False
)
handler = self.request_handlers.get(type(processed_req))
if handler:
output_batch = handler(reqs)
@@ -415,16 +456,10 @@ class Scheduler:
f"Error executing request in scheduler event loop: {e}",
exc_info=True,
)
# Determine appropriate error response format
output_batch = (
OutputBatch(error=str(e))
if reqs and isinstance(reqs[0], Req)
else OutputBatch(error=str(e))
)
output_batch = OutputBatch(error=str(e))
# 3. return results
try:
# log warmup info
is_warmup = (
processed_req.is_warmup if isinstance(processed_req, Req) else False
)
@@ -457,6 +492,7 @@ class Scheduler:
if self.receiver is not None:
self.receiver.close()
self._cleanup_disagg()
self.context.destroy(linger=0)
def _broadcast_task(self, payload: dict[str, Any]) -> None:
@@ -14,6 +14,10 @@ from typing import Any, Callable, Literal, cast
import torch
from tqdm import tqdm
from sglang.multimodal_gen.runtime.disaggregation.roles import (
RoleType,
filter_modules_for_role,
)
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
PipelineComponentLoader,
)
@@ -82,6 +86,7 @@ class ComposedPipelineBase(ABC):
use. The pipeline should be stateless and not hold any batch state.
"""
self.server_args = server_args
self._disagg_role = server_args.disagg_role
self.model_path: str = model_path
self._stages: list[PipelineStage] = []
@@ -94,6 +99,20 @@ class ComposedPipelineBase(ABC):
if self._required_config_modules is None:
raise NotImplementedError("Subclass must set _required_config_modules")
# Filter modules based on disaggregation role
if self._disagg_role != RoleType.MONOLITHIC:
original_modules = list(self._required_config_modules)
self._required_config_modules = filter_modules_for_role(
self._required_config_modules, self._disagg_role
)
skipped = set(original_modules) - set(self._required_config_modules)
if skipped:
logger.info(
"Disagg role=%s: skipping modules %s",
self._disagg_role.value,
sorted(skipped),
)
# [module_name, gpu memory usage]
self.memory_usages: dict[str, float] = {}
# Load modules directly in initialization
@@ -169,6 +188,70 @@ class ComposedPipelineBase(ABC):
"""
return
# --- Config-name → pipeline_config attribute mapping ---
_CONFIG_ATTR_MAP: dict[str, str] = {
"vae": "vae_config",
"video_vae": "vae_config",
"audio_vae": "audio_vae_config",
}
def _init_skipped_component_configs(
self,
full_model_index: dict[str, Any],
server_args: ServerArgs,
) -> None:
"""Read HF JSON configs for skipped components and run
update_model_arch + post_init so pipeline_config is fully
initialized without loading weights.
"""
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
get_diffusers_component_config,
)
required = set(self.required_config_modules)
for module_name in full_model_index:
if module_name in required:
continue # will be loaded normally
cfg_attr = self._CONFIG_ATTR_MAP.get(module_name)
if cfg_attr is None:
continue # not a config we need to patch
pipeline_cfg = getattr(server_args.pipeline_config, cfg_attr, None)
if pipeline_cfg is None:
continue
try:
component_path = self._resolve_component_path(
server_args, module_name, module_name
)
hf_config = get_diffusers_component_config(
component_path=component_path
)
hf_config.pop("_class_name", None)
hf_config.pop("_diffusers_version", None)
pipeline_cfg.update_model_arch(hf_config)
if hasattr(pipeline_cfg, "post_init"):
pipeline_cfg.post_init()
logger.info(
"Disagg role=%s: initialized %s config from HF JSON "
"(spatial_compression_ratio=%s)",
self._disagg_role.value,
module_name,
getattr(
getattr(pipeline_cfg, "arch_config", None),
"spatial_compression_ratio",
"N/A",
),
)
except Exception as e:
logger.warning(
"Disagg role=%s: failed to read HF config for skipped "
"component %s: %s",
self._disagg_role.value,
module_name,
e,
)
def _resolve_component_path(
self, server_args: ServerArgs, module_name: str, load_module_name: str
) -> str:
@@ -214,7 +297,24 @@ class ComposedPipelineBase(ABC):
"MoE pipeline detected. Adding transformer_2 to self.required_config_modules..."
)
if "transformer_2" not in self.required_config_modules:
self.required_config_modules.append("transformer_2")
# Re-apply disagg role filter: only add transformer_2 if the
# role actually needs denoising modules.
from sglang.multimodal_gen.runtime.disaggregation.roles import (
get_module_role,
)
module_role = get_module_role("transformer_2")
if (
self._disagg_role == RoleType.MONOLITHIC
or module_role is None
or module_role == self._disagg_role
):
self.required_config_modules.append("transformer_2")
else:
logger.info(
"Disagg role=%s: skipping dynamically added module transformer_2",
self._disagg_role.value,
)
else:
logger.info(
"Boundary ratio found in model_index.json without transformers; "
@@ -237,6 +337,11 @@ class ComposedPipelineBase(ABC):
len(model_index) > 1
), "model_index.json must contain at least one pipeline module"
# In disagg mode, read HF config for skipped components (e.g., VAE)
# so that update_model_arch + post_init can derive pipeline_config.
if self._disagg_role != RoleType.MONOLITHIC:
self._init_skipped_component_configs(model_index, server_args)
model_index = {
required_module: model_index[required_module]
for required_module in self.required_config_modules
@@ -341,6 +446,19 @@ class ComposedPipelineBase(ABC):
assert self.modules is not None, "No modules are registered"
# Filter stages based on disaggregation role
if self._disagg_role != RoleType.MONOLITHIC:
if stage.role_affinity != self._disagg_role:
if stage_name is None:
stage_name = self._infer_stage_name(stage)
logger.info(
"Disagg role=%s: skipping stage %s (affinity=%s)",
self._disagg_role.value,
stage_name,
stage.role_affinity.value,
)
return self
if stage_name is None:
stage_name = self._infer_stage_name(stage)
if stage_name in self._stage_name_mapping:
@@ -13,6 +13,7 @@ from enum import Enum, auto
import torch
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
VerificationResult,
@@ -103,6 +104,11 @@ class PipelineStage(ABC):
"""
pass
# Default role affinity: ENCODER. Override in subclasses for DENOISING/DECODER.
@property
def role_affinity(self) -> RoleType:
return RoleType.ENCODER
# execute on all ranks by default
@property
def parallelism_type(self) -> StageParallelismType:
@@ -56,6 +56,12 @@ class DecodingStage(PipelineStage):
output format (e.g., pixel values).
"""
@property
def role_affinity(self):
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
return RoleType.DECODER
def __init__(self, vae, pipeline=None, component_name: str = "vae") -> None:
super().__init__()
self.vae: ParallelTiledVAE = vae
@@ -35,6 +35,7 @@ from sglang.multimodal_gen.runtime.cache.cache_dit_integration import (
refresh_context_on_dual_transformer,
refresh_context_on_transformer,
)
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
from sglang.multimodal_gen.runtime.distributed import (
cfg_model_parallel_all_reduce,
get_local_torch_device,
@@ -152,6 +153,10 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
the initial noise into the final output.
"""
@property
def role_affinity(self):
return RoleType.DENOISER
def __init__(
self, transformer, scheduler, pipeline=None, transformer_2=None, vae=None
) -> None:
@@ -23,6 +23,12 @@ from sglang.multimodal_gen import envs
from sglang.multimodal_gen.configs.models.encoders import T5Config
from sglang.multimodal_gen.configs.pipeline_configs.base import PipelineConfig
from sglang.multimodal_gen.configs.quantization.nunchaku import NunchakuSVDQuantArgs
from sglang.multimodal_gen.runtime.disaggregation.disagg_args import (
DisaggArgsMixin,
add_disagg_cli_args,
convert_disagg_role_string,
)
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
from sglang.multimodal_gen.runtime.layers.quantization.configs.nunchaku_config import (
NunchakuConfig,
)
@@ -94,7 +100,7 @@ class Backend(str, Enum):
@dataclasses.dataclass
class ServerArgs:
class ServerArgs(DisaggArgsMixin):
# Model and path configuration (for convenience)
model_path: str
@@ -230,10 +236,35 @@ class ServerArgs:
# MoE parameters used by Wan2.2
boundary_ratio: float | None = None
# Disaggregation — fields defined here, methods in DisaggArgsMixin,
# CLI registration in disagg_args.add_disagg_cli_args().
base_gpu_id: int = 0
disagg_role: RoleType = RoleType.MONOLITHIC
disagg_timeout: int = 600
disagg_dispatch_policy: str = "round_robin"
disagg_mode: bool = False
disagg_server_addr: str | None = None
encoder_urls: str | None = None
denoiser_urls: str | None = None
decoder_urls: str | None = None
encoder_tp: int | None = None
denoiser_tp: int | None = None
denoiser_sp: int | None = None
denoiser_ulysses: int | None = None
denoiser_ring: int | None = None
decoder_tp: int | None = None
disagg_transfer_pool_size: int = 256 * 1024 * 1024
disagg_p2p_hostname: str = "127.0.0.1"
disagg_ib_device: str | None = None
pool_work_endpoint: str | None = None
pool_result_endpoint: str | None = None
# Logging
log_level: str = "info"
uvicorn_access_log_exclude_prefixes: list[str] = field(default_factory=list)
# get_role_parallelism, derive_pool_*_endpoint — from DisaggArgsMixin
@property
def broker_port(self) -> int:
return self.port + 1
@@ -411,12 +442,22 @@ class ServerArgs:
)
def _adjust_network_ports(self):
# Disagg role instances (encoder/denoiser/decoder) don't serve HTTP,
# so skip settling the HTTP port to avoid unnecessary port collisions.
needs_http = self.disagg_role in (
RoleType.MONOLITHIC,
RoleType.SERVER,
)
if self.strict_ports:
self._require_port(self.port, "HTTP")
if needs_http:
self._require_port(self.port, "HTTP")
self._require_port(self.scheduler_port, "Scheduler")
self._require_port(self.master_port, "Master")
if self.master_port is not None:
self._require_port(self.master_port, "Master")
else:
self.port = self.settle_port(self.port)
if needs_http:
self.port = self.settle_port(self.port)
initial_scheduler_port = self.scheduler_port + (
random.randint(0, 100) if self.scheduler_port == 5555 else 0
)
@@ -611,6 +652,9 @@ class ServerArgs:
# configure logger before use
configure_logger(server_args=self)
# Convert string disagg_role to enum (from CLI/config)
convert_disagg_role_string(self.__dict__)
# 1. adjust parameters
self._adjust_parameters()
@@ -699,6 +743,7 @@ class ServerArgs:
default=ServerArgs.num_gpus,
help="The number of GPUs to use.",
)
parser.add_argument(
"--tp-size",
type=int,
@@ -758,6 +803,9 @@ class ServerArgs:
"Increase this value if you encounter 'Connection closed by peer' errors after the service is idle. ",
)
# Disaggregated diffusion args (defined in disagg_args.py)
add_disagg_cli_args(parser)
# Prompt text file for batch processing
parser.add_argument(
"--prompt-file-path",
@@ -1139,6 +1187,9 @@ class ServerArgs:
if "backend" in kwargs and isinstance(kwargs["backend"], str):
kwargs["backend"] = Backend.from_string(kwargs["backend"])
# Convert disagg_role string to enum if necessary
convert_disagg_role_string(kwargs)
kwargs["pipeline_config"] = PipelineConfig.from_kwargs(kwargs)
return cls(**kwargs)
@@ -130,6 +130,7 @@ def get_zmq_socket(
endpoint: str,
bind: bool,
max_bind_retries: int = 10,
same_port: bool = False,
) -> tuple[zmq.Socket, str]:
"""
Create and configure a ZMQ socket.
@@ -140,10 +141,13 @@ def get_zmq_socket(
endpoint: Endpoint string (e.g., "tcp://localhost:5555")
bind: Whether to bind (True) or connect (False)
max_bind_retries: Maximum number of retries if bind fails due to address already in use
same_port: If True, retry on the same port instead of incrementing.
Useful when the port must be fixed (e.g., disagg sockets where
DiffusionServer connects to a pre-determined port).
Returns:
A tuple of (socket, actual_endpoint). The actual_endpoint may differ from the
requested endpoint if bind retry was needed.
requested endpoint if bind retry was needed (and same_port is False).
"""
mem = psutil.virtual_memory()
total_mem = mem.total / 1024**3
@@ -182,13 +186,15 @@ def get_zmq_socket(
port_match = re.search(r":(\d+)$", endpoint)
if port_match and max_bind_retries > 1:
import time as _time
original_port = int(port_match.group(1))
last_exception = None
for attempt in range(max_bind_retries):
try:
current_endpoint = endpoint
if attempt > 0:
if attempt > 0 and not same_port:
# Try next port (increment by 42 to match settle_port logic)
current_port = original_port + attempt * 42
current_endpoint = re.sub(
@@ -198,6 +204,11 @@ def get_zmq_socket(
f"ZMQ bind failed for port {original_port + (attempt - 1) * 42}, "
f"retrying with port {current_port} (attempt {attempt + 1}/{max_bind_retries})"
)
elif attempt > 0:
logger.info(
f"ZMQ bind attempt {attempt + 1}/{max_bind_retries} "
f"on same port {original_port}..."
)
socket.bind(current_endpoint)
@@ -212,7 +223,21 @@ def get_zmq_socket(
except zmq.ZMQError as e:
last_exception = e
if e.errno == zmq.EADDRINUSE and attempt < max_bind_retries - 1:
# Address already in use, try next port
# Address already in use, retry
# Longer sleep for same_port (waiting for TIME_WAIT release)
_time.sleep(1.0 if same_port else 0.5)
# Re-create socket since ZMQ socket state may be invalid after failed bind
socket.close()
socket = context.socket(socket_type)
if endpoint.find("[") != -1:
socket.setsockopt(zmq.IPV6, 1)
if socket_type == zmq.PUSH:
set_send_opt()
elif socket_type == zmq.PULL:
set_recv_opt()
elif socket_type in [zmq.DEALER, zmq.REQ, zmq.REP, zmq.ROUTER]:
set_send_opt()
set_recv_opt()
continue
elif attempt == max_bind_retries - 1:
# Last attempt failed
@@ -88,7 +88,9 @@ STANDALONE_FILES = {
"../cli/test_generate_t2i_perf.py",
"test_update_weights_from_disk.py",
],
"2-gpu": [],
"2-gpu": [
"test_disagg_server.py",
],
}
# New standalone files may omit an estimate once to learn the real CI runtime.
@@ -99,7 +101,11 @@ STANDALONE_FILE_EST_TIMES = {
"../cli/test_generate_t2i_perf.py": 240.0,
"test_update_weights_from_disk.py": 480.0,
},
"2-gpu": {},
"2-gpu": {
# Two disagg clusters × (~3 min startup + ~1 min generate) ≈ 8 min.
# Raise if CI reports a higher measured time.
"test_disagg_server.py": 600.0,
},
}
# Backward-compatible suite view for scripts that still operate on file lists.
@@ -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")
class TestPerRoleParallelism(unittest.TestCase):
"""Test per-role parallelism args and get_role_parallelism helper."""
def _from_dict(self, kwargs):
with patch.object(
PipelineConfig, "from_kwargs", return_value=QwenImagePipelineConfig()
):
return ServerArgs.from_dict(kwargs)
def test_defaults_are_none(self):
args = self._from_dict({"model_path": "/fake"})
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
for role in [RoleType.ENCODER, RoleType.DENOISER, RoleType.DECODER]:
par = args.get_role_parallelism(role)
self.assertIsNone(par["tp_size"])
self.assertIsNone(par["sp_degree"])
self.assertIsNone(par["ulysses_degree"])
self.assertIsNone(par["ring_degree"])
def test_encoder_overrides(self):
args = self._from_dict({"model_path": "/fake", "encoder_tp": 2})
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
par = args.get_role_parallelism(RoleType.ENCODER)
self.assertEqual(par["tp_size"], 2)
self.assertIsNone(par["sp_degree"])
self.assertIsNone(par["ulysses_degree"])
self.assertIsNone(par["ring_degree"])
def test_denoiser_overrides(self):
args = self._from_dict(
{
"model_path": "/fake",
"denoiser_tp": 1,
"denoiser_sp": 8,
"denoiser_ulysses": 4,
"denoiser_ring": 2,
}
)
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
par = args.get_role_parallelism(RoleType.DENOISER)
self.assertEqual(par["tp_size"], 1)
self.assertEqual(par["sp_degree"], 8)
self.assertEqual(par["ulysses_degree"], 4)
self.assertEqual(par["ring_degree"], 2)
def test_decoder_overrides(self):
args = self._from_dict({"model_path": "/fake", "decoder_tp": 2})
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
par = args.get_role_parallelism(RoleType.DECODER)
self.assertEqual(par["tp_size"], 2)
self.assertIsNone(par["sp_degree"])
self.assertIsNone(par["ulysses_degree"])
self.assertIsNone(par["ring_degree"])
def test_monolithic_returns_all_none(self):
args = self._from_dict({"model_path": "/fake", "encoder_tp": 2})
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
par = args.get_role_parallelism(RoleType.MONOLITHIC)
self.assertIsNone(par["tp_size"])
self.assertIsNone(par["sp_degree"])
def test_mixed_roles_independent(self):
"""Per-role args don't interfere with each other."""
args = self._from_dict(
{
"model_path": "/fake",
"encoder_tp": 1,
"denoiser_tp": 2,
"decoder_tp": 4,
}
)
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
self.assertEqual(args.get_role_parallelism(RoleType.ENCODER)["tp_size"], 1)
self.assertEqual(args.get_role_parallelism(RoleType.DENOISER)["tp_size"], 2)
self.assertEqual(args.get_role_parallelism(RoleType.DECODER)["tp_size"], 4)
def test_cli_args_parsed(self):
"""Per-role parallelism args are parsed from CLI."""
parser = FlexibleArgumentParser()
ServerArgs.add_cli_args(parser)
argv = [
"--model-path",
"/fake",
"--denoiser-tp",
"2",
"--denoiser-sp",
"4",
"--denoiser-ulysses",
"2",
"--denoiser-ring",
"2",
"--encoder-tp",
"1",
]
args, unknown = parser.parse_known_args(argv)
self.assertEqual(args.denoiser_tp, 2)
self.assertEqual(args.denoiser_sp, 4)
self.assertEqual(args.denoiser_ulysses, 2)
self.assertEqual(args.denoiser_ring, 2)
self.assertEqual(args.encoder_tp, 1)
self.assertIsNone(args.decoder_tp)
class TestPipelineResolutionCliOverride(unittest.TestCase):
def setUp(self):
_get_config_info.cache_clear()
@@ -102,25 +211,5 @@ class TestPipelineResolutionCliOverride(unittest.TestCase):
self.assertEqual(server_args.pipeline_config.resolution, 768)
class TestComponentPathParsing(unittest.TestCase):
def test_extract_component_paths_accepts_config_expanded_keys(self):
component_paths, remaining = ServerArgs._extract_component_paths(
[
"--component-paths.spatial-upsampler",
"/tmp/latent_upsampler",
"--component_paths.distilled-lora=/tmp/distilled.safetensors",
]
)
self.assertEqual(
component_paths,
{
"spatial_upsampler": "/tmp/latent_upsampler",
"distilled_lora": "/tmp/distilled.safetensors",
},
)
self.assertEqual(remaining, [])
if __name__ == "__main__":
unittest.main()