[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)
|
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
|
return ret
|
||||||
|
|
||||||
|
|||||||
@@ -15,12 +15,26 @@ class ImageResponseData(BaseModel):
|
|||||||
file_path: Optional[str] = None
|
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):
|
class ImageResponse(BaseModel):
|
||||||
id: str
|
id: str
|
||||||
created: int = Field(default_factory=lambda: int(time.time()))
|
created: int = Field(default_factory=lambda: int(time.time()))
|
||||||
data: List[ImageResponseData]
|
data: List[ImageResponseData]
|
||||||
peak_memory_mb: Optional[float] = None
|
peak_memory_mb: Optional[float] = None
|
||||||
inference_time_s: Optional[float] = None
|
inference_time_s: Optional[float] = None
|
||||||
|
usage: Optional[ImageUsage] = None
|
||||||
|
|
||||||
|
|
||||||
class ImageGenerationsRequest(BaseModel):
|
class ImageGenerationsRequest(BaseModel):
|
||||||
|
|||||||
@@ -434,6 +434,21 @@ def add_common_data_to_response(
|
|||||||
if result.metrics and result.metrics.total_duration_s > 0:
|
if result.metrics and result.metrics.total_duration_s > 0:
|
||||||
response["inference_time_s"] = result.metrics.total_duration_s
|
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
|
response["id"] = request_id
|
||||||
|
|
||||||
if result.action_pred is not None:
|
if result.action_pred is not None:
|
||||||
|
|||||||
@@ -817,6 +817,7 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
|
|||||||
audio=getattr(result, "audio", None),
|
audio=getattr(result, "audio", None),
|
||||||
audio_sample_rate=getattr(result, "audio_sample_rate", None),
|
audio_sample_rate=getattr(result, "audio_sample_rate", None),
|
||||||
metrics=result.metrics,
|
metrics=result.metrics,
|
||||||
|
usage=getattr(result, "usage", None),
|
||||||
trajectory_timesteps=getattr(result, "trajectory_timesteps", None),
|
trajectory_timesteps=getattr(result, "trajectory_timesteps", None),
|
||||||
trajectory_latents=getattr(result, "trajectory_latents", None),
|
trajectory_latents=getattr(result, "trajectory_latents", None),
|
||||||
rollout_trajectory_data=getattr(result, "rollout_trajectory_data", 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:
|
if output_batch.error is not None and merged.error is None:
|
||||||
merged.error = output_batch.error
|
merged.error = output_batch.error
|
||||||
merged.peak_memory_mb = max(merged.peak_memory_mb, output_batch.peak_memory_mb)
|
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 (
|
if (
|
||||||
merged.trajectory_timesteps is None
|
merged.trajectory_timesteps is None
|
||||||
and output_batch.trajectory_timesteps is not None
|
and output_batch.trajectory_timesteps is not None
|
||||||
|
|||||||
@@ -206,6 +206,7 @@ class Req:
|
|||||||
|
|
||||||
# stage logging
|
# stage logging
|
||||||
metrics: Optional[RequestMetrics] = None
|
metrics: Optional[RequestMetrics] = None
|
||||||
|
usage: dict[str, Any] | None = None
|
||||||
|
|
||||||
# tracing context (TraceReqContext or TraceNullContext)
|
# tracing context (TraceReqContext or TraceNullContext)
|
||||||
trace_ctx: Union[TraceReqContext, TraceNullContext] = field(
|
trace_ctx: Union[TraceReqContext, TraceNullContext] = field(
|
||||||
@@ -466,6 +467,7 @@ class OutputBatch:
|
|||||||
# For ComfyUI integration: noise prediction from denoising stage
|
# For ComfyUI integration: noise prediction from denoising stage
|
||||||
noise_pred: torch.Tensor | None = None
|
noise_pred: torch.Tensor | None = None
|
||||||
peak_memory_mb: float = 0.0
|
peak_memory_mb: float = 0.0
|
||||||
|
usage: dict[str, Any] | None = None
|
||||||
|
|
||||||
def drop_payload_for_warmup(self) -> None:
|
def drop_payload_for_warmup(self) -> None:
|
||||||
self.output = None
|
self.output = None
|
||||||
|
|||||||
@@ -363,6 +363,7 @@ class DecodingStage(PipelineStage):
|
|||||||
trajectory_decoded=trajectory_decoded,
|
trajectory_decoded=trajectory_decoded,
|
||||||
metrics=batch.metrics,
|
metrics=batch.metrics,
|
||||||
noise_pred=None,
|
noise_pred=None,
|
||||||
|
usage=batch.usage,
|
||||||
)
|
)
|
||||||
|
|
||||||
return output_batch
|
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)
|
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):
|
class GlmImageAR(PipelineStage):
|
||||||
r"""
|
r"""
|
||||||
Pipeline for text-to-image generation using GLM-Image.
|
Pipeline for text-to-image generation using GLM-Image.
|
||||||
@@ -304,7 +339,7 @@ class GlmImageAR(PipelineStage):
|
|||||||
image: Optional[List[PIL.Image.Image]] = None,
|
image: Optional[List[PIL.Image.Image]] = None,
|
||||||
factor: int = 32,
|
factor: int = 32,
|
||||||
seed: Optional[int] = None,
|
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.
|
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)
|
data = self._request_external_ar(payload, server_args)
|
||||||
generated_ids = data.get("output_ids")
|
generated_ids = data.get("output_ids")
|
||||||
|
usage = _extract_srt_usage(data.get("meta_info"))
|
||||||
else:
|
else:
|
||||||
if image is not None:
|
if image is not None:
|
||||||
source_grids = image_grid_thw[:-1]
|
source_grids = image_grid_thw[:-1]
|
||||||
@@ -403,6 +439,7 @@ class GlmImageAR(PipelineStage):
|
|||||||
)
|
)
|
||||||
input_len = inputs["input_ids"].shape[-1]
|
input_len = inputs["input_ids"].shape[-1]
|
||||||
generated_ids = outputs[0][input_len:]
|
generated_ids = outputs[0][input_len:]
|
||||||
|
usage = None
|
||||||
|
|
||||||
prior_token_ids = self._extract_prior_token_ids(
|
prior_token_ids = self._extract_prior_token_ids(
|
||||||
generated_ids,
|
generated_ids,
|
||||||
@@ -410,7 +447,7 @@ class GlmImageAR(PipelineStage):
|
|||||||
device,
|
device,
|
||||||
)
|
)
|
||||||
|
|
||||||
return prior_token_ids, prior_token_image_ids
|
return prior_token_ids, prior_token_image_ids, usage
|
||||||
|
|
||||||
def generate_prior_tokens_batch(
|
def generate_prior_tokens_batch(
|
||||||
self,
|
self,
|
||||||
@@ -420,7 +457,7 @@ class GlmImageAR(PipelineStage):
|
|||||||
width: int,
|
width: int,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
factor: int = 32,
|
factor: int = 32,
|
||||||
) -> list[torch.Tensor]:
|
) -> tuple[list[torch.Tensor], list[dict[str, int] | None]]:
|
||||||
device = get_local_torch_device()
|
device = get_local_torch_device()
|
||||||
height = (height // factor) * factor
|
height = (height // factor) * factor
|
||||||
width = (width // factor) * factor
|
width = (width // factor) * factor
|
||||||
@@ -472,13 +509,15 @@ class GlmImageAR(PipelineStage):
|
|||||||
)
|
)
|
||||||
|
|
||||||
prior_token_ids = []
|
prior_token_ids = []
|
||||||
|
usages = []
|
||||||
for item, generation_shape in zip(data, generation_shapes, strict=True):
|
for item, generation_shape in zip(data, generation_shapes, strict=True):
|
||||||
prior_token_ids.append(
|
prior_token_ids.append(
|
||||||
self._extract_prior_token_ids(
|
self._extract_prior_token_ids(
|
||||||
item.get("output_ids"), generation_shape, device
|
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(
|
def run_grouped_requests(
|
||||||
self,
|
self,
|
||||||
@@ -513,7 +552,7 @@ class GlmImageAR(PipelineStage):
|
|||||||
for batch, output_count in zip(batches, output_counts, strict=True)
|
for batch, output_count in zip(batches, output_counts, strict=True)
|
||||||
for output_idx in range(output_count)
|
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,
|
prompts=prompts,
|
||||||
seeds=seeds,
|
seeds=seeds,
|
||||||
height=height,
|
height=height,
|
||||||
@@ -535,6 +574,11 @@ class GlmImageAR(PipelineStage):
|
|||||||
prior_token_ids[output_offset : output_offset + output_count], dim=0
|
prior_token_ids[output_offset : output_offset + output_count], dim=0
|
||||||
)
|
)
|
||||||
batch.prior_token_image_ids = None
|
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:
|
if batch.metrics is not None:
|
||||||
batch.metrics.record_stage(stage_name, duration)
|
batch.metrics.record_stage(stage_name, duration)
|
||||||
output_offset += output_count
|
output_offset += output_count
|
||||||
@@ -585,6 +629,9 @@ class GlmImageAR(PipelineStage):
|
|||||||
output_req.seeds = None
|
output_req.seeds = None
|
||||||
output_req.generator = None
|
output_req.generator = None
|
||||||
output_req.prior_token_id = prior_token_ids[output_index : output_index + 1]
|
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:
|
if batch.request_id is not None:
|
||||||
output_req.request_id = f"{batch.request_id}:{output_index}"
|
output_req.request_id = f"{batch.request_id}:{output_index}"
|
||||||
if output_req.metrics is not None:
|
if output_req.metrics is not None:
|
||||||
@@ -633,7 +680,7 @@ class GlmImageAR(PipelineStage):
|
|||||||
and isinstance(prompt, str)
|
and isinstance(prompt, str)
|
||||||
and ar_condition_images is None
|
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,
|
prompts=[prompt] * num_outputs,
|
||||||
seeds=[_seed_for_output(seed, i) for i in range(num_outputs)],
|
seeds=[_seed_for_output(seed, i) for i in range(num_outputs)],
|
||||||
height=height,
|
height=height,
|
||||||
@@ -642,10 +689,11 @@ class GlmImageAR(PipelineStage):
|
|||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
prior_token_ids = []
|
prior_token_ids = []
|
||||||
|
output_usages = []
|
||||||
for output_idx in range(num_outputs):
|
for output_idx in range(num_outputs):
|
||||||
output_seed = _seed_for_output(seed, output_idx)
|
output_seed = _seed_for_output(seed, output_idx)
|
||||||
if output_seed is None:
|
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(
|
self.generate_prior_tokens(
|
||||||
prompt=prompt,
|
prompt=prompt,
|
||||||
image=ar_condition_images,
|
image=ar_condition_images,
|
||||||
@@ -661,7 +709,7 @@ class GlmImageAR(PipelineStage):
|
|||||||
device_type=rng_device_type,
|
device_type=rng_device_type,
|
||||||
):
|
):
|
||||||
torch.manual_seed(output_seed)
|
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(
|
self.generate_prior_tokens(
|
||||||
prompt=prompt,
|
prompt=prompt,
|
||||||
image=ar_condition_images,
|
image=ar_condition_images,
|
||||||
@@ -672,6 +720,7 @@ class GlmImageAR(PipelineStage):
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
prior_token_ids.append(prior_token_id)
|
prior_token_ids.append(prior_token_id)
|
||||||
|
output_usages.append(output_usage)
|
||||||
if prior_token_image_ids is None:
|
if prior_token_image_ids is None:
|
||||||
prior_token_image_ids = output_prior_token_image_ids
|
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.prior_token_image_ids = prior_token_image_ids
|
||||||
batch.height = height
|
batch.height = height
|
||||||
batch.width = width
|
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
|
return batch
|
||||||
|
|
||||||
|
|||||||
@@ -429,6 +429,7 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
log_requests_format: str = "text"
|
log_requests_format: str = "text"
|
||||||
log_requests_target: Optional[List[str]] = None
|
log_requests_target: Optional[List[str]] = None
|
||||||
uvicorn_access_log_exclude_prefixes: list[str] = field(default_factory=list)
|
uvicorn_access_log_exclude_prefixes: list[str] = field(default_factory=list)
|
||||||
|
enable_cache_report: bool = False
|
||||||
|
|
||||||
# Tracing
|
# Tracing
|
||||||
enable_trace: bool = False
|
enable_trace: bool = False
|
||||||
@@ -1982,6 +1983,12 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
"Defaults to empty (disabled). "
|
"Defaults to empty (disabled). "
|
||||||
"Example: --uvicorn-access-log-exclude-prefixes /metrics /health",
|
"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(
|
parser.add_argument(
|
||||||
"--backend",
|
"--backend",
|
||||||
type=str,
|
type=str,
|
||||||
|
|||||||
@@ -4,6 +4,10 @@ from unittest.mock import patch
|
|||||||
|
|
||||||
import torch
|
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 (
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.glm_image import (
|
||||||
GlmImageAR,
|
GlmImageAR,
|
||||||
)
|
)
|
||||||
@@ -29,14 +33,18 @@ class _FakeProcessor:
|
|||||||
|
|
||||||
|
|
||||||
class _FakeResponse:
|
class _FakeResponse:
|
||||||
def __init__(self, output_ids):
|
def __init__(self, output_ids, meta_info=None):
|
||||||
self._output_ids = output_ids
|
self._output_ids = output_ids
|
||||||
|
self._meta_info = meta_info
|
||||||
|
|
||||||
def raise_for_status(self):
|
def raise_for_status(self):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def json(self):
|
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):
|
class TestGlmImageARSrtBackend(unittest.TestCase):
|
||||||
@@ -63,7 +71,7 @@ class TestGlmImageARSrtBackend(unittest.TestCase):
|
|||||||
mock_post.return_value = _FakeResponse(list(range(1025)))
|
mock_post.return_value = _FakeResponse(list(range(1025)))
|
||||||
stage = GlmImageAR(processor=_FakeProcessor(), vision_language_encoder=None)
|
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",
|
prompt="A simple product sketch",
|
||||||
height=1024,
|
height=1024,
|
||||||
width=1024,
|
width=1024,
|
||||||
@@ -74,6 +82,84 @@ class TestGlmImageARSrtBackend(unittest.TestCase):
|
|||||||
self.assertTrue(payload["sampling_params"]["ignore_eos"])
|
self.assertTrue(payload["sampling_params"]["ignore_eos"])
|
||||||
self.assertEqual(payload["sampling_params"]["max_new_tokens"], 1025)
|
self.assertEqual(payload["sampling_params"]["max_new_tokens"], 1025)
|
||||||
self.assertEqual(prior_token_ids.shape, (1, 4096))
|
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(
|
@patch(
|
||||||
"sglang.multimodal_gen.runtime.pipelines_core.stages."
|
"sglang.multimodal_gen.runtime.pipelines_core.stages."
|
||||||
@@ -100,6 +186,55 @@ class TestGlmImageARSrtBackend(unittest.TestCase):
|
|||||||
server_args=self._server_args(),
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user