[diffusion] feat: disaggregated diffusion (#21701)
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user