File size: 5,438 Bytes
9589849 | 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 | #!/usr/bin/env python3
"""Summarize a verifiers traces.jsonl run without hiding errored episodes.
The scheduled-task score is the conservative primary estimate. A second score over
reward-bearing traces helps diagnose transient infrastructure failures, but is never
labelled as the run's result.
"""
from __future__ import annotations
import argparse
import json
import math
import tomllib
from collections import Counter
from pathlib import Path
def wilson(successes: int, total: int, z: float = 1.959963984540054) -> tuple[float, float]:
if total == 0:
return 0.0, 0.0
p = successes / total
denominator = 1.0 + z * z / total
center = (p + z * z / (2 * total)) / denominator
radius = z * math.sqrt(p * (1 - p) / total + z * z / (4 * total * total)) / denominator
return max(0.0, center - radius), min(1.0, center + radius)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("run", type=Path)
args = parser.parse_args()
run_dir = args.run if args.run.is_dir() else args.run.parent
path = run_dir / "traces.jsonl" if args.run.is_dir() else args.run
expected = None
config_path = run_dir / "config.toml"
if config_path.is_file():
config = tomllib.loads(config_path.read_text())
num_tasks = config.get("num_tasks")
num_rollouts = config.get("num_rollouts", 1)
if isinstance(num_tasks, int) and isinstance(num_rollouts, int):
expected = num_tasks * num_rollouts
records: list[dict] = []
traces: list[dict] = []
for line in path.read_text().splitlines():
record = json.loads(line)
records.append(record)
traces.extend(record.get("traces", []))
def has_reward(trace: dict) -> bool:
return isinstance((trace.get("rewards") or {}).get("solved"), dict)
def solved(trace: dict) -> bool:
return float(((trace.get("rewards") or {}).get("solved") or {}).get("score", 0)) > 0
clean = [
trace
for trace in traces
if trace.get("is_completed") and trace.get("ok") and not trace.get("errors")
]
scored = [trace for trace in traces if has_reward(trace)]
errored = [trace for trace in traces if trace.get("errors") or not trace.get("ok")]
successes = sum(solved(trace) for trace in traces)
scheduled = max(len(traces), expected or 0)
scheduled_low, scheduled_high = wilson(successes, scheduled)
scored_low, scored_high = wilson(successes, len(scored))
stop_counts = Counter(t.get("stop_condition", "unknown") for t in traces)
prompt_tokens = 0
completion_tokens = 0
model_calls = 0
completion_tokens_per_call: list[int] = []
calls_per_episode: list[int] = []
for trace in traces:
calls_per_episode.append(len(trace.get("calls", [])))
for call in trace.get("calls", []):
usage = call.get("usage") or {}
prompt_tokens += int(usage.get("prompt_tokens", 0))
call_completion = int(usage.get("completion_tokens", 0))
completion_tokens += call_completion
completion_tokens_per_call.append(call_completion)
model_calls += 1
def percentile(values: list[int], quantile: float) -> int | None:
if not values:
return None
ordered = sorted(values)
return ordered[round((len(ordered) - 1) * quantile)]
result = {
"trace_file": str(path),
"scheduled_episodes": scheduled,
"observed_trace_episodes": len(traces),
"missing_trace_episodes": scheduled - len(traces),
"clean_completed_episodes": len(clean),
"reward_bearing_episodes": len(scored),
"errored_or_unscored_episodes": scheduled - len(scored),
"successes": successes,
"score_scheduled": successes / scheduled if scheduled else None,
"ci95_wilson_scheduled": [scheduled_low, scheduled_high] if scheduled else None,
"diagnostic_score_reward_bearing": successes / len(scored) if scored else None,
"diagnostic_ci95_reward_bearing": [scored_low, scored_high] if scored else None,
"stop_conditions": dict(sorted(stop_counts.items())),
"model_calls": model_calls,
"calls_per_episode_mean": sum(calls_per_episode) / len(calls_per_episode)
if calls_per_episode
else None,
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"completion_tokens_per_call_mean": completion_tokens / model_calls
if model_calls
else None,
"completion_tokens_per_call_p50": percentile(completion_tokens_per_call, 0.50),
"completion_tokens_per_call_p95": percentile(completion_tokens_per_call, 0.95),
"calls_hitting_4096_token_cap": sum(
value >= 4096 for value in completion_tokens_per_call
),
"error_types": dict(
Counter(
str(error.get("type") or error.get("name") or "unknown")
for trace in errored
for error in (trace.get("errors") or [])
)
),
"episode_retry_or_env_error_types": dict(
Counter(
str(error.get("type") or error.get("name") or "unknown")
for record in records
for error in (record.get("errors") or [])
)
),
}
print(json.dumps(result, indent=2))
if __name__ == "__main__":
main()
|