sol-max-record / harness /scripts /summarize_eval.py
simonycl's picture
Upload folder using huggingface_hub
9589849 verified
Raw
History Blame Contribute Delete
5.44 kB
#!/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()