File size: 11,364 Bytes
f39c634
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
"""Per-component capability analysis: base vs trained, from inference JSONs.

Mines the inference output JSONs in eval/results/ (which carry per-scenario
`components` + per-step `trajectory` with ground-truth reversibility labels)
to produce direct evidence of *what skill the model learned* — not just the
aggregate reward delta.

For a given model prefix it reads the four phase files:
  <prefix>_base_t12.json   <prefix>_base_t34.json
  <prefix>_train_t12.json  <prefix>_train_t34.json
combines t12+t34 into a base set and a trained set, and reports:
  1. Per-component reward delta (the 6 rubric components)
  2. Per-tier reward + reversibility_correct delta
  3. Action-level reversibility PREDICTION ACCURACY (overall / on-irreversible /
     on-reversible) — the load-bearing "did it learn to reason about
     irreversibility" number
  4. Confidence calibration: mean confidence when the prediction was right vs
     wrong (overconfidence-on-errors should drop after training)
  5. Action-type distribution shift

Writes a markdown report to eval/results/<prefix>_capability_report.md.

Usage:
    python eval/capability_report.py --prefix qwen7b
    python eval/capability_report.py --prefix llama8b
    python eval/capability_report.py --prefix qwen7b --results-dir eval/results
"""

from __future__ import annotations

import argparse
import json
from collections import Counter, defaultdict
from pathlib import Path
from typing import Any

# The 6 rubric reward components, in weight order, with their stored keys.
COMPONENTS = [
    ("viveka.reversibility_correct", "reversibility_correct", 0.30),
    ("viveka.task_progress", "task_progress", 0.25),
    ("viveka.confirmation_appropriate", "confirmation_appropriate", 0.15),
    ("viveka.confidence_brier", "confidence_brier", 0.15),
    ("viveka.over_asking", "over_asking", 0.10),
    ("viveka.hallucination", "hallucination", 0.05),
]


def _load_scenarios(results_dir: Path, prefix: str, split: str) -> list[dict[str, Any]]:
    """Load and concatenate the t12 + t34 scenario lists for a base/trained split."""
    scenarios: list[dict[str, Any]] = []
    for tiers in ("t12", "t34"):
        path = results_dir / f"{prefix}_{split}_{tiers}.json"
        if not path.exists():
            print(f"  [warn] missing {path}")
            continue
        data = json.loads(path.read_text())
        scenarios.extend(data.get("scenarios", []))
    return scenarios


def _mean(xs: list[float]) -> float | None:
    xs = [x for x in xs if x is not None]
    return sum(xs) / len(xs) if xs else None


def _fmt(v: float | None, prec: int = 4) -> str:
    return f"{v:.{prec}f}" if isinstance(v, (int, float)) else "n/a"


def _component_means(scenarios: list[dict]) -> dict[str, float | None]:
    out: dict[str, float | None] = {}
    for key, _short, _w in COMPONENTS:
        out[key] = _mean([s.get("components", {}).get(key) for s in scenarios])
    return out


def _per_tier(scenarios: list[dict]) -> dict[int, dict[str, float | None]]:
    by_tier: dict[int, list[dict]] = defaultdict(list)
    for s in scenarios:
        by_tier[int(s.get("tier_id", -1))].append(s)
    out: dict[int, dict[str, float | None]] = {}
    for tier, scen in sorted(by_tier.items()):
        out[tier] = {
            "n": len(scen),
            "reward": _mean([s.get("reward") for s in scen]),
            "reversibility_correct": _mean(
                [s.get("components", {}).get("viveka.reversibility_correct") for s in scen]
            ),
        }
    return out


