mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-10 22:41:06 +00:00
fix(evaluation): preserve complete token usage metrics
This commit is contained in:
parent
dc66a0e95e
commit
793136a5cf
4 changed files with 82 additions and 11 deletions
|
|
@ -149,12 +149,14 @@ class AsAgentWrapper(BaseAgentWrapper):
|
|||
|
||||
SDK_PACKAGE = "agentscope"
|
||||
|
||||
@staticmethod
|
||||
def _agentscope_usage(usage: Any) -> TokenUsage:
|
||||
def _agentscope_usage(self, usage: Any) -> TokenUsage:
|
||||
"""Normalize AgentScope usage while preserving provider cache semantics."""
|
||||
module = type(usage).__module__ if usage is not None else ""
|
||||
model = self.as_llm.model if self.as_llm is not None else None
|
||||
module = type(model).__module__ if model is not None else ""
|
||||
# Anthropic reports normal, cache-read, and cache-write input tokens
|
||||
# separately; OpenAI-style adapters report prompt tokens inclusive.
|
||||
# separately. AgentScope normalizes every provider's usage into the
|
||||
# same ChatUsage type, so the model implementation identifies the
|
||||
# provider instead of the usage object's module.
|
||||
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):
|
||||
|
|
|
|||
|
|
@ -68,7 +68,12 @@ class TokenUsage(BaseModel):
|
|||
|
||||
@classmethod
|
||||
def combine(cls, usages: list["TokenUsage"]) -> "TokenUsage":
|
||||
"""Combine completed model calls without turning unknown into zero."""
|
||||
"""Combine completed model calls without turning partial data into a total.
|
||||
|
||||
A cache or reasoning metric is reported only when every model call
|
||||
supplied that metric. A sum over only the reporting calls would look
|
||||
complete while silently undercounting the invocation.
|
||||
"""
|
||||
optional = ("cache_read_tokens", "cache_write_tokens", "reasoning_tokens")
|
||||
values: dict[str, int | None] = {
|
||||
"input_tokens": sum(item.input_tokens for item in usages),
|
||||
|
|
@ -76,5 +81,5 @@ class TokenUsage(BaseModel):
|
|||
}
|
||||
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
|
||||
values[field] = sum(reported) if reported and len(reported) == len(usages) else None
|
||||
return cls(**values)
|
||||
|
|
|
|||
|
|
@ -1,4 +1,11 @@
|
|||
"""Read-only evaluation helpers for application job execution statistics."""
|
||||
"""Read-only evaluation helpers for application job execution statistics.
|
||||
|
||||
These helpers take before/after snapshots of application-lifetime counters.
|
||||
They are intentionally not thread-safe request attribution: overlapping calls
|
||||
in the same Application contribute to each other's deltas. They are intended
|
||||
for the benchmark utilities, where each tracked evaluation runs without other
|
||||
work sharing its Application instance.
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
|
|
@ -39,7 +46,10 @@ def check_job_count(job_name: str, app_context: "ApplicationContext") -> int:
|
|||
|
||||
|
||||
class JobCountTracker:
|
||||
"""Measure registered job calls made while this context is active."""
|
||||
"""Measure registered job calls made while this context is active.
|
||||
|
||||
Not thread-safe for per-request attribution; see the module docstring.
|
||||
"""
|
||||
|
||||
def __init__(self, job_names: list[str], app_context: "ApplicationContext") -> None:
|
||||
self.job_names = list(dict.fromkeys(job_names))
|
||||
|
|
@ -105,7 +115,10 @@ def check_agent_token_usage(agent_name: str, app_context: "ApplicationContext")
|
|||
|
||||
|
||||
class AgentTokenCountTracker:
|
||||
"""Measure one token metric for agent wrappers during a context block."""
|
||||
"""Measure one token metric for agent wrappers during a context block.
|
||||
|
||||
Not thread-safe for per-request attribution; see the module docstring.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -154,7 +167,10 @@ def track_agent_token_counts(
|
|||
|
||||
|
||||
class AgentTokenUsageTracker:
|
||||
"""Measure all token metrics for agent wrappers during a context block."""
|
||||
"""Measure all token metrics for agent wrappers during a context block.
|
||||
|
||||
Not thread-safe for per-request attribution; see the module docstring.
|
||||
"""
|
||||
|
||||
def __init__(self, agent_names: list[str], app_context: "ApplicationContext") -> None:
|
||||
self.agent_names = list(dict.fromkeys(agent_names))
|
||||
|
|
|
|||
|
|
@ -1,6 +1,8 @@
|
|||
"""Tests for unified agent token accounting."""
|
||||
|
||||
from reme.components.agent_wrapper import BaseAgentWrapper
|
||||
from agentscope.model._model_usage import ChatUsage
|
||||
|
||||
from reme.components.agent_wrapper import AsAgentWrapper, BaseAgentWrapper
|
||||
from reme.components.application_context import ApplicationContext
|
||||
from reme.schema import TokenUsage
|
||||
from reme.utils import global_counter_get_all
|
||||
|
|
@ -51,6 +53,52 @@ def test_codex_style_usage_does_not_double_count_cached_input():
|
|||
assert usage.total_tokens == 64
|
||||
|
||||
|
||||
def test_agentscope_anthropic_usage_includes_cache_tokens(tmp_path):
|
||||
"""AgentScope uses one usage type, so provider identity comes from its model."""
|
||||
context = ApplicationContext(workspace_dir=str(tmp_path))
|
||||
wrapper = AsAgentWrapper(name="research", as_llm="", app_context=context)
|
||||
wrapper.as_llm = type(
|
||||
"AnthropicLLM",
|
||||
(),
|
||||
{"model": type("AnthropicModel", (), {"__module__": "agentscope.model._anthropic._model"})()},
|
||||
)()
|
||||
usage = ChatUsage(
|
||||
input_tokens=10,
|
||||
output_tokens=4,
|
||||
time=0.0,
|
||||
cache_input_tokens=20,
|
||||
cache_creation_input_tokens=30,
|
||||
)
|
||||
|
||||
assert wrapper._agentscope_usage(usage).model_dump() == { # pylint: disable=protected-access
|
||||
"input_tokens": 60,
|
||||
"output_tokens": 4,
|
||||
"cache_read_tokens": 20,
|
||||
"cache_write_tokens": 30,
|
||||
"reasoning_tokens": None,
|
||||
"total_tokens": 64,
|
||||
}
|
||||
|
||||
|
||||
def test_combined_usage_marks_partially_reported_metrics_as_unknown():
|
||||
"""A partial cache/reasoning sum must not be presented as a full total."""
|
||||
usage = TokenUsage.combine(
|
||||
[
|
||||
TokenUsage(input_tokens=10, output_tokens=4, cache_read_tokens=6),
|
||||
TokenUsage(input_tokens=5, output_tokens=2),
|
||||
],
|
||||
)
|
||||
|
||||
assert usage.model_dump() == {
|
||||
"input_tokens": 15,
|
||||
"output_tokens": 6,
|
||||
"cache_read_tokens": None,
|
||||
"cache_write_tokens": None,
|
||||
"reasoning_tokens": None,
|
||||
"total_tokens": 21,
|
||||
}
|
||||
|
||||
|
||||
def test_token_counter_is_a_per_agent_metric_tree(tmp_path):
|
||||
"""Recorded usage accumulates per agent, and optional metrics track reported calls."""
|
||||
context = ApplicationContext(workspace_dir=str(tmp_path))
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue