File size: 6,430 Bytes
a2ffd07 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 | #!/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
|