#!/usr/bin/env python3 from __future__ import annotations import argparse import json import re from pathlib import Path TASK_ORDER = [ "alzheimer-mouse", "comparative-genomics", "cystic-fibrosis", "deseq", "evolution", "giab", "metagenomics", "single-cell", "transcript-quant", "viral-metagenomics", ] def load_json(path: Path, default=None): if not path.exists(): return default return json.loads(path.read_text(encoding="utf-8")) def as_int(value) -> int: if isinstance(value, bool): return int(value) if isinstance(value, int): return value if isinstance(value, float): return int(value) if isinstance(value, str) and value.strip(): try: return int(float(value)) except ValueError: return 0 return 0 def latest_run_dir(runs_root: Path, task_id: str) -> Path | None: pattern = re.compile(rf"^{re.escape(task_id)}_\d{{8}}_\d{{6}}$") candidates = sorted(run_dir for run_dir in runs_root.iterdir() if run_dir.is_dir() and pattern.match(run_dir.name)) return candidates[-1] if candidates else None def summarize_run(task_id: str, run_dir: Path | None) -> dict: row = { "task": task_id, "run_dir": str(run_dir) if run_dir else None, "cases": 1 if run_dir else 0, "cases_with_usage": 0, "total_prompt_tokens": 0, "total_completion_tokens": 0, "total_tokens": 0, "total_planning_context_tokens": 0, "avg_prompt_tokens_per_case": 0.0, "avg_completion_tokens_per_case": 0.0, "avg_total_tokens_per_case": 0.0, "avg_planning_context_tokens_per_case": 0.0, "avg_llm_calls_per_case": 0.0, "planning_latency_seconds": None, "total_runtime_seconds": None, } if run_dir is None: return row summary = load_json(run_dir / "run_summary.json", default={}) or {} retrieval = load_json(run_dir / "retrieval_plan.json", default={}) or {} token_usage = retrieval.get("token_usage") if not isinstance(token_usage, dict): token_usage = summary prompt_tokens = as_int(token_usage.get("prompt_tokens")) completion_tokens = as_int(token_usage.get("completion_tokens")) total_tokens = as_int(token_usage.get("total_tokens")) llm_call_count = as_int(token_usage.get("llm_call_count")) planning_context_tokens = as_int(retrieval.get("planning_context_tokens", summary.get("planning_context_tokens"))) row.update( { "cases_with_usage": 1 if any(x > 0 for x in (prompt_tokens, completion_tokens, total_tokens)) else 0, "total_prompt_tokens": prompt_tokens, "total_completion_tokens": completion_tokens, "total_tokens": total_tokens, "total_planning_context_tokens": planning_context_tokens, "avg_prompt_tokens_per_case": round(float(prompt_tokens), 2), "avg_completion_tokens_per_case": round(float(completion_tokens), 2), "avg_total_tokens_per_case": round(float(total_tokens), 2), "avg_planning_context_tokens_per_case": round(float(planning_context_tokens), 2), "avg_llm_calls_per_case": round(float(llm_call_count), 2), "planning_latency_seconds": retrieval.get("planning_latency_seconds", summary.get("planning_latency_seconds")), "total_runtime_seconds": retrieval.get("total_runtime_seconds", summary.get("total_runtime_seconds")), } ) return row def summarize_overall(rows: list[dict]) -> dict: present = [row for row in rows if row.get("cases")] n = len(present) if not n: return { "task": "__overall__", "cases": 0, "cases_with_usage": 0, "total_prompt_tokens": 0, "total_completion_tokens": 0, "total_tokens": 0, "total_planning_context_tokens": 0, "avg_prompt_tokens_per_case": 0.0, "avg_completion_tokens_per_case": 0.0, "avg_total_tokens_per_case": 0.0, "avg_planning_context_tokens_per_case": 0.0, "avg_llm_calls_per_case": 0.0, } total_prompt = sum(as_int(row.get("total_prompt_tokens")) for row in present) total_completion = sum(as_int(row.get("total_completion_tokens")) for row in present) total_tokens = sum(as_int(row.get("total_tokens")) for row in present) total_context = sum(as_int(row.get("total_planning_context_tokens")) for row in present) total_calls = sum(as_int(row.get("avg_llm_calls_per_case")) for row in present) return { "task": "__overall__", "cases": n, "cases_with_usage": sum(as_int(row.get("cases_with_usage")) for row in present), "total_prompt_tokens": total_prompt, "total_completion_tokens": total_completion, "total_tokens": total_tokens, "total_planning_context_tokens": total_context, "avg_prompt_tokens_per_case": round(total_prompt / n, 2), "avg_completion_tokens_per_case": round(total_completion / n, 2), "avg_total_tokens_per_case": round(total_tokens / n, 2), "avg_planning_context_tokens_per_case": round(total_context / n, 2), "avg_llm_calls_per_case": round(total_calls / n, 2), } def main() -> int: parser = argparse.ArgumentParser(description="Summarize BioAgentBench real token usage and planning-context tokens.") parser.add_argument("--results-dir", type=Path, required=True) parser.add_argument("--out-json", type=Path, default=None) args = parser.parse_args() rows = [summarize_run(task_id, latest_run_dir(args.results_dir, task_id)) for task_id in TASK_ORDER] rows.append(summarize_overall(rows)) text = json.dumps(rows, indent=2, ensure_ascii=False) print(text) if args.out_json: args.out_json.parent.mkdir(parents=True, exist_ok=True) args.out_json.write_text(text + "\n", encoding="utf-8") return 0 if __name__ == "__main__": raise SystemExit(main())