hallucination / mechanistic_interp /scripts /plot_probe_latent_metrics.py
ToiTenBao's picture
Upload hallucination folder
a2ffd07 verified
Raw
History Blame Contribute Delete
3.17 kB
#!/usr/bin/env python
"""Plot AUC / F1 vs layer for the latent probes (toilet & bathroom).
Parses the per-layer validation blocks written by train_probe_latent.py to the
run logs and renders a single figure with two panels (AUC, F1), one line per
object. Best layer per object is annotated.
Driver: mechanistic_interp/scripts/plot_probe_latent_metrics.sh
"""
import argparse
import re
from pathlib import Path
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
def parse_log(path):
"""Return {layer: {'acc','auc','f1'}} parsed from a training log."""
rows, cur, cur_metrics = {}, None, {}
keys = {"acc": "Accuracy", "auc": "AUC-ROC", "f1": "F1 "}
for line in Path(path).read_text().splitlines():
layer = re.search(r"Layer (\d+) —", line)
if layer:
if cur is not None and cur_metrics:
rows[cur] = cur_metrics
cur, cur_metrics = int(layer.group(1)), {}
for key, label in keys.items():
hit = re.search(label + r"\s*:\s*([0-9.]+)", line)
if hit:
cur_metrics[key] = float(hit.group(1))
if cur is not None and cur_metrics:
rows[cur] = cur_metrics
return dict(sorted(rows.items()))
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--log_dir", default="mechanistic_interp/logs")
ap.add_argument("--out", default="mechanistic_interp/graph/probe_latent_metrics.png")
ap.add_argument("--objects", nargs="+", default=["toilet", "bathroom"])
ap.add_argument("--title", default="Latent probe — validation metrics per layer",
help="Figure suptitle.")
args = ap.parse_args()
colors = {"toilet": "#d1495b", "bathroom": "#2e86ab"}
data = {obj: parse_log(f"{args.log_dir}/probe_latent_{obj}.log") for obj in args.objects}
fig, axes = plt.subplots(1, 2, figsize=(13, 5), sharex=True)
for metric, ax, title in [("auc", axes[0], "AUC-ROC"), ("f1", axes[1], "F1")]:
for obj in args.objects:
rows = data[obj]
layers = list(rows)
ys = [rows[l][metric] for l in layers]
color = colors.get(obj, None)
ax.plot(layers, ys, marker="o", ms=4, lw=1.8, color=color, label=obj)
best_l = max(layers, key=lambda l: rows[l][metric])
best_y = rows[best_l][metric]
ax.scatter([best_l], [best_y], s=120, facecolors="none",
edgecolors=color, linewidths=2, zorder=5)
ax.annotate(f"L{best_l}\n{best_y:.4f}", (best_l, best_y),
textcoords="offset points", xytext=(0, -28),
ha="center", fontsize=8, color=color)
ax.set_title(f"{title} vs layer", fontsize=12)
ax.set_xlabel("residual-stream layer (hook_resid_post)")
ax.set_ylabel(title)
ax.grid(True, alpha=0.3)
ax.legend(title="object")
fig.suptitle(args.title, fontsize=14)
fig.tight_layout()
out = Path(args.out)
out.parent.mkdir(parents=True, exist_ok=True)
fig.savefig(out, dpi=150, bbox_inches="tight")
print(f"Saved → {out}")
if __name__ == "__main__":
main()