hallucination / experiment /probing /plot_exemplar_heatmaps.py
ToiTenBao's picture
Upload hallucination folder
a2ffd07 verified
Raw
History Blame Contribute Delete
4.97 kB
# experiment/probing/plot_exemplar_heatmaps.py
"""Render the heatmap / trajectory / per-feature figures per image_id.
Pure parquet reader — no model loading, no GPU required.
"""
from __future__ import annotations
import argparse
import json
import os
import sys
import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "../.."))
from experiment.probing._helpers import heatmap_matrix_from_parquet
METHODS = ["base", "adv", "efuf", "nullu", "lora_finetune"]
def _load_tau(tau_path: str, alpha: float) -> dict[str, float]:
with open(tau_path) as f:
d = json.load(f)
return {m: d[m][f"{alpha:.2f}"] for m in METHODS if m in d}
def _flatten_feature_ids(features_path: str, set_key: str) -> list[int]:
with open(features_path) as f:
feats = json.load(f)
ids: set[int] = set()
for _, v in feats[set_key].items():
ids.update(v)
return sorted(ids)
def render_heatmap(parquet_path: str, output_dir: str,
feature_ids: list[int], tau: dict[str, float]) -> None:
fig, axes = plt.subplots(1, len(METHODS), figsize=(4 * len(METHODS), 6), sharey=True)
vmin, vmax = 0.0, max(tau.values()) * 2.0
im = None
for ax, m in zip(axes, METHODS):
try:
mat, tokens, is_tt = heatmap_matrix_from_parquet(
parquet_path, method=m, image_id=_iid_from_parquet(parquet_path),
pass_="teacher", feature_ids=feature_ids,
)
except ValueError:
ax.set_title(f"{m}\n(no data)")
ax.axis("off")
continue
im = ax.imshow(mat, aspect="auto", origin="lower", cmap="hot", vmin=vmin, vmax=vmax)
ax.set_title(f"{m} τ_c={tau.get(m, float('nan')):.2f}")
ax.set_xlabel("token")
ax.set_xticks(range(len(tokens)))
ax.set_xticklabels(
[t.strip() for t in tokens], rotation=90, fontsize=6,
)
for i, tt in enumerate(is_tt):
if tt:
ax.get_xticklabels()[i].set_color("red")
axes[0].set_ylabel("layer")
if im is not None:
fig.colorbar(im, ax=axes.tolist(), shrink=0.6)
fig.suptitle(f"image_id={_iid_from_parquet(parquet_path)} (teacher-forced)")
fig.savefig(os.path.join(output_dir, "heatmap.png"), dpi=140, bbox_inches="tight")
plt.close(fig)
def render_trajectory(parquet_path: str, output_dir: str,
feature_ids: list[int], tau: dict[str, float]) -> None:
fig, ax = plt.subplots(figsize=(8, 4))
for m in METHODS:
try:
mat, _, is_tt = heatmap_matrix_from_parquet(
parquet_path, method=m, image_id=_iid_from_parquet(parquet_path),
pass_="teacher", feature_ids=feature_ids,
)
except ValueError:
continue
toilet_cols = [i for i, tt in enumerate(is_tt) if tt]
if not toilet_cols:
traj = mat.max(axis=1)
else:
traj = mat[:, toilet_cols].max(axis=1)
ax.plot(traj, label=f"{m} (τ_c={tau.get(m, float('nan')):.2f})")
ax.axhline(tau.get(m, 0.0), linestyle=":", alpha=0.4)
ax.set_xlabel("layer")
ax.set_ylabel("max_{f ∈ Φ_toilet, t ∈ toilet-tokens} z")
ax.set_title(f"image_id={_iid_from_parquet(parquet_path)}")
ax.legend(fontsize=8)
fig.savefig(os.path.join(output_dir, "trajectory.png"), dpi=140, bbox_inches="tight")
plt.close(fig)
def _iid_from_parquet(parquet_path: str) -> str:
return os.path.basename(os.path.dirname(parquet_path))
def main():
p = argparse.ArgumentParser()
p.add_argument("--toilet_features", required=True)
p.add_argument("--tau_c", required=True)
p.add_argument("--feature_set", choices=["A", "B", "C"], default="A")
p.add_argument("--alpha", type=float, default=0.05)
p.add_argument("--image_ids", default="")
p.add_argument("--image_ids_file", default="")
p.add_argument("--output_root", default="outputs/feature_ks")
args = p.parse_args()
feature_ids = _flatten_feature_ids(args.toilet_features, args.feature_set)
tau = _load_tau(args.tau_c, args.alpha)
if args.image_ids:
ids = [s.strip() for s in args.image_ids.split(",") if s.strip()]
elif args.image_ids_file:
with open(args.image_ids_file) as f:
ids = [line.strip() for line in f if line.strip()]
else:
ids = [d for d in os.listdir(args.output_root)
if os.path.isdir(os.path.join(args.output_root, d))]
for iid in ids:
d = os.path.join(args.output_root, iid)
parquet = os.path.join(d, "activations.parquet")
if not os.path.exists(parquet):
print(f" SKIP {iid}: no parquet")
continue
render_heatmap(parquet, d, feature_ids, tau)
render_trajectory(parquet, d, feature_ids, tau)
print(f" rendered {iid}")
if __name__ == "__main__":
main()