burtenshaw's picture
burtenshaw HF Staff
feat: publish beam pi study source
5741b22 verified
Raw History Blame Contribute Delete
14.1 kB
#!/usr/bin/env python3
"""Post-hoc, task-weighted ProgramBench pilot curves. Never uses best-so-far scores."""
from __future__ import annotations
import argparse
from bisect import bisect_right
from contextlib import closing
import csv
from datetime import datetime, timezone
import json
import math
from pathlib import Path
import sqlite3
SCHEDULE = [0, 30, 60, 120, 300, 600, 1200, 2400, 4800, 7200]
CONDITIONS = ["single", "peers", "async"]
AXES = ["wall_seconds", "reference_seconds", "api_prefill_reference_seconds"]
def finite(value):
return isinstance(value, (int, float)) and not isinstance(value, bool) and math.isfinite(value)
def read_json(path, default):
return json.loads(path.read_text()) if path.exists() else default
def epoch(value):
return datetime.fromisoformat(value.replace("Z", "+00:00")).timestamp()
def clock_trace(events, requests, prefill_tps=10_000, decode_tps=265):
"""Two causal clocks from gateway operations and measured tool durations.
Primary prefill counts estimated new context input; sensitivity prefill counts
every billed prompt token. Gateway start/end times also capture compaction calls
that bypass Pi's regular generation hooks. Only completed operations add time.
"""
timeline, warnings = [], []
origin = next((epoch(e["utc"]) - e.get("elapsed_ms", 0) / 1000
for e in events if e.get("utc")), None)
for index, event in enumerate(events):
if finite(event.get("elapsed_ms")):
timeline.append((event["elapsed_ms"] / 1000, index, event))
ledger_ids = {r.get("client_request_id") for r in requests if r.get("client_request_id")}
observed_ids = {e.get("request_id") for e in events if e.get("type") == "model_request_end"}
if observed_ids - ledger_ids:
warnings.append("model_events_missing_from_gateway_ledger")
for index, row in enumerate(requests):
if row.get("status") == "rejected":
continue
output, prompt, context = row.get("completion_tokens"), row.get("prompt_tokens"), row.get("context_charged")
if not all(finite(v) and v >= 0 for v in (output, prompt, context)) or not row.get("ended_at"):
warnings.append("request_has_unknown_usage_or_end_time")
continue
if origin is None:
warnings.append("missing_episode_clock_origin")
continue
start, end = epoch(row["started_at"]) - origin, epoch(row["ended_at"]) - origin
if end < start:
warnings.append("invalid_gateway_timestamps")
continue
ident = row["id"]
durations = (max(0, context - output) / prefill_tps + output / decode_tps,
prompt / prefill_tps + output / decode_tps)
timeline.extend([(start, -2, {"type": "clock_model_start", "agent_id": row["agent_id"], "id": ident}),
(end, -1, {"type": "clock_model_end", "agent_id": row["agent_id"], "id": ident,
"durations": durations})])
agents, sent, starts = {}, {}, {}
trace = [(0.0, 0.0, 0.0)]
snapshot_clocks = {}
for wall, _, event in sorted(timeline, key=lambda row: (row[0], row[1])):
kind, agent = event.get("type"), event.get("agent_id")
if agent:
agents.setdefault(agent, (0.0, 0.0))
if kind == "agent_created":
agents[agent] = agents.get(event.get("parent_id"), (0.0, 0.0))
elif kind == "message_sent":
sent[event.get("sequence")] = agents.get(agent, (0.0, 0.0))
elif kind == "message_delivered":
inherited = sent.get(event.get("causal_event_id"))
if inherited is None:
warnings.append("missing_message_send_event")
else:
agents[agent] = tuple(max(a, b) for a, b in zip(agents[agent], inherited))
elif kind in ("clock_model_start", "tool_start"):
starts[(kind, event.get("id", event.get("sequence")))] = agents.get(agent, (0.0, 0.0))
elif kind in ("clock_model_end", "tool_end"):
if kind == "clock_model_end":
key, durations = ("clock_model_start", event["id"]), event["durations"]
else:
duration = event.get("duration_ms")
if not finite(duration) or duration < 0:
warnings.append("invalid_tool_duration")
continue
key, durations = ("tool_start", event.get("causal_event_id")), (duration / 1000,) * 2
anchor = starts.get(key, agents.get(agent, (0.0, 0.0)))
agents[agent] = tuple(max(current, begin + duration)
for current, begin, duration in zip(agents[agent], anchor, durations))
current = tuple(max((clock[i] for clock in agents.values()), default=0.0) for i in (0, 1))
trace.append((max(0, wall), *current))
if kind == "snapshot" and event.get("path"):
snapshot_clocks[str(event["path"])] = current
if not events:
warnings.append("missing_runner_events")
return {"trace": trace, "snapshot_clocks": snapshot_clocks,
"valid": not warnings, "warnings": sorted(set(warnings))}
def clocks_at(clocks, wall, snapshot_path=None):
if snapshot_path in clocks["snapshot_clocks"]:
return clocks["snapshot_clocks"][snapshot_path]
rows = clocks["trace"]
index = bisect_right([row[0] for row in rows], wall) - 1
return rows[max(0, index)][1:]
def score_at(points, time_value, axis):
eligible = [point for point in points if point[axis] <= time_value]
if not eligible:
return 0.0, False
# Ordering includes original wall time for tied normalized clocks. Regressions stay.
latest = max(eligible, key=lambda point: (point[axis], point["wall_seconds"], point["order"]))
return latest["score"], True
def episode_data(root, run, requests, prefill_tps, decode_tps):
directory = root / "episodes" / run["run_id"]
result = read_json(directory / "runner" / "result.json", {})
controller_status = read_json(directory / "status.json", {})
event_path = directory / "runner" / "events.jsonl"
events = [json.loads(line) for line in event_path.read_text().splitlines() if line.strip()] if event_path.exists() else []
clocks = clock_trace(events, requests, prefill_tps, decode_tps)
grades = read_json(directory / "grades.json", [])
if isinstance(grades, dict):
grades = grades.get("grades", [])
points, invalid_grades = [], 0
for order, grade in enumerate(grades):
score = grade.get("score")
valid = grade.get("valid", True)
if score is None and grade.get("sha256"):
summary = read_json(directory / "evaluations" / grade["sha256"] / run["task"] / "summary.json", {})
score = summary.get("analysis_score", summary.get("score"))
valid = valid and summary.get("valid", summary.get("complete", False))
wall = grade.get("elapsed_seconds")
if not valid or not finite(score) or not 0 <= score <= 1 or not finite(wall) or wall < 0:
invalid_grades += 1
continue
primary, sensitivity = clocks_at(clocks, wall, grade.get("snapshot_path"))
points.append({"order": order, "wall_seconds": wall, "reference_seconds": primary,
"api_prefill_reference_seconds": sensitivity, "score": score})
points.sort(key=lambda point: (point["wall_seconds"], point["order"]))
last_score = points[-1]["score"] if points else 0.0
stop_reason = controller_status.get("stop_reason", result.get("stop_reason", "missing_result"))
status = controller_status.get("status")
if status not in {"completed", "failed", "interrupted", "budget_exhausted", "running", "pending"}:
status = ("completed" if stop_reason in {"completed", "wall_time_limit", "wall_time_exhausted"}
else "budget_exhausted" if "budget" in stop_reason or stop_reason == "daily_quota_headroom"
else "interrupted" if stop_reason == "interrupted" else "failed")
return {"run_id": run["run_id"], "task": run["task"], "condition": run["condition"],
"status": status, "stop_reason": stop_reason, "elapsed_seconds": result.get("elapsed_seconds"),
"final_score": last_score, "has_valid_grade": bool(points), "invalid_grade_count": invalid_grades,
"clock_valid": clocks["valid"], "clock_warnings": clocks["warnings"], "points": points,
"api_tokens_charged": sum(row.get("charged", 0) for row in requests)}
def aggregate_at(episodes, tasks, condition, time_value, axis):
selected = [row for row in episodes if row["condition"] == condition]
clock_valid = axis == "wall_seconds" or all(row["clock_valid"] for row in selected)
task_scores, observed = [], 0
for task in tasks:
repetitions = [row for row in selected if row["task"] == task]
values = [score_at(row["points"], time_value, axis) for row in repetitions]
task_scores.append(sum(value for value, _ in values) / len(values) if values else 0.0)
observed += int(any(valid for _, valid in values))
return {"axis": axis, "time_seconds": time_value, "condition": condition,
"mean_hidden_test_fraction": sum(task_scores) / len(tasks) if clock_valid else None,
"tasks_total": len(tasks), "tasks_with_valid_grade": observed,
"episodes_configured": len(selected), "clock_valid": clock_valid}
def analyze(experiment_dir, config, prefill_tps=10_000, decode_tps=265):
if prefill_tps <= 0 or decode_tps <= 0:
raise ValueError("Reference rates must be positive")
root = Path(experiment_dir)
tasks = [task["instance_id"] for task in config["tasks"]]
if not tasks or len(tasks) != len(set(tasks)):
raise ValueError("Selected task IDs must be nonempty and unique")
runs = [row for row in config["runs"] if row.get("category", "benchmark") == "benchmark"]
if len({run["run_id"] for run in runs}) != len(runs):
raise ValueError("Duplicate run ID")
if any(run["task"] not in tasks or run["condition"] not in CONDITIONS for run in runs):
raise ValueError("Unexpected task or condition")
requests = []
ledger = root / "budget.sqlite"
if ledger.exists():
with closing(sqlite3.connect(f"file:{ledger.resolve()}?mode=ro", uri=True)) as db:
db.row_factory = sqlite3.Row
requests = [dict(row) for row in db.execute("SELECT * FROM requests")]
episodes = [episode_data(root, run, [row for row in requests if row["run_id"] == run["run_id"]],
prefill_tps, decode_tps) for run in runs]
curves = []
for axis in AXES:
grid = set(SCHEDULE if axis == "wall_seconds" else [0.0])
grid.update(point[axis] for row in episodes for point in row["points"])
for instant in sorted(grid):
curves.extend(aggregate_at(episodes, tasks, condition, instant, axis) for condition in CONDITIONS)
final = {}
for condition in CONDITIONS:
scores = []
selected = [row for row in episodes if row["condition"] == condition]
for task in tasks:
values = [row["final_score"] for row in selected if row["task"] == task]
scores.append(sum(values) / len(values) if values else 0.0)
final[condition] = {"mean_hidden_test_fraction": sum(scores) / len(tasks),
"tasks_total": len(tasks),
"tasks_with_valid_grade": len({row["task"] for row in selected if row["has_valid_grade"]}),
"episodes_configured": len(selected),
"episodes_completed": sum(row["status"] == "completed" for row in selected)}
summary = {"generated_at_utc": datetime.now(timezone.utc).isoformat(), "tasks": tasks,
"aggregation": "mean hidden-test fraction within each task, then equal-weight mean over every selected task",
"failure_policy": "latest valid score carried forward; zero before any valid grade; no best-so-far filtering",
"reference_clock": {"decode_tokens_per_second": decode_tps, "prefill_tokens_per_second": prefill_tps,
"primary_prefill": "estimated newly entering context tokens",
"sensitivity_prefill": "all API prompt tokens",
"label": "operational normalization; assumed reference rates; not Anthropic's unknown constants",
"includes": "completed model operations, actual tool durations, parent/message causal handoffs",
"limitation": "no partial in-flight decode/tool progress at snapshots"},
"uncertainty": "exploratory five-task convenience sample; one repetition; no confidence intervals",
"final": final, "episodes": episodes}
root.mkdir(parents=True, exist_ok=True)
(root / "summary.json").write_text(json.dumps(summary, indent=2) + "\n")
with (root / "curves.csv").open("w", newline="") as stream:
writer = csv.DictWriter(stream, fieldnames=list(curves[0]))
writer.writeheader()
writer.writerows(curves)
return summary, curves
def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--experiment-dir", type=Path, required=True)
parser.add_argument("--config", type=Path, required=True)
parser.add_argument("--reference-prefill-tps", type=float, default=10_000)
parser.add_argument("--reference-decode-tps", type=float, default=265)
args = parser.parse_args()
summary, curves = analyze(args.experiment_dir, json.loads(args.config.read_text()),
args.reference_prefill_tps, args.reference_decode_tps)
print(json.dumps({"summary": str(args.experiment_dir / "summary.json"),
"curves": str(args.experiment_dir / "curves.csv"), "final": summary["final"]}))
if __name__ == "__main__":
main()