benchmark输出完整token消耗统计

This commit is contained in:
酱牛肉 2026-07-30 14:43:33 +08:00
parent 26f221e765
commit a028cae33d
6 changed files with 182 additions and 29 deletions

View file

@ -252,12 +252,12 @@ async def answer_question_agentic(app, question: str) -> tuple[str, dict]:
Returns (answer, metadata)
"""
from reme.utils.evaluation_interface import track_agent_token_counts, track_job_counts
from reme.utils.evaluation_interface import track_agent_token_usage, track_job_counts
with track_job_counts(["search"], app.context) as counts, track_agent_token_counts(
with track_job_counts(["search"], app.context) as tool_counts, track_agent_token_usage(
["bench"],
app.context,
) as token_counts:
) as token_usages:
query_resp = await app.run_job(
"agentic_answer",
query=question,
@ -266,8 +266,8 @@ async def answer_question_agentic(app, question: str) -> tuple[str, dict]:
return answer, {
"mode": "agentic",
"search_calls": counts["search"],
"token_count": token_counts["bench"],
"tool_counts": tool_counts,
"token_usage": token_usages["bench"],
}
@ -459,9 +459,9 @@ async def evaluate_case(eval_config: dict, case_id: str, eval_only: bool = False
agentic_answer = "(no answer generated)"
logger.info(f"[Case {case_id}] Agentic answer: {agentic_answer[:200]}...")
logger.info(
f"[Case {case_id}] Agentic search calls: {agentic_meta.get('search_calls', 0)}",
f"[Case {case_id}] Agentic tool calls: {agentic_meta.get('tool_counts', {})}",
)
logger.info(f"[Case {case_id}] Bench token usage: {agentic_meta.get('token_count', 0)}")
logger.info(f"[Case {case_id}] Bench token usage: {agentic_meta.get('token_usage', {})}")
# Judge agentic answer
logger.info(f"[Case {case_id}] Judging agentic ({q_type})...")
@ -692,8 +692,8 @@ def main( # pylint: disable=too-many-statements
type_binary_scores: dict[str, list[float]] = {}
all_scores: list[float] = []
all_binary_scores: list[float] = []
all_search_calls: list[int] = []
all_token_counts: list[int] = []
all_tool_call_totals: list[int] = []
all_token_usages: list[dict[str, int | None]] = []
for case_result in results:
if "error" in case_result:
@ -716,8 +716,9 @@ def main( # pylint: disable=too-many-statements
type_binary_scores[qtype].append(binary_score)
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))
metadata = q.get("agentic_metadata", {})
all_tool_call_totals.append(sum(metadata.get("tool_counts", {}).values()))
all_token_usages.append(metadata.get("token_usage", {}))
print("\n ── AGENTIC ──")
if all_scores:
@ -731,10 +732,16 @@ def main( # pylint: disable=too-many-statements
binary_overall = sum(all_binary_scores) / len(all_binary_scores) if all_binary_scores else 0
print(f" {'-'*38}")
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}")
tool_call_mean, tool_call_variance = _mean_and_variance(all_tool_call_totals)
print(f" Tool calls/query: mean={tool_call_mean:.2f} variance={tool_call_variance:.2f}")
print(" Bench tokens/query:")
for metric in _TOKEN_USAGE_METRICS:
values = [usage[metric] for usage in all_token_usages if usage.get(metric) is not None]
if values:
mean, variance = _mean_and_variance(values)
print(f" {metric}: mean={mean:.2f} variance={variance:.2f}")
else:
print(f" {metric}: unavailable")
else:
print(" (no results)")
@ -774,6 +781,24 @@ def main( # pylint: disable=too-many-statements
print("=" * 70 + "\n")
_TOKEN_USAGE_METRICS = (
"input_tokens",
"output_tokens",
"cache_read_tokens",
"cache_write_tokens",
"reasoning_tokens",
"total_tokens",
)
def _mean_and_variance(values: list[int]) -> tuple[float, float]:
"""Return population mean and variance for one per-question metric."""
if not values:
return 0.0, 0.0
mean = sum(values) / len(values)
return mean, sum((value - mean) ** 2 for value in values) / len(values)
if __name__ == "__main__":
import argparse

View file

@ -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_agent_token_counts, track_job_counts
from reme.utils.evaluation_interface import track_agent_token_usage, track_job_counts
reme_cfg = eval_config["reme"]
dream_trigger_hour = reme_cfg.get("dream_trigger_hour", 23)
@ -433,20 +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, track_agent_token_counts(
with track_job_counts(["search"], app.context) as tool_counts, track_agent_token_usage(
["bench"],
app.context,
) as token_counts:
) as token_usages:
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_tool_counts = tool_counts
agentic_token_usage = token_usages["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}")
logger.info(f"[Item {item_index}] Agentic tool calls: {agentic_tool_counts}")
logger.info(f"[Item {item_index}] Bench token usage: {agentic_token_usage}")
# ── Phase 5: Judge agentic response (via answer_judge_step) ──────────
logger.info(f"[Item {item_index}] Judging agentic (binary, type={item['question_type']})...")
@ -469,8 +469,8 @@ async def evaluate_item(item: dict, eval_config: dict, item_index: int, eval_onl
"ground_truth": item["answer"],
"agentic_response": agentic_response,
"agentic_judgment": agentic_judgment,
"agentic_search_calls": agentic_search_calls,
"agentic_token_count": agentic_token_count,
"agentic_tool_counts": agentic_tool_counts,
"agentic_token_usage": agentic_token_usage,
"sessions_ingested": len(sorted_sessions),
"dreams_triggered": len(dream_dates_triggered),
}
@ -718,10 +718,18 @@ def _print_summary(results: list[dict], start_time: float) -> None:
# Agentic stats
print("\n ── Agentic (ReAct) ──")
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}")
tool_call_totals = [sum(r.get("agentic_tool_counts", {}).values()) for r in results]
tool_call_mean, tool_call_variance = _mean_and_variance(tool_call_totals)
print(f" Tool calls/query: mean={tool_call_mean:.2f} variance={tool_call_variance:.2f}")
token_usages = [r.get("agentic_token_usage", {}) for r in results]
print(" Bench tokens/query:")
for metric in _TOKEN_USAGE_METRICS:
values = [usage[metric] for usage in token_usages if usage.get(metric) is not None]
if values:
mean, variance = _mean_and_variance(values)
print(f" {metric}: mean={mean:.2f} variance={variance:.2f}")
else:
print(f" {metric}: unavailable")
print(" Per-type accuracy:")
for qtype, stats in sorted(agentic_type_stats.items()):
acc = 100 * stats["correct"] / stats["total"] if stats["total"] else 0
@ -735,6 +743,24 @@ def _print_summary(results: list[dict], start_time: float) -> None:
print("=" * 60 + "\n")
_TOKEN_USAGE_METRICS = (
"input_tokens",
"output_tokens",
"cache_read_tokens",
"cache_write_tokens",
"reasoning_tokens",
"total_tokens",
)
def _mean_and_variance(values: list[int]) -> tuple[float, float]:
"""Return population mean and variance for one per-query metric."""
if not values:
return 0.0, 0.0
mean = sum(values) / len(values)
return mean, sum((value - mean) ** 2 for value in values) / len(values)
if __name__ == "__main__":
import argparse

View file

@ -246,6 +246,11 @@ class BaseAgentWrapper(BaseComponent):
value = getattr(usage, field)
if value is not None:
global_counter_add(metadata, [self.TOKEN_COUNTER_PREFIX, self.name, field], value)
global_counter_add(
metadata,
[self.TOKEN_COUNTER_PREFIX, self.name, f"{field}_reported_calls"],
1,
)
@abstractmethod
async def reply(self, inputs: Any, **kwargs) -> dict:

View file

@ -3,13 +3,21 @@
from typing import TYPE_CHECKING
from ..components.job import BackgroundJob, BaseJob, CronJob, StreamJob
from .counter import global_counter_get
from .counter import global_counter_get, global_counter_get_all
if TYPE_CHECKING:
from ..components.application_context import ApplicationContext
_JOB_ENTRY_CLASSES = {CronJob, StreamJob, BackgroundJob, BaseJob}
_TOKEN_METRICS = (
"input_tokens",
"output_tokens",
"cache_read_tokens",
"cache_write_tokens",
"reasoning_tokens",
"total_tokens",
)
def check_job_count(job_name: str, app_context: "ApplicationContext") -> int:
@ -81,6 +89,21 @@ def check_agent_token_count(
return global_counter_get(app_context.metadata, ["__token_counter", agent_name, metric])
def check_agent_token_usage(agent_name: str, app_context: "ApplicationContext") -> dict[str, int | None]:
"""Return all token metrics currently accumulated for one agent wrapper.
Optional cache and reasoning metrics remain ``None`` until the backend has
reported them at least once. This keeps unknown usage distinct from zero.
"""
tree = global_counter_get_all(app_context.metadata, ["__token_counter", agent_name])
children = tree.get("children", {}) if tree is not None else {}
usage: dict[str, int | None] = {}
for metric in _TOKEN_METRICS:
node = children.get(metric)
usage[metric] = node["value"] if node is not None else None
return usage
class AgentTokenCountTracker:
"""Measure one token metric for agent wrappers during a context block."""
@ -129,3 +152,51 @@ def track_agent_token_counts(
assert counts["bench"] > 0
"""
return AgentTokenCountTracker(agent_names, app_context, metric)
class AgentTokenUsageTracker:
"""Measure all token metrics for agent wrappers during a context block."""
def __init__(self, agent_names: list[str], app_context: "ApplicationContext") -> None:
self.agent_names = list(dict.fromkeys(agent_names))
self.app_context = app_context
self._start_usage: dict[str, dict[str, int | None]] = {}
self._start_report_counts: dict[str, dict[str, int]] = {}
self.usages: dict[str, dict[str, int | None]] = {}
def __enter__(self) -> dict[str, dict[str, int | None]]:
self._start_usage = {
name: check_agent_token_usage(name, self.app_context) for name in self.agent_names
}
self._start_report_counts = {
name: {
metric: check_agent_token_count(name, self.app_context, f"{metric}_reported_calls")
for metric in ("cache_read_tokens", "cache_write_tokens", "reasoning_tokens")
}
for name in self.agent_names
}
return self.usages
def __exit__(self, exc_type, exc_value, traceback) -> bool:
for name in self.agent_names:
end_usage = check_agent_token_usage(name, self.app_context)
delta: dict[str, int | None] = {}
for metric in _TOKEN_METRICS:
current = end_usage[metric]
start = self._start_usage[name][metric]
if metric in self._start_report_counts[name]:
end_reports = check_agent_token_count(name, self.app_context, f"{metric}_reported_calls")
if end_reports == self._start_report_counts[name][metric]:
delta[metric] = None
continue
delta[metric] = (current or 0) - (start or 0)
self.usages[name] = delta
return False
def track_agent_token_usage(
agent_names: list[str],
app_context: "ApplicationContext",
) -> AgentTokenUsageTracker:
"""Return a context manager that reports full per-agent token usage deltas."""
return AgentTokenUsageTracker(agent_names, app_context)

