feat: report multimodal (image/audio/video) token counts in usage.prompt_tokens_details (#27122)

This commit is contained in:
Zyann
2026-06-15 11:04:10 +08:00
committed by GitHub
parent c0dfe4c8ec
commit 37505eca27
11 changed files with 178 additions and 2 deletions
@@ -1546,6 +1546,11 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
decode_req.req.cached_tokens_device = cached_tokens[1].item()
decode_req.req.cached_tokens_host = cached_tokens[2].item()
decode_req.req.cached_tokens_storage = cached_tokens[3].item()
# Multimodal prompt token counts packed into cached_tokens slots 4-6
# by the prefill node (see MetadataBuffers.set_buf).
decode_req.req.mm_image_tokens = cached_tokens[4].item()
decode_req.req.mm_audio_tokens = cached_tokens[5].item()
decode_req.req.mm_video_tokens = cached_tokens[6].item()
if not self.spec_algorithm.is_none():
decode_req.req.output_topk_p = output_topk_p
decode_req.req.output_topk_index = output_topk_index
+14
View File
@@ -312,10 +312,24 @@ class MetadataBuffers:
def set_buf(self, req: Req):
self.output_ids[req.metadata_buffer_index][0] = req.output_ids[0]
# The cached_tokens buffer is (size, 16); slots 0-3 hold cached token
# counts and slots 4-6 are reused for multimodal prompt token counts
# (slots 7-15 remain spare). This avoids adding new RDMA buffers.
# Slot map: 0=cached 1=device 2=host 3=storage 4=image 5=audio 6=video.
self.cached_tokens[req.metadata_buffer_index][0] = req.cached_tokens
self.cached_tokens[req.metadata_buffer_index][1] = req.cached_tokens_device
self.cached_tokens[req.metadata_buffer_index][2] = req.cached_tokens_host
self.cached_tokens[req.metadata_buffer_index][3] = req.cached_tokens_storage
# Compute multimodal prompt token counts on the prefill node so decode
# can report them in usage.
if req.multimodal_inputs:
image_t, audio_t, video_t = req.multimodal_inputs.compute_mm_token_counts()
else:
image_t = audio_t = video_t = 0
self.cached_tokens[req.metadata_buffer_index][4] = image_t
self.cached_tokens[req.metadata_buffer_index][5] = audio_t
self.cached_tokens[req.metadata_buffer_index][6] = video_t
if req.return_logprob:
if req.logprob.output_token_logprobs_val: # not none or empty list
self.output_token_logprobs_val[req.metadata_buffer_index][0] = (
@@ -177,6 +177,20 @@ class PromptTokensDetails(BaseModel):
"""Details about prompt tokens."""
cached_tokens: int = 0
# Multimodal prompt token counts (only populated when present in the prompt)
image_tokens: Optional[int] = None
audio_tokens: Optional[int] = None
video_tokens: Optional[int] = None
@model_serializer(mode="wrap")
def _serialize(self, handler):
data = handler(self)
# Drop multimodal fields when absent so text-only/cache-only responses
# keep the original {"cached_tokens": N} shape.
for key in ("image_tokens", "audio_tokens", "video_tokens"):
if data.get(key) is None:
data.pop(key, None)
return data
class UsageInfo(BaseModel):
@@ -974,6 +974,9 @@ class OpenAIServingChat(OpenAIServingBase):
hidden_states = {}
routed_experts = {}
cached_tokens_details = {}
image_tokens = {}
audio_tokens = {}
video_tokens = {}
stream_started = False
try:
@@ -1000,6 +1003,9 @@ class OpenAIServingChat(OpenAIServingBase):
cached_tokens_details[index] = content["meta_info"].get(
"cached_tokens_details", None
)
image_tokens[index] = content["meta_info"].get("image_tokens", 0)
audio_tokens[index] = content["meta_info"].get("audio_tokens", 0)
video_tokens[index] = content["meta_info"].get("video_tokens", 0)
# Handle logprobs
choice_logprobs = None
@@ -1142,6 +1148,17 @@ class OpenAIServingChat(OpenAIServingBase):
# Additional usage chunk
if include_usage:
# Multimodal tokens are per-prompt (input side), so aggregate
# once per prompt (first choice), matching prompt/cached semantics.
total_image_tokens = sum(
tok for idx, tok in image_tokens.items() if idx % request.n == 0
)
total_audio_tokens = sum(
tok for idx, tok in audio_tokens.items() if idx % request.n == 0
)
total_video_tokens = sum(
tok for idx, tok in video_tokens.items() if idx % request.n == 0
)
usage = UsageProcessor.calculate_streaming_usage(
prompt_tokens,
reasoning_tokens,
@@ -1149,6 +1166,9 @@ class OpenAIServingChat(OpenAIServingBase):
cached_tokens=cached_tokens,
n_choices=request.n,
enable_cache_report=self.tokenizer_manager.server_args.enable_cache_report,
image_tokens=total_image_tokens,
audio_tokens=total_audio_tokens,
video_tokens=total_video_tokens,
)
usage_chunk = ChatCompletionStreamResponse(
id=content["meta_info"]["id"],
@@ -1304,11 +1324,27 @@ class OpenAIServingChat(OpenAIServingBase):
)
choices.append(choice_data)
# Calculate usage
# Calculate usage. Multimodal tokens are per-prompt (input side), so
# aggregate once per prompt (stride by n), matching prompt/cached semantics.
image_tokens = sum(
ret[i]["meta_info"].get("image_tokens", 0)
for i in range(0, len(ret), request.n)
)
audio_tokens = sum(
ret[i]["meta_info"].get("audio_tokens", 0)
for i in range(0, len(ret), request.n)
)
video_tokens = sum(
ret[i]["meta_info"].get("video_tokens", 0)
for i in range(0, len(ret), request.n)
)
usage = UsageProcessor.calculate_response_usage(
ret,
n_choices=request.n,
enable_cache_report=self.tokenizer_manager.server_args.enable_cache_report,
image_tokens=image_tokens,
audio_tokens=audio_tokens,
video_tokens=video_tokens,
)
return ChatCompletionResponse(
@@ -19,6 +19,9 @@ class UsageProcessor:
responses: List[Dict[str, Any]],
n_choices: int = 1,
enable_cache_report: bool = False,
image_tokens: int = 0,
audio_tokens: int = 0,
video_tokens: int = 0,
) -> UsageInfo:
completion_tokens = sum(
r["meta_info"].get("completion_tokens", 0) for r in responses
@@ -46,6 +49,9 @@ class UsageProcessor:
reasoning_tokens=reasoning_tokens,
completion_tokens=completion_tokens,
cached_tokens=cached_details,
image_tokens=image_tokens,
audio_tokens=audio_tokens,
video_tokens=video_tokens,
)
@staticmethod
@@ -56,6 +62,9 @@ class UsageProcessor:
cached_tokens: Mapping[int, int],
n_choices: int,
enable_cache_report: bool = False,
image_tokens: int = 0,
audio_tokens: int = 0,
video_tokens: int = 0,
) -> UsageInfo:
# index % n_choices == 0 marks the first choice of a prompt
total_prompt_tokens = sum(
@@ -77,6 +86,9 @@ class UsageProcessor:
reasoning_tokens=total_reasoning_tokens,
completion_tokens=total_completion_tokens,
cached_tokens=cached_details,
image_tokens=image_tokens,
audio_tokens=audio_tokens,
video_tokens=video_tokens,
)
@staticmethod
@@ -85,12 +97,30 @@ class UsageProcessor:
completion_tokens: int,
reasoning_tokens: Optional[int] = 0,
cached_tokens: Optional[PromptTokensDetails] = None,
image_tokens: int = 0,
audio_tokens: int = 0,
video_tokens: int = 0,
) -> UsageInfo:
"""Calculate token usage information"""
# `cached_tokens` is already a PromptTokensDetails (or None) carrying the
# cached count. Attach multimodal counts to the same object, creating one
# only when there is something to report so plain-text requests keep
# prompt_tokens_details=None (backward compatible).
details = cached_tokens
if image_tokens or audio_tokens or video_tokens:
if details is None:
details = PromptTokensDetails()
if image_tokens:
details.image_tokens = image_tokens
if audio_tokens:
details.audio_tokens = audio_tokens
if video_tokens:
details.video_tokens = video_tokens
return UsageInfo(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=prompt_tokens + completion_tokens,
prompt_tokens_details=cached_tokens,
prompt_tokens_details=details,
reasoning_tokens=reasoning_tokens,
)
@@ -421,6 +421,9 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
completion_tokens=recv_obj.completion_tokens,
cached_tokens=recv_obj.cached_tokens,
cached_tokens_details=recv_obj.cached_tokens_details,
image_tokens=recv_obj.image_tokens,
audio_tokens=recv_obj.audio_tokens,
video_tokens=recv_obj.video_tokens,
spec_verify_ct=recv_obj.spec_verify_ct,
spec_num_correct_drafts=recv_obj.spec_num_correct_drafts,
spec_correct_drafts_histogram=recv_obj.spec_correct_drafts_histogram,
+10
View File
@@ -1191,6 +1191,11 @@ class BatchTokenIDOutput(BaseBatchReq, SpeculativeDecodingMetricsMixin):
# For observability
time_stats: Optional[List[SchedulerReqTimeStats]] = None
# Multimodal prompt token counts (image/audio/video). None when not applicable.
image_tokens: Optional[List[int]] = None
audio_tokens: Optional[List[int]] = None
video_tokens: Optional[List[int]] = None
@dataclass
class BatchStrOutput(BaseBatchReq, SpeculativeDecodingMetricsMixin):
@@ -1257,6 +1262,11 @@ class BatchStrOutput(BaseBatchReq, SpeculativeDecodingMetricsMixin):
# For observability
time_stats: Optional[List[SchedulerReqTimeStats]] = None
# Multimodal prompt token counts (image/audio/video). None when not applicable.
image_tokens: Optional[List[int]] = None
audio_tokens: Optional[List[int]] = None
video_tokens: Optional[List[int]] = None
@dataclass
class BatchEmbeddingOutput(BaseBatchReq):
@@ -163,6 +163,9 @@ def _handle_output_by_index(output, i):
cached_tokens_details=_extract_field_by_index(
output, "cached_tokens_details", i
),
image_tokens=_extract_field_by_index(output, "image_tokens", i),
audio_tokens=_extract_field_by_index(output, "audio_tokens", i),
video_tokens=_extract_field_by_index(output, "video_tokens", i),
input_token_logprobs_val=_extract_field_by_index(
output, "input_token_logprobs_val", i, check_length=False
),
@@ -253,6 +256,9 @@ def _handle_output_by_index(output, i):
cached_tokens_details=_extract_field_by_index(
output, "cached_tokens_details", i
),
image_tokens=_extract_field_by_index(output, "image_tokens", i),
audio_tokens=_extract_field_by_index(output, "audio_tokens", i),
video_tokens=_extract_field_by_index(output, "video_tokens", i),
input_token_logprobs_val=_extract_field_by_index(
output, "input_token_logprobs_val", i, check_length=False
),
@@ -579,6 +579,25 @@ class MultimodalInputs:
def contains_mm_input(self) -> bool:
return any(True for item in self.mm_items if item.is_valid())
def compute_mm_token_counts(self) -> Tuple[int, int, int]:
"""Count prompt tokens consumed by each modality (image, audio, video).
A modality's token count is the total span covered by its items'
offsets. Returns a (image_tokens, audio_tokens, video_tokens) tuple.
"""
image_tokens = audio_tokens = video_tokens = 0
for item in self.mm_items:
if not item.offsets:
continue
num_tokens = sum(end - start + 1 for start, end in item.offsets)
if item.is_image():
image_tokens += num_tokens
elif item.is_audio():
audio_tokens += num_tokens
elif item.is_video():
video_tokens += num_tokens
return image_tokens, audio_tokens, video_tokens
def merge(self, other: MultimodalInputs):
"""
merge image inputs when requests are being merged
@@ -811,6 +830,11 @@ class Req(ReqDllmMixin):
# For multimodal inputs
self.multimodal_inputs: Optional[MultimodalInputs] = None
# Pre-computed multimodal prompt token counts; populated on the prefill
# node and transferred to decode via the metadata buffer in disagg (PD) mode.
self.mm_image_tokens: int = 0
self.mm_audio_tokens: int = 0
self.mm_video_tokens: int = 0
# Prefix info
# The indices to kv cache for the shared prefix.
@@ -270,6 +270,9 @@ class _GenerationStreamAccumulator:
cached_tokens_details: list = field(
default_factory=list
) # Detailed breakdown by cache source
image_tokens: list = field(default_factory=list)
audio_tokens: list = field(default_factory=list)
video_tokens: list = field(default_factory=list)
spec_verify_ct: list = field(default_factory=list)
spec_num_correct_drafts: list = field(default_factory=list)
spec_correct_drafts_histogram: list = field(default_factory=list)
@@ -377,6 +380,22 @@ class _GenerationStreamAccumulator:
# Collect detailed cache breakdown if available
self.cached_tokens_details.append(self.get_cached_tokens_details(req))
# Multimodal prompt token counts. In disagg decode mode the prefill node
# already computed these and transferred them via the metadata buffer
# (req.mm_*), so prefer the pre-stored values; otherwise compute them
# from the request's multimodal items.
if req.mm_image_tokens or req.mm_audio_tokens or req.mm_video_tokens:
image_t = req.mm_image_tokens
audio_t = req.mm_audio_tokens
video_t = req.mm_video_tokens
elif req.multimodal_inputs:
image_t, audio_t, video_t = req.multimodal_inputs.compute_mm_token_counts()
else:
image_t = audio_t = video_t = 0
self.image_tokens.append(image_t)
self.audio_tokens.append(audio_t)
self.video_tokens.append(video_t)
self.retraction_counts.append(req.retraction_count)
self.time_stats.append(req.time_stats)
@@ -504,6 +523,9 @@ class _GenerationStreamAccumulator:
completion_tokens=self.completion_tokens,
cached_tokens=self.cached_tokens,
cached_tokens_details=self.cached_tokens_details,
image_tokens=self.image_tokens,
audio_tokens=self.audio_tokens,
video_tokens=self.video_tokens,
input_token_logprobs_val=self.input_token_logprobs_val,
input_token_logprobs_idx=self.input_token_logprobs_idx,
output_token_logprobs_val=self.output_token_logprobs_val,
@@ -1907,6 +1907,18 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
state.customized_info_accumulated[k].extend(v[i])
meta_info[k] = state.customized_info_accumulated[k]
# Add multimodal prompt token counts only for requests that
# actually consumed them, so plain-text meta_info stays unchanged.
image_tokens_list = getattr(recv_obj, "image_tokens", None)
audio_tokens_list = getattr(recv_obj, "audio_tokens", None)
video_tokens_list = getattr(recv_obj, "video_tokens", None)
if image_tokens_list and image_tokens_list[i]:
meta_info["image_tokens"] = image_tokens_list[i]
if audio_tokens_list and audio_tokens_list[i]:
meta_info["audio_tokens"] = audio_tokens_list[i]
if video_tokens_list and video_tokens_list[i]:
meta_info["video_tokens"] = video_tokens_list[i]
if getattr(recv_obj, "output_hidden_states", None):
hidden_states = recv_obj.output_hidden_states[i]
if hidden_states is not None: