hallucination / mechanistic_interp /scripts /write_seqprobe_metrics_md.py
ToiTenBao's picture
Upload hallucination folder
a2ffd07 verified
Raw
History Blame Contribute Delete
2.28 kB
"""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()