From 0801cc05ed202544e2e5c67d6b0a1d454de54fca Mon Sep 17 00:00:00 2001 From: Makcum888e <79456407+Makcum888e@users.noreply.github.com> Date: Mon, 25 May 2026 13:51:25 +0300 Subject: [PATCH] [Diffusion][NPU] Disaggregation diffusion stages support for NPU (#25895) Co-authored-by: ronnie_zheng --- .../ascend-npus/diffusion/disaggregation.mdx | 33 ++++++++++++++ .../runtime/disaggregation/scheduler_mixin.py | 27 ++++++++---- .../disaggregation/transport/buffer.py | 16 +++---- .../runtime/disaggregation/transport/codec.py | 2 +- .../disaggregation/transport/manager.py | 44 ++++++++++--------- 5 files changed, 84 insertions(+), 38 deletions(-) create mode 100644 docs_new/docs/hardware-platforms/ascend-npus/diffusion/disaggregation.mdx diff --git a/docs_new/docs/hardware-platforms/ascend-npus/diffusion/disaggregation.mdx b/docs_new/docs/hardware-platforms/ascend-npus/diffusion/disaggregation.mdx new file mode 100644 index 000000000..d3adf0bb2 --- /dev/null +++ b/docs_new/docs/hardware-platforms/ascend-npus/diffusion/disaggregation.mdx @@ -0,0 +1,33 @@ +--- +title: "Disaggregation of Diffusion Pipeline on Ascend NPU" +--- + +## Quick Start + +Please follow the [NPU installation guide](../ascend_npu.mdx) first. + +Then install Mooncake from sources. + +```bash +git clone https://github.com/kvcache-ai/Mooncake.git +cd Mooncake +git checkout v0.3.10.post2 +bash dependencies.sh +mkdir build +cd build +export GLOG_logtostderr=1 +cmake -DUSE_ASCEND_DIRECT=ON .. +make -j +make install +cd ../mooncake-wheel/ +pip install . +``` + +Before run servers ensure that *.so files of mooncake in LD_LIBRARY_PATH +```bash +export LD_LIBRARY_PATH=/usr/local/python3.11.14/lib/python3.11/site-packages/mooncake:$LD_LIBRARY_PATH +``` + +## Examples + +For usage examples please follow the [main disaggregation guide](../../../sglang-diffusion/disaggregation.mdx) diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py b/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py index 0cb1b0aac..24dc62806 100644 --- a/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py +++ b/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py @@ -51,6 +51,7 @@ from sglang.multimodal_gen.runtime.pipelines_core import Req from sglang.multimodal_gen.runtime.pipelines_core.diffusion_scheduler_utils import ( clone_scheduler_runtime, ) +from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.utils.common import get_zmq_socket from sglang.multimodal_gen.runtime.utils.distributed import broadcast_pyobj from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger @@ -205,7 +206,7 @@ def _init_disagg_request_scheduler(self: Scheduler, req: Req) -> None: scheduler_template = self.worker.pipeline.get_module("scheduler") if scheduler_template is None: return - device = torch.device(f"cuda:{self.worker.local_rank}") + device = torch.device(f"{current_platform.device_type}:{self.worker.local_rank}") _init_request_scheduler_from_template(scheduler_template, req, device) @@ -384,8 +385,8 @@ class SchedulerDisaggMixin: if self._disagg_role != RoleType.MONOLITHIC: self._disagg_metrics = DisaggMetrics(role=self._disagg_role.value) - device = torch.device(f"cuda:{local_rank}") - self._transfer_stream = torch.cuda.Stream(device=device) + device = torch.device(f"{current_platform.device_type}:{local_rank}") + self._transfer_stream = torch.get_device_module().Stream(device=device) self._init_disagg_sockets() self._init_disagg_transfer_manager() @@ -455,7 +456,11 @@ class SchedulerDisaggMixin: ) # Use GPU buffer when engine supports GPUDirect RDMA, CPU pinned otherwise - device = f"cuda:{physical_gpu_id}" if engine.supports_gpu_direct else "cpu" + device = ( + f"{current_platform.device_type}:{physical_gpu_id}" + if engine.supports_gpu_direct + else "cpu" + ) buffer = TransferTensorBuffer( pool_size=pool_size, device=device, role_name=self._disagg_role.value ) @@ -651,7 +656,7 @@ class SchedulerDisaggMixin: self._transfer_manager.register_prealloc_as_receive(request_id, slot) # Load tensors on transfer_stream (non-blocking) - local_device = f"cuda:{self.worker.local_rank}" + local_device = f"{current_platform.device_type}:{self.worker.local_rank}" tensors, load_event = self._transfer_manager.load_tensors_async( request_id, manifest, @@ -778,7 +783,9 @@ class SchedulerDisaggMixin: # (set via torch.cuda.set_device(local_rank) during init), which is # already the right physical GPU — the .to() is effectively a no-op # but makes the invariant explicit for future readers. - local_device = torch.device(f"cuda:{self.worker.local_rank}") + local_device = torch.device( + f"{current_platform.device_type}:{self.worker.local_rank}" + ) for key, value in list(tensor_fields.items()): if isinstance(value, torch.Tensor): tensor_fields[key] = value.to(local_device, non_blocking=True) @@ -847,7 +854,9 @@ class SchedulerDisaggMixin: ) # Wait for load to complete on compute stream if load_event is not None: - torch.cuda.current_stream().wait_event(load_event) + torch.get_device_module().current_stream().wait_event( + load_event + ) # Now safe to free the receive slot if prealloc_slot_id is not None: with self._transfer_manager._lock: @@ -1222,7 +1231,7 @@ class SchedulerDisaggMixin: self._transfer_manager.register_prealloc_as_receive(request_id, slot) # 1. Start load on transfer_stream (non-blocking) - local_device = f"cuda:{self.worker.local_rank}" + local_device = f"{current_platform.device_type}:{self.worker.local_rank}" tensors, load_event = self._transfer_manager.load_tensors_async( request_id, manifest, @@ -1239,7 +1248,7 @@ class SchedulerDisaggMixin: # 4. Wait for load before compute (GPU must see the data) if load_event is not None: - torch.cuda.current_stream().wait_event(load_event) + torch.get_device_module().current_stream().wait_event(load_event) # 5. Free receive slot after load completes (data is on compute GPU) if prealloc_slot_id is not None: diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/transport/buffer.py b/python/sglang/multimodal_gen/runtime/disaggregation/transport/buffer.py index e7a8fae8a..49c3faeb5 100644 --- a/python/sglang/multimodal_gen/runtime/disaggregation/transport/buffer.py +++ b/python/sglang/multimodal_gen/runtime/disaggregation/transport/buffer.py @@ -104,7 +104,7 @@ class TransferTensorBuffer: name: str, tensor: torch.Tensor, byte_offset: int = 0, - stream: torch.cuda.Stream | None = None, + stream: torch.Stream | None = None, ) -> int: """Copy a tensor into the pool slot. Returns bytes written.""" src_tensor = tensor.contiguous() @@ -122,7 +122,7 @@ class TransferTensorBuffer: src_bytes = src_tensor.view(torch.uint8).reshape(-1) if stream is not None: - with torch.cuda.stream(stream): + with torch.get_device_module().stream(stream): dst.copy_(src_bytes, non_blocking=True) else: dst.copy_(src_bytes, non_blocking=True) @@ -136,7 +136,7 @@ class TransferTensorBuffer: dtype: torch.dtype, byte_offset: int = 0, device: torch.device | str = "cpu", - stream: torch.cuda.Stream | None = None, + stream: torch.Stream | None = None, ) -> torch.Tensor: """Read a tensor from the pool slot. Returns a clone on target device.""" nbytes = 1 @@ -157,12 +157,12 @@ class TransferTensorBuffer: if same_device: # Clone to decouple tensor lifetime from pool slot if stream is not None: - with torch.cuda.stream(stream): + with torch.get_device_module().stream(stream): return src.clone() return src.clone() if stream is not None: - with torch.cuda.stream(stream): + with torch.get_device_module().stream(stream): return src.to(device, non_blocking=True) return src.to(device, non_blocking=True) @@ -170,7 +170,7 @@ class TransferTensorBuffer: self, handle: SlotHandle, tensors: dict[str, torch.Tensor | list[torch.Tensor] | None], - stream: torch.cuda.Stream | None = None, + stream: torch.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]] = {} @@ -178,7 +178,7 @@ class TransferTensorBuffer: # Ensure copy stream sees all prior compute kernels if stream is not None: - stream.wait_stream(torch.cuda.current_stream()) + stream.wait_stream(torch.get_device_module().current_stream()) for name, value in tensors.items(): if value is None: @@ -225,7 +225,7 @@ class TransferTensorBuffer: handle: SlotHandle, manifest: dict[str, list[dict]], device: torch.device | str = "cpu", - stream: torch.cuda.Stream | None = None, + stream: torch.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]] = {} diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/transport/codec.py b/python/sglang/multimodal_gen/runtime/disaggregation/transport/codec.py index 51664b6b4..14667348a 100644 --- a/python/sglang/multimodal_gen/runtime/disaggregation/transport/codec.py +++ b/python/sglang/multimodal_gen/runtime/disaggregation/transport/codec.py @@ -48,7 +48,7 @@ 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: + if tensor.is_cuda or tensor.is_npu: tensor = tensor.cpu() if not tensor.is_contiguous(): tensor = tensor.contiguous() diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/transport/manager.py b/python/sglang/multimodal_gen/runtime/disaggregation/transport/manager.py index 9647c0bb4..b9241710b 100644 --- a/python/sglang/multimodal_gen/runtime/disaggregation/transport/manager.py +++ b/python/sglang/multimodal_gen/runtime/disaggregation/transport/manager.py @@ -14,6 +14,7 @@ from sglang.multimodal_gen.runtime.disaggregation.transport.buffer import ( from sglang.multimodal_gen.runtime.disaggregation.transport.engine import ( BaseTransferEngine, ) +from sglang.multimodal_gen.runtime.platforms import current_platform logger = logging.getLogger(__name__) @@ -75,7 +76,7 @@ class DiffusionTransferManager: request_id: str, tensor_fields: dict[str, torch.Tensor | list[torch.Tensor] | None], scalar_fields: dict | None = None, - stream: torch.cuda.Stream | None = None, + stream: torch.Stream | None = None, ) -> StagedTransfer | None: """Stage GPU tensors into the local TransferBuffer. Returns None on allocation failure.""" total_size = 0 @@ -112,8 +113,8 @@ class DiffusionTransferManager: if stream is not None: stream.synchronize() - elif torch.cuda.is_available(): - torch.cuda.synchronize() + elif torch.get_device_module().is_available(): + torch.get_device_module().synchronize() staged = StagedTransfer( request_id=request_id, @@ -137,8 +138,8 @@ class DiffusionTransferManager: 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]: + stream: torch.Stream | None = None, + ) -> tuple[StagedTransfer | None, torch.Event | None]: """Stage GPU tensors, returning a CUDA event instead of blocking. Caller MUST wait on the event before reading buffer data. @@ -177,11 +178,11 @@ class DiffusionTransferManager: d2h_event = None if stream is not None: - d2h_event = torch.cuda.Event() + d2h_event = torch.get_device_module().Event() d2h_event.record(stream) - elif torch.cuda.is_available(): - d2h_event = torch.cuda.Event() - d2h_event.record(torch.cuda.current_stream()) + elif torch.get_device_module().is_available(): + d2h_event = torch.get_device_module().Event() + d2h_event.record(torch.get_device_module().current_stream()) staged = StagedTransfer( request_id=request_id, @@ -204,9 +205,12 @@ class DiffusionTransferManager: 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]: + device: torch.device | str = current_platform.device_type, + stream: torch.Stream | None = None, + ) -> tuple[ + dict[str, torch.Tensor | list[torch.Tensor]], + torch.get_device_module().Event | None, + ]: """Load tensors from receive slot to GPU, returning a CUDA event. Caller MUST wait on the event before using the returned tensors. @@ -225,11 +229,11 @@ class DiffusionTransferManager: load_event = None if stream is not None: - load_event = torch.cuda.Event() + load_event = torch.get_device_module().Event() load_event.record(stream) - elif torch.cuda.is_available(): - load_event = torch.cuda.Event() - load_event.record(torch.cuda.current_stream()) + elif torch.get_device_module().is_available(): + load_event = torch.get_device_module().Event() + load_event.record(torch.get_device_module().current_stream()) logger.debug( "TransferManager: loaded_async %d tensor fields for %s to %s", @@ -315,8 +319,8 @@ class DiffusionTransferManager: self, request_id: str, manifest: dict, - device: torch.device | str = "cuda", - stream: torch.cuda.Stream | None = None, + 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: @@ -333,8 +337,8 @@ class DiffusionTransferManager: if stream is not None: stream.synchronize() - elif torch.cuda.is_available(): - torch.cuda.synchronize() + elif torch.get_device_module().is_available(): + torch.get_device_module().synchronize() logger.debug( "TransferManager: loaded %d tensor fields for %s to %s",