[NPU] [Diffusion] support distributed inference pipeline for GLM-Image (#31320)

Co-authored-by: Xiaoyu Zhang <1182563586@qq.com>
This commit is contained in:
Артем Савкин
2026-08-28 15:39:05 +03:00
committed by GitHub
co-authored by Xiaoyu Zhang
parent 803b4fb31c
commit ecbadf0b4b
14 changed files with 1166 additions and 118 deletions
@@ -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
@@ -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},
@@ -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.
@@ -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)
@@ -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]
@@ -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 {
@@ -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
@@ -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)
@@ -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,
@@ -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"<no log at {path}>"
try:
return "\n".join(path.read_text(errors="ignore").splitlines()[-lines:])
except OSError as error:
return f"<log read failed: {error}>"
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()
@@ -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
@@ -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(
@@ -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())