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