From ecbadf0b4bd57c67b4c7eecfcf4b9ae6b6dde56b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=D0=90=D1=80=D1=82=D0=B5=D0=BC=20=D0=A1=D0=B0=D0=B2=D0=BA?= =?UTF-8?q?=D0=B8=D0=BD?= <58187114+OrangeRedeng@users.noreply.github.com> Date: Fri, 28 Aug 2026 15:39:05 +0300 Subject: [PATCH] [NPU] [Diffusion] support distributed inference pipeline for GLM-Image (#31320) Co-authored-by: Xiaoyu Zhang <1182563586@qq.com> --- docs/docs/sglang-diffusion/disaggregation.mdx | 115 +++++ .../runtime/disaggregation/orchestrator.py | 424 +++++++++++++++++- .../runtime/disaggregation/request_state.py | 6 +- .../runtime/disaggregation/scheduler_mixin.py | 150 ++++++- .../multimodal_gen/runtime/launch_server.py | 116 ++--- .../runtime/pipelines/glm_image.py | 51 ++- .../pipelines_core/composed_pipeline_base.py | 6 + .../runtime/pipelines_core/schedule_batch.py | 4 + .../stages/model_specific_stages/glm_image.py | 53 ++- .../sglang/multimodal_gen/test/run_suite.py | 1 + .../ascend/test_glm_image_distributed.py | 315 +++++++++++++ .../server/ascend/testcase_configs_npu.py | 17 +- .../test/unit/test_glm_image_ar.py | 18 +- .../test/unit/test_glm_image_multi_output.py | 8 +- 14 files changed, 1166 insertions(+), 118 deletions(-) create mode 100644 python/sglang/multimodal_gen/test/server/ascend/test_glm_image_distributed.py diff --git a/docs/docs/sglang-diffusion/disaggregation.mdx b/docs/docs/sglang-diffusion/disaggregation.mdx index ccc9b292c..b8d70d5a6 100644 --- a/docs/docs/sglang-diffusion/disaggregation.mdx +++ b/docs/docs/sglang-diffusion/disaggregation.mdx @@ -153,6 +153,118 @@ sglang serve --model-path ... --disagg-role server \ --decoder-urls "tcp://10.0.0.5:35000" ``` +### GLM-Image distributed mode + +GLM-Image can batch AR generation in the head and dispatch the resulting prior +tokens to distributed denoiser workers. Each worker runs prompt/glyph preparation, +DiT, and VAE decoding locally; no latent or embedding tensors are transferred. + +The following 16-device deployment uses devices 0-1 for the external AR server +and devices 2-15 for 14 independent batch-1 Cache-DiT denoisers. + +```bash +# Run 14 distributed denoisers +DISAGG_SERVER="tcp://127.0.0.1:19655" +MODEL_PATH="zai-org/GLM-Image" +BASE_MASTER_PORT=29005 + +export SGLANG_CACHE_DIT_FN=2 +export SGLANG_CACHE_DIT_BN=1 +export SGLANG_CACHE_DIT_WARMUP=4 +export SGLANG_CACHE_DIT_RDT=0.4 +export SGLANG_CACHE_DIT_MC=4 +export SGLANG_CACHE_DIT_TAYLORSEER=true +export SGLANG_CACHE_DIT_TS_ORDER=2 +export SGLANG_CACHE_DIT_ENABLED=true + +worker_pids=() + +cleanup() { + trap - EXIT + if ((${#worker_pids[@]})); then + kill "${worker_pids[@]}" 2>/dev/null || true + wait "${worker_pids[@]}" 2>/dev/null || true + fi +} + +trap cleanup EXIT +trap 'exit 130' INT +trap 'exit 143' TERM + +for i in $(seq 2 15); do + scheduler_port=$((19001 + i)) + master_port=$((BASE_MASTER_PORT + i)) + + sglang serve \ + --model-path "$MODEL_PATH" \ + --disagg-role denoiser \ + --disagg-server-addr "$DISAGG_SERVER" \ + --srt-encoder-url http://127.0.0.1:30020 \ + --scheduler-port "$scheduler_port" \ + --master-port "$master_port" \ + --num-gpus 1 \ + --base-gpu-id "$i" \ + --denoiser-sp 1 \ + --cfg-parallel-size 1 \ + --batching-max-size 1 \ + --dit-cpu-offload false \ + --attention-backend fa & + worker_pids+=("$!") +done + +# Stop all denoisers if any worker exits or fails during startup. +wait -n "${worker_pids[@]}" +``` + +Run the external AR server: + +```bash +sglang serve \ + --model-path ./zai-org/GLM-Image/vision_language_encoder/ \ + --tokenizer-path ./zai-org/GLM-Image/processor/ \ + --enable-multimodal \ + --cuda-graph-max-bs 28 \ + --device npu \ + --attention-backend ascend \ + --disable-fast-image-processor \ + --tp-size 2 \ + --host 0.0.0.0 \ + --port 30020 \ + --mem-fraction-static 0.8 +``` + +Run the public head. `--encoder-urls` and `--decoder-urls` are intentionally +omitted for this topology. + +```bash +sglang serve \ + --model-path zai-org/GLM-Image \ + --disagg-role server \ + --srt-encoder-url http://127.0.0.1:30020 \ + --srt-encoder-timeout 300 \ + --denoiser-urls "tcp://127.0.0.1:19003;tcp://127.0.0.1:19004;tcp://127.0.0.1:19005;tcp://127.0.0.1:19006;tcp://127.0.0.1:19007;tcp://127.0.0.1:19008;tcp://127.0.0.1:19009;tcp://127.0.0.1:19010;tcp://127.0.0.1:19011;tcp://127.0.0.1:19012;tcp://127.0.0.1:19013;tcp://127.0.0.1:19014;tcp://127.0.0.1:19015;tcp://127.0.0.1:19016" \ + --batching-mode dynamic \ + --batching-max-size 28 \ + --batching-delay-ms 30 \ + --enable-batching-metrics \ + --host 0.0.0.0 \ + --port 30052 \ + --scheduler-port 19655 \ + --output-path ./outputs +``` + +Workers return decoded pixels to the head, which saves and serves the final +files from `--output-path`. The PR benchmark used +[longtext-bench.zip](https://github.com/user-attachments/files/29779516/longtext-bench.zip): + +```bash +python fetch_images.py \ + --base-url http://localhost:30052/v1 \ + --model GLM-Image-distributed-test \ + --output-dir generated_images \ + --max-concurrency 56 +``` + ## Port Convention Result endpoints are derived deterministically from the head node's `--scheduler-port` (default: 5555): @@ -196,6 +308,9 @@ Tensor data between roles (encoder→denoiser, denoiser→decoder) is transferre **mooncake-transfer-engine** is required for disaggregated diffusion. It provides RDMA for direct GPU-to-GPU data movement. +The GLM-Image distributed mode is an exception: it relays only prior token IDs +and request metadata over ZMQ and does not require Mooncake. + ```bash pip install mooncake-transfer-engine ``` diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/orchestrator.py b/python/sglang/multimodal_gen/runtime/disaggregation/orchestrator.py index 05a158def..638a2fc64 100644 --- a/python/sglang/multimodal_gen/runtime/disaggregation/orchestrator.py +++ b/python/sglang/multimodal_gen/runtime/disaggregation/orchestrator.py @@ -7,9 +7,14 @@ import pickle import threading import time from collections import deque -from dataclasses import dataclass +from concurrent.futures import Future, ThreadPoolExecutor +from dataclasses import dataclass, field +from typing import TYPE_CHECKING +import torch import zmq +from transformers import AutoProcessor +from zmq.utils.monitor import recv_monitor_message from sglang.multimodal_gen.runtime.disaggregation.dispatch_policy import ( PoolDispatcher, @@ -20,6 +25,7 @@ from sglang.multimodal_gen.runtime.disaggregation.request_state import ( ) from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType from sglang.multimodal_gen.runtime.disaggregation.transport.codec import ( + send_tensors, unpack_tensors, ) from sglang.multimodal_gen.runtime.disaggregation.transport.protocol import ( @@ -31,11 +37,70 @@ from sglang.multimodal_gen.runtime.disaggregation.transport.protocol import ( encode_transfer_msg, is_transfer_message, ) +from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import ( + OutputBatch, + Req, +) from sglang.multimodal_gen.runtime.utils.common import get_zmq_socket +from sglang.multimodal_gen.runtime.utils.perf_logger import ( + MemorySnapshot, + RequestMetrics, +) + +if TYPE_CHECKING: + from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.glm_image import ( + GlmImageAR, + ) + from sglang.multimodal_gen.runtime.server_args import ServerArgs logger = logging.getLogger(__name__) +def _deserialize_request_metrics(data: dict | None) -> RequestMetrics | None: + if data is None: + return None + + metrics = RequestMetrics(request_id=data["request_id"]) + metrics.stages = data.get("stages", {}) + metrics.steps = data.get("steps", []) + metrics.total_duration_ms = data.get("total_duration_ms", 0.0) + for name, snapshot in data.get("memory_snapshots", {}).items(): + metrics.memory_snapshots[name] = MemorySnapshot( + allocated_mb=snapshot.get("allocated_mb", 0.0), + reserved_mb=snapshot.get("reserved_mb", 0.0), + peak_allocated_mb=snapshot.get("peak_allocated_mb", 0.0), + peak_reserved_mb=snapshot.get("peak_reserved_mb", 0.0), + peak_host_anon_mb=snapshot.get("peak_host_anon_mb", 0.0), + ) + return metrics + + +@dataclass +class _GlmDistributedRequest: + """Track one client request as it moves from AR to a denoiser.""" + + client_request_id: str + req: Req + enqueue_time: float + worker_idx: int | None = None + + +@dataclass +class _GlmDistributedModeState: + """State used only by the GLM external-AR distributed topology.""" + + server_args: "ServerArgs" + ar_stage: "GlmImageAR" + executor: ThreadPoolExecutor + denoiser_worker_available: list[bool] + pending_ar_requests: deque[_GlmDistributedRequest] = field(default_factory=deque) + pending_denoiser_requests: deque[_GlmDistributedRequest] = field( + default_factory=deque + ) + denoiser_requests: dict[str, _GlmDistributedRequest] = field(default_factory=dict) + active_ar_batch: tuple[Future, list[_GlmDistributedRequest]] | None = None + + @dataclass class _EncoderTTAEntry: request_id: str @@ -89,9 +154,11 @@ class DiffusionServer: dispatch_policy_name: str = "round_robin", timeout_s: float = 600.0, encoder_capacity: int = 4, - denoiser_capacity: int = 2, + denoiser_capacity_per_worker: int = 2, decoder_capacity: int = 4, p2p_mode: bool = True, + server_args=None, + glm_distributed_mode_enabled: bool = False, ): self._frontend_endpoint = frontend_endpoint self._encoder_work_endpoints = encoder_work_endpoints @@ -108,9 +175,9 @@ class DiffusionServer: self._tracker = RequestTracker() self._dispatcher = PoolDispatcher( - num_encoders=self._num_encoders, + num_encoders=max(1, self._num_encoders), num_denoisers=self._num_denoisers, - num_decoders=self._num_decoders, + num_decoders=max(1, self._num_decoders), policy_name=dispatch_policy_name, ) @@ -124,7 +191,7 @@ class DiffusionServer: # FreeBufferSlots per instance self._encoder_free_slots = [encoder_capacity] * self._num_encoders - self._denoiser_free_slots = [denoiser_capacity] * self._num_denoisers + self._denoiser_free_slots = [denoiser_capacity_per_worker] * self._num_denoisers self._decoder_free_slots = [decoder_capacity] * self._num_decoders # TTA queues per role type @@ -133,6 +200,24 @@ class DiffusionServer: self._decoder_tta: deque[_RoleTTAEntry] = deque() self._transfer_mode = p2p_mode + self._glm_distributed_state: _GlmDistributedModeState | None = None + + if glm_distributed_mode_enabled: + from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.glm_image import ( + GlmImageAR, + ) + + processor = AutoProcessor.from_pretrained( + server_args.model_path, subfolder="processor" + ) + self._glm_distributed_state = _GlmDistributedModeState( + server_args=server_args, + ar_stage=GlmImageAR(processor=processor, vision_language_encoder=None), + executor=ThreadPoolExecutor( + max_workers=1, thread_name_prefix="glm-distributed-ar" + ), + denoiser_worker_available=[False] * self._num_denoisers, + ) self._transfer_state: dict[str, _TransferRequestState] = {} # Per-instance registration: instance_idx -> {session_id, pool_ptr, pool_size} @@ -197,6 +282,10 @@ class DiffusionServer: if self._thread is not None: self._thread.join(timeout=5.0) self._thread = None + if self._glm_distributed_state is not None: + self._glm_distributed_state.executor.shutdown( + wait=False, cancel_futures=True + ) def _event_loop(self) -> None: frontend, _ = get_zmq_socket( @@ -209,9 +298,16 @@ class DiffusionServer: encoder_pushes.append(sock) denoiser_pushes: list[zmq.Socket] = [] + denoiser_monitors: list[zmq.Socket] = [] for i, ep in enumerate(self._denoiser_work_endpoints): sock, _ = get_zmq_socket(self._context, zmq.PUSH, ep, bind=False) denoiser_pushes.append(sock) + if self._glm_distributed_state is not None: + denoiser_monitors.append( + sock.get_monitor_socket( + events=zmq.EVENT_CONNECTED | zmq.EVENT_DISCONNECTED + ) + ) decoder_pushes: list[zmq.Socket] = [] for i, ep in enumerate(self._decoder_work_endpoints): @@ -233,6 +329,8 @@ class DiffusionServer: poller.register(encoder_result_pull, zmq.POLLIN) poller.register(denoiser_result_pull, zmq.POLLIN) poller.register(decoder_result_pull, zmq.POLLIN) + for monitor in denoiser_monitors: + poller.register(monitor, zmq.POLLIN) self._encoder_pushes = encoder_pushes self._denoiser_pushes = denoiser_pushes @@ -246,6 +344,7 @@ class DiffusionServer: + encoder_pushes + denoiser_pushes + decoder_pushes + + denoiser_monitors ) try: @@ -266,7 +365,16 @@ class DiffusionServer: if decoder_result_pull in events: self._handle_role_result(decoder_result_pull, RoleType.DECODER) - self._drain_all_queues() + for worker_idx, monitor in enumerate(denoiser_monitors): + if monitor in events: + self._handle_glm_denoiser_monitor_event(worker_idx, monitor) + + if self._glm_distributed_state is not None: + self._process_glm_ar_batch_result_if_ready() + self._dispatch_glm_ar_batch_if_ready() + self._dispatch_glm_denoiser_requests_if_ready() + else: + self._drain_all_queues() except Exception: logger.exception("DiffusionServer event loop error") @@ -285,7 +393,9 @@ class DiffusionServer: self._handle_transfer_result(frames, role) return - if role == RoleType.DECODER: + if self._glm_distributed_state is not None and role == RoleType.DENOISER: + self._handle_glm_denoiser_result_frames(frames) + elif role == RoleType.DECODER: self._handle_decoder_result_frames(frames) else: # Non-transfer frames from encoder/denoiser are error results @@ -378,6 +488,22 @@ class DiffusionServer: self._tracker.transition(request_id, RequestState.ENCODER_WAITING) except ValueError: pass + if self._glm_distributed_state is not None: + if ( + not isinstance(req.prompt, str) + or getattr(req, "image_path", None) is not None + ): + self._complete_with_error( + request_id, + "GLM distributed mode supports one text prompt without image input", + ) + return + now = time.monotonic() + self._glm_distributed_state.pending_ar_requests.append( + _GlmDistributedRequest(request_id, req, now) + ) + return + self._encoder_tta.append( _EncoderTTAEntry( request_id=request_id, @@ -390,11 +516,245 @@ class DiffusionServer: request_id, ) - def _handle_decoder_result_frames(self, frames: list) -> None: - from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import ( - OutputBatch, + def _dispatch_glm_ar_batch_if_ready(self) -> None: + """Dispatch one compatible batch to the external AR server.""" + state = self._glm_distributed_state + assert state is not None + if state.active_ar_batch is not None or not state.pending_ar_requests: + return + + batch_max_size = ( + max(1, state.server_args.batching_max_size) + if state.server_args.batching_mode == "dynamic" + else 1 + ) + base = state.pending_ar_requests[0] + indices = [0] + output_slots = max(1, int(base.req.num_outputs_per_prompt or 1)) + for index in range(1, len(state.pending_ar_requests)): + if output_slots >= batch_max_size: + break + candidate = state.pending_ar_requests[index] + if (base.req.height, base.req.width) == ( + candidate.req.height, + candidate.req.width, + ): + candidate_outputs = max( + 1, int(candidate.req.num_outputs_per_prompt or 1) + ) + if output_slots + candidate_outputs > batch_max_size: + continue + indices.append(index) + output_slots += candidate_outputs + + waited = time.monotonic() - base.enqueue_time + batch_delay_s = state.server_args.batching_delay_ms / 1000.0 + if output_slots < batch_max_size and waited < batch_delay_s: + return + + requests = [state.pending_ar_requests[index] for index in indices] + for index in reversed(indices): + del state.pending_ar_requests[index] + + for request in requests: + try: + self._tracker.transition( + request.client_request_id, RequestState.ENCODER_RUNNING + ) + except ValueError: + pass + state.active_ar_batch = ( + state.executor.submit( + state.ar_stage.generate_and_assign_prior_tokens, + [request.req for request in requests], + state.server_args, + device=torch.device("cpu"), + ), + requests, + ) + logger.info( + "GLM distributed AR dispatched batch size=%d requests, %d outputs", + len(requests), + output_slots, ) + def _process_glm_ar_batch_result_if_ready(self) -> None: + """Move a completed AR batch into the denoiser dispatch queue.""" + state = self._glm_distributed_state + assert state is not None + active_batch = state.active_ar_batch + if active_batch is None or not active_batch[0].done(): + return + future, requests = active_batch + state.active_ar_batch = None + try: + future.result() + except Exception as error: + for request in requests: + self._complete_with_error( + request.client_request_id, f"GLM AR error: {error}" + ) + return + + group_id = f"glm-distributed::{time.monotonic_ns()}" + for request_index, request in enumerate(requests): + if self._tracker.get(request.client_request_id) is None: + continue + try: + self._tracker.transition( + request.client_request_id, RequestState.ENCODER_DONE + ) + self._tracker.transition( + request.client_request_id, RequestState.DENOISING_WAITING + ) + except ValueError: + pass + denoiser_req = request.req + denoiser_req.request_id = f"{group_id}::request::{request_index}" + state.denoiser_requests[denoiser_req.request_id] = request + state.pending_denoiser_requests.append(request) + + def _dispatch_glm_denoiser_requests_if_ready(self) -> None: + """Dispatch AR-complete requests to available denoisers. + + Each GLM denoiser accepts one request at a time. Dispatch continues until + either the pending queue is empty or every connected worker is busy. + """ + from sglang.multimodal_gen.runtime.disaggregation.scheduler_mixin import ( + extract_transfer_fields, + ) + + state = self._glm_distributed_state + assert state is not None + while state.pending_denoiser_requests: + available_slots = [ + slots if state.denoiser_worker_available[index] else 0 + for index, slots in enumerate(self._denoiser_free_slots) + ] + worker_idx = self._dispatcher.select_denoiser_with_capacity(available_slots) + if worker_idx is None: + return + request = state.pending_denoiser_requests.popleft() + self._denoiser_free_slots[worker_idx] -= 1 + request.worker_idx = worker_idx + tensor_fields, scalar_fields = extract_transfer_fields(request.req) + scalar_fields["request_id"] = request.req.request_id + send_tensors( + self._denoiser_pushes[worker_idx], tensor_fields, scalar_fields + ) + try: + self._tracker.transition( + request.client_request_id, + RequestState.DENOISING_RUNNING, + denoiser_instance=worker_idx, + ) + except ValueError: + pass + logger.info( + "GLM distributed mode dispatched request with %d output(s) " + "to denoiser[%d]", + request.req.num_outputs_per_prompt, + worker_idx, + ) + + def _handle_glm_denoiser_monitor_event( + self, worker_idx: int, monitor: zmq.Socket + ) -> None: + """Update dispatch eligibility when a denoiser connects or disconnects.""" + state = self._glm_distributed_state + assert state is not None + event = recv_monitor_message(monitor, flags=zmq.NOBLOCK)["event"] + if event == zmq.EVENT_DISCONNECTED: + state.denoiser_worker_available[worker_idx] = False + self._denoiser_free_slots[worker_idx] = 0 + replayed_requests = [ + request + for request in state.denoiser_requests.values() + if request.worker_idx == worker_idx + ] + for request in replayed_requests: + request.worker_idx = None + try: + self._tracker.transition( + request.client_request_id, RequestState.DENOISING_WAITING + ) + except ValueError: + pass + state.pending_denoiser_requests.extendleft(reversed(replayed_requests)) + logger.warning( + "GLM denoiser[%d] disconnected; requeued %d request(s)", + worker_idx, + len(replayed_requests), + ) + elif event == zmq.EVENT_CONNECTED: + registered = worker_idx in self._denoiser_peers + state.denoiser_worker_available[worker_idx] = True + if not any( + request.worker_idx == worker_idx + for request in state.denoiser_requests.values() + ): + self._denoiser_free_slots[worker_idx] = 1 + logger.info( + "GLM denoiser[%d] connected (registered=%s)", + worker_idx, + registered, + ) + + def _handle_glm_denoiser_result_frames(self, frames: list) -> None: + """Return decoded denoiser output to the originating HTTP request.""" + state = self._glm_distributed_state + assert state is not None + tensor_fields, scalar_fields = unpack_tensors(frames, device="cpu") + denoiser_request_id = scalar_fields.get("request_id") + request = state.denoiser_requests.pop(denoiser_request_id, None) + if request is None: + logger.warning( + "Unknown GLM distributed denoiser result: %s", denoiser_request_id + ) + return + if request.worker_idx is not None: + self._denoiser_free_slots[request.worker_idx] = 1 + + error = scalar_fields.get("error") + output = tensor_fields.get("output") + total = max(1, int(request.req.num_outputs_per_prompt or 1)) + output_size = len(output) if output is not None else None + if output_size is not None and output_size != total: + error = ( + f"GLM distributed output size mismatch: got {output_size}, " + f"expected {total}" + ) + output = None + result = OutputBatch( + output=output, + error=error, + metrics=_deserialize_request_metrics(scalar_fields.get("metrics")), + metrics_list=[ + _deserialize_request_metrics(metrics) + for metrics in scalar_fields.get("metrics_list", []) + ] + or None, + peak_memory_mb=scalar_fields.get("peak_memory_mb", 0.0), + usage=scalar_fields.get("usage"), + ) + with self._lock: + identity = self._pending.pop(request.client_request_id, None) + if identity is not None: + self._frontend.send_multipart([identity, b"", pickle.dumps(result)]) + try: + self._tracker.transition( + request.client_request_id, RequestState.DENOISING_DONE + ) + self._tracker.transition( + request.client_request_id, + RequestState.FAILED if error else RequestState.DONE, + error=error, + ) + except ValueError: + pass + self._tracker.remove(request.client_request_id) + + def _handle_decoder_result_frames(self, frames: list) -> None: request_id = self._extract_request_id(frames) if request_id is None: logger.warning("DiffusionServer: decoder result missing request_id") @@ -521,10 +881,6 @@ class DiffusionServer: return None def _complete_with_error(self, request_id: str, error_msg: str) -> None: - from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import ( - OutputBatch, - ) - logger.error("DiffusionServer: %s — %s", request_id, error_msg) try: @@ -578,6 +934,25 @@ class DiffusionServer: self._decoder_tta = deque( e for e in self._decoder_tta if e.request_id not in timed_set ) + if self._glm_distributed_state is not None: + state = self._glm_distributed_state + state.pending_ar_requests = deque( + request + for request in state.pending_ar_requests + if request.client_request_id not in timed_set + ) + state.pending_denoiser_requests = deque( + request + for request in state.pending_denoiser_requests + if request.client_request_id not in timed_set + ) + for denoiser_request_id, request in list( + state.denoiser_requests.items() + ): + if request.client_request_id in timed_set: + state.denoiser_requests.pop(denoiser_request_id) + if request.worker_idx is not None: + self._denoiser_free_slots[request.worker_idx] = 1 def _free_slot_for_record(self, record) -> None: if ( @@ -667,6 +1042,8 @@ class DiffusionServer: prealloc = msg.get("preallocated_slots", []) info["free_preallocated_slots"] = list(prealloc) peers[idx] = info + if role == RoleType.DENOISER and self._glm_distributed_state is not None: + self._glm_distributed_state.denoiser_worker_available[idx] = True logger.info( "DiffusionServer transfer: registered %s[%d] work_endpoint=%s " @@ -1026,10 +1403,6 @@ class DiffusionServer: ) def _transfer_return_to_client_from_msg(self, request_id: str, msg: dict) -> None: - from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import ( - OutputBatch, - ) - with self._lock: client_identity = self._pending.pop(request_id, None) @@ -1054,7 +1427,7 @@ class DiffusionServer: def get_stats(self) -> dict: with self._lock: pending_count = len(self._pending) - return { + stats = { "role": "diffusion_server", "transfer_mode": self._transfer_mode, "num_encoders": self._num_encoders, @@ -1074,3 +1447,16 @@ class DiffusionServer: "decoder_peers": len(self._decoder_peers), "tracker": self._tracker.snapshot(), } + if self._glm_distributed_state is not None: + state = self._glm_distributed_state + stats.update( + { + "glm_ar_queue_depth": len(state.pending_ar_requests), + "glm_ar_in_flight": state.active_ar_batch is not None, + "glm_denoiser_queue_depth": len(state.pending_denoiser_requests), + "glm_denoiser_worker_available": list( + state.denoiser_worker_available + ), + } + ) + return stats diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/request_state.py b/python/sglang/multimodal_gen/runtime/disaggregation/request_state.py index 7c80906bc..f09f61a04 100644 --- a/python/sglang/multimodal_gen/runtime/disaggregation/request_state.py +++ b/python/sglang/multimodal_gen/runtime/disaggregation/request_state.py @@ -45,10 +45,14 @@ _VALID_TRANSITIONS: dict[RequestState, set[RequestState]] = { RequestState.DENOISING_RUNNING, }, RequestState.DENOISING_WAITING: {RequestState.DENOISING_RUNNING}, - RequestState.DENOISING_RUNNING: {RequestState.DENOISING_DONE}, + RequestState.DENOISING_RUNNING: { + RequestState.DENOISING_WAITING, + RequestState.DENOISING_DONE, + }, RequestState.DENOISING_DONE: { RequestState.DECODER_WAITING, RequestState.DECODER_RUNNING, + RequestState.DONE, }, RequestState.DECODER_WAITING: {RequestState.DECODER_RUNNING}, RequestState.DECODER_RUNNING: {RequestState.DONE}, diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py b/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py index 27c9300e7..3f868280f 100644 --- a/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py +++ b/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py @@ -30,6 +30,7 @@ from sglang.multimodal_gen.runtime.disaggregation.transport.buffer import ( ) from sglang.multimodal_gen.runtime.disaggregation.transport.codec import ( send_tensors, + unpack_tensors, ) from sglang.multimodal_gen.runtime.disaggregation.transport.engine import ( create_transfer_engine, @@ -49,6 +50,7 @@ from sglang.multimodal_gen.runtime.disaggregation.transport.protocol import ( encode_transfer_msg, is_transfer_message, ) +from sglang.multimodal_gen.runtime.entrypoints.utils import expand_request_outputs from sglang.multimodal_gen.runtime.pipelines_core import Req from sglang.multimodal_gen.runtime.pipelines_core.diffusion_scheduler_utils import ( clone_scheduler_runtime, @@ -65,6 +67,44 @@ if TYPE_CHECKING: logger = init_logger(__name__) + +def _advertised_pool_work_endpoint(server_args) -> str: + host = server_args.disagg_p2p_hostname or server_args.host or "127.0.0.1" + if host == "0.0.0.0": + host = server_args.disagg_p2p_hostname or "127.0.0.1" + return server_args.pool_work_endpoint.replace("0.0.0.0", host) + + +def _expand_glm_distributed_outputs(req: Req) -> list[Req]: + """Split external-AR tokens into requests for sequential denoising.""" + output_count = max(1, int(req.num_outputs_per_prompt or 1)) + if output_count == 1: + return [req] + + prior_token_ids = req.prior_token_id + if not isinstance(prior_token_ids, torch.Tensor) or ( + prior_token_ids.shape[0] != output_count + ): + actual_count = ( + prior_token_ids.shape[0] + if isinstance(prior_token_ids, torch.Tensor) + else type(prior_token_ids).__name__ + ) + raise RuntimeError( + "Cannot split GLM-Image AR output for distributed inference: " + f"expected {output_count} token rows, got {actual_count}." + ) + + usage_by_output = req.extra.get("usage_by_output") + output_reqs = expand_request_outputs(req) + for output_index, output_req in enumerate(output_reqs): + output_req.prior_token_id = prior_token_ids[output_index : output_index + 1] + output_req.extra.pop("usage_by_output", None) + if usage_by_output is not None and output_index < len(usage_by_output): + output_req.usage = usage_by_output[output_index] + return output_reqs + + # --------------------------------------------------------------------------- # Field extraction: split Req into tensors (transfer buffer) and scalars (JSON) # --------------------------------------------------------------------------- @@ -116,6 +156,9 @@ _SAMPLING_PARAMS_EXCLUDE_FIELDS = frozenset( } ) +# Receivers reconstruct base SamplingParams, so only base defaults can be omitted. +_BASE_SAMPLING_PARAM_FIELDS = {f.name: f for f in dataclasses.fields(SamplingParams)} + def _is_tensor_like(value) -> bool: if isinstance(value, torch.Tensor): @@ -278,7 +321,8 @@ def extract_transfer_fields(req) -> tuple[dict, dict]: value = getattr(sp, name, None) if value is None: continue - if _is_default(value, f): + base_field = _BASE_SAMPLING_PARAM_FIELDS.get(name) + if base_field is not None and _is_default(value, base_field): continue try: scalar_fields[name] = _to_json_serializable(value) @@ -458,6 +502,23 @@ class SchedulerDisaggMixin: sa = self.server_args + if self._is_glm_distributed_mode(): + self._preallocated_slots = {} + register_msg = TransferRegisterMsg( + role=self._disagg_role.value, + work_endpoint=_advertised_pool_work_endpoint(sa), + ) + self._pool_result_push.send_multipart(encode_transfer_msg(register_msg)) + self._compute_ready_queue = queue.Queue(maxsize=4) + self._recv_prefetch_thread = threading.Thread( + target=self._recv_prefetch_loop, + daemon=True, + name="recv-prefetch-glm-distributed-denoiser", + ) + self._recv_prefetch_thread.start() + logger.info("GLM distributed denoiser registered") + return + # Pool size: configurable, default 256 MiB pool_size = getattr(sa, "disagg_transfer_pool_size", 256 * 1024 * 1024) @@ -522,7 +583,7 @@ class SchedulerDisaggMixin: session_id=self._transfer_manager.session_id, pool_ptr=self._transfer_manager.pool_data_ptr, pool_size=self._transfer_manager.pool_size, - work_endpoint=sa.pool_work_endpoint, + work_endpoint=_advertised_pool_work_endpoint(sa), preallocated_slots=preallocated_slot_info, ) self._pool_result_push.send_multipart(encode_transfer_msg(register_msg)) @@ -623,6 +684,10 @@ class SchedulerDisaggMixin: raw_frames = self._pool_work_pull.recv_multipart() frames = [bytes(f) for f in raw_frames] + if not is_transfer_message(frames): + self._compute_ready_queue.put(("relay_compute", frames)) + continue + msg = decode_transfer_msg(frames) msg_type = msg.get("msg_type", "") @@ -903,6 +968,18 @@ class SchedulerDisaggMixin: self._broadcast_to_all_ranks(("skip",)) self._handle_transfer_msg(data) + elif msg_type == "relay_compute": + local_device = ( + f"{current_platform.device_type}:{self.worker.local_rank}" + ) + tensors, scalar_fields = unpack_tensors(data, device=local_device) + request_id = scalar_fields.get("request_id", "unknown") + req = self._build_disagg_req(scalar_fields, tensors) + if is_multi_rank: + self._broadcast_to_all_ranks(("compute",)) + self._broadcast_req_to_all_ranks(req) + self._execute_glm_distributed_denoiser_request(req, request_id) + self._consecutive_error_count = 0 except Exception as e: @@ -1298,11 +1375,16 @@ class SchedulerDisaggMixin: (:meth:`_disagg_non_rank0_event_loop`). """ if self._disagg_role == RoleType.DENOISER: - # Initialize scheduler timesteps (same as rank 0) - _init_disagg_request_scheduler(self, req) + if not self._is_glm_distributed_mode(): + _init_disagg_request_scheduler(self, req) with self._disagg_trace_dispatch(req): - self.worker.execute_forward([req], return_req=True) + if self._is_glm_distributed_mode(): + req.save_output = False + req.return_file_paths_only = False + self.worker.execute_forward([req]) + else: + self.worker.execute_forward([req], return_req=True) elif self._disagg_role == RoleType.DECODER: req.save_output = False @@ -1456,6 +1538,64 @@ class SchedulerDisaggMixin: duration_s, ) + def _is_glm_distributed_mode(self: Scheduler) -> bool: + """Return whether this scheduler is a GLM distributed denoiser.""" + return ( + self._disagg_role == RoleType.DENOISER + and self.worker.pipeline.pipeline_name == "GlmImagePipeline" + and self.server_args.srt_encoder_url is not None + ) + + def _execute_glm_distributed_denoiser_request( + self: Scheduler, req: Req, request_id: str + ) -> None: + """Run local preparation, DiT, and VAE, then return decoded pixels.""" + req.save_output = False + start_time = time.monotonic() + with self._disagg_trace_dispatch(req): + if ( + self.server_args.pipeline_config.supports_sequential_multi_output_inference() + and max(1, int(req.num_outputs_per_prompt or 1)) > 1 + ): + output_reqs = _expand_glm_distributed_outputs(req) + output_batches = list( + self.worker.execute_forward_sequentially(output_reqs) + ) + output_batch = self.worker._merge_expanded_output_batches( + output_batches + ) + else: + output_batch = self.worker.execute_forward([req]) + + tensor_fields = {} + scalar_fields = {"request_id": request_id} + if output_batch.output is not None: + tensor_fields["output"] = output_batch.output + if output_batch.error is not None: + scalar_fields["error"] = output_batch.error + if output_batch.usage is not None: + scalar_fields["usage"] = output_batch.usage + if output_batch.metrics is not None: + scalar_fields["metrics"] = output_batch.metrics.to_dict() + if output_batch.metrics_list is not None: + scalar_fields["metrics_list"] = [ + metrics.to_dict() if metrics is not None else None + for metrics in output_batch.metrics_list + ] + scalar_fields["peak_memory_mb"] = output_batch.peak_memory_mb + send_tensors(self._pool_result_push, tensor_fields, scalar_fields) + + if self._disagg_metrics: + if output_batch.error: + self._disagg_metrics.record_request_failed(request_id) + else: + self._disagg_metrics.record_request_complete(request_id) + logger.debug( + "GLM distributed denoiser: processed %s in %.2f s", + request_id, + time.monotonic() - start_time, + ) + def _disagg_decoder_compute(self: Scheduler, req: Req, request_id: str) -> None: """Run decoder compute in transfer mode, send result to DS. diff --git a/python/sglang/multimodal_gen/runtime/launch_server.py b/python/sglang/multimodal_gen/runtime/launch_server.py index af3d09c4c..df6fed265 100644 --- a/python/sglang/multimodal_gen/runtime/launch_server.py +++ b/python/sglang/multimodal_gen/runtime/launch_server.py @@ -3,12 +3,9 @@ import dataclasses import multiprocessing as mp import os -import signal import sys -import threading import time -import psutil import uvicorn from sglang.multimodal_gen.runtime.disaggregation.orchestrator import ( @@ -24,7 +21,10 @@ from sglang.multimodal_gen.runtime.server_args import ( prepare_server_args, set_global_server_args, ) -from sglang.multimodal_gen.runtime.utils.common import is_port_available +from sglang.multimodal_gen.runtime.utils.common import ( + is_port_available, + kill_process_tree, +) from sglang.multimodal_gen.runtime.utils.logging_utils import configure_logger, logger from sglang.multimodal_gen.runtime.utils.trace_wrapper import init_diffusion_tracing from sglang.multimodal_gen.utils import kill_itself_when_parent_died @@ -53,45 +53,6 @@ def _find_available_port( ) -def kill_process_tree(parent_pid, include_parent: bool = True, skip_pid: int = None): - """Kill the process and all its child processes.""" - # Remove sigchld handler to avoid spammy logs. - if threading.current_thread() is threading.main_thread(): - signal.signal(signal.SIGCHLD, signal.SIG_DFL) - - if parent_pid is None: - parent_pid = os.getpid() - include_parent = False - - try: - itself = psutil.Process(parent_pid) - except psutil.NoSuchProcess: - return - - children = itself.children(recursive=True) - for child in children: - if child.pid == skip_pid: - continue - try: - child.kill() - except psutil.NoSuchProcess: - pass - - if include_parent: - try: - if parent_pid == os.getpid(): - itself.kill() - sys.exit(0) - - itself.kill() - - # Sometime processes cannot be killed with SIGKILL (e.g, PID=1 launched by kubernetes), - # so we send an additional signal to kill them. - itself.send_signal(signal.SIGQUIT) - except psutil.NoSuchProcess: - pass - - def _process_names(processes) -> str: return ", ".join(getattr(p, "name", repr(p)) for p in processes) @@ -585,12 +546,12 @@ def launch_http_server_only(server_args): ) -def parse_url_string(url_str: str) -> list[str]: +def parse_url_string(url_str: str | None) -> 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()] + return [u.strip() for u in (url_str or "").split(";") if u.strip()] def launch_disagg_server(server_args: ServerArgs): @@ -605,12 +566,23 @@ def launch_disagg_server(server_args: ServerArgs): decoder result: scheduler_port + 3 """ configure_logger(server_args) + set_global_server_args(server_args) - for name, val in [ - ("--encoder-urls", server_args.encoder_urls), - ("--denoiser-urls", server_args.denoiser_urls), - ("--decoder-urls", server_args.decoder_urls), - ]: + glm_distributed_mode_enabled = ( + type(server_args.pipeline_config).__name__ == "GlmImagePipelineConfig" + and server_args.srt_encoder_url is not None + and server_args.encoder_urls is None + and server_args.decoder_urls is None + ) + required_urls = [("--denoiser-urls", server_args.denoiser_urls)] + if not glm_distributed_mode_enabled: + required_urls.extend( + [ + ("--encoder-urls", server_args.encoder_urls), + ("--decoder-urls", server_args.decoder_urls), + ] + ) + for name, val in required_urls: if val is None: raise ValueError(f"{name} is required for --disagg-role server") @@ -644,6 +616,9 @@ def launch_disagg_server(server_args: ServerArgs): decoder_result_ep, ) + denoiser_options = ( + {"denoiser_capacity_per_worker": 1} if glm_distributed_mode_enabled else {} + ) diffusion_server = DiffusionServer( frontend_endpoint=frontend_endpoint, encoder_work_endpoints=encoder_work_endpoints, @@ -654,6 +629,9 @@ def launch_disagg_server(server_args: ServerArgs): decoder_result_endpoint=decoder_result_ep, dispatch_policy_name=server_args.disagg_dispatch_policy, timeout_s=float(server_args.disagg_timeout), + server_args=server_args, + glm_distributed_mode_enabled=glm_distributed_mode_enabled, + **denoiser_options, ) diffusion_server.start() @@ -726,6 +704,44 @@ def launch_disagg_role(server_args: ServerArgs): "ulysses_degree": role_par["ulysses_degree"], "ring_degree": role_par["ring_degree"], } + role_tp = role_par["tp_size"] or 1 + role_sp = role_par["sp_degree"] or 1 + cfg_degree = ( + server_args.cfg_parallel_degree if server_args.enable_cfg_parallel else 1 + ) + cfg_parallel_explicit = server_args.is_arg_explicitly_set( + "enable_cfg_parallel" + ) or server_args.is_arg_explicitly_set("cfg_parallel_degree") + required_devices = role_tp * role_sp * cfg_degree * server_args.dp_size + if not cfg_parallel_explicit and ( + required_devices > server_args.num_gpus + or server_args.num_gpus % required_devices != 0 + ): + logger.warning( + "Disabling auto-enabled CFG parallel for %s role because tp=%d, " + "sp=%d, cfg=%d, dp=%d is incompatible with %d devices", + role_type.value, + role_tp, + role_sp, + cfg_degree, + server_args.dp_size, + server_args.num_gpus, + ) + role_overrides["enable_cfg_parallel"] = False + role_overrides["cfg_parallel_degree"] = 1 + cfg_degree = 1 + required_devices = role_tp * role_sp * server_args.dp_size + + if ( + required_devices > server_args.num_gpus + or server_args.num_gpus % required_devices != 0 + ): + raise ValueError( + f"Invalid parallelism for {role_type.value} role: " + f"tp={role_tp}, sp={role_sp}, cfg={cfg_degree}, " + f"dp={server_args.dp_size} requires groups of {required_devices} " + f"devices, but num_gpus={server_args.num_gpus}" + ) base_dict = { f.name: getattr(server_args, f.name) diff --git a/python/sglang/multimodal_gen/runtime/pipelines/glm_image.py b/python/sglang/multimodal_gen/runtime/pipelines/glm_image.py index 4b72b5a84..57e75959d 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines/glm_image.py +++ b/python/sglang/multimodal_gen/runtime/pipelines/glm_image.py @@ -12,6 +12,22 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.g from sglang.multimodal_gen.runtime.server_args import ServerArgs +class GlmImageDenoiserDecodingStage(GlmImageDecodingStage): + """Run VAE decoding on the denoiser because this topology has no decoder worker.""" + + @property + def role_affinity(self) -> RoleType: + return RoleType.DENOISER + + +class GlmImageDenoiserPreparationStage(GlmImageBeforeDenoisingStage): + """Run DiT preparation on the denoiser because AR runs in the server head.""" + + @property + def role_affinity(self) -> RoleType: + return RoleType.DENOISER + + class GlmImagePipeline(LoRAPipeline, ComposedPipelineBase): pipeline_name = "GlmImagePipeline" @@ -26,6 +42,10 @@ class GlmImagePipeline(LoRAPipeline, ComposedPipelineBase): ] def create_pipeline_stages(self, server_args: ServerArgs): + is_glm_distributed_mode = ( + self._disagg_role == RoleType.DENOISER + and server_args.srt_encoder_url is not None + ) self.add_stage( GlmImageAR( processor=self.get_module("processor"), @@ -34,8 +54,13 @@ class GlmImagePipeline(LoRAPipeline, ComposedPipelineBase): "glm_image_ar", ) + before_denoising_stage_cls = ( + GlmImageDenoiserPreparationStage + if is_glm_distributed_mode + else GlmImageBeforeDenoisingStage + ) self.add_stage( - GlmImageBeforeDenoisingStage( + before_denoising_stage_cls( vae=self.get_module("vae"), text_encoder=self.get_module("text_encoder"), tokenizer=self.get_module("tokenizer"), @@ -52,14 +77,22 @@ class GlmImagePipeline(LoRAPipeline, ComposedPipelineBase): ), ) - self.add_stage_factory( - RoleType.DECODER, - lambda: GlmImageDecodingStage( - vae=self.get_module("vae"), - pipeline=self, - ), - "decoding_stage", - ) + if is_glm_distributed_mode: + self.add_stage( + GlmImageDenoiserDecodingStage( + vae=self.get_module("vae"), pipeline=self + ), + "decoding_stage", + ) + else: + self.add_stage_factory( + RoleType.DECODER, + lambda: GlmImageDecodingStage( + vae=self.get_module("vae"), + pipeline=self, + ), + "decoding_stage", + ) EntryClass = [GlmImagePipeline] diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py b/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py index 3f0ada45f..edcc7a08e 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py @@ -270,6 +270,12 @@ class ComposedPipelineBase(ABC): extra_allowed_modules = set( role_to_pipeline_modules.get(role, {}).get(self.pipeline_name, set()) ) + if ( + role == RoleType.DENOISER + and self.pipeline_name == "GlmImagePipeline" + and getattr(self.server_args, "srt_encoder_url", None) is not None + ): + extra_allowed_modules.update({"text_encoder", "tokenizer", "vae"}) if role == RoleType.DENOISER and task_name == "ti2v": if self.pipeline_name in { diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py b/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py index bc36f0d96..0e217caeb 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py @@ -108,6 +108,10 @@ class Req: pooled_embeds: list[torch.Tensor] = field(default_factory=list) neg_pooled_embeds: list[torch.Tensor] = field(default_factory=list) + # GLM-Image autoregressive prior tokens + prior_token_id: torch.Tensor | None = None + prior_token_image_ids: torch.Tensor | list[torch.Tensor] | None = None + # Additional text-related parameters max_sequence_length: int | None = None prompt_template: dict[str, Any] | None = None diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/glm_image.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/glm_image.py index 3faaa54da..5455eb488 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/glm_image.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/glm_image.py @@ -15,7 +15,6 @@ from sglang.multimodal_gen.configs.sample.glmimage import ( GLM_IMAGE_RESOLUTION_ALIGNMENT, align_glm_image_resolution, ) -from sglang.multimodal_gen.runtime.distributed import get_local_torch_device from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import ( ComponentUse, @@ -27,6 +26,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.base import ( StageParallelismType, ) from sglang.multimodal_gen.runtime.pipelines_core.stages.decoding import DecodingStage +from sglang.multimodal_gen.runtime.platforms import current_platform 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.precision import ( @@ -414,7 +414,7 @@ class GlmImageAR(PipelineStage): Tuple of the D16 prior token IDs, optional source-image token IDs, and optional usage statistics returned by an external AR server. """ - device = get_local_torch_device() + device = current_platform.get_local_torch_device() _validate_glm_image_resolution_alignment(width, height) is_text_to_image = image is None or len(image) == 0 @@ -516,8 +516,9 @@ class GlmImageAR(PipelineStage): height: int, width: int, server_args: ServerArgs, + device: Optional[torch.device] = None, ) -> tuple[list[torch.Tensor], list[dict[str, int] | None]]: - device = get_local_torch_device() + device = device or current_platform.get_local_torch_device() _validate_glm_image_resolution_alignment(width, height) input_ids = [] @@ -577,27 +578,15 @@ class GlmImageAR(PipelineStage): usages.append(_extract_srt_usage(item.get("meta_info"))) return prior_token_ids, usages - def run_grouped_requests( + def generate_and_assign_prior_tokens( self, batches: list[Req], server_args: ServerArgs, + device: Optional[torch.device] = None, ) -> list[Req]: - can_batch_ar = ( - len(batches) > 1 - and server_args.srt_encoder_url is not None - and all( - isinstance(batch.prompt, str) and batch.image_path is None - for batch in batches - ) - ) - if not can_batch_ar: - return super().run_grouped_requests(batches, server_args) - + """Generate one AR batch and assign its tokens and usage to each request.""" height = batches[0].height width = batches[0].width - if any(batch.height != height or batch.width != width for batch in batches[1:]): - return super().run_grouped_requests(batches, server_args) - start_time = time.time() output_counts = [_num_outputs_per_prompt(batch) for batch in batches] prompts = [ @@ -616,6 +605,7 @@ class GlmImageAR(PipelineStage): height=height, width=width, server_args=server_args, + device=device, ) duration = time.time() - start_time logger.info( @@ -642,6 +632,29 @@ class GlmImageAR(PipelineStage): output_offset += output_count return batches + def run_grouped_requests( + self, + batches: list[Req], + server_args: ServerArgs, + ) -> list[Req]: + can_batch_ar = ( + len(batches) > 1 + and server_args.srt_encoder_url is not None + and all( + isinstance(batch.prompt, str) and batch.image_path is None + for batch in batches + ) + ) + if not can_batch_ar: + return super().run_grouped_requests(batches, server_args) + + height = batches[0].height + width = batches[0].width + if any(batch.height != height or batch.width != width for batch in batches[1:]): + return super().run_grouped_requests(batches, server_args) + + return self.generate_and_assign_prior_tokens(batches, server_args) + def iter_sequential_requests( self, batch: Req, server_args: ServerArgs ) -> Iterator[Req]: @@ -714,7 +727,7 @@ class GlmImageAR(PipelineStage): else: ar_condition_images = None - device = get_local_torch_device() + device = current_platform.get_local_torch_device() if ar_condition_images is not None: height = height or ar_condition_images[0].height @@ -1158,7 +1171,7 @@ class GlmImageBeforeDenoisingStage(PipelineStage): height = batch.height width = batch.width - device = get_local_torch_device() + device = current_platform.get_local_torch_device() batch_size = _num_outputs_per_prompt(batch) max_sequence_length = 1024 seed = getattr(batch, "seed", None) diff --git a/python/sglang/multimodal_gen/test/run_suite.py b/python/sglang/multimodal_gen/test/run_suite.py index 5acc334bc..aa8f33a02 100644 --- a/python/sglang/multimodal_gen/test/run_suite.py +++ b/python/sglang/multimodal_gen/test/run_suite.py @@ -39,6 +39,7 @@ if current_platform.is_npu(): DEFAULT_STANDALONE_EST_TIME_SECONDS, FILE_SUITES, PARAMETRIZED_CASE_GROUPS, + STANDALONE_FILE_EST_TIMES, STANDALONE_FILES, STARTUP_OVERHEAD_SECONDS, SUITES, diff --git a/python/sglang/multimodal_gen/test/server/ascend/test_glm_image_distributed.py b/python/sglang/multimodal_gen/test/server/ascend/test_glm_image_distributed.py new file mode 100644 index 000000000..d8fad5fe1 --- /dev/null +++ b/python/sglang/multimodal_gen/test/server/ascend/test_glm_image_distributed.py @@ -0,0 +1,315 @@ +"""NPU smoke test for the GLM-Image external-AR distributed topology.""" + +from __future__ import annotations + +import base64 +import os +import signal +import subprocess +import time +import unittest +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path + +import requests +import torch + +from sglang.multimodal_gen.test.server.ascend.testcase_configs_npu import ( + GLM_IMAGE_WEIGHTS_PATH, +) +from sglang.multimodal_gen.test.test_utils import 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")) +_STARTUP_TIMEOUT_S = float(os.environ.get("SGLANG_GLM_AR_STARTUP_TIMEOUT", "600")) +# A3 warm AR (26.711s) plus six sequential denoiser outputs (24.1945s each). +_EXPECTED_MAKESPAN_S = 172.0 +_PERFORMANCE_TOLERANCE = 0.25 + + +def _kill_process_tree(proc: subprocess.Popen) -> None: + try: + os.killpg(os.getpgid(proc.pid), signal.SIGKILL) + except (ProcessLookupError, PermissionError): + pass + + +def _tail_log(path: Path, lines: int = 80) -> str: + if not path.exists(): + return f"" + try: + return "\n".join(path.read_text(errors="ignore").splitlines()[-lines:]) + except OSError as error: + return f"" + + +def _wait_for_log(path: Path, message: str, timeout: float) -> None: + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + if path.exists(): + try: + if message in path.read_text(errors="ignore"): + return + except OSError: + pass + time.sleep(2) + raise TimeoutError(f"Missing {message!r} in {path}:\n{_tail_log(path)}") + + +class _GlmDistributedCluster: + def __init__(self) -> None: + self.model_path = Path(GLM_IMAGE_WEIGHTS_PATH) + self.ar_port = find_free_port(HOST) + self.denoiser_port = find_free_port(HOST) + self.head_port = find_free_port(HOST) + self.head_scheduler_port = find_free_port(HOST) + self.master_port = find_free_port(HOST) + self.processes: list[subprocess.Popen] = [] + self.log_paths = { + "ar": _LOG_DIR / "glm_image_distributed_ar.log", + "denoiser": _LOG_DIR / "glm_image_distributed_denoiser.log", + "head": _LOG_DIR / "glm_image_distributed_head.log", + } + self._log_handles: list = [] + + def __enter__(self) -> _GlmDistributedCluster: + if not self.model_path.is_dir(): + raise RuntimeError( + f"GLM-Image ModelScope cache is missing: {self.model_path}" + ) + try: + self._start_ar() + self._start_denoiser() + self._start_head() + except Exception: + self.stop() + raise + return self + + def __exit__(self, *exc) -> None: + self.stop() + + def _start_process(self, command: list[str], log_name: str) -> None: + log_handle = open(self.log_paths[log_name], "w") + self._log_handles.append(log_handle) + self.processes.append( + subprocess.Popen( + command, + stdout=log_handle, + stderr=subprocess.STDOUT, + preexec_fn=os.setsid, + env=os.environ.copy(), + ) + ) + + def _start_ar(self) -> None: + self._start_process( + [ + "sglang", + "serve", + "--model-path", + str(self.model_path / "vision_language_encoder"), + "--tokenizer-path", + str(self.model_path / "processor"), + "--enable-multimodal", + "--device", + "npu", + "--attention-backend", + "ascend", + "--disable-fast-image-processor", + "--tp-size", + "1", + "--cuda-graph-bs", + "2", + "--base-gpu-id", + "0", + "--host", + HOST, + "--port", + str(self.ar_port), + ], + "ar", + ) + self._wait_for_health("ar", self.ar_port) + + def _start_denoiser(self) -> None: + self._start_process( + [ + "sglang", + "serve", + "--model-path", + str(self.model_path), + "--disagg-role", + "denoiser", + "--disagg-server-addr", + f"tcp://{HOST}:{self.head_scheduler_port}", + "--srt-encoder-url", + f"http://{HOST}:{self.ar_port}", + "--scheduler-port", + str(self.denoiser_port), + "--master-port", + str(self.master_port), + "--num-gpus", + "1", + "--base-gpu-id", + "1", + "--denoiser-sp", + "1", + "--cfg-parallel-size", + "1", + "--batching-max-size", + "1", + "--dit-cpu-offload", + "false", + "--attention-backend", + "fa", + ], + "denoiser", + ) + + def _start_head(self) -> None: + self._start_process( + [ + "sglang", + "serve", + "--model-path", + str(self.model_path), + "--disagg-role", + "server", + "--denoiser-urls", + f"tcp://{HOST}:{self.denoiser_port}", + "--srt-encoder-url", + f"http://{HOST}:{self.ar_port}", + "--batching-mode", + "dynamic", + "--batching-max-size", + "2", + "--batching-delay-ms", + "100", + "--scheduler-port", + str(self.head_scheduler_port), + "--host", + HOST, + "--port", + str(self.head_port), + ], + "head", + ) + _wait_for_log( + self.log_paths["denoiser"], "Role DENOISER ready", _STARTUP_TIMEOUT_S + ) + self._wait_for_health("head", self.head_port) + + def _wait_for_health(self, name: str, port: int) -> None: + try: + wait_for_server_health( + f"http://{HOST}:{port}", path="/v1/models", timeout=_STARTUP_TIMEOUT_S + ) + except Exception as error: + raise RuntimeError( + f"{name} failed to become healthy:\n{_tail_log(self.log_paths[name])}" + ) from error + + def stop(self) -> None: + for process in self.processes: + _kill_process_tree(process) + for log_handle in self._log_handles: + log_handle.close() + self.processes.clear() + self._log_handles.clear() + + +class TestGlmImageDistributedNpu(CustomTestCase): + @classmethod + def setUpClass(cls) -> None: + super().setUpClass() + if not hasattr(torch, "npu") or torch.npu.device_count() < 2: + raise unittest.SkipTest("requires two Ascend NPUs") + cls.cluster = _GlmDistributedCluster() + cls.cluster.__enter__() + + @classmethod + def tearDownClass(cls) -> None: + if hasattr(cls, "cluster"): + for name, path in cls.cluster.log_paths.items(): + print(f"\n=== [glm-image-distributed] {name} log tail ===") + print(_tail_log(path)) + cls.cluster.stop() + super().tearDownClass() + + def _generate(self, prompt: str, n: int = 1) -> list[bytes]: + response = requests.post( + f"http://{HOST}:{self.cluster.head_port}/v1/images/generations", + json={ + "model": str(self.cluster.model_path), + "prompt": prompt, + "n": n, + "size": "1024x1024", + "response_format": "b64_json", + }, + timeout=600, + ) + response.raise_for_status() + images = response.json()["data"] + self.assertEqual(len(images), n) + return [base64.b64decode(image["b64_json"]) for image in images] + + def test_external_ar_batching_multi_output_disaggregation_performance( + self, + ) -> None: + self._generate("A warmup landscape") + requests_to_generate = [ + ("A mountain sunrise", 2), + ("A city at night", 1), + ("A forest lake", 2), + ("A desert sunset", 1), + ] + + def generate(request: tuple[str, int]) -> tuple[list[bytes], float]: + prompt, output_count = request + start_time = time.perf_counter() + images = self._generate(prompt, n=output_count) + return images, time.perf_counter() - start_time + + start_time = time.perf_counter() + with ThreadPoolExecutor(max_workers=4) as executor: + results = list(executor.map(generate, requests_to_generate)) + makespan_s = time.perf_counter() - start_time + request_latencies_s = [latency_s for _, latency_s in results] + images = [image for request_images, _ in results for image in request_images] + + print( + "GLM distributed performance: " + f"makespan={makespan_s:.2f}s, " + f"request_latencies={[round(value, 2) for value in request_latencies_s]}" + ) + self.assertLessEqual( + makespan_s, + _EXPECTED_MAKESPAN_S * (1 + _PERFORMANCE_TOLERANCE), + "GLM external-AR and DiT overlap performance regressed", + ) + + self.assertTrue( + all(images), + "each requested output must produce a decoded image", + ) + head_log = self.cluster.log_paths["head"].read_text(errors="ignore") + self.assertGreaterEqual( + head_log.count( + "GLM distributed AR dispatched batch size=2 requests, 2 outputs" + ), + 1, + "the two n=1 requests must share an external-AR batch", + ) + self.assertGreaterEqual( + head_log.count( + "GLM distributed AR dispatched batch size=1 requests, 2 outputs" + ), + 2, + "each n=2 request must preserve both outputs through disaggregation", + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/python/sglang/multimodal_gen/test/server/ascend/testcase_configs_npu.py b/python/sglang/multimodal_gen/test/server/ascend/testcase_configs_npu.py index d3745d0f5..afccd1008 100644 --- a/python/sglang/multimodal_gen/test/server/ascend/testcase_configs_npu.py +++ b/python/sglang/multimodal_gen/test/server/ascend/testcase_configs_npu.py @@ -237,14 +237,26 @@ DEFAULT_EST_TIME_SECONDS = 300.0 STARTUP_OVERHEAD_SECONDS = 120.0 DEFAULT_STANDALONE_EST_TIME_SECONDS = 300.0 +STANDALONE_FILES = { + "2-npu": [ + "ascend/test_glm_image_distributed.py", + ], +} + +STANDALONE_FILE_EST_TIMES = { + "2-npu": { + "ascend/test_glm_image_distributed.py": 900.0, + }, +} + SUITES = { "1-npu": [ "ascend/test_server_1_npu.py", - # add new 1-npu test files here + *STANDALONE_FILES.get("1-npu", []), ], "2-npu": [ "ascend/test_server_2_npu.py", - # add new 2-npu test files here + *STANDALONE_FILES.get("2-npu", []), ], } @@ -258,6 +270,5 @@ PARAMETRIZED_CASE_GROUPS = { } FILE_SUITES = {} -STANDALONE_FILES = {} COMPONENT_ACCURACY_SUITES = {} _UPDATE_WEIGHTS_FROM_DISK_TEST_FILE = None diff --git a/python/sglang/multimodal_gen/test/unit/test_glm_image_ar.py b/python/sglang/multimodal_gen/test/unit/test_glm_image_ar.py index 9f11c8e78..bf21e70e6 100644 --- a/python/sglang/multimodal_gen/test/unit/test_glm_image_ar.py +++ b/python/sglang/multimodal_gen/test/unit/test_glm_image_ar.py @@ -73,7 +73,7 @@ class TestGlmImageARSrtBackend(unittest.TestCase): @patch( "sglang.multimodal_gen.runtime.pipelines_core.stages." - "model_specific_stages.glm_image.get_local_torch_device", + "model_specific_stages.glm_image.current_platform.get_local_torch_device", return_value=torch.device("cpu"), ) @patch( @@ -102,7 +102,7 @@ class TestGlmImageARSrtBackend(unittest.TestCase): @patch( "sglang.multimodal_gen.runtime.pipelines_core.stages." - "model_specific_stages.glm_image.get_local_torch_device", + "model_specific_stages.glm_image.current_platform.get_local_torch_device", return_value=torch.device("cpu"), ) @patch( @@ -142,7 +142,7 @@ class TestGlmImageARSrtBackend(unittest.TestCase): @patch( "sglang.multimodal_gen.runtime.pipelines_core.stages." - "model_specific_stages.glm_image.get_local_torch_device", + "model_specific_stages.glm_image.current_platform.get_local_torch_device", return_value=torch.device("cpu"), ) @patch( @@ -182,7 +182,7 @@ class TestGlmImageARSrtBackend(unittest.TestCase): @patch( "sglang.multimodal_gen.runtime.pipelines_core.stages." - "model_specific_stages.glm_image.get_local_torch_device", + "model_specific_stages.glm_image.current_platform.get_local_torch_device", return_value=torch.device("cpu"), ) @patch( @@ -207,7 +207,7 @@ class TestGlmImageARSrtBackend(unittest.TestCase): @patch( "sglang.multimodal_gen.runtime.pipelines_core.stages." - "model_specific_stages.glm_image.get_local_torch_device", + "model_specific_stages.glm_image.current_platform.get_local_torch_device", return_value=torch.device("cpu"), ) def test_forward_aligns_runtime_dimensions_before_ar_generation(self, _mock_device): @@ -248,7 +248,7 @@ class TestGlmImageARSrtBackend(unittest.TestCase): @patch( "sglang.multimodal_gen.runtime.pipelines_core.stages." - "model_specific_stages.glm_image.get_local_torch_device", + "model_specific_stages.glm_image.current_platform.get_local_torch_device", return_value=torch.device("cpu"), ) @patch( @@ -282,7 +282,7 @@ class TestGlmImageARSrtBackend(unittest.TestCase): @patch( "sglang.multimodal_gen.runtime.pipelines_core.stages." - "model_specific_stages.glm_image.get_local_torch_device", + "model_specific_stages.glm_image.current_platform.get_local_torch_device", return_value=torch.device("cpu"), ) @patch( @@ -349,7 +349,7 @@ class TestGlmImageARSrtBackend(unittest.TestCase): @patch( "sglang.multimodal_gen.runtime.pipelines_core.stages." - "model_specific_stages.glm_image.get_local_torch_device", + "model_specific_stages.glm_image.current_platform.get_local_torch_device", return_value=torch.device("cpu"), ) def test_generate_prior_tokens_rejects_unaligned_internal_dimensions( @@ -370,7 +370,7 @@ class TestGlmImageARSrtBackend(unittest.TestCase): @patch( "sglang.multimodal_gen.runtime.pipelines_core.stages." - "model_specific_stages.glm_image.get_local_torch_device", + "model_specific_stages.glm_image.current_platform.get_local_torch_device", return_value=torch.device("cpu"), ) def test_generate_prior_tokens_batch_rejects_unaligned_internal_dimensions( diff --git a/python/sglang/multimodal_gen/test/unit/test_glm_image_multi_output.py b/python/sglang/multimodal_gen/test/unit/test_glm_image_multi_output.py index 8469d8a24..a8bf09ac6 100644 --- a/python/sglang/multimodal_gen/test/unit/test_glm_image_multi_output.py +++ b/python/sglang/multimodal_gen/test/unit/test_glm_image_multi_output.py @@ -110,7 +110,9 @@ def test_ar_stage_generates_one_prior_per_requested_output(): ) with patch.object( - glm_stage, "get_local_torch_device", return_value=torch.device("cpu") + glm_stage.current_platform, + "get_local_torch_device", + return_value=torch.device("cpu"), ): result = stage.forward(batch, SimpleNamespace()) @@ -135,7 +137,9 @@ def test_before_denoising_expands_latents_and_conditions_for_requested_outputs() ) with patch.object( - glm_stage, "get_local_torch_device", return_value=torch.device("cpu") + glm_stage.current_platform, + "get_local_torch_device", + return_value=torch.device("cpu"), ): result = stage.forward(batch, SimpleNamespace())