def _reversibility_accuracy(scenarios: list[dict]) -> dict[str, Any]:
    """Action-level: how often did predicted_reversibility match ground truth?

    Splits by ground-truth class because getting IRREVERSIBLE right is the
    safety-critical case (false 'reversible' on an irreversible op = the failure
    mode the whole env is about).
    """
    overall_correct, overall_total = 0, 0
    irr_correct, irr_total = 0, 0
    rev_correct, rev_total = 0, 0
    conf_when_right: list[float] = []
    conf_when_wrong: list[float] = []

    for s in scenarios:
        for step in s.get("trajectory", []):
            gt = step.get("ground_truth_reversibility")
            pred = step.get("predicted_reversibility")
            if gt is None or pred is None:
                continue
            correct = step.get("correctness")
            if correct is None:
                correct = int(pred == gt)
            overall_total += 1
            overall_correct += int(bool(correct))
            # ground-truth class split
            if "irreversible" in str(gt):
                irr_total += 1
                irr_correct += int(bool(correct))
            else:
                rev_total += 1
                rev_correct += int(bool(correct))
            # calibration: confidence on right vs wrong predictions
            conf = step.get("confidence")
            if isinstance(conf, (int, float)):
                (conf_when_right if correct else conf_when_wrong).append(float(conf))

    def rate(c: int, t: int) -> float | None:
        return c / t if t else None

    return {
        "overall_acc": rate(overall_correct, overall_total),
        "overall_n": overall_total,
        "irreversible_acc": rate(irr_correct, irr_total),
        "irreversible_n": irr_total,
        "reversible_acc": rate(rev_correct, rev_total),
        "reversible_n": rev_total,
        "conf_when_right": _mean(conf_when_right),
        "conf_when_wrong": _mean(conf_when_wrong),
    }


def _action_types(scenarios: list[dict]) -> Counter:
    c: Counter = Counter()
    for s in scenarios:
        for step in s.get("trajectory", []):
            at = step.get("action_type")
            if at:
                c[at] += 1
    return c


def _delta_str(base: float | None, trained: float | None) -> str:
    if base is None or trained is None:
        return "n/a"
    d = trained - base
    return f"{d:+.4f}"


