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