[NPU] [Diffusion] support distributed inference pipeline for GLM-Image (#31320)
Co-authored-by: Xiaoyu Zhang <1182563586@qq.com>
This commit is contained in:
co-authored by
Xiaoyu Zhang
parent
803b4fb31c
commit
ecbadf0b4b
@@ -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
|
||||
|
||||
+33
-20
@@ -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())
|
||||
|
||||
|
||||
Reference in New Issue
Block a user