Beyond_Prompt-based_Retrieval / Biomanus_upload /summarize_bioagent_bench_token_usage.py
czty's picture
Add files using upload-large-folder tool
b2c86fd verified
Raw
History Blame Contribute Delete
6.08 kB
#!/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())