[Feature] Add GLM Image usage report (#33378)

Co-authored-by: wuyuefeng <wuyuefeng@noreply.gitcode.com>
Co-authored-by: ronnie_zheng <zl19940307@163.com>
This commit is contained in:
Yuefeng Wu
2026-08-05 13:59:49 +03:00
committed by GitHub
co-authored by wuyuefeng ronnie_zheng
parent 8279702e0b
commit 22d558b103
9 changed files with 249 additions and 11 deletions
@@ -225,6 +225,8 @@ def _build_image_response_kwargs(
)
ret = add_common_data_to_response(ret, request_id=request_id, result=result)
if ret.get("usage") is not None:
ret["usage"]["image_count"] = len(save_file_path_list)
return ret
@@ -15,12 +15,26 @@ class ImageResponseData(BaseModel):
file_path: Optional[str] = None
class ImagePromptTokensDetails(BaseModel):
cached_tokens: int = 0
class ImageUsage(BaseModel):
prompt_tokens: Optional[int] = None
total_tokens: Optional[int] = None
completion_tokens: Optional[int] = None
prompt_tokens_details: Optional[ImagePromptTokensDetails] = None
reasoning_tokens: Optional[int] = 0
image_count: Optional[int] = None
class ImageResponse(BaseModel):
id: str
created: int = Field(default_factory=lambda: int(time.time()))
data: List[ImageResponseData]
peak_memory_mb: Optional[float] = None
inference_time_s: Optional[float] = None
usage: Optional[ImageUsage] = None
class ImageGenerationsRequest(BaseModel):
@@ -434,6 +434,21 @@ def add_common_data_to_response(
if result.metrics and result.metrics.total_duration_s > 0:
response["inference_time_s"] = result.metrics.total_duration_s
if result.usage is not None:
usage = dict(result.usage)
cached_tokens = usage.pop("cached_tokens", None)
enable_cache_report = getattr(
get_global_server_args(), "enable_cache_report", False
)
if (
enable_cache_report
and cached_tokens is not None
and int(cached_tokens) > 0
and usage.get("prompt_tokens_details") is None
):
usage["prompt_tokens_details"] = {"cached_tokens": int(cached_tokens)}
response["usage"] = usage
response["id"] = request_id
if result.action_pred is not None:
@@ -817,6 +817,7 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
audio=getattr(result, "audio", None),
audio_sample_rate=getattr(result, "audio_sample_rate", None),
metrics=result.metrics,
usage=getattr(result, "usage", None),
trajectory_timesteps=getattr(result, "trajectory_timesteps", None),
trajectory_latents=getattr(result, "trajectory_latents", None),
rollout_trajectory_data=getattr(result, "rollout_trajectory_data", None),
@@ -851,6 +852,14 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
if output_batch.error is not None and merged.error is None:
merged.error = output_batch.error
merged.peak_memory_mb = max(merged.peak_memory_mb, output_batch.peak_memory_mb)
if output_batch.usage is not None:
if merged.usage is None:
merged.usage = {}
for key, value in output_batch.usage.items():
if isinstance(value, int):
merged.usage[key] = int(merged.usage.get(key, 0)) + value
else:
merged.usage[key] = value
if (
merged.trajectory_timesteps is None
and output_batch.trajectory_timesteps is not None
@@ -206,6 +206,7 @@ class Req:
# stage logging
metrics: Optional[RequestMetrics] = None
usage: dict[str, Any] | None = None
# tracing context (TraceReqContext or TraceNullContext)
trace_ctx: Union[TraceReqContext, TraceNullContext] = field(
@@ -466,6 +467,7 @@ class OutputBatch:
# For ComfyUI integration: noise prediction from denoising stage
noise_pred: torch.Tensor | None = None
peak_memory_mb: float = 0.0
usage: dict[str, Any] | None = None
def drop_payload_for_warmup(self) -> None:
self.output = None
@@ -363,6 +363,7 @@ class DecodingStage(PipelineStage):
trajectory_decoded=trajectory_decoded,
metrics=batch.metrics,
noise_pred=None,
usage=batch.usage,
)
return output_batch
@@ -154,6 +154,41 @@ def _repeat_to_batch(tensor: Optional[torch.Tensor], batch_size: int):
return tensor.repeat(*repeat_shape)
def _extract_srt_usage(meta_info: dict[str, Any] | None) -> dict[str, int] | None:
if not isinstance(meta_info, dict):
return None
usage = {
"prompt_tokens": int(meta_info.get("prompt_tokens", 0) or 0),
"completion_tokens": int(meta_info.get("completion_tokens", 0) or 0),
"reasoning_tokens": int(meta_info.get("reasoning_tokens", 0) or 0),
"cached_tokens": int(meta_info.get("cached_tokens", 0) or 0),
}
usage["total_tokens"] = usage["prompt_tokens"] + usage["completion_tokens"]
return usage
def _merge_srt_usage(
total_usage: dict[str, Any] | None, usage: dict[str, int] | None
) -> dict[str, Any] | None:
if usage is None:
return total_usage
if total_usage is None:
total_usage = {}
for key, value in usage.items():
total_usage[key] = int(total_usage.get(key, 0)) + int(value)
return total_usage
def _merge_srt_usages(
usages: list[dict[str, int] | None],
) -> dict[str, Any] | None:
total_usage = None
for usage in usages:
total_usage = _merge_srt_usage(total_usage, usage)
return total_usage
class GlmImageAR(PipelineStage):
r"""
Pipeline for text-to-image generation using GLM-Image.
@@ -304,7 +339,7 @@ class GlmImageAR(PipelineStage):
image: Optional[List[PIL.Image.Image]] = None,
factor: int = 32,
seed: Optional[int] = None,
) -> Tuple[torch.Tensor, int, int]:
) -> Tuple[torch.Tensor, Optional[List[torch.Tensor]], Optional[dict[str, int]]]:
"""
Generate prior tokens using the AR (vision_language_encoder) model.
@@ -369,6 +404,7 @@ class GlmImageAR(PipelineStage):
}
data = self._request_external_ar(payload, server_args)
generated_ids = data.get("output_ids")
usage = _extract_srt_usage(data.get("meta_info"))
else:
if image is not None:
source_grids = image_grid_thw[:-1]
@@ -403,6 +439,7 @@ class GlmImageAR(PipelineStage):
)
input_len = inputs["input_ids"].shape[-1]
generated_ids = outputs[0][input_len:]
usage = None
prior_token_ids = self._extract_prior_token_ids(
generated_ids,
@@ -410,7 +447,7 @@ class GlmImageAR(PipelineStage):
device,
)
return prior_token_ids, prior_token_image_ids
return prior_token_ids, prior_token_image_ids, usage
def generate_prior_tokens_batch(
self,
@@ -420,7 +457,7 @@ class GlmImageAR(PipelineStage):
width: int,
server_args: ServerArgs,
factor: int = 32,
) -> list[torch.Tensor]:
) -> tuple[list[torch.Tensor], list[dict[str, int] | None]]:
device = get_local_torch_device()
height = (height // factor) * factor
width = (width // factor) * factor
@@ -472,13 +509,15 @@ class GlmImageAR(PipelineStage):
)
prior_token_ids = []
usages = []
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
usages.append(_extract_srt_usage(item.get("meta_info")))
return prior_token_ids, usages
def run_grouped_requests(
self,
@@ -513,7 +552,7 @@ class GlmImageAR(PipelineStage):
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(
prior_token_ids, usages = self.generate_prior_tokens_batch(
prompts=prompts,
seeds=seeds,
height=height,
@@ -535,6 +574,11 @@ class GlmImageAR(PipelineStage):
prior_token_ids[output_offset : output_offset + output_count], dim=0
)
batch.prior_token_image_ids = None
output_usages = usages[output_offset : output_offset + output_count]
batch.extra["usage_by_output"] = output_usages
usage = _merge_srt_usages(output_usages)
if usage is not None:
batch.usage = usage
if batch.metrics is not None:
batch.metrics.record_stage(stage_name, duration)
output_offset += output_count
@@ -585,6 +629,9 @@ class GlmImageAR(PipelineStage):
output_req.seeds = None
output_req.generator = None
output_req.prior_token_id = prior_token_ids[output_index : output_index + 1]
usage_by_output = 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]
if batch.request_id is not None:
output_req.request_id = f"{batch.request_id}:{output_index}"
if output_req.metrics is not None:
@@ -633,7 +680,7 @@ class GlmImageAR(PipelineStage):
and isinstance(prompt, str)
and ar_condition_images is None
):
prior_token_ids = self.generate_prior_tokens_batch(
prior_token_ids, output_usages = self.generate_prior_tokens_batch(
prompts=[prompt] * num_outputs,
seeds=[_seed_for_output(seed, i) for i in range(num_outputs)],
height=height,
@@ -642,10 +689,11 @@ class GlmImageAR(PipelineStage):
)
else:
prior_token_ids = []
output_usages = []
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 = (
prior_token_id, output_prior_token_image_ids, output_usage = (
self.generate_prior_tokens(
prompt=prompt,
image=ar_condition_images,
@@ -661,7 +709,7 @@ class GlmImageAR(PipelineStage):
device_type=rng_device_type,
):
torch.manual_seed(output_seed)
prior_token_id, output_prior_token_image_ids = (
prior_token_id, output_prior_token_image_ids, output_usage = (
self.generate_prior_tokens(
prompt=prompt,
image=ar_condition_images,
@@ -672,6 +720,7 @@ class GlmImageAR(PipelineStage):
)
)
prior_token_ids.append(prior_token_id)
output_usages.append(output_usage)
if prior_token_image_ids is None:
prior_token_image_ids = output_prior_token_image_ids
@@ -684,6 +733,10 @@ class GlmImageAR(PipelineStage):
batch.prior_token_image_ids = prior_token_image_ids
batch.height = height
batch.width = width
batch.extra["usage_by_output"] = output_usages
usage = _merge_srt_usages(output_usages)
if usage is not None:
batch.usage = usage
return batch
@@ -429,6 +429,7 @@ class ServerArgs(DisaggServerArgsMixin):
log_requests_format: str = "text"
log_requests_target: Optional[List[str]] = None
uvicorn_access_log_exclude_prefixes: list[str] = field(default_factory=list)
enable_cache_report: bool = False
# Tracing
enable_trace: bool = False
@@ -1982,6 +1983,12 @@ class ServerArgs(DisaggServerArgsMixin):
"Defaults to empty (disabled). "
"Example: --uvicorn-access-log-exclude-prefixes /metrics /health",
)
parser.add_argument(
"--enable-cache-report",
action="store_true",
default=ServerArgs.enable_cache_report,
help="Return number of cached tokens in usage.prompt_tokens_details for each OpenAI-compatible request.",
)
parser.add_argument(
"--backend",
type=str,
@@ -4,6 +4,10 @@ from unittest.mock import patch
import torch
from sglang.multimodal_gen.runtime.entrypoints.openai.image_api import (
_build_image_response_kwargs,
)
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.glm_image import (
GlmImageAR,
)
@@ -29,14 +33,18 @@ class _FakeProcessor:
class _FakeResponse:
def __init__(self, output_ids):
def __init__(self, output_ids, meta_info=None):
self._output_ids = output_ids
self._meta_info = meta_info
def raise_for_status(self):
return None
def json(self):
return {"output_ids": self._output_ids}
data = {"output_ids": self._output_ids}
if self._meta_info is not None:
data["meta_info"] = self._meta_info
return data
class TestGlmImageARSrtBackend(unittest.TestCase):
@@ -63,7 +71,7 @@ class TestGlmImageARSrtBackend(unittest.TestCase):
mock_post.return_value = _FakeResponse(list(range(1025)))
stage = GlmImageAR(processor=_FakeProcessor(), vision_language_encoder=None)
prior_token_ids, _ = stage.generate_prior_tokens(
prior_token_ids, _, usage = stage.generate_prior_tokens(
prompt="A simple product sketch",
height=1024,
width=1024,
@@ -74,6 +82,84 @@ class TestGlmImageARSrtBackend(unittest.TestCase):
self.assertTrue(payload["sampling_params"]["ignore_eos"])
self.assertEqual(payload["sampling_params"]["max_new_tokens"], 1025)
self.assertEqual(prior_token_ids.shape, (1, 4096))
self.assertIsNone(usage)
@patch(
"sglang.multimodal_gen.runtime.pipelines_core.stages."
"model_specific_stages.glm_image.get_local_torch_device",
return_value=torch.device("cpu"),
)
@patch(
"sglang.multimodal_gen.runtime.pipelines_core.stages."
"model_specific_stages.glm_image.requests.post"
)
def test_srt_ar_extracts_usage_from_meta_info(self, mock_post, _mock_device):
set_global_server_args(self._server_args())
mock_post.return_value = _FakeResponse(
list(range(1025)),
meta_info={
"prompt_tokens": 13,
"completion_tokens": 25,
"reasoning_tokens": 0,
"cached_tokens": 5,
},
)
stage = GlmImageAR(processor=_FakeProcessor(), vision_language_encoder=None)
_, _, usage = stage.generate_prior_tokens(
prompt="A simple product sketch",
height=1024,
width=1024,
server_args=self._server_args(),
)
self.assertEqual(
usage,
{
"prompt_tokens": 13,
"completion_tokens": 25,
"reasoning_tokens": 0,
"cached_tokens": 5,
"total_tokens": 38,
},
)
@patch(
"sglang.multimodal_gen.runtime.pipelines_core.stages."
"model_specific_stages.glm_image.get_local_torch_device",
return_value=torch.device("cpu"),
)
@patch(
"sglang.multimodal_gen.runtime.pipelines_core.stages."
"model_specific_stages.glm_image.requests.post"
)
def test_srt_ar_forward_aggregates_usage(self, mock_post, _mock_device):
set_global_server_args(self._server_args())
mock_post.side_effect = [
_FakeResponse(
list(range(1025)),
meta_info={"prompt_tokens": 13, "completion_tokens": 25},
),
_FakeResponse(
list(range(1025)),
meta_info={"prompt_tokens": 13, "completion_tokens": 25},
),
]
stage = GlmImageAR(processor=_FakeProcessor(), vision_language_encoder=None)
batch = SimpleNamespace(
prompt="A simple product sketch",
height=1025,
width=1001,
image_path=None,
num_outputs_per_prompt=2,
seed=None,
)
batch = stage.forward(batch, self._server_args())
self.assertEqual(batch.usage["prompt_tokens"], 26)
self.assertEqual(batch.usage["completion_tokens"], 50)
self.assertEqual(batch.usage["total_tokens"], 76)
@patch(
"sglang.multimodal_gen.runtime.pipelines_core.stages."
@@ -100,6 +186,55 @@ class TestGlmImageARSrtBackend(unittest.TestCase):
server_args=self._server_args(),
)
def test_image_response_adds_image_count_to_usage(self):
set_global_server_args(SimpleNamespace(enable_cache_report=False))
response = _build_image_response_kwargs(
["/tmp/glm-image-0.jpg", "/tmp/glm-image-1.jpg"],
"b64_json",
"A simple product sketch",
"req-0",
OutputBatch(
usage={
"prompt_tokens": 13,
"completion_tokens": 25,
"total_tokens": 38,
"reasoning_tokens": 0,
"cached_tokens": 5,
}
),
b64_list=["aGVsbG8=", "d29ybGQ="],
is_persistent=False,
)
self.assertEqual(response["usage"]["image_count"], 2)
self.assertNotIn("prompt_tokens_details", response["usage"])
self.assertNotIn("cached_tokens", response["usage"])
def test_image_response_reports_cached_tokens_when_cache_report_enabled(self):
set_global_server_args(SimpleNamespace(enable_cache_report=True))
response = _build_image_response_kwargs(
["/tmp/glm-image-0.jpg"],
"b64_json",
"A simple product sketch",
"req-0",
OutputBatch(
usage={
"prompt_tokens": 13,
"completion_tokens": 25,
"total_tokens": 38,
"reasoning_tokens": 0,
"cached_tokens": 5,
}
),
b64_list=["aGVsbG8="],
is_persistent=False,
)
self.assertEqual(
response["usage"]["prompt_tokens_details"], {"cached_tokens": 5}
)
self.assertNotIn("cached_tokens", response["usage"])
if __name__ == "__main__":
unittest.main()