[diffusion] refactor: simplify disaggregation transport hygiene (#33844)

This commit is contained in:
Mick
2026-08-06 21:50:32 +08:00
committed by GitHub
parent 1a15cf1536
commit 45dfd80674
5 changed files with 29 additions and 205 deletions
@@ -44,6 +44,7 @@ from sglang.multimodal_gen.runtime.disaggregation.transport.protocol import (
TransferMsgType, TransferMsgType,
TransferPushedMsg, TransferPushedMsg,
TransferRegisterMsg, TransferRegisterMsg,
TransferStagedMsg,
decode_transfer_msg, decode_transfer_msg,
encode_transfer_msg, encode_transfer_msg,
is_transfer_message, is_transfer_message,
@@ -106,15 +107,7 @@ _EXCLUDE_FIELDS = frozenset(
} }
) )
# Sampling-params fields that should never be transferred across roles: # SamplingParams fields that are reconstructed locally or not JSON-safe.
# - data_type / supported_resolutions: enums / non-JSON classvars reconstructed on the receiver
# - teacache_params: model-specific object, not JSON-safe
# - output_* / save_output / return_*: output-side concerns owned by the decoder role
#
# Everything else on SamplingParams is forwarded automatically via a field-walk
# below; this keeps new request-level features (e.g. Qwen-Image's
# true_cfg_scale, guidance_rescale, cfg_normalization, ...) from silently
# getting dropped just because nobody remembered to add them to a whitelist.
_SAMPLING_PARAMS_EXCLUDE_FIELDS = frozenset( _SAMPLING_PARAMS_EXCLUDE_FIELDS = frozenset(
{ {
"data_type", "data_type",
@@ -123,11 +116,6 @@ _SAMPLING_PARAMS_EXCLUDE_FIELDS = frozenset(
} }
) )
_BASE_SP_DEFAULTS: dict[str, Any] = {}
for _f in dataclasses.fields(SamplingParams):
if _f.default is not dataclasses.MISSING:
_BASE_SP_DEFAULTS[_f.name] = _f.default
def _is_tensor_like(value) -> bool: def _is_tensor_like(value) -> bool:
if isinstance(value, torch.Tensor): if isinstance(value, torch.Tensor):
@@ -290,8 +278,7 @@ def extract_transfer_fields(req) -> tuple[dict, dict]:
value = getattr(sp, name, None) value = getattr(sp, name, None)
if value is None: if value is None:
continue continue
base_default = _BASE_SP_DEFAULTS.get(name, dataclasses.MISSING) if _is_default(value, f):
if base_default is not dataclasses.MISSING and value == base_default:
continue continue
try: try:
scalar_fields[name] = _to_json_serializable(value) scalar_fields[name] = _to_json_serializable(value)
@@ -669,12 +656,11 @@ class SchedulerDisaggMixin:
Called from the recv prefetch thread. Loads on _transfer_stream Called from the recv prefetch thread. Loads on _transfer_stream
and builds the Req, so the main thread can start compute immediately. and builds the Req, so the main thread can start compute immediately.
Returns (req, load_event, request_id, role_name, prealloc_slot_id). Returns (req, load_event, request_id, prealloc_slot_id).
""" """
request_id = msg["request_id"] request_id = msg["request_id"]
manifest = msg.get("manifest", {}) manifest = msg.get("manifest", {})
scalar_fields = msg.get("scalar_fields", {}) scalar_fields = msg.get("scalar_fields", {})
role_name = self._disagg_role.value.upper()
if self._disagg_metrics: if self._disagg_metrics:
self._disagg_metrics.record_request_start(request_id) self._disagg_metrics.record_request_start(request_id)
@@ -710,7 +696,7 @@ class SchedulerDisaggMixin:
# running denoising loop on the main thread. Deferred to main thread # running denoising loop on the main thread. Deferred to main thread
# in _disagg_prefetch_event_loop, right before compute. # in _disagg_prefetch_event_loop, right before compute.
return (req, load_event, request_id, role_name, prealloc_slot_id, scalar_fields) return req, load_event, request_id, prealloc_slot_id
# ------------------------------------------------------------------ # ------------------------------------------------------------------
# Broadcast # Broadcast
@@ -865,11 +851,7 @@ class SchedulerDisaggMixin:
- queue timeout: broadcast "skip" - queue timeout: broadcast "skip"
- shutdown: broadcast None - shutdown: broadcast None
""" """
is_multi_rank = ( is_multi_rank = self._is_multi_rank()
self.server_args.sp_degree != 1
or self.server_args.tp_size > 1
or self.server_args.enable_cfg_parallel
)
while self._running: while self._running:
try: try:
@@ -882,9 +864,7 @@ class SchedulerDisaggMixin:
if msg_type == "transfer_compute": if msg_type == "transfer_compute":
# Load already done by recv thread # Load already done by recv thread
req, load_event, request_id, rn, prealloc_slot_id, scalar_fields = ( req, load_event, request_id, prealloc_slot_id = data
data
)
# Wait for load to complete on compute stream # Wait for load to complete on compute stream
if load_event is not None: if load_event is not None:
torch.get_device_module().current_stream().wait_event( torch.get_device_module().current_stream().wait_event(
@@ -913,9 +893,9 @@ class SchedulerDisaggMixin:
_init_disagg_request_scheduler(self, req) _init_disagg_request_scheduler(self, req)
# Run compute # Run compute
if self._disagg_role == RoleType.DENOISER: if self._disagg_role == RoleType.DENOISER:
self._disagg_denoiser_compute(req, request_id, rn) self._disagg_denoiser_compute(req, request_id)
elif self._disagg_role == RoleType.DECODER: elif self._disagg_role == RoleType.DECODER:
self._disagg_decoder_compute(req, request_id, rn) self._disagg_decoder_compute(req, request_id)
elif msg_type == "transfer_control": elif msg_type == "transfer_control":
# alloc, push messages — handle on main thread (rank 0 only) # alloc, push messages — handle on main thread (rank 0 only)
@@ -1249,8 +1229,6 @@ class SchedulerDisaggMixin:
request_id = msg["request_id"] request_id = msg["request_id"]
manifest = msg.get("manifest", {}) manifest = msg.get("manifest", {})
scalar_fields = msg.get("scalar_fields", {}) scalar_fields = msg.get("scalar_fields", {})
role_name = self._disagg_role.value.upper()
if self._disagg_metrics: if self._disagg_metrics:
self._disagg_metrics.record_request_start(request_id) self._disagg_metrics.record_request_start(request_id)
@@ -1299,9 +1277,9 @@ class SchedulerDisaggMixin:
# 7. Run compute # 7. Run compute
if self._disagg_role == RoleType.DENOISER: if self._disagg_role == RoleType.DENOISER:
self._disagg_denoiser_compute(req, request_id, role_name) self._disagg_denoiser_compute(req, request_id)
elif self._disagg_role == RoleType.DECODER: elif self._disagg_role == RoleType.DECODER:
self._disagg_decoder_compute(req, request_id, role_name) self._disagg_decoder_compute(req, request_id)
# ------------------------------------------------------------------ # ------------------------------------------------------------------
# Compute # Compute
@@ -1407,9 +1385,7 @@ class SchedulerDisaggMixin:
with trace_slice(ctx, DiffStage.SCHEDULER_DISPATCH, thread_finish_flag=True): with trace_slice(ctx, DiffStage.SCHEDULER_DISPATCH, thread_finish_flag=True):
yield yield
def _disagg_denoiser_compute( def _disagg_denoiser_compute(self: Scheduler, req: Req, request_id: str) -> None:
self: Scheduler, req: Req, request_id: str, role_name: str
) -> None:
"""Run denoiser compute in transfer mode, then stage output for decoder. """Run denoiser compute in transfer mode, then stage output for decoder.
Note: Scheduler timestep init is done in _handle_transfer_ready Note: Scheduler timestep init is done in _handle_transfer_ready
@@ -1480,9 +1456,7 @@ class SchedulerDisaggMixin:
duration_s, duration_s,
) )
def _disagg_decoder_compute( def _disagg_decoder_compute(self: Scheduler, req: Req, request_id: str) -> None:
self: Scheduler, req: Req, request_id: str, role_name: str
) -> None:
"""Run decoder compute in transfer mode, send result to DS. """Run decoder compute in transfer mode, send result to DS.
Decoder result is sent as raw ZMQ multipart frames (same format as Decoder result is sent as raw ZMQ multipart frames (same format as
@@ -1622,22 +1596,20 @@ class SchedulerDisaggMixin:
self._disagg_metrics.record_request_failed(request_id) self._disagg_metrics.record_request_failed(request_id)
return return
# 2. Build transfer metadata dict while staging runs (CPU work, overlapped) # 2. Build transfer metadata while staging runs (CPU work, overlapped)
staged_data = { staged_msg = TransferStagedMsg(
"msg_type": "transfer_staged", request_id=request_id,
"request_id": request_id, data_size=staged.slot.size if staged.slot else 0,
"data_size": staged.slot.size if staged.slot else 0, manifest=staged.manifest,
"manifest": staged.manifest, session_id=self._transfer_manager.session_id,
"session_id": self._transfer_manager.session_id, pool_ptr=self._transfer_manager.pool_data_ptr,
"pool_ptr": self._transfer_manager.pool_data_ptr, slot_offset=staged.slot.offset if staged.slot else 0,
"slot_offset": staged.slot.offset if staged.slot else 0, scalar_fields=staged.scalar_fields,
"scalar_fields": staged.scalar_fields, )
}
msg_bytes = json.dumps(staged_data, separators=(",", ":")).encode("utf-8")
# 3. Wait for staging to complete before sending (buffer must be ready) # 3. Wait for staging to complete before sending (buffer must be ready)
if stage_event is not None: if stage_event is not None:
stage_event.synchronize() stage_event.synchronize()
# 4. Send transfer staged message # 4. Send transfer staged message
self._pool_result_push.send_multipart([TRANSFER_MAGIC, msg_bytes]) self._pool_result_push.send_multipart(encode_transfer_msg(staged_msg))
@@ -93,20 +93,6 @@ class BuddyAllocator:
"free_blocks_by_size": free_blocks_by_order, "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) --- # --- Internal (caller must hold self._lock) ---
def _allocate_locked(self, target_order: int, request_id: str | None) -> int | None: def _allocate_locked(self, target_order: int, request_id: str | None) -> int | None:
@@ -101,7 +101,6 @@ class TransferTensorBuffer:
def write_tensor( def write_tensor(
self, self,
handle: SlotHandle, handle: SlotHandle,
name: str,
tensor: torch.Tensor, tensor: torch.Tensor,
byte_offset: int = 0, byte_offset: int = 0,
stream: torch.Stream | None = None, stream: torch.Stream | None = None,
@@ -186,7 +185,7 @@ class TransferTensorBuffer:
entries = [] entries = []
if isinstance(value, torch.Tensor): if isinstance(value, torch.Tensor):
nbytes = self.write_tensor(handle, name, value, byte_offset, stream) nbytes = self.write_tensor(handle, value, byte_offset, stream)
entries.append( entries.append(
{ {
"offset": byte_offset, "offset": byte_offset,
@@ -201,9 +200,7 @@ class TransferTensorBuffer:
for i, t in enumerate(value): for i, t in enumerate(value):
if t is None: if t is None:
continue continue
nbytes = self.write_tensor( nbytes = self.write_tensor(handle, t, byte_offset, stream)
handle, f"{name}[{i}]", t, byte_offset, stream
)
entries.append( entries.append(
{ {
"offset": byte_offset, "offset": byte_offset,
@@ -261,12 +258,3 @@ class TransferTensorBuffer:
) )
return result 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
@@ -22,7 +22,7 @@ logger = logging.getLogger(__name__)
@dataclass @dataclass
class StagedTransfer: class StagedTransfer:
request_id: str request_id: str
slot: SlotHandle slot: SlotHandle | None
manifest: dict manifest: dict
scalar_fields: dict = field(default_factory=dict) scalar_fields: dict = field(default_factory=dict)
@@ -71,68 +71,6 @@ class DiffusionTransferManager:
def pool_size(self) -> int: def pool_size(self) -> int:
return self._buffer.pool_size 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.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.get_device_module().is_available():
torch.get_device_module().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( def stage_tensors_async(
self, self,
request_id: str, request_id: str,
@@ -315,39 +253,6 @@ class DiffusionTransferManager:
) )
return pending return pending
def load_tensors(
self,
request_id: str,
manifest: dict,
device: torch.device | str = current_platform.device_type,
stream: torch.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.get_device_module().is_available():
torch.get_device_module().synchronize()
logger.debug(
"TransferManager: loaded %d tensor fields for %s to %s",
len(tensors),
request_id,
device,
)
return tensors
def register_prealloc_as_receive( def register_prealloc_as_receive(
self, request_id: str, slot: "SlotHandle" self, request_id: str, slot: "SlotHandle"
) -> "PendingReceive": ) -> "PendingReceive":
@@ -365,27 +270,6 @@ class DiffusionTransferManager:
self._buffer.free(pending.slot) self._buffer.free(pending.slot)
logger.debug("TransferManager: freed receive slot for %s", request_id) 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: def cleanup(self) -> None:
self._engine.deregister_buffer(self._buffer.pool_data_ptr) self._engine.deregister_buffer(self._buffer.pool_data_ptr)
logger.info("DiffusionTransferManager cleaned up") logger.info("DiffusionTransferManager cleaned up")
@@ -6,12 +6,9 @@ in frame[0] and JSON payload in frame[1].
""" """
import json import json
import logging
from dataclasses import asdict, dataclass, field from dataclasses import asdict, dataclass, field
from typing import Any from typing import Any
logger = logging.getLogger(__name__)
TRANSFER_MAGIC = b"__transfer__" TRANSFER_MAGIC = b"__transfer__"
@@ -37,14 +34,11 @@ class TransferStagedMsg:
msg_type: str = TransferMsgType.STAGED msg_type: str = TransferMsgType.STAGED
request_id: str = "" request_id: str = ""
data_size: int = 0 data_size: int = 0
manifest: dict = None manifest: dict = field(default_factory=dict)
session_id: str = "" session_id: str = ""
pool_ptr: int = 0 pool_ptr: int = 0
slot_offset: int = 0 slot_offset: int = 0
scalar_fields: dict = field(default_factory=dict)
def __post_init__(self):
if self.manifest is None:
self.manifest = {}
@dataclass @dataclass