File size: 6,075 Bytes
b2c86fd | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 | #!/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())
|