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()