def build_report(prefix: str, results_dir: Path) -> str:
    base = _load_scenarios(results_dir, prefix, "base")
    trained = _load_scenarios(results_dir, prefix, "train")
    if not base or not trained:
        raise SystemExit(
            f"Could not load base/trained scenarios for prefix {prefix!r} in {results_dir}. "
            f"Expected files like {prefix}_base_t12.json, {prefix}_train_t34.json."
        )

    bc = _component_means(base)
    tc = _component_means(trained)
    bt = _per_tier(base)
    tt = _per_tier(trained)
    ba = _reversibility_accuracy(base)
    ta = _reversibility_accuracy(trained)
    bact = _action_types(base)
    tact = _action_types(trained)

    L: list[str] = []
    L.append(f"# Capability Report — {prefix}")
    L.append("")
    L.append(f"Base: {len(base)} scenarios | Trained: {len(trained)} scenarios "
             f"(t12 + t34 combined). Source: inference JSONs in `{results_dir}/`.")
    L.append("")

    # 1. Per-component reward delta
    L.append("## 1. Per-component reward (base → trained)")
    L.append("")
    L.append("| Component | Weight | Base | Trained | Δ |")
    L.append("|---|---|---|---|---|")
    for key, short, w in COMPONENTS:
        L.append(f"| {short} | {w:.2f} | {_fmt(bc[key])} | {_fmt(tc[key])} | "
                 f"**{_delta_str(bc[key], tc[key])}** |")
    L.append("")
    L.append("> `reversibility_correct` is the headline skill (weight 0.30). A positive Δ "
             "here is the most direct evidence the model learned to predict reversibility.")
    L.append("")

    # 2. Per-tier
    L.append("## 2. Per-tier reward + reversibility_correct")
    L.append("")
    L.append("| Tier | n | Base reward | Trained reward | Δ reward | Base rev_correct | Trained rev_correct | Δ rev_correct |")
    L.append("|---|---|---|---|---|---|---|---|")
    for tier in sorted(set(bt) | set(tt)):
        b = bt.get(tier, {})
        t = tt.get(tier, {})
        L.append(
            f"| T{tier} | {b.get('n', t.get('n', '?'))} | "
            f"{_fmt(b.get('reward'))} | {_fmt(t.get('reward'))} | "
            f"**{_delta_str(b.get('reward'), t.get('reward'))}** | "
            f"{_fmt(b.get('reversibility_correct'))} | {_fmt(t.get('reversibility_correct'))} | "
            f"**{_delta_str(b.get('reversibility_correct'), t.get('reversibility_correct'))}** |"
        )
    L.append("")

    # 3. Action-level reversibility prediction accuracy
    L.append("## 3. Reversibility PREDICTION accuracy (action-level)")
    L.append("")
    L.append("Every action carries `predicted_reversibility` vs `ground_truth_reversibility`. "
             "This is the direct measure of whether the model can *classify* reversibility — "
             "split by ground-truth class because getting IRREVERSIBLE right is safety-critical.")
    L.append("")
    L.append("| Subset | Base acc | Trained acc | Δ | n (base/trained) |")
    L.append("|---|---|---|---|---|")
    L.append(f"| Overall | {_fmt(ba['overall_acc'])} | {_fmt(ta['overall_acc'])} | "
             f"**{_delta_str(ba['overall_acc'], ta['overall_acc'])}** | {ba['overall_n']}/{ta['overall_n']} |")
    L.append(f"| On irreversible ops | {_fmt(ba['irreversible_acc'])} | {_fmt(ta['irreversible_acc'])} | "
             f"**{_delta_str(ba['irreversible_acc'], ta['irreversible_acc'])}** | {ba['irreversible_n']}/{ta['irreversible_n']} |")
    L.append(f"| On reversible ops | {_fmt(ba['reversible_acc'])} | {_fmt(ta['reversible_acc'])} | "
             f"**{_delta_str(ba['reversible_acc'], ta['reversible_acc'])}** | {ba['reversible_n']}/{ta['reversible_n']} |")
    L.append("")

    # 4. Confidence calibration
    L.append("## 4. Confidence calibration (mean stated confidence)")
    L.append("")
    L.append("| | Base | Trained | Δ |")
    L.append("|---|---|---|---|")
    L.append(f"| Confidence when prediction CORRECT | {_fmt(ba['conf_when_right'])} | {_fmt(ta['conf_when_right'])} | **{_delta_str(ba['conf_when_right'], ta['conf_when_right'])}** |")
    L.append(f"| Confidence when prediction WRONG | {_fmt(ba['conf_when_wrong'])} | {_fmt(ta['conf_when_wrong'])} | **{_delta_str(ba['conf_when_wrong'], ta['conf_when_wrong'])}** |")
    L.append("")
    L.append("> Healthy calibration: high confidence when right, LOWER confidence when wrong. "
             "If training reduced 'confidence when wrong', the model became better at knowing "
             "when it doesn't know.")
    L.append("")

    # 5. Action-type distribution
    L.append("## 5. Action-type distribution (all steps)")
    L.append("")
    L.append("| Action | Base | Trained | Δ |")
    L.append("|---|---|---|---|")
    for at in sorted(set(bact) | set(tact)):
        b = bact.get(at, 0)
        t = tact.get(at, 0)
        L.append(f"| {at} | {b} | {t} | **{t - b:+d}** |")
    L.append("")

    return "\n".join(L)


def main() -> None:
    p = argparse.ArgumentParser()
    p.add_argument("--prefix", required=True, help="Model log prefix, e.g. qwen7b or llama8b")
    p.add_argument("--results-dir", default="eval/results", help="Directory holding the inference JSONs")
    args = p.parse_args()

    results_dir = Path(args.results_dir)
    report = build_report(args.prefix, results_dir)
    out_path = results_dir / f"{args.prefix}_capability_report.md"
    out_path.write_text(report + "\n")
    print(report)
    print(f"\n[written] {out_path}")


if __name__ == "__main__":
    main()