add token usage stats

This commit is contained in:
Ryuichi Leo Takashige
2026-01-22 21:04:56 +00:00
parent 4591301767
commit d9ddf90575
3 changed files with 53 additions and 0 deletions
+46
View File
@@ -42,6 +42,7 @@ from exo.shared.types.api import (
CompletionLogprobs,
CompletionResponse,
CompletionTaskParams,
CompletionTokensDetails,
CreateInstanceParams,
CreateInstanceResponse,
DeleteInstanceResponse,
@@ -65,6 +66,7 @@ from exo.shared.types.api import (
PlacementPreviewResponse,
StreamingChoiceResponse,
ToolCall,
Usage,
)
from exo.shared.types.chunks import (
CompletionChunk,
@@ -231,6 +233,12 @@ class API:
self._image_store = ImageStore(EXO_IMAGE_CACHE_DIR)
self._tg: TaskGroup | None = None
# Accumulated usage stats across all requests
self._total_prompt_tokens: int = 0
self._total_completion_tokens: int = 0
self._total_reasoning_tokens: int = 0
self._total_requests: int = 0
def reset(self, new_session_id: SessionId, result_clock: int):
logger.info("Resetting API State")
self.state = State()
@@ -296,6 +304,17 @@ class API:
self.app.get("/images/{image_id}")(self.get_image)
self.app.get("/state")(lambda: self.state)
self.app.get("/events")(lambda: self._event_log)
self.app.get("/v1/usage")(self.get_usage)
def get_usage(self) -> dict[str, int]:
"""Return accumulated token usage across all requests."""
return {
"total_requests": self._total_requests,
"total_prompt_tokens": self._total_prompt_tokens,
"total_completion_tokens": self._total_completion_tokens,
"total_reasoning_tokens": self._total_reasoning_tokens,
"total_tokens": self._total_prompt_tokens + self._total_completion_tokens,
}
async def place_instance(self, payload: PlaceInstanceParams):
command = PlaceInstance(
@@ -563,6 +582,13 @@ class API:
yield f"data: {chunk_response.model_dump_json()}\n\n"
if chunk.finish_reason is not None:
# Accumulate usage stats from the final chunk
if isinstance(chunk, TokenChunk) and chunk.stats is not None:
s = chunk.stats
self._total_prompt_tokens += s.prompt_tokens
self._total_completion_tokens += s.generation_tokens
self._total_reasoning_tokens += s.reasoning_tokens
self._total_requests += 1
yield "data: [DONE]\n\n"
async def _collect_chat_completion(
@@ -575,6 +601,7 @@ class API:
logprobs_items: list[LogprobsContentItem] = []
model: str | None = None
finish_reason: FinishReason | None = None
stats: GenerationStats | None = None
async for chunk in self._chat_chunk_stream(command_id):
# Skip CompletionChunk - it's for the legacy completions API
@@ -592,6 +619,8 @@ class API:
if isinstance(chunk, TokenChunk):
text_parts.append(chunk.text)
if chunk.stats is not None:
stats = chunk.stats
if chunk.logprob is not None:
lp = _build_logprobs(chunk)
if lp.content:
@@ -621,6 +650,22 @@ class API:
if logprobs_items:
logprobs = Logprobs(content=logprobs_items)
usage: Usage | None = None
if stats is not None:
completion_tokens = stats.generation_tokens
usage = Usage(
prompt_tokens=stats.prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=stats.prompt_tokens + completion_tokens,
completion_tokens_details=CompletionTokensDetails(
reasoning_tokens=stats.reasoning_tokens,
) if stats.reasoning_tokens > 0 else None,
)
self._total_prompt_tokens += stats.prompt_tokens
self._total_completion_tokens += completion_tokens
self._total_reasoning_tokens += stats.reasoning_tokens
self._total_requests += 1
return ChatCompletionResponse(
id=command_id,
created=int(time.time()),
@@ -637,6 +682,7 @@ class API:
finish_reason=finish_reason,
)
],
usage=usage,
)
async def _collect_chat_completion_with_stats(
+1
View File
@@ -151,6 +151,7 @@ class GenerationStats(BaseModel):
generation_tps: float
prompt_tokens: int
generation_tokens: int
reasoning_tokens: int = 0
peak_memory_usage: Memory
+6
View File
@@ -59,6 +59,7 @@ class ActiveRequest:
gpt_oss_parser: Any | None = None # StreamableParser for GPT-OSS models
gpt_oss_thinking: bool = False
tokens_generated: int = 0
reasoning_tokens: int = 0
prompt_tokens: int = 0
start_time: float = field(default_factory=time.perf_counter)
@@ -290,6 +291,10 @@ class BatchedInferenceHandler:
delta: str | None = parser.last_content_delta # pyright: ignore[reportAny]
channel: str = parser.current_channel # pyright: ignore[reportAny]
# Track reasoning tokens (analysis channel = thinking)
if channel == "analysis":
active_request.reasoning_tokens += 1
# Handle thinking tag transitions
prefix = ""
if channel == "analysis" and not active_request.gpt_oss_thinking:
@@ -339,6 +344,7 @@ class BatchedInferenceHandler:
generation_tps=generation_tps,
prompt_tokens=active_request.prompt_tokens,
generation_tokens=active_request.tokens_generated,
reasoning_tokens=active_request.reasoning_tokens,
peak_memory_usage=Memory.from_bytes(peak_memory_bytes),
)