mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-11 03:40:03 +00:00
token消耗量统计
This commit is contained in:
parent
e14ff16cac
commit
26f221e765
11 changed files with 370 additions and 31 deletions
|
|
@ -252,16 +252,23 @@ async def answer_question_agentic(app, question: str) -> tuple[str, dict]:
|
|||
|
||||
Returns (answer, metadata)
|
||||
"""
|
||||
from reme.utils.evaluation_interface import track_job_counts
|
||||
from reme.utils.evaluation_interface import track_agent_token_counts, track_job_counts
|
||||
|
||||
with track_job_counts(["search"], app.context) as counts:
|
||||
with track_job_counts(["search"], app.context) as counts, track_agent_token_counts(
|
||||
["bench"],
|
||||
app.context,
|
||||
) as token_counts:
|
||||
query_resp = await app.run_job(
|
||||
"agentic_answer",
|
||||
query=question,
|
||||
)
|
||||
answer = (query_resp.answer or "").strip()
|
||||
|
||||
return answer, {"mode": "agentic", "search_calls": counts["search"]}
|
||||
return answer, {
|
||||
"mode": "agentic",
|
||||
"search_calls": counts["search"],
|
||||
"token_count": token_counts["bench"],
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
@ -454,6 +461,7 @@ async def evaluate_case(eval_config: dict, case_id: str, eval_only: bool = False
|
|||
logger.info(
|
||||
f"[Case {case_id}] Agentic search calls: {agentic_meta.get('search_calls', 0)}",
|
||||
)
|
||||
logger.info(f"[Case {case_id}] Bench token usage: {agentic_meta.get('token_count', 0)}")
|
||||
|
||||
# Judge agentic answer
|
||||
logger.info(f"[Case {case_id}] Judging agentic ({q_type})...")
|
||||
|
|
@ -685,6 +693,7 @@ def main( # pylint: disable=too-many-statements
|
|||
all_scores: list[float] = []
|
||||
all_binary_scores: list[float] = []
|
||||
all_search_calls: list[int] = []
|
||||
all_token_counts: list[int] = []
|
||||
|
||||
for case_result in results:
|
||||
if "error" in case_result:
|
||||
|
|
@ -708,6 +717,7 @@ def main( # pylint: disable=too-many-statements
|
|||
all_scores.append(score)
|
||||
all_binary_scores.append(binary_score)
|
||||
all_search_calls.append(q.get("agentic_metadata", {}).get("search_calls", 0))
|
||||
all_token_counts.append(q.get("agentic_metadata", {}).get("token_count", 0))
|
||||
|
||||
print("\n ── AGENTIC ──")
|
||||
if all_scores:
|
||||
|
|
@ -723,6 +733,8 @@ def main( # pylint: disable=too-many-statements
|
|||
print(f" {'OVERALL':<40s}: {overall:.3f} binary={binary_overall:.3f} ({len(all_scores)} Qs)")
|
||||
avg_search_calls = sum(all_search_calls) / len(all_search_calls)
|
||||
print(f" Average search calls/query: {avg_search_calls:.2f}")
|
||||
avg_token_count = sum(all_token_counts) / len(all_token_counts)
|
||||
print(f" Average bench tokens/query: {avg_token_count:.2f}")
|
||||
else:
|
||||
print(" (no results)")
|
||||
|
||||
|
|
|
|||
|
|
@ -259,7 +259,7 @@ async def evaluate_item(item: dict, eval_config: dict, item_index: int, eval_onl
|
|||
"""
|
||||
from reme import Application
|
||||
from reme.config import resolve_app_config
|
||||
from reme.utils.evaluation_interface import track_job_counts
|
||||
from reme.utils.evaluation_interface import track_agent_token_counts, track_job_counts
|
||||
|
||||
reme_cfg = eval_config["reme"]
|
||||
dream_trigger_hour = reme_cfg.get("dream_trigger_hour", 23)
|
||||
|
|
@ -433,19 +433,20 @@ async def evaluate_item(item: dict, eval_config: dict, item_index: int, eval_onl
|
|||
f"[Item {item_index}] Asking (agentic): {question[:80]}... query_time={query_time}",
|
||||
)
|
||||
|
||||
with track_job_counts(["search"], app.context) as counts:
|
||||
query_resp = await app.run_job(
|
||||
"agentic_answer",
|
||||
query=question,
|
||||
query_time=query_time,
|
||||
)
|
||||
with track_job_counts(["search"], app.context) as counts, track_agent_token_counts(
|
||||
["bench"],
|
||||
app.context,
|
||||
) as token_counts:
|
||||
query_resp = await app.run_job("agentic_answer", query=question, query_time=query_time)
|
||||
agentic_search_calls = counts["search"]
|
||||
agentic_token_count = token_counts["bench"]
|
||||
agentic_response = (query_resp.answer or "").strip()
|
||||
if not agentic_response:
|
||||
agentic_response = "(no answer generated)"
|
||||
|
||||
logger.info(f"[Item {item_index}] Agentic response: {agentic_response[:200]}...")
|
||||
logger.info(f"[Item {item_index}] Agentic search calls: {agentic_search_calls}")
|
||||
logger.info(f"[Item {item_index}] Bench token usage: {agentic_token_count}")
|
||||
|
||||
# ── Phase 5: Judge agentic response (via answer_judge_step) ──────────
|
||||
logger.info(f"[Item {item_index}] Judging agentic (binary, type={item['question_type']})...")
|
||||
|
|
@ -469,6 +470,7 @@ async def evaluate_item(item: dict, eval_config: dict, item_index: int, eval_onl
|
|||
"agentic_response": agentic_response,
|
||||
"agentic_judgment": agentic_judgment,
|
||||
"agentic_search_calls": agentic_search_calls,
|
||||
"agentic_token_count": agentic_token_count,
|
||||
"sessions_ingested": len(sorted_sessions),
|
||||
"dreams_triggered": len(dream_dates_triggered),
|
||||
}
|
||||
|
|
@ -718,6 +720,8 @@ def _print_summary(results: list[dict], start_time: float) -> None:
|
|||
print(f" Overall accuracy: {agentic_correct}/{total} ({100*agentic_correct/total:.1f}%)")
|
||||
avg_search_calls = sum(r.get("agentic_search_calls", 0) for r in results) / total if total else 0
|
||||
print(f" Average search calls/query: {avg_search_calls:.2f}")
|
||||
avg_token_count = sum(r.get("agentic_token_count", 0) for r in results) / total if total else 0
|
||||
print(f" Average bench tokens/query: {avg_token_count:.2f}")
|
||||
print(" Per-type accuracy:")
|
||||
for qtype, stats in sorted(agentic_type_stats.items()):
|
||||
acc = 100 * stats["correct"] / stats["total"] if stats["total"] else 0
|
||||
|
|
|
|||
|
|
@ -59,7 +59,7 @@ from .base_agent_wrapper import BaseAgentWrapper
|
|||
from ..as_llm import BaseAsLLM
|
||||
from ..component_registry import R
|
||||
from ...enumeration import ChunkEnum
|
||||
from ...schema import StreamChunk
|
||||
from ...schema import StreamChunk, TokenUsage
|
||||
from ...utils import AsStateHandler
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -149,6 +149,14 @@ class AsAgentWrapper(BaseAgentWrapper):
|
|||
|
||||
SDK_PACKAGE = "agentscope"
|
||||
|
||||
@staticmethod
|
||||
def _agentscope_usage(usage: Any) -> TokenUsage:
|
||||
"""Normalize AgentScope usage while preserving provider cache semantics."""
|
||||
module = type(usage).__module__ if usage is not None else ""
|
||||
# Anthropic reports normal, cache-read, and cache-write input tokens
|
||||
# separately; OpenAI-style adapters report prompt tokens inclusive.
|
||||
return TokenUsage.from_provider(usage, input_includes_cache="_anthropic" not in module)
|
||||
|
||||
def __init__(self, as_llm: str = "default", session_retention_days: int = 10, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.as_llm = self.bind(as_llm, BaseAsLLM, optional=False)
|
||||
|
|
@ -343,8 +351,18 @@ class AsAgentWrapper(BaseAgentWrapper):
|
|||
kwargs = self._merged_kwargs(kwargs)
|
||||
agent, inputs = await self._build_agent(inputs, **kwargs)
|
||||
|
||||
usages: list[TokenUsage] = []
|
||||
await agent.observe(inputs)
|
||||
await agent.reply()
|
||||
async for event in agent.reply_stream():
|
||||
if isinstance(event, ModelCallEndEvent):
|
||||
# AgentScope's event intentionally contains only the portable
|
||||
# input/output pair. Cache dimensions are unavailable here.
|
||||
usages.append(
|
||||
TokenUsage(
|
||||
input_tokens=event.input_tokens,
|
||||
output_tokens=event.output_tokens,
|
||||
),
|
||||
)
|
||||
await self._dump_state(agent.state)
|
||||
last_msg = agent.state.context[-1]
|
||||
|
||||
|
|
@ -365,6 +383,12 @@ class AsAgentWrapper(BaseAgentWrapper):
|
|||
tool_choice=ToolChoice(mode="auto"),
|
||||
)
|
||||
result["structured_output"] = res.content
|
||||
if res.usage is not None:
|
||||
usages.append(self._agentscope_usage(res.usage))
|
||||
|
||||
usage = TokenUsage.combine(usages)
|
||||
result["usage"] = usage.model_dump()
|
||||
self._record_token_usage(usage)
|
||||
|
||||
return result
|
||||
|
||||
|
|
@ -453,13 +477,16 @@ class AsAgentWrapper(BaseAgentWrapper):
|
|||
if isinstance(event, ModelCallStartEvent):
|
||||
return cls._chunk(ChunkEnum.USAGE, chunk="", metadata={"model_name": getattr(event, "model_name", None)})
|
||||
if isinstance(event, ModelCallEndEvent):
|
||||
usage = {"input_tokens": event.input_tokens, "output_tokens": event.output_tokens}
|
||||
usage = TokenUsage(input_tokens=event.input_tokens, output_tokens=event.output_tokens)
|
||||
return cls._chunk(
|
||||
ChunkEnum.USAGE,
|
||||
chunk=json.dumps(usage),
|
||||
input_tokens=event.input_tokens,
|
||||
output_tokens=event.output_tokens,
|
||||
metadata={"model_name": getattr(event, "model_name", None)},
|
||||
chunk=json.dumps(usage.model_dump()),
|
||||
input_tokens=usage.input_tokens,
|
||||
output_tokens=usage.output_tokens,
|
||||
metadata={
|
||||
"model_name": getattr(event, "model_name", None),
|
||||
"usage": usage.model_dump(),
|
||||
},
|
||||
)
|
||||
if isinstance(event, ExceedMaxItersEvent):
|
||||
return cls._chunk(ChunkEnum.ERROR, chunk="Exceeded max iterations")
|
||||
|
|
@ -469,11 +496,20 @@ class AsAgentWrapper(BaseAgentWrapper):
|
|||
"""Stream agent events as unified StreamChunk objects."""
|
||||
kwargs = self._merged_stream_kwargs(kwargs)
|
||||
agent, inputs = await self._build_agent(inputs, **kwargs)
|
||||
usages: list[TokenUsage] = []
|
||||
|
||||
async for event in agent.reply_stream(inputs):
|
||||
if isinstance(event, ModelCallEndEvent):
|
||||
usages.append(
|
||||
TokenUsage(
|
||||
input_tokens=event.input_tokens,
|
||||
output_tokens=event.output_tokens,
|
||||
),
|
||||
)
|
||||
chunk = self._event_to_chunk(event)
|
||||
if chunk is not None:
|
||||
chunk.session_id = chunk.session_id or agent.state.session_id
|
||||
yield chunk
|
||||
|
||||
await self._dump_state(agent.state)
|
||||
self._record_token_usage(TokenUsage.combine(usages))
|
||||
|
|
|
|||
|
|
@ -11,7 +11,8 @@ from pydantic import BaseModel
|
|||
from ..base_component import BaseComponent
|
||||
from ..outbound_proxy import BaseOutboundProxy
|
||||
from ...enumeration import ChunkEnum, ComponentEnum
|
||||
from ...schema import StreamChunk
|
||||
from ...schema import StreamChunk, TokenUsage
|
||||
from ...utils import global_counter_add
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..job.base_job import BaseJob
|
||||
|
|
@ -22,6 +23,7 @@ class BaseAgentWrapper(BaseComponent):
|
|||
|
||||
component_type = ComponentEnum.AGENT_WRAPPER
|
||||
SDK_PACKAGE: ClassVar[str | None] = None
|
||||
TOKEN_COUNTER_PREFIX: ClassVar[str] = "__token_counter"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -227,6 +229,24 @@ class BaseAgentWrapper(BaseComponent):
|
|||
"""Create a StreamChunk with a short backend-friendly call site."""
|
||||
return StreamChunk(chunk_type=chunk_type, **kwargs)
|
||||
|
||||
def _record_token_usage(self, usage: TokenUsage) -> None:
|
||||
"""Add one completed invocation to the application token tree."""
|
||||
if self.app_context is None:
|
||||
return
|
||||
metadata = getattr(self.app_context, "metadata", None)
|
||||
if not isinstance(metadata, dict):
|
||||
return
|
||||
for field in ("input_tokens", "output_tokens", "total_tokens"):
|
||||
global_counter_add(
|
||||
metadata,
|
||||
[self.TOKEN_COUNTER_PREFIX, self.name, field],
|
||||
getattr(usage, field),
|
||||
)
|
||||
for field in ("cache_read_tokens", "cache_write_tokens", "reasoning_tokens"):
|
||||
value = getattr(usage, field)
|
||||
if value is not None:
|
||||
global_counter_add(metadata, [self.TOKEN_COUNTER_PREFIX, self.name, field], value)
|
||||
|
||||
@abstractmethod
|
||||
async def reply(self, inputs: Any, **kwargs) -> dict:
|
||||
"""Send inputs to the agent and return a dict with session_id and last_message."""
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ from typing import Any, TYPE_CHECKING
|
|||
from .base_agent_wrapper import BaseAgentWrapper
|
||||
from ..component_registry import R
|
||||
from ...enumeration import ChunkEnum
|
||||
from ...schema import StreamChunk
|
||||
from ...schema import StreamChunk, TokenUsage
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from claude_agent_sdk import AssistantMessage, ResultMessage, UserMessage
|
||||
|
|
@ -36,6 +36,11 @@ class CcAgentWrapper(BaseAgentWrapper):
|
|||
DEFAULT_DISALLOWED_TOOLS = ["WebSearch"]
|
||||
MCP_SERVER_NAME = "mcp_server"
|
||||
|
||||
@staticmethod
|
||||
def _claude_usage(usage: dict[str, Any] | None) -> TokenUsage:
|
||||
"""Normalize Claude CLI usage, whose cache dimensions are separate."""
|
||||
return TokenUsage.from_provider(usage or {}, input_includes_cache=False)
|
||||
|
||||
@property
|
||||
def session_path(self) -> Path:
|
||||
"""Directory used for persisted Claude Code sessions."""
|
||||
|
|
@ -248,12 +253,17 @@ class CcAgentWrapper(BaseAgentWrapper):
|
|||
if event_type == "message_delta":
|
||||
delta = raw.get("delta", {})
|
||||
usage = raw.get("usage", {})
|
||||
normalized = cls._claude_usage(usage)
|
||||
return cls._chunk(
|
||||
ChunkEnum.USAGE,
|
||||
session_id=session_id,
|
||||
chunk=json.dumps(usage),
|
||||
output_tokens=usage.get("output_tokens"),
|
||||
metadata={"stop_reason": delta.get("stop_reason")},
|
||||
chunk=json.dumps(normalized.model_dump()),
|
||||
input_tokens=normalized.input_tokens,
|
||||
output_tokens=normalized.output_tokens,
|
||||
metadata={
|
||||
"stop_reason": delta.get("stop_reason"),
|
||||
"usage": normalized.model_dump(),
|
||||
},
|
||||
)
|
||||
|
||||
if event_type == "message_stop":
|
||||
|
|
@ -389,14 +399,16 @@ class CcAgentWrapper(BaseAgentWrapper):
|
|||
"""Convert the SDK terminal result into usage and error chunks."""
|
||||
session_id = msg.session_id or ""
|
||||
usage = msg.usage or {}
|
||||
normalized = cls._claude_usage(usage)
|
||||
chunks = [
|
||||
cls._chunk(
|
||||
ChunkEnum.USAGE,
|
||||
session_id=session_id,
|
||||
chunk=json.dumps(usage),
|
||||
input_tokens=usage.get("input_tokens"),
|
||||
output_tokens=usage.get("output_tokens"),
|
||||
chunk=json.dumps(normalized.model_dump()),
|
||||
input_tokens=normalized.input_tokens,
|
||||
output_tokens=normalized.output_tokens,
|
||||
metadata={
|
||||
"usage": normalized.model_dump(),
|
||||
"duration_ms": msg.duration_ms,
|
||||
"duration_api_ms": msg.duration_api_ms,
|
||||
"stop_reason": msg.stop_reason,
|
||||
|
|
@ -447,7 +459,9 @@ class CcAgentWrapper(BaseAgentWrapper):
|
|||
"session_id": last_msg.session_id or "",
|
||||
"last_message": asdict(last_msg),
|
||||
"result": last_msg.result,
|
||||
"usage": self._claude_usage(last_msg.usage).model_dump(),
|
||||
}
|
||||
self._record_token_usage(TokenUsage.model_validate(result["usage"]))
|
||||
if kwargs.get("output_schema") is not None:
|
||||
result["structured_output"] = last_msg.structured_output
|
||||
return result
|
||||
|
|
@ -523,6 +537,7 @@ class CcAgentWrapper(BaseAgentWrapper):
|
|||
)
|
||||
for chunk in self._result_message_to_chunks(msg):
|
||||
yield chunk
|
||||
self._record_token_usage(self._claude_usage(msg.usage))
|
||||
if reply_open or not emitted_reply_end:
|
||||
emitted_reply_end = True
|
||||
reply_open = False
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@ from typing import Any, TYPE_CHECKING
|
|||
from .base_agent_wrapper import BaseAgentWrapper
|
||||
from ..component_registry import R
|
||||
from ...enumeration import ChunkEnum
|
||||
from ...schema import StreamChunk
|
||||
from ...schema import StreamChunk, TokenUsage
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from openai_codex import AsyncCodex, AsyncThread, CodexConfig, RunInput
|
||||
|
|
@ -79,6 +79,11 @@ class CodexAgentWrapper(BaseAgentWrapper):
|
|||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _codex_usage(usage: Any) -> TokenUsage:
|
||||
"""Normalize Codex's per-turn token usage snapshot."""
|
||||
return TokenUsage.from_provider(usage, input_includes_cache=True)
|
||||
|
||||
# pylint: disable=too-many-arguments
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -421,7 +426,12 @@ class CodexAgentWrapper(BaseAgentWrapper):
|
|||
"last_message": final_response,
|
||||
"result": final_response,
|
||||
"turn": self._serialize(result),
|
||||
"usage": TokenUsage().model_dump(),
|
||||
}
|
||||
if (raw_usage := getattr(result, "usage", None)) is not None:
|
||||
usage = self._codex_usage(raw_usage.last)
|
||||
response["usage"] = usage.model_dump()
|
||||
self._record_token_usage(usage)
|
||||
if kwargs.get("output_schema") is not None:
|
||||
try:
|
||||
response["structured_output"] = json.loads(final_response)
|
||||
|
|
@ -508,13 +518,14 @@ class CodexAgentWrapper(BaseAgentWrapper):
|
|||
]
|
||||
if method == "thread/tokenUsage/updated":
|
||||
usage = payload.token_usage.last
|
||||
data = cls._serialize(usage)
|
||||
normalized = cls._codex_usage(usage)
|
||||
return [
|
||||
make_chunk(
|
||||
ChunkEnum.USAGE,
|
||||
chunk=data,
|
||||
input_tokens=usage.input_tokens,
|
||||
output_tokens=usage.output_tokens,
|
||||
chunk=normalized.model_dump(),
|
||||
input_tokens=normalized.input_tokens,
|
||||
output_tokens=normalized.output_tokens,
|
||||
metadata={"usage": normalized.model_dump()},
|
||||
),
|
||||
]
|
||||
if method == "error":
|
||||
|
|
@ -564,10 +575,13 @@ class CodexAgentWrapper(BaseAgentWrapper):
|
|||
turn = await thread.turn(inputs, **self._turn_kwargs(kwargs))
|
||||
stream = turn.stream()
|
||||
completed = False
|
||||
final_usage: TokenUsage | None = None
|
||||
try:
|
||||
async for event in stream:
|
||||
if event.method == "turn/completed":
|
||||
completed = True
|
||||
if event.method == "thread/tokenUsage/updated":
|
||||
final_usage = self._codex_usage(event.payload.token_usage.last)
|
||||
for chunk in self._event_to_chunks(event, thread.id):
|
||||
yield chunk
|
||||
finally:
|
||||
|
|
@ -577,3 +591,5 @@ class CodexAgentWrapper(BaseAgentWrapper):
|
|||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
self.logger.warning(f"Failed to interrupt Codex turn {turn.id}: {exc}")
|
||||
await stream.aclose()
|
||||
if final_usage is not None:
|
||||
self._record_token_usage(final_usage)
|
||||
|
|
|
|||
|
|
@ -40,6 +40,7 @@ from .file_node import FileNode
|
|||
from .request import Request
|
||||
from .response import Response
|
||||
from .stream_chunk import StreamChunk
|
||||
from .token_usage import TokenUsage
|
||||
|
||||
__all__ = [
|
||||
"ApplicationConfig",
|
||||
|
|
@ -83,5 +84,6 @@ __all__ = [
|
|||
"Response",
|
||||
"SelectedPaper",
|
||||
"StreamChunk",
|
||||
"TokenUsage",
|
||||
"TopicSelectionOutput",
|
||||
]
|
||||
|
|
|
|||
84
reme/schema/token_usage.py
Normal file
84
reme/schema/token_usage.py
Normal file
|
|
@ -0,0 +1,84 @@
|
|||
"""Backend-neutral token accounting contracts."""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
|
||||
|
||||
class TokenUsage(BaseModel):
|
||||
"""Token usage for one completed agent invocation.
|
||||
|
||||
``input_tokens`` is the complete prompt size. Cache-read and cache-write
|
||||
tokens are included in it, but are also exposed separately when the
|
||||
backend reports them. ``reasoning_tokens`` is a subset of output tokens.
|
||||
A ``None`` cache field means that the backend did not report it; it must
|
||||
not be interpreted as zero.
|
||||
"""
|
||||
|
||||
input_tokens: int = Field(default=0, ge=0)
|
||||
output_tokens: int = Field(default=0, ge=0)
|
||||
cache_read_tokens: int | None = Field(default=None, ge=0)
|
||||
cache_write_tokens: int | None = Field(default=None, ge=0)
|
||||
reasoning_tokens: int | None = Field(default=None, ge=0)
|
||||
total_tokens: int = Field(default=0, ge=0)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _set_total(self) -> "TokenUsage":
|
||||
self.total_tokens = self.input_tokens + self.output_tokens
|
||||
return self
|
||||
|
||||
@classmethod
|
||||
def from_provider(
|
||||
cls,
|
||||
usage: Any,
|
||||
*,
|
||||
input_includes_cache: bool,
|
||||
) -> "TokenUsage":
|
||||
"""Normalize a provider usage object or mapping.
|
||||
|
||||
Providers that report cache tokens separately from normal input (for
|
||||
example Claude) pass ``False``. Providers whose input count already
|
||||
includes cached input (for example Codex/OpenAI) pass ``True``.
|
||||
"""
|
||||
|
||||
def get(*names: str) -> int | None:
|
||||
for name in names:
|
||||
value = usage.get(name) if isinstance(usage, dict) else getattr(usage, name, None)
|
||||
if value is not None:
|
||||
return int(value)
|
||||
return None
|
||||
|
||||
reported_input = get("input_tokens", "prompt_tokens") or 0
|
||||
cache_read = get(
|
||||
"cache_read_input_tokens",
|
||||
"cache_input_tokens",
|
||||
"cached_input_tokens",
|
||||
)
|
||||
cache_write = get("cache_creation_input_tokens", "cache_write_input_tokens")
|
||||
input_tokens = reported_input
|
||||
if not input_includes_cache:
|
||||
input_tokens += (cache_read or 0) + (cache_write or 0)
|
||||
return cls(
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=get("output_tokens", "completion_tokens") or 0,
|
||||
cache_read_tokens=cache_read,
|
||||
cache_write_tokens=cache_write,
|
||||
reasoning_tokens=get("reasoning_output_tokens", "reasoning_tokens"),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def combine(cls, usages: list["TokenUsage"]) -> "TokenUsage":
|
||||
"""Combine completed model calls without turning unknown into zero."""
|
||||
optional = ("cache_read_tokens", "cache_write_tokens", "reasoning_tokens")
|
||||
values: dict[str, int | None] = {
|
||||
"input_tokens": sum(item.input_tokens for item in usages),
|
||||
"output_tokens": sum(item.output_tokens for item in usages),
|
||||
}
|
||||
for field in optional:
|
||||
reported = [
|
||||
getattr(item, field)
|
||||
for item in usages
|
||||
if getattr(item, field) is not None
|
||||
]
|
||||
values[field] = sum(reported) if reported else None
|
||||
return cls(**values)
|
||||
|
|
@ -65,3 +65,67 @@ def track_job_counts(job_names: list[str], app_context: "ApplicationContext") ->
|
|||
assert counts == {"search": 2}
|
||||
"""
|
||||
return JobCountTracker(job_names, app_context)
|
||||
|
||||
|
||||
def check_agent_token_count(
|
||||
agent_name: str,
|
||||
app_context: "ApplicationContext",
|
||||
metric: str = "total_tokens",
|
||||
) -> int:
|
||||
"""Return one application-lifetime token metric for an agent wrapper.
|
||||
|
||||
The agent name is the configured ``agent_wrapper`` component name (for
|
||||
example ``"bench"``), and ``metric`` is one leaf in ReMe's token counter
|
||||
tree, such as ``input_tokens`` or ``total_tokens``.
|
||||
"""
|
||||
return global_counter_get(app_context.metadata, ["__token_counter", agent_name, metric])
|
||||
|
||||
|
||||
class AgentTokenCountTracker:
|
||||
"""Measure one token metric for agent wrappers during a context block."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
agent_names: list[str],
|
||||
app_context: "ApplicationContext",
|
||||
metric: str = "total_tokens",
|
||||
) -> None:
|
||||
self.agent_names = list(dict.fromkeys(agent_names))
|
||||
self.app_context = app_context
|
||||
self.metric = metric
|
||||
self._start_counts: dict[str, int] = {}
|
||||
self.counts: dict[str, int] = {}
|
||||
|
||||
def __enter__(self) -> dict[str, int]:
|
||||
self._start_counts = {
|
||||
name: check_agent_token_count(name, self.app_context, self.metric)
|
||||
for name in self.agent_names
|
||||
}
|
||||
return self.counts
|
||||
|
||||
def __exit__(self, exc_type, exc_value, traceback) -> bool:
|
||||
self.counts.update(
|
||||
{
|
||||
name: check_agent_token_count(name, self.app_context, self.metric) - start_count
|
||||
for name, start_count in self._start_counts.items()
|
||||
},
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
def track_agent_token_counts(
|
||||
agent_names: list[str],
|
||||
app_context: "ApplicationContext",
|
||||
metric: str = "total_tokens",
|
||||
) -> AgentTokenCountTracker:
|
||||
"""Return a context manager that reports agent token deltas.
|
||||
|
||||
Example:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
with track_agent_token_counts(["bench"], app.context) as counts:
|
||||
await app.run_job("agentic_answer", query="...")
|
||||
assert counts["bench"] > 0
|
||||
"""
|
||||
return AgentTokenCountTracker(agent_names, app_context, metric)
|
||||
|
|
|
|||
|
|
@ -6,7 +6,13 @@ from types import SimpleNamespace
|
|||
import pytest
|
||||
|
||||
from reme.components.job import BaseJob, StreamJob
|
||||
from reme.utils.evaluation_interface import check_job_count, track_job_counts
|
||||
from reme.utils import global_counter_add
|
||||
from reme.utils.evaluation_interface import (
|
||||
check_agent_token_count,
|
||||
check_job_count,
|
||||
track_agent_token_counts,
|
||||
track_job_counts,
|
||||
)
|
||||
|
||||
|
||||
def test_check_job_count_reads_registered_base_job_count():
|
||||
|
|
@ -86,3 +92,15 @@ def test_track_job_counts_updates_results_when_body_raises():
|
|||
assert counts == {"search": 1}
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
|
||||
def test_track_agent_token_counts_returns_delta_for_one_agent():
|
||||
"""Token tracking mirrors job-count tracking over the token counter tree."""
|
||||
app_context = SimpleNamespace(metadata={}, jobs={})
|
||||
global_counter_add(app_context.metadata, ["__token_counter", "bench", "total_tokens"], 10)
|
||||
|
||||
with track_agent_token_counts(["bench"], app_context) as counts:
|
||||
global_counter_add(app_context.metadata, ["__token_counter", "bench", "total_tokens"], 25)
|
||||
|
||||
assert counts == {"bench": 25}
|
||||
assert check_agent_token_count("bench", app_context) == 35
|
||||
|
|
|
|||
68
tests/unit/test_token_usage.py
Normal file
68
tests/unit/test_token_usage.py
Normal file
|
|
@ -0,0 +1,68 @@
|
|||
"""Tests for unified agent token accounting."""
|
||||
|
||||
from reme.components.agent_wrapper import BaseAgentWrapper
|
||||
from reme.components.application_context import ApplicationContext
|
||||
from reme.schema import TokenUsage
|
||||
from reme.utils import global_counter_get_all
|
||||
|
||||
|
||||
class _UsageWrapper(BaseAgentWrapper):
|
||||
async def reply(self, inputs, **kwargs):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
def test_claude_style_usage_normalizes_cache_into_complete_input():
|
||||
usage = TokenUsage.from_provider(
|
||||
{
|
||||
"input_tokens": 10,
|
||||
"output_tokens": 4,
|
||||
"cache_read_input_tokens": 20,
|
||||
"cache_creation_input_tokens": 30,
|
||||
},
|
||||
input_includes_cache=False,
|
||||
)
|
||||
|
||||
assert usage.model_dump() == {
|
||||
"input_tokens": 60,
|
||||
"output_tokens": 4,
|
||||
"cache_read_tokens": 20,
|
||||
"cache_write_tokens": 30,
|
||||
"reasoning_tokens": None,
|
||||
"total_tokens": 64,
|
||||
}
|
||||
|
||||
|
||||
def test_codex_style_usage_does_not_double_count_cached_input():
|
||||
usage = TokenUsage.from_provider(
|
||||
{
|
||||
"input_tokens": 60,
|
||||
"output_tokens": 4,
|
||||
"cached_input_tokens": 20,
|
||||
"reasoning_output_tokens": 2,
|
||||
},
|
||||
input_includes_cache=True,
|
||||
)
|
||||
|
||||
assert usage.input_tokens == 60
|
||||
assert usage.cache_read_tokens == 20
|
||||
assert usage.reasoning_tokens == 2
|
||||
assert usage.total_tokens == 64
|
||||
|
||||
|
||||
def test_token_counter_is_a_per_agent_metric_tree(tmp_path):
|
||||
context = ApplicationContext(workspace_dir=str(tmp_path))
|
||||
wrapper = _UsageWrapper(name="research", app_context=context)
|
||||
wrapper._record_token_usage( # pylint: disable=protected-access
|
||||
TokenUsage(input_tokens=10, output_tokens=4, cache_read_tokens=6),
|
||||
)
|
||||
wrapper._record_token_usage(TokenUsage(input_tokens=3, output_tokens=2)) # pylint: disable=protected-access
|
||||
|
||||
assert global_counter_get_all(context.metadata, ["__token_counter", "research"]) == {
|
||||
"value": 0,
|
||||
"children": {
|
||||
"input_tokens": {"value": 13, "children": {}},
|
||||
"output_tokens": {"value": 6, "children": {}},
|
||||
"total_tokens": {"value": 19, "children": {}},
|
||||
"cache_read_tokens": {"value": 6, "children": {}},
|
||||
},
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue