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