diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/image_api.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/image_api.py index 2f6e58e7e..e523ab3aa 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/image_api.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/image_api.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py index 0517ee733..ff4fed8f1 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py @@ -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): diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py index 880fe0a8d..6b5ef964d 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py @@ -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: diff --git a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py index e25a8921a..19fe3b1a0 100644 --- a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py +++ b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py @@ -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 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 a868d2d9b..bc36f0d96 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py @@ -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 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 8d48d967f..b94e2334d 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py @@ -363,6 +363,7 @@ class DecodingStage(PipelineStage): trajectory_decoded=trajectory_decoded, metrics=batch.metrics, noise_pred=None, + usage=batch.usage, ) return output_batch 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 b59fb67cb..01e4d84ca 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 @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/server_args/server_args.py b/python/sglang/multimodal_gen/runtime/server_args/server_args.py index b27f445a0..9e54713d5 100644 --- a/python/sglang/multimodal_gen/runtime/server_args/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args/server_args.py @@ -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, 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 45a973b3a..5c1f89545 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 @@ -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()