[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:
co-authored by
wuyuefeng
ronnie_zheng
parent
8279702e0b
commit
22d558b103
@@ -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
|
||||
|
||||
+61
-8
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user