View file

@ -9,7 +9,9 @@ from reme.components.job import BaseJob, StreamJob
from reme.utils import global_counter_add
from reme.utils.evaluation_interface import (
check_agent_token_count,
check_agent_token_usage,
check_job_count,
track_agent_token_usage,
track_agent_token_counts,
track_job_counts,
)
@ -104,3 +106,26 @@ def test_track_agent_token_counts_returns_delta_for_one_agent():
assert counts == {"bench": 25}
assert check_agent_token_count("bench", app_context) == 35
def test_track_agent_token_usage_preserves_unreported_cache_as_none():
"""Detailed usage tracking does not turn an unknown cache value into zero."""
app_context = SimpleNamespace(metadata={}, jobs={})
for metric, value in (("input_tokens", 10), ("output_tokens", 5), ("total_tokens", 15)):
global_counter_add(app_context.metadata, ["__token_counter", "bench", metric], value)
with track_agent_token_usage(["bench"], app_context) as usages:
for metric, value in (("input_tokens", 20), ("output_tokens", 7), ("total_tokens", 27)):
global_counter_add(app_context.metadata, ["__token_counter", "bench", metric], value)
assert usages == {
"bench": {
"input_tokens": 20,
"output_tokens": 7,
"cache_read_tokens": None,
"cache_write_tokens": None,
"reasoning_tokens": None,
"total_tokens": 27,
},
}
assert check_agent_token_usage("bench", app_context)["cache_read_tokens"] is None

View file

@ -64,5 +64,6 @@ def test_token_counter_is_a_per_agent_metric_tree(tmp_path):
"output_tokens": {"value": 6, "children": {}},
"total_tokens": {"value": 19, "children": {}},
"cache_read_tokens": {"value": 6, "children": {}},
"cache_read_tokens_reported_calls": {"value": 1, "children": {}},
},
}