hallucination / mechanistic_interp /scripts /d_sae_correlation_patching.sh
ToiTenBao's picture
Upload hallucination folder
a2ffd07 verified
Raw
History Blame Contribute Delete
6.43 kB
#!/bin/bash
# Top-k d_sae features -> toilet SAE features (per-feature gradient attribution).
#
# For a chosen readout layer L_inspect with toilet feature set T, find WHICH
# features anywhere in the d_sae at the layers BEFORE L_inspect most affect T:
# M_{L_inspect} = mean_tokens( sum_{f in T} z[L_inspect][.., f] )
# a[L_patch][f] = sum_tokens z_clean[L_patch][t,f] * dM_{L_inspect}/dz[L_patch][t,f]
# normalised by the inspect-layer metric M_{L_inspect} (fraction of the toilet
# readout the feature explains) so inspect layers at different depths compare
# (raw |a| grows ~exponentially with depth: ~500 at L_inspect=1 vs ~4e7 at 31):
# score[L_inspect][L_patch][f] = a[L_inspect][L_patch][f] / M_{L_inspect}
#
# Output: a JSON (OUT) holding, per L_inspect, the top-K features per upstream
# layer AND a GLOBAL cross-layer ranking; plus an INTERACTIVE Plotly HTML (GRAPH)
# with a dropdown over L_inspect showing a horizontal bar chart of the global
# top-N features ranked ACROSS all upstream layers (one comparable ranking — a
# heatmap is NOT used, since its per-row top-k can't compare ranks between layers).
# Each bar is labelled "L<patch>·#<feat>", coloured by signed normalised attr
# (red promotes T, blue suppresses T), hover reveals the raw value.
#
# Tunables (override via env):
# DEVICE target GPU (default cuda:2)
# DTYPE bfloat16 | float16 | float32 (default bfloat16)
# SAE_CKPT BatchTopK SAE checkpoint (threshold buffer required)
# TOILET_FEATS JSON {layer_<L>: {features: [...]}} - readout set T
# K features kept per upstream layer in the JSON (default 10)
# TOP_N features in the global cross-layer ranking / bars (default 30)
# VERIFY 1 = re-rank the shortlist by FAITHFUL ablation ΔM (comparable
# across layers); 0 = first-order score only (default 0). First-
# order inflates early layers by orders of magnitude, so it is
# NOT comparable across L_patch - use VERIFY=1 to compare layers.
# Ablates ONLY the top_n shortlist, never the whole d_sae.
# VERIFY_BATCH features ablated per forward when VERIFY=1 (default 8)
# RANK_BY abs | pos | neg (default abs)
# abs = strongest either way, pos = promoters, neg = suppressors
# INSPECT_LAYERS space-separated readout layers (default: all with T; e.g. "24 31")
# OUT_TAG extra suffix on output filenames (e.g. top10)
# SAMPLES samples.json to filter images (base mentioned object, LoRA
# suppressed it). Set SAMPLES="" to use the IMAGES list below.
# SAMPLES_PROMPT prompt_results key driving the filter (default "Describe this image.")
# IMAGE_ROOT directory holding <image_id>.jpg for SAMPLES filtering.
# PROMPT text fed to the model for the forward/backward pass.
#
# Usage:
# bash mechanistic_interp/scripts/sae_attribution_patching.sh
# INSPECT_LAYERS="24 31" K=15 bash mechanistic_interp/scripts/sae_attribution_patching.sh
# RANK_BY=pos bash mechanistic_interp/scripts/sae_attribution_patching.sh
set -euo pipefail
REPO_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/../.." && pwd)"
MODEL_NAME="llava-hf/llava-1.5-7b-hf"
SAE_CKPT="${SAE_CKPT:-/data/caotue/SAE_checkpoints/16_d_model/last.ckpt}"
TOILET_FEATS="${TOILET_FEATS:-${REPO_ROOT}/mechanistic_interp/probes/selected_features_f1_post_toilet_top10.json}"
DEVICE="${DEVICE:-cuda:5}"
DTYPE="${DTYPE:-bfloat16}"
K="${K:-20}"
TOP_N="${TOP_N:-30}"
VERIFY="${VERIFY:-1}"
VERIFY_BATCH="${VERIFY_BATCH:-16}"
RANK_BY="${RANK_BY:-pos}"
INSPECT_LAYERS="${INSPECT_LAYERS:-}"
OUT_TAG="${OUT_TAG:-}"
TAG_SUFFIX="${OUT_TAG:+_${OUT_TAG}}"
OUT="${REPO_ROOT}/mechanistic_interp/graph/d_sae_correlation_patching${TAG_SUFFIX}.json"
GRAPH="${REPO_ROOT}/mechanistic_interp/graph/d_sae_correlation_patching${TAG_SUFFIX}.html"
SAMPLES="${SAMPLES:-${REPO_ROOT}/mechanistic_interp/toilet_bathroom/samples.json}"
SAMPLES_PROMPT="${SAMPLES_PROMPT:-Describe this image.}"
IMAGE_ROOT="${IMAGE_ROOT:-/data/caotue/CC3M-Dataset/cc3m_images/train}"
PROMPT="${PROMPT:-USER: <image> Describe this image. ASSISTANT: In this image, there is a bathroom and a}"
# Fallback image when SAMPLES="" and no --images is passed.
IMAGES=(
/data/caotue/CC3M-Dataset/cc3m_images/train/001607409.jpg
)
# ------------------------------------------------------------------ #
mkdir -p "$(dirname "${OUT}")"
echo "========================================================"
echo " top-k d_sae features -> toilet features (per-feature attribution)"
echo "========================================================"
echo " Toilet feats : ${TOILET_FEATS}"
echo " SAE ckpt : ${SAE_CKPT}"
echo " Device / dtype : ${DEVICE} / ${DTYPE}"
echo " k / top_n : ${K} / ${TOP_N}"
echo " rank_by : ${RANK_BY}"
echo " verify : ${VERIFY} (faithful ΔM ablation; batch ${VERIFY_BATCH})"
echo " Inspect layers : ${INSPECT_LAYERS:-<all with toilet feats>}"
echo " Prompt : ${PROMPT}"
echo " Out JSON : ${OUT}"
echo " Out graph : ${GRAPH}"
echo "========================================================"
cd "${REPO_ROOT}"
export PYTHONPATH="$(cd .. && pwd):$(pwd):${PYTHONPATH:-}"
ARGS=(
--toilet_feats "${TOILET_FEATS}"
--sae_ckpt "${SAE_CKPT}"
--model_name "${MODEL_NAME}"
--device "${DEVICE}"
--dtype "${DTYPE}"
--prompt "${PROMPT}"
--k "${K}"
--top_n "${TOP_N}"
--verify_batch "${VERIFY_BATCH}"
--rank_by "${RANK_BY}"
--out "${OUT}"
--graph "${GRAPH}"
)
[ -n "${INSPECT_LAYERS}" ] && ARGS+=(--inspect_layers ${INSPECT_LAYERS})
[ "${VERIFY}" = "1" ] && ARGS+=(--verify)
# Image source precedence: explicit --images on the CLI > SAMPLES filter > IMAGES list.
if [[ "$*" == *"--images"* ]]; then
python -m mechanistic_interp.d_sae_correlation_patching "${ARGS[@]}" "$@"
elif [ -n "${SAMPLES}" ]; then
python -m mechanistic_interp.d_sae_correlation_patching \
--samples "${SAMPLES}" \
--samples_prompt "${SAMPLES_PROMPT}" \
--image_root "${IMAGE_ROOT}" \
"${ARGS[@]}" \
"$@"
else
python -m mechanistic_interp.d_sae_correlation_patching \
--images "${IMAGES[@]}" \
"${ARGS[@]}" \
"$@"
fi