[Front-end] [Messages] Populate num_cache_creation_tokens (#48535)

Signed-off-by: Yifan Zong <[email protected]>
This commit is contained in:
yzong-rh
2026-07-18 13:04:35 -04:00
committed by GitHub
parent e94243893d
commit a287eb163f
15 changed files with 201 additions and 127 deletions
@@ -103,6 +103,9 @@ pub struct PrefillStats {
/// Tokens to be prefilled from external KV transfer.
#[serde(default)]
pub num_external_cached_tokens: u32,
/// Prompt tokens newly admitted into the local prefix cache.
#[serde(default)]
pub num_cache_creation_tokens: u32,
}
/// Stats for debugging the metrics calculation.
+1
View File
@@ -375,6 +375,7 @@ mod tests {
num_cached_tokens: 4,
num_local_cached_tokens: 4,
num_external_cached_tokens: 0,
..Default::default()
}),
..Default::default()
},
@@ -23,7 +23,6 @@ from vllm.entrypoints.anthropic.protocol import (
from vllm.entrypoints.anthropic.serving import (
AnthropicServingMessages,
_build_anthropic_usage,
_get_cached_tokens,
)
from vllm.entrypoints.openai.chat_completion.protocol import (
ChatCompletionResponse,
@@ -668,42 +667,6 @@ class TestThinkingBlockConversion:
# ======================================================================
class TestGetCachedTokens:
"""Tests for _get_cached_tokens helper."""
def test_none_usage(self):
assert _get_cached_tokens(None) is None
def test_no_prompt_tokens_details(self):
usage = UsageInfo(prompt_tokens=100, completion_tokens=10)
assert _get_cached_tokens(usage) is None
def test_cached_tokens_present(self):
usage = UsageInfo(
prompt_tokens=100,
completion_tokens=10,
prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=80),
)
assert _get_cached_tokens(usage) == 80
def test_cached_tokens_zero(self):
"""Zero cached tokens should return 0, not None."""
usage = UsageInfo(
prompt_tokens=100,
completion_tokens=10,
prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=0),
)
assert _get_cached_tokens(usage) == 0
def test_cached_tokens_none_in_details(self):
usage = UsageInfo(
prompt_tokens=100,
completion_tokens=10,
prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=None),
)
assert _get_cached_tokens(usage) is None
class TestBuildAnthropicUsage:
"""Tests for _build_anthropic_usage helper.
@@ -711,36 +674,32 @@ class TestBuildAnthropicUsage:
vLLM's prompt_tokens is the total.
"""
def test_no_cache_info(self):
"""When cache info is unavailable, return raw prompt_tokens."""
result = _build_anthropic_usage(100, 10, None)
assert result.input_tokens == 100
assert result.output_tokens == 10
assert result.cache_read_input_tokens is None
assert result.cache_creation_input_tokens is None
def test_cache_hit(self):
"""When cache is hit, input_tokens excludes cached tokens."""
usage = UsageInfo(
prompt_tokens=100,
completion_tokens=10,
prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=80),
prompt_tokens_details=PromptTokenUsageInfo(
cached_tokens=80, created_cache_tokens=10
),
)
result = _build_anthropic_usage(100, 10, usage)
assert result.input_tokens == 20 # 100 - 80
result = _build_anthropic_usage(usage)
assert result.input_tokens == 10 # 100 - 80 - 10
assert result.output_tokens == 10
assert result.cache_read_input_tokens == 80
assert result.cache_creation_input_tokens == 0
assert result.cache_creation_input_tokens == 10
def test_zero_cached_tokens(self):
"""Zero cached tokens should still set cache_creation to 0."""
usage = UsageInfo(
prompt_tokens=100,
completion_tokens=10,
prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=0),
prompt_tokens_details=PromptTokenUsageInfo(
cached_tokens=0, created_cache_tokens=0
),
)
result = _build_anthropic_usage(100, 10, usage)
assert result.input_tokens == 100 # 100 - 0
result = _build_anthropic_usage(usage)
assert result.input_tokens == 100 # 100 - 0 - 0
assert result.cache_read_input_tokens == 0
assert result.cache_creation_input_tokens == 0
@@ -749,9 +708,11 @@ class TestBuildAnthropicUsage:
usage = UsageInfo(
prompt_tokens=100,
completion_tokens=10,
prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=100),
prompt_tokens_details=PromptTokenUsageInfo(
cached_tokens=100, created_cache_tokens=0
),
)
result = _build_anthropic_usage(100, 10, usage)
result = _build_anthropic_usage(usage)
assert result.input_tokens == 0
assert result.cache_read_input_tokens == 100
assert result.cache_creation_input_tokens == 0
@@ -759,7 +720,7 @@ class TestBuildAnthropicUsage:
def test_no_prompt_tokens_details(self):
"""UsageInfo without prompt_tokens_details returns no cache info."""
usage = UsageInfo(prompt_tokens=100, completion_tokens=10)
result = _build_anthropic_usage(100, 10, usage)
result = _build_anthropic_usage(usage)
assert result.input_tokens == 100
assert result.cache_read_input_tokens is None
assert result.cache_creation_input_tokens is None
@@ -1241,7 +1202,9 @@ class TestStreamingCacheUsageSemantics:
prompt_tokens=100,
completion_tokens=5,
total_tokens=105,
prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=80),
prompt_tokens_details=PromptTokenUsageInfo(
cached_tokens=80, created_cache_tokens=10
),
),
)
yield "data: [DONE]"
@@ -1263,9 +1226,9 @@ class TestStreamingCacheUsageSemantics:
delta_usage = next(
data["usage"] for ev, data in events if ev == "message_delta"
)
assert delta_usage["input_tokens"] == 20 # 100 - 80
assert delta_usage["input_tokens"] == 10 # 100 - 80 - 10
assert delta_usage["cache_read_input_tokens"] == 80
assert delta_usage["cache_creation_input_tokens"] == 0
assert delta_usage["cache_creation_input_tokens"] == 10
@pytest.mark.asyncio
async def test_streaming_no_cache_hit(self):
@@ -1284,7 +1247,9 @@ class TestStreamingCacheUsageSemantics:
prompt_tokens=50,
completion_tokens=5,
total_tokens=55,
prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=0),
prompt_tokens_details=PromptTokenUsageInfo(
cached_tokens=0, created_cache_tokens=0
),
),
)
yield "data: [DONE]"
@@ -1302,7 +1267,7 @@ class TestStreamingCacheUsageSemantics:
assert start_usage["input_tokens"] == 50
assert "cache_read_input_tokens" not in start_usage
assert "cache_creation_input_tokens" not in start_usage
assert delta_usage["input_tokens"] == 50 # 50 - 0
assert delta_usage["input_tokens"] == 50 # 50 - 0 - 0
assert delta_usage["cache_read_input_tokens"] == 0
assert delta_usage["cache_creation_input_tokens"] == 0
@@ -18,6 +18,7 @@ def server():
"--max-model-len",
"2048",
"--enforce-eager",
"--enable-prompt-tokens-details",
"--enable-auto-tool-choice",
"--tool-call-parser",
"hermes",
@@ -191,3 +192,54 @@ async def test_anthropic_structured_output(client: anthropic.AsyncAnthropic):
json_obj = json.loads(response.content[0].text)
for key in ["name", "email", "plan_interest", "demo_requested"]:
assert key in json_obj, f"Missing key in output: {key}"
@pytest.mark.asyncio
async def test_anthropic_streaming_cache_usage(client: anthropic.AsyncAnthropic):
async def get_stream_usage(resp):
prompt_tokens = None
usage = None
async for chunk in resp:
if (
chunk.type == "message_start"
and chunk.message is not None
and chunk.message.usage is not None
):
prompt_tokens = chunk.message.usage.input_tokens
elif chunk.type == "message_delta" and chunk.usage is not None:
usage = chunk.usage
assert usage is not None
assert usage.input_tokens >= 0
assert usage.output_tokens >= 0
cache_created = usage.cache_creation_input_tokens
cache_read = usage.cache_read_input_tokens
assert cache_read is not None
assert cache_created is not None
assert cache_created >= 0
assert cache_read >= 0
assert prompt_tokens == usage.input_tokens + cache_created + cache_read
return usage
request = dict(
model="claude-3-7-sonnet-latest",
max_tokens=1,
temperature=0.0,
messages=[
{
"role": "user",
"content": "Cache coverage sentinel. " * 256
+ "Answer with exactly one word: ok.",
}
],
stream=True,
)
cold_usage = await get_stream_usage(await client.messages.create(**request))
assert cold_usage.cache_read_input_tokens == 0
assert cold_usage.cache_creation_input_tokens is not None
assert cold_usage.cache_creation_input_tokens > 0
warm_usage = await get_stream_usage(await client.messages.create(**request))
assert warm_usage.cache_read_input_tokens is not None
assert warm_usage.cache_read_input_tokens > 0
@@ -831,16 +831,23 @@ def test_mm_prompt_tokens_details():
assert counts == {"image": 600, "video": 1200}
# Gated off, or nothing to report -> no details.
assert _make_prompt_tokens_details(False, 5, counts) is None
assert _make_prompt_tokens_details(True, None, None) is None
assert _make_prompt_tokens_details(False, 5, 0, counts) is None
assert _make_prompt_tokens_details(True, None, None, None) is None
# Zero cached_tokens is still reported (not None), matching the cached-only
# behavior; multimodal counts ride alongside even when cached_tokens is None.
assert _make_prompt_tokens_details(True, 0, None).cached_tokens == 0
details = _make_prompt_tokens_details(True, None, counts)
details = _make_prompt_tokens_details(True, 0, 0, None)
assert details.cached_tokens == 0
assert details.created_cache_tokens == 0
assert details.multimodal_tokens is None
details = _make_prompt_tokens_details(True, None, None, counts)
assert details.cached_tokens is None
assert details.created_cache_tokens is None
assert details.multimodal_tokens == {"image": 600, "video": 1200}
details = _make_prompt_tokens_details(True, 3, 0, counts)
assert details.cached_tokens == 3
assert details.created_cache_tokens == 0
assert details.multimodal_tokens == {"image": 600, "video": 1200}
assert _make_prompt_tokens_details(True, 3, counts).cached_tokens == 3
@pytest.mark.asyncio
+1
View File
@@ -280,6 +280,7 @@ def test_abort_request_when_structured_output_fsm_cannot_advance():
scheduler.waiting = Mock()
scheduler.kv_cache_manager = Mock()
scheduler.kv_cache_manager.take_events.return_value = None
scheduler.kv_cache_manager.estimate_cached_tokens.return_value = 0
scheduler.kv_event_publisher = Mock()
scheduler.finished_req_ids = set()
scheduler.finished_req_ids_dict = None
+1
View File
@@ -3036,6 +3036,7 @@ def test_abort_request_when_structured_output_fsm_cannot_advance():
scheduler.waiting = Mock()
scheduler.kv_cache_manager = Mock()
scheduler.kv_cache_manager.take_events.return_value = None
scheduler.kv_cache_manager.estimate_cached_tokens.return_value = 0
scheduler.kv_event_publisher = Mock()
scheduler.finished_req_ids = set()
scheduler.finished_req_ids_dict = None
+29 -51
View File
@@ -53,47 +53,43 @@ from vllm.renderers.online_renderer import OnlineRenderer
logger = logging.getLogger(__name__)
def _get_cached_tokens(usage: UsageInfo | None) -> int | None:
"""Extract cached token count from OpenAI UsageInfo."""
if usage is None or usage.prompt_tokens_details is None:
return None
return usage.prompt_tokens_details.cached_tokens
def _build_anthropic_usage(
prompt_tokens: int,
completion_tokens: int | None,
usage: UsageInfo | None,
) -> AnthropicUsage:
"""Build an AnthropicUsage from OpenAI-style token counts.
"""Build an AnthropicUsage from UsageInfo.
Anthropic defines ``total_input == input_tokens + cache_read +
cache_creation``. vLLM's ``prompt_tokens`` is the total, so
``input_tokens = prompt_tokens - cached_tokens``.
``input_tokens = prompt_tokens - cache_read - cache_creation``.
OpenAI usage only exposes ``cached_tokens`` (hits); there is no
cache-creation analog, so ``cache_creation_input_tokens`` is ``0``
when cache info is present. When cache info is absent (e.g.
``--enable-prompt-tokens-details`` off, or a streaming chunk that
hasn't carried it yet), cache fields are left **unset** so
``exclude_unset=True`` serialization omits them entirely.
Cache fields are taken from ``UsageInfo.prompt_tokens_details``.
When cache info is absent (e.g. ``--enable-prompt-tokens-details``
off, or a streaming chunk that hasn't carried it yet), cache fields
are left **unset** so ``exclude_unset=True`` serialization omits them
entirely.
``completion_tokens`` follows ``UsageInfo`` and may be ``None`` on
intermediate stream chunks; we coerce to ``0`` for the wire format.
"""
output_tokens = completion_tokens or 0
cached = _get_cached_tokens(usage)
if cached is not None:
return AnthropicUsage(
input_tokens=prompt_tokens - cached,
output_tokens=output_tokens,
cache_read_input_tokens=cached,
cache_creation_input_tokens=0,
)
return AnthropicUsage(
input_tokens=prompt_tokens,
output_tokens=output_tokens,
)
kwargs = {}
if usage is None:
kwargs["input_tokens"] = 0
kwargs["output_tokens"] = 0
else:
kwargs["output_tokens"] = usage.completion_tokens or 0
input_tokens = usage.prompt_tokens
if (details := usage.prompt_tokens_details) is not None:
if (cache_read := details.cached_tokens) is not None:
input_tokens -= cache_read
kwargs["cache_read_input_tokens"] = cache_read
if (cache_creation := details.created_cache_tokens) is not None:
input_tokens -= cache_creation
kwargs["cache_creation_input_tokens"] = cache_creation
kwargs["input_tokens"] = max(0, input_tokens)
return AnthropicUsage(**kwargs)
def wrap_data_with_event(data: str, event: str):
@@ -625,11 +621,7 @@ class AnthropicServingMessages(OpenAIServingChat):
id=generator.id,
content=[],
model=generator.model,
usage=_build_anthropic_usage(
generator.usage.prompt_tokens,
generator.usage.completion_tokens,
generator.usage,
),
usage=_build_anthropic_usage(generator.usage),
kv_transfer_params=generator.kv_transfer_params,
ec_transfer_params=generator.ec_transfer_params,
)
@@ -816,13 +808,7 @@ class AnthropicServingMessages(OpenAIServingChat):
model=origin_chunk.model,
stop_reason=None,
stop_sequence=None,
usage=_build_anthropic_usage(
origin_chunk.usage.prompt_tokens
if origin_chunk.usage
else 0,
0,
origin_chunk.usage,
),
usage=_build_anthropic_usage(origin_chunk.usage),
),
)
first_item = False
@@ -840,15 +826,7 @@ class AnthropicServingMessages(OpenAIServingChat):
chunk = AnthropicStreamEvent(
type="message_delta",
delta=AnthropicDelta(stop_reason=stop_reason),
usage=_build_anthropic_usage(
origin_chunk.usage.prompt_tokens
if origin_chunk.usage
else 0,
origin_chunk.usage.completion_tokens
if origin_chunk.usage
else 0,
origin_chunk.usage,
),
usage=_build_anthropic_usage(origin_chunk.usage),
)
data = chunk.model_dump_json(exclude_unset=True)
yield wrap_data_with_event(data, "message_delta")
@@ -90,15 +90,21 @@ def _get_mm_token_counts(engine_input: EngineInput) -> dict[str, int]:
def _make_prompt_tokens_details(
enable_prompt_tokens_details: bool,
num_cached_tokens: int | None,
num_cache_creation_tokens: int | None,
mm_token_counts: dict[str, int] | None,
) -> PromptTokenUsageInfo | None:
"""Build ``prompt_tokens_details`` from cached + multimodal token counts."""
if not enable_prompt_tokens_details:
return None
if num_cached_tokens is None and not mm_token_counts:
if (
num_cached_tokens is None
and num_cache_creation_tokens is None
and not mm_token_counts
):
return None
return PromptTokenUsageInfo(
cached_tokens=num_cached_tokens,
created_cache_tokens=num_cache_creation_tokens,
multimodal_tokens=mm_token_counts or None,
)
@@ -427,6 +433,7 @@ class OpenAIServingChat(GenerateBaseServing):
finish_reason_sent = [False] * num_choices
num_prompt_tokens = 0
num_cached_tokens = None
num_cache_creation_tokens = None
tools_streamed = [False] * num_choices
if isinstance(request.tool_choice, ChatCompletionNamedToolChoiceParam):
@@ -479,6 +486,7 @@ class OpenAIServingChat(GenerateBaseServing):
# response (by the try...catch).
if first_iteration:
num_cached_tokens = res.num_cached_tokens
num_cache_creation_tokens = res.num_cache_creation_tokens
# Send first response for each request.n (index) with
# the role
role = self.get_chat_request_role(request)
@@ -756,6 +764,7 @@ class OpenAIServingChat(GenerateBaseServing):
final_usage.prompt_tokens_details = _make_prompt_tokens_details(
self.enable_prompt_tokens_details,
num_cached_tokens,
num_cache_creation_tokens,
mm_token_counts,
)
@@ -1030,6 +1039,7 @@ class OpenAIServingChat(GenerateBaseServing):
usage.prompt_tokens_details = _make_prompt_tokens_details(
self.enable_prompt_tokens_details,
final_res.num_cached_tokens,
final_res.num_cache_creation_tokens,
mm_token_counts,
)
@@ -101,6 +101,7 @@ class ModelList(OpenAIBaseModel):
class PromptTokenUsageInfo(OpenAIBaseModel):
cached_tokens: int | None = None
created_cache_tokens: int | None = None
multimodal_tokens: dict[str, int] | None = None
"""Prompt tokens contributed by each input modality, keyed by modality name
(e.g. `image`, `audio`, `video`). A breakdown of the multimodal
+6 -1
View File
@@ -103,6 +103,8 @@ class RequestOutput:
encoder_prompt_token_ids: The token IDs of the encoder prompt.
None if decoder-only.
num_cached_tokens: The number of tokens with prefix cache hit.
num_cache_creation_tokens: Prompt tokens currently counted as local
prefix-cache writes for this request.
kv_transfer_params: The params for remote K/V transfer.
ec_transfer_params: The params for remote encoder-cache transfer.
"""
@@ -120,6 +122,7 @@ class RequestOutput:
encoder_prompt: str | None = None,
encoder_prompt_token_ids: list[int] | None = None,
num_cached_tokens: int | None = None,
num_cache_creation_tokens: int | None = None,
*,
kv_transfer_params: dict[str, Any] | None = None,
ec_transfer_params: dict[str, Any] | None = None,
@@ -142,6 +145,7 @@ class RequestOutput:
self.encoder_prompt = encoder_prompt
self.encoder_prompt_token_ids = encoder_prompt_token_ids
self.num_cached_tokens = num_cached_tokens
self.num_cache_creation_tokens = num_cache_creation_tokens
self.kv_transfer_params = kv_transfer_params
self.ec_transfer_params = ec_transfer_params
@@ -188,7 +192,8 @@ class RequestOutput:
f"finished={self.finished}, "
f"metrics={self.metrics}, "
f"lora_request={self.lora_request}, "
f"num_cached_tokens={self.num_cached_tokens})"
f"num_cached_tokens={self.num_cached_tokens}, "
f"num_cache_creation_tokens={self.num_cache_creation_tokens})"
)
+29
View File
@@ -657,6 +657,35 @@ class KVCacheManager:
clipped_block_ids.append(ids[:num_valid_blocks])
return tuple(clipped_block_ids)
def estimate_cached_tokens(self, request: Request) -> int:
"""Estimate the number of tokens cached by the request."""
cached_tokens: int | None = None
for group, blocks in zip(
self.kv_cache_config.kv_cache_groups,
self.get_blocks(request.request_id).blocks,
):
if isinstance(
group.kv_cache_spec,
(CrossAttentionSpec, EncoderOnlyAttentionSpec),
):
# Cross-attention and encoder-only groups are not prefix cached.
continue
group_cached_tokens = 0
for block in blocks:
group_cached_tokens = max(
group_cached_tokens,
block.block_hash_num_tokens or 0,
)
cached_tokens = (
group_cached_tokens
if cached_tokens is None
else min(cached_tokens, group_cached_tokens)
)
return cached_tokens or 0
def cache_blocks(self, request: Request, num_computed_tokens: int) -> None:
"""Cache the blocks for the request, if enabled.
+17 -9
View File
@@ -799,7 +799,10 @@ class Scheduler(SchedulerInterface):
continue
# Track first scheduled prefill, not post-preemption repeat prefills
if request.prefill_stats is not None:
if (
request.prefill_stats is not None
and request.num_preemptions <= 0
):
assert num_computed_tokens <= request.num_prompt_tokens
request.prefill_stats.set(
num_prompt_tokens=request.num_prompt_tokens,
@@ -1708,6 +1711,7 @@ class Scheduler(SchedulerInterface):
pooler_output = pooler_outputs[req_index] if pooler_outputs else None
kv_transfer_params = None
ec_transfer_params = None
prefill_stats = None
status_before_stop = request.status
num_output_tokens_before = len(request._output_token_ids)
@@ -1787,6 +1791,16 @@ class Scheduler(SchedulerInterface):
# Normal decode / re-prefill: token(s) at the END.
routed_experts = routing_data[end - len(new_token_ids) : end]
should_emit_output = bool(
new_token_ids or pooler_output is not None or stopped
)
if should_emit_output:
prefill_stats = request.take_prefill_stats()
if prefill_stats is not None:
prefill_stats.finalize(
self.kv_cache_manager.estimate_cached_tokens(request)
)
finish_reason = None
if stopped:
# Capture finish_reason BEFORE _handle_stopped_request, which may
@@ -1814,13 +1828,7 @@ class Scheduler(SchedulerInterface):
# Get prompt logprobs for this request.
prompt_logprobs_tensors = prompt_logprobs_dict.get(req_id)
if (
new_token_ids
or pooler_output is not None
or kv_transfer_params
or ec_transfer_params
or stopped
):
if should_emit_output:
# Add EngineCoreOutput for this Request.
outputs[request.client_index].append(
EngineCoreOutput(
@@ -1832,7 +1840,7 @@ class Scheduler(SchedulerInterface):
pooling_output=pooler_output,
stop_reason=request.stop_reason,
events=request.take_events(),
prefill_stats=request.take_prefill_stats(),
prefill_stats=prefill_stats,
kv_transfer_params=kv_transfer_params,
ec_transfer_params=ec_transfer_params,
trace_headers=request.trace_headers,
+5
View File
@@ -172,6 +172,7 @@ class RequestState:
self.is_prefilling = True
self.queue = queue
self.num_cached_tokens = 0
self.num_cache_creation_tokens = 0
self.stats = RequestStateStats(arrival_time=arrival_time) if log_stats else None
@@ -377,6 +378,7 @@ class RequestState:
kv_transfer_params=kv_transfer_params,
ec_transfer_params=ec_transfer_params,
num_cached_tokens=self.num_cached_tokens,
num_cache_creation_tokens=self.num_cache_creation_tokens,
metrics=self.stats,
)
@@ -639,6 +641,9 @@ class OutputProcessor:
req_state.num_cached_tokens = (
engine_core_output.prefill_stats.num_cached_tokens
)
req_state.num_cache_creation_tokens = (
engine_core_output.prefill_stats.num_cache_creation_tokens
)
req_state.is_prefilling = False
if pooling_output is None:
+8
View File
@@ -265,6 +265,7 @@ class PrefillStats:
num_cached_tokens: Tokens to be prefilled without actual compute work.
num_local_cached_tokens: Tokens to be prefilled from local prefix cache.
num_external_cached_tokens: Tokens to be prefilled from external KV transfer.
num_cache_creation_tokens: Tokens computed and written to the prefix cache.
"""
num_prompt_tokens: int = 0
@@ -272,6 +273,7 @@ class PrefillStats:
num_cached_tokens: int = 0
num_local_cached_tokens: int = 0
num_external_cached_tokens: int = 0
num_cache_creation_tokens: int = 0
def set(
self,
@@ -288,6 +290,12 @@ class PrefillStats:
self.num_local_cached_tokens = num_local_cached_tokens
self.num_external_cached_tokens = num_external_cached_tokens
def finalize(self, num_cached_tokens: int) -> None:
assert num_cached_tokens >= 0
self.num_cache_creation_tokens = max(
0, min(num_cached_tokens, self.num_prompt_tokens) - self.num_cached_tokens
)
@dataclass
class PromptTokenStats: