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