diff --git a/benchmark/beam/run.py b/benchmark/beam/run.py index 44dff90f..cc278fd0 100644 --- a/benchmark/beam/run.py +++ b/benchmark/beam/run.py @@ -254,10 +254,13 @@ async def answer_question_agentic(app, question: str) -> tuple[str, dict]: """ from reme.utils.evaluation_interface import track_agent_token_usage, track_job_counts - with track_job_counts(["search"], app.context) as tool_counts, track_agent_token_usage( - ["bench"], - app.context, - ) as token_usages: + with ( + track_job_counts(["search"], app.context) as tool_counts, + track_agent_token_usage( + ["bench"], + app.context, + ) as token_usages, + ): query_resp = await app.run_job( "agentic_answer", query=question, @@ -732,14 +735,14 @@ 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)") - 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}") + tool_call_mean, tool_call_std = _mean_and_std(all_tool_call_totals) + print(f" Tool calls/query: mean={tool_call_mean:.2f} std={tool_call_std:.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}") + mean, std = _mean_and_std(values) + print(f" {metric}: mean={mean:.2f} std={std:.2f}") else: print(f" {metric}: unavailable") else: @@ -791,12 +794,12 @@ _TOKEN_USAGE_METRICS = ( ) -def _mean_and_variance(values: list[int]) -> tuple[float, float]: - """Return population mean and variance for one per-question metric.""" +def _mean_and_std(values: list[int]) -> tuple[float, float]: + """Return population mean and standard deviation 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) + return mean, (sum((value - mean) ** 2 for value in values) / len(values)) ** 0.5 if __name__ == "__main__": diff --git a/benchmark/longmemeval/run.py b/benchmark/longmemeval/run.py index 1c8e1708..b7047595 100644 --- a/benchmark/longmemeval/run.py +++ b/benchmark/longmemeval/run.py @@ -433,10 +433,13 @@ 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 tool_counts, track_agent_token_usage( - ["bench"], - app.context, - ) as token_usages: + with ( + track_job_counts(["search"], app.context) as tool_counts, + track_agent_token_usage( + ["bench"], + app.context, + ) as token_usages, + ): query_resp = await app.run_job("agentic_answer", query=question, query_time=query_time) agentic_tool_counts = tool_counts agentic_token_usage = token_usages["bench"] @@ -719,15 +722,15 @@ def _print_summary(results: list[dict], start_time: float) -> None: print("\n ── Agentic (ReAct) ──") print(f" Overall accuracy: {agentic_correct}/{total} ({100*agentic_correct/total:.1f}%)") 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}") + tool_call_mean, tool_call_std = _mean_and_std(tool_call_totals) + print(f" Tool calls/query: mean={tool_call_mean:.2f} std={tool_call_std:.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}") + mean, std = _mean_and_std(values) + print(f" {metric}: mean={mean:.2f} std={std:.2f}") else: print(f" {metric}: unavailable") print(" Per-type accuracy:") @@ -753,12 +756,12 @@ _TOKEN_USAGE_METRICS = ( ) -def _mean_and_variance(values: list[int]) -> tuple[float, float]: - """Return population mean and variance for one per-query metric.""" +def _mean_and_std(values: list[int]) -> tuple[float, float]: + """Return population mean and standard deviation 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) + return mean, (sum((value - mean) ** 2 for value in values) / len(values)) ** 0.5 if __name__ == "__main__": diff --git a/reme/components/agent_wrapper/base_agent_wrapper.py b/reme/components/agent_wrapper/base_agent_wrapper.py index 09ce7c22..46006122 100644 --- a/reme/components/agent_wrapper/base_agent_wrapper.py +++ b/reme/components/agent_wrapper/base_agent_wrapper.py @@ -233,21 +233,21 @@ class BaseAgentWrapper(BaseComponent): """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): + counters = getattr(self.app_context, "metadata", None) + if not isinstance(counters, dict): return for field in ("input_tokens", "output_tokens", "total_tokens"): global_counter_add( - metadata, + counters, [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) + global_counter_add(counters, [self.TOKEN_COUNTER_PREFIX, self.name, field], value) global_counter_add( - metadata, + counters, [self.TOKEN_COUNTER_PREFIX, self.name, f"{field}_reported_calls"], 1, ) diff --git a/reme/schema/token_usage.py b/reme/schema/token_usage.py index 86557b08..27ad9518 100644 --- a/reme/schema/token_usage.py +++ b/reme/schema/token_usage.py @@ -75,10 +75,6 @@ class TokenUsage(BaseModel): "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 - ] + 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) diff --git a/reme/utils/evaluation_interface.py b/reme/utils/evaluation_interface.py index 0e87d2db..ea5e4306 100644 --- a/reme/utils/evaluation_interface.py +++ b/reme/utils/evaluation_interface.py @@ -121,8 +121,7 @@ class AgentTokenCountTracker: 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 + name: check_agent_token_count(name, self.app_context, self.metric) for name in self.agent_names } return self.counts @@ -165,9 +164,7 @@ class AgentTokenUsageTracker: 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_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") diff --git a/tests/unit/test_token_usage.py b/tests/unit/test_token_usage.py index 5affd28f..d0dd5b04 100644 --- a/tests/unit/test_token_usage.py +++ b/tests/unit/test_token_usage.py @@ -12,6 +12,7 @@ class _UsageWrapper(BaseAgentWrapper): def test_claude_style_usage_normalizes_cache_into_complete_input(): + """Cache-excluded provider input is normalized into the complete input total.""" usage = TokenUsage.from_provider( { "input_tokens": 10, @@ -33,6 +34,7 @@ def test_claude_style_usage_normalizes_cache_into_complete_input(): def test_codex_style_usage_does_not_double_count_cached_input(): + """Cache-inclusive provider input is preserved without adding cached tokens twice.""" usage = TokenUsage.from_provider( { "input_tokens": 60, @@ -50,6 +52,7 @@ def test_codex_style_usage_does_not_double_count_cached_input(): 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)) wrapper = _UsageWrapper(name="research", app_context=context) wrapper._record_token_usage( # pylint: disable=protected-access