From d96df7bed5566ef507faf8bccb715ea2fe6cd934 Mon Sep 17 00:00:00 2001
From: =?UTF-8?q?=D0=90=D1=80=D1=82=D0=B5=D0=BC=20=D0=A1=D0=B0=D0=B2=D0=BA?=
=?UTF-8?q?=D0=B8=D0=BD?= <58187114+OrangeRedeng@users.noreply.github.com>
Date: Wed, 5 Aug 2026 08:47:06 +0300
Subject: [PATCH] [Diffusion] Batch GLM-Image AR requests (#30683)
Co-authored-by: Xiaoyu Zhang <1182563586@qq.com>
---
.../sglang-diffusion/dynamic_batching.mdx | 6 +-
.../configs/pipeline_configs/base.py | 8 +
.../configs/pipeline_configs/glm_image.py | 15 +
.../managers/dynamic_batch_admission.py | 10 +-
.../runtime/managers/gpu_worker.py | 79 +++-
.../runtime/managers/scheduler.py | 223 ++++++++--
.../pipelines_core/composed_pipeline_base.py | 26 +-
.../executors/parallel_executor.py | 12 +-
.../executors/pipeline_executor.py | 53 +++
.../runtime/pipelines_core/schedule_batch.py | 5 +-
.../runtime/pipelines_core/stages/base.py | 7 +
.../runtime/pipelines_core/stages/decoding.py | 17 +-
.../stages/model_specific_stages/glm_image.py | 398 ++++++++++++++----
.../test/unit/test_glm_image_ar.py | 3 +
14 files changed, 708 insertions(+), 154 deletions(-)
diff --git a/docs/docs/sglang-diffusion/dynamic_batching.mdx b/docs/docs/sglang-diffusion/dynamic_batching.mdx
index a4d5d2640..3fbd54089 100644
--- a/docs/docs/sglang-diffusion/dynamic_batching.mdx
+++ b/docs/docs/sglang-diffusion/dynamic_batching.mdx
@@ -12,6 +12,10 @@ Use it for concurrent T2I or T2V traffic with the same model and sampling shape.
Dynamic batching is disabled by default with `--batching-max-size 1`.
+For GLM-Image T2I with an external AR server, compatible requests share one
+batched AR `/generate` call. DiT denoising and VAE decoding remain independent
+per-request executions; GLM-Image does not batch DiT work.
+
```bash Command
sglang serve \
--model-path black-forest-labs/FLUX.1-dev \
@@ -70,7 +74,7 @@ An initial implementation of dynamic batching for T2I and T2V models can be foun
| FLUX.2-Klein-Base-9B | ? | ? |
| Z-Image | ? | - |
| Z-Image-Turbo | ✅ | - |
- | GLM-Image | ❌ | - |
+ | GLM-Image (external AR) | ✅ | - |
| Qwen Image | ✅ | - |
| Qwen Image 2512 | ✅ | - |
| Qwen Image Edit | - | ❌ |
diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py
index 616bed20f..887d87ed1 100644
--- a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py
+++ b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py
@@ -408,6 +408,14 @@ class PipelineConfig:
"""Return whether dynamic batches should run as grouped Req lists."""
return False
+ def supports_sequential_dit_inference(self):
+ """Return whether batched AR is followed by per-request DiT inference."""
+ return False
+
+ def supports_sequential_multi_output_inference(self):
+ """Return whether one request's outputs run through DiT/VAE sequentially."""
+ return False
+
def estimate_request_cost(self, batch) -> float:
"""Return the relative cost used for batching admission caps.
diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/glm_image.py b/python/sglang/multimodal_gen/configs/pipeline_configs/glm_image.py
index 32b4874a4..a58fd8fec 100644
--- a/python/sglang/multimodal_gen/configs/pipeline_configs/glm_image.py
+++ b/python/sglang/multimodal_gen/configs/pipeline_configs/glm_image.py
@@ -13,6 +13,8 @@ from sglang.multimodal_gen.configs.pipeline_configs.base import (
SpatialImagePipelineConfig,
shard_rotary_emb_for_sp,
)
+from sglang.multimodal_gen.runtime.platforms import current_platform
+from sglang.multimodal_gen.runtime.server_args import get_global_server_args
@dataclass
@@ -48,6 +50,19 @@ class GlmImagePipelineConfig(SpatialImagePipelineConfig):
self.vae_scale_factor = self.vae_config.get_vae_scale_factor()
self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor)
+ def supports_dynamic_batching(self):
+ server_args = get_global_server_args()
+ return server_args.srt_encoder_url is not None
+
+ def supports_native_grouped_requests(self):
+ return True
+
+ def supports_sequential_dit_inference(self):
+ return True
+
+ def supports_sequential_multi_output_inference(self):
+ return current_platform.is_npu()
+
def get_freqs_cis(self, batch, device, rotary_emb, dtype):
height = batch.height // self.vae_scale_factor
width = batch.width // self.vae_scale_factor
diff --git a/python/sglang/multimodal_gen/runtime/managers/dynamic_batch_admission.py b/python/sglang/multimodal_gen/runtime/managers/dynamic_batch_admission.py
index 528bf86c4..c398251ea 100644
--- a/python/sglang/multimodal_gen/runtime/managers/dynamic_batch_admission.py
+++ b/python/sglang/multimodal_gen/runtime/managers/dynamic_batch_admission.py
@@ -185,7 +185,7 @@ class BatchAdmissionController:
proposed = current_reqs + [candidate_req]
limit = self.limit_for(proposed[0])
return limit.reject_reason(
- batch_size=len(proposed),
+ batch_size=self._effective_batch_size(proposed),
batch_cost=self.estimate_batch_cost(proposed),
)
@@ -195,7 +195,7 @@ class BatchAdmissionController:
return len(reqs) >= self._user_max_batch_size
limit = self.limit_for(reqs[0])
- if len(reqs) >= limit.max_batch_size:
+ if self._effective_batch_size(reqs) >= limit.max_batch_size:
return True
next_cost = self.estimate_batch_cost(reqs + [reqs[0]])
@@ -206,7 +206,7 @@ class BatchAdmissionController:
return None
limit = self.limit_for(reqs[0])
- if len(reqs) >= limit.max_batch_size:
+ if self._effective_batch_size(reqs) >= limit.max_batch_size:
return limit.cap_reason or f"config_cap:{limit.max_batch_size}"
next_cost = self.estimate_batch_cost(reqs + [reqs[0]])
@@ -240,6 +240,10 @@ class BatchAdmissionController:
float(self._pipeline_config.estimate_request_cost(req)) for req in reqs
)
+ @staticmethod
+ def _effective_batch_size(reqs: list[Req]) -> int:
+ return sum(max(1, int(req.num_outputs_per_prompt or 1)) for req in reqs)
+
def _matching_rules(self, req: Req) -> list[BatchingRule]:
return [
rule
diff --git a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py
index aaf14edd3..e25a8921a 100644
--- a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py
+++ b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py
@@ -9,7 +9,7 @@ import tempfile
import time
from contextlib import ExitStack
from dataclasses import dataclass, field
-from typing import Any, Callable, List, Union
+from typing import Any, Callable, Iterator, List, Union
import numpy as np
import torch
@@ -371,6 +371,54 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
error_context=f"request {req.request_id}",
)
+ def execute_forward_sequentially(self, batch: list[Req]) -> Iterator[OutputBatch]:
+ """Yield grouped results after each request finishes its terminal stage."""
+ assert self.pipeline is not None
+ results = self.pipeline.forward_batch_sequentially(batch, self.server_args)
+ group_start_time = time.monotonic()
+
+ try:
+ for req in batch:
+ output_count = (
+ max(1, int(req.num_outputs_per_prompt or 1))
+ if self.server_args.pipeline_config.supports_sequential_multi_output_inference()
+ else 1
+ )
+ output_batch = self._execute_forward_common(
+ req,
+ forward_fn=lambda results=results, output_count=output_count: (
+ self._collect_sequential_outputs(results, output_count)
+ ),
+ log_reqs=[req],
+ return_req=False,
+ save_output_paths=lambda output_batch, req=req: self._save_output_paths(
+ req, output_batch
+ ),
+ error_context=f"grouped request {req.request_id}",
+ execution_start_time=group_start_time,
+ propagate_forward_errors=True,
+ )
+ assert isinstance(output_batch, OutputBatch)
+ yield output_batch
+ del output_batch
+ finally:
+ close = getattr(results, "close", None)
+ if close is not None:
+ close()
+
+ def _collect_sequential_outputs(
+ self,
+ results: Iterator[OutputBatch | Req],
+ output_count: int,
+ ) -> OutputBatch | Req:
+ if output_count == 1:
+ return next(results)
+
+ output_batches = [
+ self._to_output_batch(next(results)) for _ in range(output_count)
+ ]
+ return self._merge_expanded_output_batches(output_batches)
+
def _execute_forward_batch(self, batch: list[Req]) -> OutputBatch | Req:
"""Execute expanded multi-output requests as one grouped forward."""
# TODO: support early return or mix-stage execution for reqs in a group
@@ -396,17 +444,24 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
return_req: bool,
save_output_paths: Callable[[OutputBatch], None],
error_context: str,
+ execution_start_time: float | None = None,
+ propagate_forward_errors: bool = False,
) -> OutputBatch | Req:
"""
Args:
forward_fn: the actual forward function for reqs
"""
output_batch = None
+ forward_failed = False
try:
if self.rank == 0 and not current_platform.is_cpu():
torch.get_device_module().reset_peak_memory_stats()
- start_time = time.monotonic()
+ start_time = (
+ execution_start_time
+ if execution_start_time is not None
+ else time.monotonic()
+ )
self._realtime_sessions.attach(req)
# capture memory baseline for each req in grouped forward on rank-0
@@ -425,7 +480,11 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
stack.enter_context(
trace_slice(item.trace_ctx, DiffStage.GPU_FORWARD)
)
- result = forward_fn()
+ try:
+ result = forward_fn()
+ except Exception:
+ forward_failed = True
+ raise
# disagg roles return raw Req so callers can keep and transfer intermediate tensors
# before converting it to OutputBatch
@@ -456,11 +515,11 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
self._materialize_output_transport(output_batch, req, save_output_paths)
if (
- torch.cuda.is_initialized()
+ not current_platform.is_cpu()
and output_batch.output is None
and not req.return_raw_frames
):
- torch.cuda.empty_cache()
+ torch.get_device_module().empty_cache()
if req.perf_dump_path is not None or envs.SGLANG_DIFFUSION_STAGE_LOGGING:
if not req.is_warmup:
@@ -479,6 +538,12 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
tag="server_perf_dump",
)
except Exception as e:
+ if propagate_forward_errors and forward_failed:
+ if isinstance(e, StopIteration):
+ raise RuntimeError(
+ "Grouped pipeline returned fewer outputs than requests."
+ ) from e
+ raise
logger.error(
f"Error executing {error_context}: {e}",
exc_info=True,
@@ -490,8 +555,8 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
output_batch.error = f"Error executing {error_context}: {e}"
self._record_output_peak_memory(output_batch)
# clean cache if OOM
- if torch.cuda.is_initialized():
- torch.cuda.empty_cache()
+ if not current_platform.is_cpu():
+ torch.get_device_module().empty_cache()
return output_batch
def _materialize_output_transport(
diff --git a/python/sglang/multimodal_gen/runtime/managers/scheduler.py b/python/sglang/multimodal_gen/runtime/managers/scheduler.py
index cf64593e5..9ec0f7001 100644
--- a/python/sglang/multimodal_gen/runtime/managers/scheduler.py
+++ b/python/sglang/multimodal_gen/runtime/managers/scheduler.py
@@ -72,6 +72,11 @@ _MAX_RECV_REQS_PER_POLL = 1024
_BATCH_METRICS_LOG_INTERVAL = 5
+@dataclasses.dataclass(frozen=True)
+class _SequentiallyReturnedOutputs:
+ outputs: Iterator[OutputBatch]
+
+
class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisaggMixin):
"""
Runs the main event loop for the rank 0 worker.
@@ -249,7 +254,7 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
def _dispatch_items(
self, items: list[tuple[bytes | None, Any]]
- ) -> OutputBatch | list[OutputBatch]:
+ ) -> OutputBatch | list[OutputBatch] | _SequentiallyReturnedOutputs:
"""Dispatch ready queue items; several plain `Req`s form one dynamic batch."""
reqs = [item[1] for item in items]
if len(reqs) > 1 and all(isinstance(req, Req) for req in reqs):
@@ -279,6 +284,15 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
DiffStage.SCHEDULER_DISPATCH,
thread_finish_flag=True,
):
+ if (
+ len(reqs) == 1
+ and self.server_args.pipeline_config.supports_sequential_multi_output_inference()
+ and max(1, int(req.num_outputs_per_prompt or 1)) > 1
+ ):
+ return _SequentiallyReturnedOutputs(
+ self._iter_grouped_outputs_sequentially(reqs)
+ )
+
if len(reqs) == 1 or not allow_dynamic_batching:
return self.worker.execute_forward(reqs)
@@ -330,8 +344,15 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
error_msg=f"Dynamic batching failed: {e}",
)
- def _execute_generation_grouped(self, reqs: List[Req]) -> List[OutputBatch]:
+ def _execute_generation_grouped(
+ self, reqs: List[Req]
+ ) -> List[OutputBatch] | _SequentiallyReturnedOutputs:
batch_size = len(reqs)
+ if self.server_args.pipeline_config.supports_sequential_dit_inference():
+ return _SequentiallyReturnedOutputs(
+ self._iter_grouped_outputs_sequentially(reqs)
+ )
+
try:
output_batch = self.worker.execute_forward(reqs)
if output_batch.error:
@@ -372,6 +393,18 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
error_msg=f"Native grouped execution failed: {e}",
)
+ def _iter_grouped_outputs_sequentially(
+ self, reqs: List[Req]
+ ) -> Iterator[OutputBatch]:
+ yield from self.worker.execute_forward_sequentially(reqs)
+ logger.info(
+ "Processed native grouped batch sequentially: %d/%d request(s) "
+ "with max_delay=%.2fms",
+ len(reqs),
+ self._batching_max_size,
+ self._batching_delay_s * 1000.0,
+ )
+
def _execute_generation_sequential(self, reqs: List[Req]) -> List[OutputBatch]:
return [self.worker.execute_forward([req]) for req in reqs]
@@ -412,10 +445,14 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
except Exception:
return None
+ exclude_num_outputs = (
+ self.server_args.pipeline_config.supports_sequential_dit_inference()
+ )
return [
(f.name, self._freeze_signature_value(getattr(sp, f.name, None)))
for f in sp_fields
if not f.metadata.get("batch_sig_exclude", False)
+ and not (exclude_num_outputs and f.name == "num_outputs_per_prompt")
]
def _diffusers_kwargs_signature_value(self, req: Req) -> Any:
@@ -424,14 +461,21 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
def _build_dynamic_batch_signature(self, req: Req) -> tuple[Any, ...] | None:
"""Build the request compatibility signature for dynamic batching.
- The signature is built from `SamplingParams` fields, excluding fields
- marked with `batch_sig_exclude`, plus generation-affecting
- `extra.diffusers_kwargs`.
+ The signature is built from batch-shared `SamplingParams` fields, plus
+ generation-affecting `extra.diffusers_kwargs` and profiling settings
+ used by grouped execution.
"""
signature_items = self._sampling_param_signature_items(req)
if signature_items is None:
return None
+ profile_signature = (
+ (True, req.profile_all_stages, req.num_profiled_timesteps)
+ if req.profile
+ else (False,)
+ )
+ signature_items.append(("profiling", profile_signature))
+
if req.extra:
diffusers_kwargs = req.extra.get("diffusers_kwargs")
if diffusers_kwargs:
@@ -478,6 +522,26 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
if base_diffusers_kwargs != candidate_diffusers_kwargs:
return "extra.diffusers_kwargs"
+ if base_req.profile:
+ base_profile = (
+ True,
+ base_req.profile_all_stages,
+ base_req.num_profiled_timesteps,
+ )
+ else:
+ base_profile = (False,)
+
+ if candidate_req.profile:
+ candidate_profile = (
+ True,
+ candidate_req.profile_all_stages,
+ candidate_req.num_profiled_timesteps,
+ )
+ else:
+ candidate_profile = (False,)
+ if base_profile != candidate_profile:
+ return "profiling"
+
return None
def _get_dynamic_batch_reject_reason(
@@ -548,20 +612,23 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
def _record_batch_dispatch_metrics(
self,
- batch_size: int,
+ request_count: int,
+ output_count: int,
queue_wait_ms: float,
- effective_max_batch_size: int,
+ effective_max_output_count: int,
reject_reasons: list[str] | None = None,
stop_reason: str | None = None,
) -> None:
if not self._batch_metrics_enabled:
return
- effective_max_batch_size = max(1, effective_max_batch_size)
+ effective_max_output_count = max(1, effective_max_output_count)
logger.info(
- "Dynamic batch dispatch: size=%d/%d, user_max=%d, queue_wait=%.2fms, stop_reason=%s",
- batch_size,
- effective_max_batch_size,
+ "Dynamic batch dispatch: requests=%d, outputs=%d/%d, "
+ "user_max_outputs=%d, queue_wait=%.2fms, stop_reason=%s",
+ request_count,
+ output_count,
+ effective_max_output_count,
self._batching_max_size,
max(queue_wait_ms, 0.0),
stop_reason or "unspecified",
@@ -569,11 +636,15 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
window = self._batch_metrics_window
window.dispatches += 1
- window.total_requests += batch_size
- window.total_capacity += effective_max_batch_size
- if batch_size > 1:
+ window.total_requests += request_count
+ window.total_outputs += output_count
+ window.total_capacity += effective_max_output_count
+ if request_count > 1:
window.merged_dispatches += 1
- if self._dynamic_batching_enabled() and batch_size >= effective_max_batch_size:
+ if (
+ self._dynamic_batching_enabled()
+ and output_count >= effective_max_output_count
+ ):
window.full_dispatches += 1
window.wait_times_ms.append(max(queue_wait_ms, 0.0))
if reject_reasons:
@@ -590,8 +661,9 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
if window.dispatches == 0:
return
- avg_size = window.total_requests / window.dispatches
- utilization = window.total_requests / max(1, window.total_capacity)
+ avg_requests = window.total_requests / window.dispatches
+ avg_outputs = window.total_outputs / window.dispatches
+ utilization = window.total_outputs / max(1, window.total_capacity)
avg_wait_ms = sum(window.wait_times_ms) / len(window.wait_times_ms)
p95_wait_ms = self._percentile(window.wait_times_ms, 95.0)
merged_rate = window.merged_dispatches / window.dispatches
@@ -604,9 +676,13 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
top_rejects = "none"
logger.info(
- "Dynamic batch stats (last %d dispatches): avg_size=%.2f, merged_rate=%.1f%%, full_rate=%.1f%%, utilization=%.1f%%, wait_avg=%.2fms, wait_p95=%.2fms, top_rejects=%s",
+ "Dynamic batch stats (last %d dispatches): avg_requests=%.2f, "
+ "avg_outputs=%.2f, merged_rate=%.1f%%, full_rate=%.1f%%, "
+ "utilization=%.1f%%, wait_avg=%.2fms, wait_p95=%.2fms, "
+ "top_rejects=%s",
window.dispatches,
- avg_size,
+ avg_requests,
+ avg_outputs,
merged_rate * 100.0,
full_rate * 100.0,
utilization * 100.0,
@@ -659,6 +735,68 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
):
self.receiver.send_multipart([identity, b"", payload])
+ def _return_item_result(
+ self,
+ item: tuple[bytes | None, Any],
+ output_batch: OutputBatch,
+ ) -> None:
+ identity, processed_req = item
+ is_warmup = is_warmup_req(processed_req)
+ self._log_warmup_result(output_batch, processed_req, is_warmup)
+
+ if self._should_return_lightweight_warmup_result(processed_req):
+ output_batch.drop_payload_for_warmup()
+ self.return_result(output_batch, identity, should_not_return=False)
+ else:
+ self.return_result(output_batch, identity, should_not_return=is_warmup)
+
+ def _return_results_sequentially(
+ self,
+ items: list[tuple[bytes | None, Any]],
+ outputs: Iterator[OutputBatch],
+ ) -> None:
+ output_iter = iter(outputs)
+ try:
+ for index, item in enumerate(items):
+ output_batch, error = self._fetch_next_output(output_iter)
+ if error is not None:
+ self._return_sequential_errors(items[index:], error)
+ return
+
+ assert output_batch is not None
+ self._return_item_result(item, output_batch)
+ del output_batch
+ finally:
+ close = getattr(output_iter, "close", None)
+ if close is not None:
+ close()
+
+ @staticmethod
+ def _fetch_next_output(
+ output_iter: Iterator[OutputBatch],
+ ) -> tuple[OutputBatch | None, str | None]:
+ try:
+ return next(output_iter), None
+ except StopIteration:
+ error = (
+ "Grouped execution returned fewer outputs than requests "
+ "while processing sequentially."
+ )
+ logger.error(error)
+ return None, error
+ except Exception as e:
+ error = f"Failed to execute grouped requests sequentially: {e}"
+ logger.error(error, exc_info=True)
+ return None, error
+
+ def _return_sequential_errors(
+ self,
+ items: list[tuple[bytes | None, Any]],
+ error: str,
+ ) -> None:
+ for item in items:
+ self._return_item_result(item, OutputBatch(error=error))
+
@contextmanager
def _record_return_stage(
self, output_batch: OutputBatch, stage_name: str
@@ -844,10 +982,12 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
if not self._dynamic_batching_enabled():
identity, req, enqueue_time = self.waiting_queue.popleft()
if isinstance(req, Req):
+ output_count = max(1, int(req.num_outputs_per_prompt or 1))
self._record_batch_dispatch_metrics(
- batch_size=1,
+ request_count=1,
+ output_count=output_count,
queue_wait_ms=(time.monotonic() - enqueue_time) * 1000.0,
- effective_max_batch_size=1,
+ effective_max_output_count=output_count,
stop_reason="dynamic_disabled",
)
return [(identity, req)]
@@ -866,10 +1006,12 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
reason = self._get_dynamic_batch_reject_reason(req, req)
if reason is not None:
reject_reasons.append(f"head:{reason}")
+ output_count = max(1, int(req.num_outputs_per_prompt or 1))
self._record_batch_dispatch_metrics(
- batch_size=1,
+ request_count=1,
+ output_count=output_count,
queue_wait_ms=(time.monotonic() - head_enqueue_time) * 1000.0,
- effective_max_batch_size=1,
+ effective_max_output_count=output_count,
reject_reasons=reject_reasons,
stop_reason=reject_reasons[0] if reject_reasons else "head_ineligible",
)
@@ -930,9 +1072,12 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
else:
stop_reason = "ready"
self._record_batch_dispatch_metrics(
- batch_size=batch_len,
+ request_count=batch_len,
+ output_count=sum(
+ max(1, int(req.num_outputs_per_prompt or 1)) for req in compatible_reqs
+ ),
queue_wait_ms=oldest_wait_s * 1000.0,
- effective_max_batch_size=self._batch_admission.max_admissible_batch_size(
+ effective_max_output_count=self._batch_admission.max_admissible_batch_size(
compatible_reqs[0]
),
reject_reasons=reject_reasons,
@@ -1086,6 +1231,13 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
)
handler_result = OutputBatch(error=str(e))
+ if isinstance(handler_result, _SequentiallyReturnedOutputs):
+ try:
+ self._return_results_sequentially(items, handler_result.outputs)
+ except zmq.ZMQError as e:
+ logger.error(f"ZMQ error sending replies sequentially: {e}")
+ continue
+
if isinstance(handler_result, list):
output_batches = handler_result
else:
@@ -1109,25 +1261,8 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
# 3. return results
try:
- for (identity, processed_req), output_batch in zip(
- items, output_batches, strict=True
- ):
- is_warmup = is_warmup_req(processed_req)
- self._log_warmup_result(output_batch, processed_req, is_warmup)
-
- should_return_lightweight_warmup_result = (
- self._should_return_lightweight_warmup_result(processed_req)
- )
- if should_return_lightweight_warmup_result:
- # internal prewarm is a real-path request; reply but drop payloads
- output_batch.drop_payload_for_warmup()
- self.return_result(
- output_batch, identity, should_not_return=False
- )
- else:
- self.return_result(
- output_batch, identity, should_not_return=is_warmup
- )
+ for item, output_batch in zip(items, output_batches, strict=True):
+ self._return_item_result(item, output_batch)
except zmq.ZMQError as e:
# Reply failed; log and keep loop alive to accept future requests
logger.error(f"ZMQ error sending reply: {e}")
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py b/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py
index 095454229..0aa5eb792 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py
@@ -9,7 +9,7 @@ This module defines the base class for pipelines that are composed of multiple s
import os
from abc import ABC, abstractmethod
-from typing import Any, Callable, Literal, cast
+from typing import Any, Callable, Iterator, Literal, cast
import torch
from tqdm import tqdm
@@ -1065,3 +1065,27 @@ class ComposedPipelineBase(ABC):
return self.executor.execute_group_with_profiling(
self.stages, batches, server_args
)
+
+ @torch.no_grad()
+ def forward_batch_sequentially(
+ self,
+ batches: list[Req],
+ server_args: ServerArgs,
+ ) -> Iterator[OutputBatch]:
+ """Yield grouped outputs as each terminal-stage invocation completes."""
+ if len(batches) == 1 and (
+ not server_args.pipeline_config.supports_sequential_multi_output_inference()
+ or max(1, int(batches[0].num_outputs_per_prompt or 1)) == 1
+ ):
+ yield self.forward(batches[0], server_args)
+ return
+
+ self.component_residency_manager = get_global_component_residency_manager(
+ self, server_args
+ )
+ self.executor.component_residency_manager = self.component_residency_manager
+ yield from self.executor.execute_group_sequentially_with_profiling(
+ self.stages,
+ batches,
+ server_args,
+ )
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/executors/parallel_executor.py b/python/sglang/multimodal_gen/runtime/pipelines_core/executors/parallel_executor.py
index 25e171aab..f3ff019ee 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/executors/parallel_executor.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/executors/parallel_executor.py
@@ -141,13 +141,15 @@ class ParallelExecutor(PipelineExecutor):
broadcasted_list = broadcast_pyobj(
obj_list, rank=rank, dist_group=group.cpu_group, src=0
)
- if rank != 0:
- success, batch = broadcasted_list[0], broadcasted_list[1]
- else:
- success = obj_list[0]
+ success, broadcasted_batch = broadcasted_list
if not success:
- raise RuntimeError(f"Error on rank 0") from batch
+ if isinstance(broadcasted_batch, BaseException):
+ raise RuntimeError("Error on rank 0") from broadcasted_batch
+ raise RuntimeError(f"Error on rank 0: {broadcasted_batch}")
+
+ if rank != 0:
+ batch = broadcasted_batch
torch.distributed.barrier()
return batch
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/executors/pipeline_executor.py b/python/sglang/multimodal_gen/runtime/pipelines_core/executors/pipeline_executor.py
index cc5e22738..5d70b628c 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/executors/pipeline_executor.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/executors/pipeline_executor.py
@@ -6,6 +6,7 @@ Base class for all pipeline executors.
"""
import contextlib
+import time
from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Any, Callable, List
@@ -152,6 +153,21 @@ class PipelineExecutor(ABC):
batches = self.execute_group(stages, batches, server_args)
return batches
+ def execute_group_sequentially_with_profiling(
+ self,
+ stages: List["PipelineStage"],
+ batches: list[Req],
+ server_args: ServerArgs,
+ ):
+ """Run the AR stage as a group, then yield each completed DiT request."""
+ with self.profile_execution(batches[0], dump_rank=0):
+ with current_platform.inference_mode():
+ yield from self.execute_group_sequentially(
+ stages,
+ batches,
+ server_args,
+ )
+
@staticmethod
@contextlib.contextmanager
def _stage_execution_context(stage: "PipelineStage", server_args: ServerArgs):
@@ -245,6 +261,43 @@ class PipelineExecutor(ABC):
batches = stage.run_grouped_requests(batches, server_args)
return batches
+ def execute_group_sequentially(
+ self,
+ stages: List["PipelineStage"],
+ batches: list[Req],
+ server_args: ServerArgs,
+ ):
+ """Yield outputs after batched AR and sequential DiT/VAE inference."""
+ batches = self.execute_group(stages[:1], batches, server_args)
+
+ remaining_stages = stages[1:]
+ sequential_start_time = time.monotonic()
+ for parent_batch in batches:
+ for batch in stages[0].iter_sequential_requests(parent_batch, server_args):
+ if batch.metrics is not None:
+ batch.metrics.record_stage(
+ "PipelineExecutor.sequential_wait",
+ time.monotonic() - sequential_start_time,
+ )
+ try:
+ output = self.execute(remaining_stages, batch, server_args)
+ except Exception as e:
+ logger.error(
+ "Sequential DiT/VAE inference failed for request %s: %s",
+ batch.request_id,
+ e,
+ exc_info=True,
+ )
+ output = OutputBatch(
+ error=f"Error executing grouped request {batch.request_id}: {e}",
+ metrics=batch.metrics,
+ )
+ yield output
+ del output
+ del batch
+ if current_platform.is_npu():
+ torch.get_device_module().empty_cache()
+
@contextlib.contextmanager
def profile_execution(self, batch: Req, dump_rank: int = 0):
"""
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py b/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py
index 59166a632..a868d2d9b 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py
@@ -50,12 +50,13 @@ SAMPLING_PARAMS_FIELDS = {f.name for f in fields(SamplingParams)}
class BatchMetricsWindow:
"""Counters accumulated between dynamic batching metric logs.
- `total_capacity` uses each dispatch's effective admission cap, so
- utilization reflects model/config limits instead of only the user max.
+ `total_outputs` and `total_capacity` use output slots, so utilization
+ reflects model/config limits even when one request asks for many outputs.
"""
dispatches: int = 0
total_requests: int = 0
+ total_outputs: int = 0
total_capacity: int = 0
merged_dispatches: int = 0
full_dispatches: int = 0
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/base.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/base.py
index 5eb76fa9f..0e05c7bcc 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/base.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/base.py
@@ -144,6 +144,13 @@ class PipelineStage(StageDedupMixin, ABC):
"""
pass
+ def iter_sequential_requests(
+ self, batch: Req, server_args: ServerArgs
+ ) -> Iterator[Req]:
+ """Expand one post-stage request into sequential downstream requests."""
+ del server_args
+ return iter((batch,))
+
def set_component_residency_manager(self, manager) -> None:
self._component_residency_manager = manager
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py
index 782fa4494..8d48d967f 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py
@@ -250,7 +250,22 @@ class DecodingStage(PipelineStage):
with temporary_module_dtype(
self.vae, vae_dtype, enabled=should_cast_vae
) as vae:
- decode_output = self._get_vae_decode_fn(vae, server_args)(latents)
+ try:
+ decode_output = self._get_vae_decode_fn(vae, server_args)(latents)
+ except Exception as error:
+ if "out of memory" in str(error).lower():
+ if not server_args.pipeline_config.vae_tiling:
+ logger.warning(
+ "OOM detected during VAE decoding. Please enable "
+ "--vae-tiling to reduce peak memory usage."
+ )
+ else:
+ logger.warning(
+ "OOM detected during VAE decoding with tiling enabled. "
+ "Please reduce the resolution or enable "
+ "--vae-cpu-offload."
+ )
+ raise
image = _ensure_tensor_decode_output(decode_output)
# De-normalize image to [0, 1] range
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/glm_image.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/glm_image.py
index 517d92e92..b59fb67cb 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/glm_image.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/glm_image.py
@@ -1,7 +1,8 @@
import inspect
import re
import time
-from typing import List, Optional, Tuple, Union
+from copy import copy, deepcopy
+from typing import Any, Iterator, List, Optional, Tuple, Union
import numpy as np
import PIL
@@ -214,6 +215,86 @@ class GlmImageAR(PipelineStage):
token_ids = token_ids.reshape(1, -1)
return token_ids
+ @staticmethod
+ def _external_ar_sampling_params(max_new_tokens: int, seed: Optional[int]):
+ sampling_params = {
+ "temperature": 1.0,
+ "max_new_tokens": max_new_tokens,
+ "ignore_eos": True,
+ }
+ if seed is not None:
+ sampling_params["sampling_seed"] = seed
+ return sampling_params
+
+ @staticmethod
+ def _request_external_ar(payload: dict, server_args: ServerArgs):
+ try:
+ response = requests.post(
+ server_args.srt_encoder_url + "/generate",
+ json=payload,
+ timeout=(
+ server_args.srt_encoder_connect_timeout,
+ server_args.srt_encoder_timeout,
+ ),
+ )
+ response.raise_for_status()
+ except requests.ConnectTimeout as e:
+ logger.error(
+ "Connection timeout to SGLang encoder (%s). Try to increase "
+ "--srt-encoder-connection-timeout (current: %s sec). Details: %s",
+ server_args.srt_encoder_url,
+ server_args.srt_encoder_connect_timeout,
+ e,
+ )
+ raise
+ except requests.ReadTimeout as e:
+ logger.error(
+ "Read timeout from SGLang encoder (%s). Try to increase "
+ "--srt-encoder-timeout (current: %s sec). Details: %s",
+ server_args.srt_encoder_url,
+ server_args.srt_encoder_timeout,
+ e,
+ )
+ raise
+ except requests.ConnectionError as e:
+ logger.error(
+ "Failed to connect to SGLang encoder at %s: %s",
+ server_args.srt_encoder_url,
+ e,
+ )
+ raise
+ except requests.RequestException as e:
+ logger.error(
+ "SGLang encoder request to %s failed: %s",
+ server_args.srt_encoder_url,
+ e,
+ )
+ raise
+ return response.json()
+
+ def _extract_prior_token_ids(
+ self,
+ generated_ids: Any,
+ generation_shape: tuple[int, int, int],
+ device: torch.device,
+ ) -> torch.Tensor:
+ large_image_offset, token_h, token_w = generation_shape
+ expected_output_len = large_image_offset + token_h * token_w
+ actual_output_len = 0 if generated_ids is None else len(generated_ids)
+ if actual_output_len < expected_output_len:
+ raise RuntimeError(
+ "GLM-Image AR returned too few output_ids: "
+ f"got {actual_output_len}, need at least {expected_output_len} "
+ f"(large_image_offset={large_image_offset}, "
+ f"token_h={token_h}, token_w={token_w})."
+ )
+
+ prior_token_ids_d32 = torch.tensor(
+ generated_ids[large_image_offset : large_image_offset + token_h * token_w],
+ device=device,
+ )
+ return self._upsample_token_ids(prior_token_ids_d32, token_h, token_w)
+
def generate_prior_tokens(
self,
prompt: str,
@@ -222,6 +303,7 @@ class GlmImageAR(PipelineStage):
server_args: ServerArgs,
image: Optional[List[PIL.Image.Image]] = None,
factor: int = 32,
+ seed: Optional[int] = None,
) -> Tuple[torch.Tensor, int, int]:
"""
Generate prior tokens using the AR (vision_language_encoder) model.
@@ -281,56 +363,11 @@ class GlmImageAR(PipelineStage):
payload = {
"input_ids": inputs["input_ids"][0].tolist(),
"image_data": [{"image_grid_thw": image_grid_thw.tolist()}],
- "sampling_params": {
- "temperature": 1.0,
- "max_new_tokens": max_new_tokens,
- "ignore_eos": True,
- },
+ "sampling_params": self._external_ar_sampling_params(
+ max_new_tokens, seed
+ ),
}
- try:
- response = requests.post(
- server_args.srt_encoder_url + "/generate",
- json=payload,
- timeout=(
- server_args.srt_encoder_connect_timeout,
- server_args.srt_encoder_timeout,
- ),
- )
- except requests.ConnectionError as e:
- logger.error(
- "Failed to establish a connection to SGLang encoder server at %s. "
- "Verify that the AR model server is running and accessible. Error details: %s",
- server_args.srt_encoder_url,
- e,
- )
- raise
- except requests.ConnectTimeout as e:
- logger.error(
- "Connection timeout to SGLang encoder (%s). Try to increase --srt-encoder-connection-timeout (current: %s sec). Details: %s",
- server_args.srt_encoder_url,
- server_args.srt_encoder_connect_timeout,
- e,
- )
- raise
- except requests.ReadTimeout as e:
- logger.error(
- "Read timeout from SGLang encoder (%s). Try to increase --srt-encoder-timeout (current: %s sec). Details: %s",
- server_args.srt_encoder_url,
- server_args.srt_encoder_timeout,
- e,
- )
- raise
- except requests.RequestException as e:
- logger.error(
- "An error occurred during communication with SGLang encoder server at %s. "
- "The server is reachable, but the request failed. Error type: %s, Details: %s",
- server_args.srt_encoder_url,
- type(e).__name__,
- e,
- )
- raise
-
- data = response.json()
+ data = self._request_external_ar(payload, server_args)
generated_ids = data.get("output_ids")
else:
if image is not None:
@@ -367,27 +404,193 @@ class GlmImageAR(PipelineStage):
input_len = inputs["input_ids"].shape[-1]
generated_ids = outputs[0][input_len:]
- expected_output_len = large_image_offset + token_h * token_w
- actual_output_len = 0 if generated_ids is None else len(generated_ids)
- if actual_output_len < expected_output_len:
- raise RuntimeError(
- "GLM-Image AR returned too few output_ids: "
- f"got {actual_output_len}, need at least {expected_output_len} "
- f"(large_image_offset={large_image_offset}, "
- f"token_h={token_h}, token_w={token_w})."
- )
-
- # Extract large image tokens + upsample D32→D16
- prior_token_ids_d32 = torch.tensor(
- generated_ids[large_image_offset : large_image_offset + token_h * token_w],
- device=device,
- )
- prior_token_ids = self._upsample_token_ids(
- prior_token_ids_d32, token_h, token_w
+ prior_token_ids = self._extract_prior_token_ids(
+ generated_ids,
+ (large_image_offset, token_h, token_w),
+ device,
)
return prior_token_ids, prior_token_image_ids
+ def generate_prior_tokens_batch(
+ self,
+ prompts: list[str],
+ seeds: list[Optional[int]],
+ height: int,
+ width: int,
+ server_args: ServerArgs,
+ factor: int = 32,
+ ) -> list[torch.Tensor]:
+ device = get_local_torch_device()
+ height = (height // factor) * factor
+ width = (width // factor) * factor
+
+ input_ids = []
+ image_data = []
+ sampling_params = []
+ generation_shapes = []
+ for prompt, seed in zip(prompts, seeds, strict=True):
+ messages = [
+ {
+ "role": "user",
+ "content": [{"type": "text", "text": prompt}],
+ }
+ ]
+ inputs = self.processor.apply_chat_template(
+ messages,
+ tokenize=True,
+ target_h=height,
+ target_w=width,
+ return_dict=True,
+ return_tensors="pt",
+ )
+ image_grid_thw = inputs.get("image_grid_thw")
+ max_new_tokens, large_image_offset, token_h, token_w = (
+ self._compute_generation_params(
+ image_grid_thw=image_grid_thw,
+ is_text_to_image=True,
+ )
+ )
+ input_ids.append(inputs["input_ids"][0].tolist())
+ image_data.append([{"image_grid_thw": image_grid_thw.tolist()}])
+ sampling_params.append(
+ self._external_ar_sampling_params(max_new_tokens, seed)
+ )
+ generation_shapes.append((large_image_offset, token_h, token_w))
+
+ payload = {
+ "input_ids": input_ids,
+ "image_data": image_data,
+ "sampling_params": sampling_params,
+ }
+ data = self._request_external_ar(payload, server_args)
+ if not isinstance(data, list) or len(data) != len(prompts):
+ raise RuntimeError(
+ "GLM-Image AR batch returned an unexpected response: "
+ f"expected {len(prompts)} outputs, got "
+ f"{len(data) if isinstance(data, list) else type(data).__name__}."
+ )
+
+ prior_token_ids = []
+ for item, generation_shape in zip(data, generation_shapes, strict=True):
+ prior_token_ids.append(
+ self._extract_prior_token_ids(
+ item.get("output_ids"), generation_shape, device
+ )
+ )
+ return prior_token_ids
+
+ 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)
+
+ start_time = time.time()
+ output_counts = [_num_outputs_per_prompt(batch) for batch in batches]
+ prompts = [
+ batch.prompt
+ for batch, output_count in zip(batches, output_counts, strict=True)
+ for _ in range(output_count)
+ ]
+ seeds = [
+ _seed_for_output(batch.seed, output_idx)
+ for batch, output_count in zip(batches, output_counts, strict=True)
+ for output_idx in range(output_count)
+ ]
+ prior_token_ids = self.generate_prior_tokens_batch(
+ prompts=prompts,
+ seeds=seeds,
+ height=height,
+ width=width,
+ server_args=server_args,
+ )
+ duration = time.time() - start_time
+ logger.info(
+ "generate_prior_tokens_batch time: %.3fs for %d requests (%d outputs)",
+ duration,
+ len(batches),
+ len(prior_token_ids),
+ )
+
+ stage_name = self._active_profile_stage_name()
+ output_offset = 0
+ for batch, output_count in zip(batches, output_counts, strict=True):
+ batch.prior_token_id = torch.cat(
+ prior_token_ids[output_offset : output_offset + output_count], dim=0
+ )
+ batch.prior_token_image_ids = None
+ if batch.metrics is not None:
+ batch.metrics.record_stage(stage_name, duration)
+ output_offset += output_count
+ return batches
+
+ def iter_sequential_requests(
+ self, batch: Req, server_args: ServerArgs
+ ) -> Iterator[Req]:
+ if not server_args.pipeline_config.supports_sequential_multi_output_inference():
+ return iter((batch,))
+
+ output_count = _num_outputs_per_prompt(batch)
+ if output_count == 1:
+ return iter((batch,))
+
+ prior_token_ids = batch.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 sequential inference: "
+ f"expected {output_count} token rows, got {actual_count}."
+ )
+
+ return map(
+ lambda output_index: self._make_sequential_request(
+ batch, prior_token_ids, output_index
+ ),
+ range(output_count),
+ )
+
+ @staticmethod
+ def _make_sequential_request(
+ batch: Req, prior_token_ids: torch.Tensor, output_index: int
+ ) -> Req:
+ output_req = copy(batch)
+ output_req.sampling_params = copy(batch.sampling_params)
+ output_req.extra = dict(batch.extra)
+ output_req.condition_inputs = dict(batch.condition_inputs)
+ output_req.metrics = deepcopy(batch.metrics)
+ output_req.num_outputs_per_prompt = 1
+ output_req.seed = _seed_for_output(batch.seed, output_index)
+ output_req.seeds = None
+ output_req.generator = None
+ output_req.prior_token_id = prior_token_ids[output_index : output_index + 1]
+ if batch.request_id is not None:
+ output_req.request_id = f"{batch.request_id}:{output_index}"
+ if output_req.metrics is not None:
+ output_req.metrics.request_id = output_req.request_id
+ return output_req
+
@torch.no_grad()
def forward(
self,
@@ -423,27 +626,25 @@ class GlmImageAR(PipelineStage):
rng_devices.append(torch.npu.current_device())
rng_device_type = "npu"
- prior_token_ids = []
prior_token_image_ids = None
- for output_idx in range(num_outputs):
- output_seed = _seed_for_output(seed, output_idx)
- if output_seed is None:
- prior_token_id, output_prior_token_image_ids = (
- self.generate_prior_tokens(
- prompt=prompt,
- image=ar_condition_images,
- height=height,
- width=width,
- server_args=server_args,
- )
- )
- else:
- with torch.random.fork_rng(
- devices=rng_devices,
- enabled=True,
- device_type=rng_device_type,
- ):
- torch.manual_seed(output_seed)
+ if (
+ num_outputs > 1
+ and getattr(server_args, "srt_encoder_url", None) is not None
+ and isinstance(prompt, str)
+ and ar_condition_images is None
+ ):
+ prior_token_ids = self.generate_prior_tokens_batch(
+ prompts=[prompt] * num_outputs,
+ seeds=[_seed_for_output(seed, i) for i in range(num_outputs)],
+ height=height,
+ width=width,
+ server_args=server_args,
+ )
+ else:
+ prior_token_ids = []
+ for output_idx in range(num_outputs):
+ output_seed = _seed_for_output(seed, output_idx)
+ if output_seed is None:
prior_token_id, output_prior_token_image_ids = (
self.generate_prior_tokens(
prompt=prompt,
@@ -453,9 +654,26 @@ class GlmImageAR(PipelineStage):
server_args=server_args,
)
)
- prior_token_ids.append(prior_token_id)
- if prior_token_image_ids is None:
- prior_token_image_ids = output_prior_token_image_ids
+ else:
+ with torch.random.fork_rng(
+ devices=rng_devices,
+ enabled=True,
+ device_type=rng_device_type,
+ ):
+ torch.manual_seed(output_seed)
+ prior_token_id, output_prior_token_image_ids = (
+ self.generate_prior_tokens(
+ prompt=prompt,
+ image=ar_condition_images,
+ height=height,
+ width=width,
+ server_args=server_args,
+ seed=output_seed,
+ )
+ )
+ prior_token_ids.append(prior_token_id)
+ if prior_token_image_ids is None:
+ prior_token_image_ids = output_prior_token_image_ids
prior_token_id = torch.cat(prior_token_ids, dim=0)
prior_token_id = prior_token_id.to(device=device)
diff --git a/python/sglang/multimodal_gen/test/unit/test_glm_image_ar.py b/python/sglang/multimodal_gen/test/unit/test_glm_image_ar.py
index eabe11fe8..45a973b3a 100644
--- a/python/sglang/multimodal_gen/test/unit/test_glm_image_ar.py
+++ b/python/sglang/multimodal_gen/test/unit/test_glm_image_ar.py
@@ -32,6 +32,9 @@ class _FakeResponse:
def __init__(self, output_ids):
self._output_ids = output_ids
+ def raise_for_status(self):
+ return None
+
def json(self):
return {"output_ids": self._output_ids}