File size: 2,279 Bytes
a2ffd07
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Combine per-layer metric JSONs (from eval_seqprobe_bce.py) into one markdown table."""
import argparse
import json
import os


def table(obj, d):
    lines = [f"## {obj}  (val n={d['n']}, pos={d['n_pos']})", "",
             "| layer | Accuracy | AUC | F1 | Precision | Recall | BCE |",
             "|------:|---------:|----:|---:|----------:|-------:|----:|"]
    pl = d["per_layer"]
    best_auc = max(pl, key=lambda l: pl[l]["auc"])
    best_bce = min(pl, key=lambda l: pl[l]["bce"])
    for l in sorted(pl, key=int):
        m = pl[l]
        star = ""
        if l == best_auc:
            star += " ★AUC"
        if l == best_bce:
            star += " ★BCE"
        lines.append(f"| {l}{star} | {m['accuracy']:.4f} | {m['auc']:.4f} | {m['f1']:.4f} | "
                     f"{m['precision']:.4f} | {m['recall']:.4f} | {m['bce']:.4f} |")
    lines.append("")
    lines.append(f"Best AUC: L{best_auc} ({pl[best_auc]['auc']:.4f}) · "
                 f"Lowest BCE: L{best_bce} ({pl[best_bce]['bce']:.4f})")
    lines.append("")
    return "\n".join(lines)


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--jsons", nargs="+", required=True, help="per-layer metric JSON files")
    ap.add_argument("--out", default="mechanistic_interp/graph/seqprobe_metrics.md")
    ap.add_argument("--title", default="SequenceLayerProbes — per-layer validation metrics")
    args = ap.parse_args()
    parts = [f"# {args.title}", "",
             "Validation metrics per layer (hook_resid_post). "
             "Threshold = logit>0. BCE = BCEWithLogits.", ""]
    for jp in args.jsons:
        d = json.load(open(jp))
        # Two schemas: (a) eval_seqprobe_bce {object,n,n_pos,per_layer};
        # (b) train_probe_latent val_metrics_byvariant {probe_type, per_group_per_layer.overall}.
        if "per_group_per_layer" in d:
            ov = d["per_group_per_layer"]["overall"]
            l0 = ov[next(iter(ov))]
            d = {"object": d.get("probe_type", "?"),
                 "n": l0["n"], "n_pos": l0["pos"], "per_layer": ov}
        parts.append(table(d["object"], d))
    os.makedirs(os.path.dirname(args.out), exist_ok=True)
    open(args.out, "w").write("\n".join(parts))
    print(f"saved {args.out}")


if __name__ == "__main__":
    main()