File size: 5,850 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 138 139 140 141 142 143 144 145 146 147 148 149 150 | #!/bin/bash
# =============================================================================
# select_features.sh — pick top-k SAE features per language-model layer.
#
# Two modes:
# MODE=probe (default) — top-k by largest positive probe weight.
# PROBES_PATH can be an all-layers state-dict file
# OR a directory containing one .pth file per layer.
# (e.g. mechanistic_interp/probes/probes_gen_bathroom_toilet.pt
# or mechanistic_interp/probes/toilet/mid/)
# One run, one probe → one OBJECT.
# MODE=f1 — top-k by best per-feature F1 against an HF column
# (0/1 label).
# Positives : HF rows where OBJECT column == 1
# (includes rows where OTHER_OBJECT=1 too).
# Negatives : HF rows where OTHER_OBJECT=1 & OBJECT=0
# (contrastive, if OTHER_OBJECT is set)
# + random CC3M ids from NEG_JSONL (if set).
# At least one of OTHER_OBJECT or NEG_JSONL is required.
# To process multiple objects in one shot, set
# OBJECTS="bathroom toilet ..." — one JSON per object.
#
# Examples:
# # probe mode (single object: whatever the state-dict was trained for)
# OBJECT=toilet \
# PROBES_PATH=mechanistic_interp/probes/probes_gen_bathroom_toilet.pt \
# bash mechanistic_interp/scripts/select_features.sh
#
# # probe mode using per-layer .pth files from Train_Probe_SAE.py output
# # Directory structure: probes/toilet/mid/probe_mid_0.pth, probe_mid_1.pth, ...
# OBJECT=toilet \
# PROBES_PATH=mechanistic_interp/probes/toilet/mid \
# bash mechanistic_interp/scripts/select_features.sh
#
# # f1: contrastive HF negatives only (no random CC3M)
# MODE=f1 OBJECT=toilet OTHER_OBJECT=bathroom \
# HF_DATASET=pbcong/bathroom-toilet \
# bash mechanistic_interp/scripts/select_features.sh
#
# # f1: contrastive HF negatives + random CC3M negatives
# MODE=f1 OBJECT=toilet OTHER_OBJECT=bathroom \
# HF_DATASET=pbcong/bathroom-toilet \
# NEG_JSONL=mechanistic_interp/neg_cc3m_5k.json \
# bash mechanistic_interp/scripts/select_features.sh
#
# # f1 on a different HF dataset with two new columns (json negatives only)
# MODE=f1 OBJECTS="kitchen fridge" \
# HF_DATASET=youruser/kitchen-fridge \
# NEG_JSONL=mechanistic_interp/neg_cc3m_5k.json \
# bash mechanistic_interp/scripts/select_features.sh
#
# Output:
# ${OUT_DIR}/selected_features_${MODE}_${OBJ}_top${TOP_K}.json (one per object)
# =============================================================================
set -euo pipefail
cd "$(dirname "$0")/../.."
export PYTHONPATH="$(cd .. && pwd):$(pwd):${PYTHONPATH:-}"
MODE="${MODE:-probe}"
TOP_K="${TOP_K:-20}"
HOOK_TYPE="${HOOK_TYPE:-post}"
DEVICE="${DEVICE:-cuda:4}"
# Object selection. OBJECTS (plural) wins when set; otherwise fall back to
# OBJECT (singular). OBJECT is also what the legacy single-object call uses.
OBJECT="${OBJECT:-bathroom}"
OBJECTS="${OBJECTS:-}"
# Common paths.
OUT_DIR="${OUT_DIR:-mechanistic_interp/probes}"
# probe-mode
PROBES_PATH="${PROBES_PATH:-/data/caotue/SAE_checkpoints/16_d_model/probes_bathroom_epoch10}"
# f1-mode
SAE_CKPT="${SAE_CKPT:-/data/caotue/SAE_checkpoints/16_d_model/last.ckpt}"
HF_DATASET="${HF_DATASET:-pbcong/bathroom-toilet}"
OTHER_OBJECT="${OTHER_OBJECT:-toilet}" # contrastive column (e.g. 'bathroom' when OBJECT=toilet)
NEG_JSONL="${NEG_JSONL:-mechanistic_interp/neg_cc3m_5k.json}" # optional random-CC3M negatives JSON
IMG_BATCH="${IMG_BATCH:-16}"
DTYPE="${DTYPE:-bfloat16}"
# Resolve object list.
if [ -n "${OBJECTS}" ]; then
LIST=( ${OBJECTS} )
else
LIST=( "${OBJECT}" )
fi
echo "MODE=${MODE} TOP_K=${TOP_K} objects: ${LIST[*]}"
OK_LIST=()
FAIL_LIST=()
for OBJ in "${LIST[@]}"; do
OUT="${OUT_DIR}/selected_features_${MODE}_${OBJ}_top${TOP_K}.json"
ARGS=(
-m mechanistic_interp.select_features
--mode "${MODE}"
--top_k "${TOP_K}"
--out "${OUT}"
--hook_type "${HOOK_TYPE}"
--label "${OBJ}"
)
if [ "${MODE}" = "probe" ]; then
ARGS+=( --probes_path "${PROBES_PATH}" )
else
# f1 mode: --probe_type drives positives (rows where the named HF
# column == 1), --label is just the JSON tag, OUT path is unique per
# object, and each object runs in its own python process so a failure
# on one (e.g. OOM) doesn't block the others.
ARGS+=(
--sae_ckpt "${SAE_CKPT}"
--hf_dataset "${HF_DATASET}"
--probe_type "${OBJ}"
--device "${DEVICE}"
--img_batch "${IMG_BATCH}"
--dtype "${DTYPE}"
)
if [ -n "${OTHER_OBJECT}" ]; then
ARGS+=( --other_object "${OTHER_OBJECT}" )
fi
if [ -n "${NEG_JSONL}" ]; then
ARGS+=( --neg_jsonl "${NEG_JSONL}" )
fi
fi
echo
echo "[$(date +%T)] OBJECT=${OBJ} → ${OUT}"
echo "Running: python ${ARGS[*]}"
if python "${ARGS[@]}"; then
echo "[$(date +%T)] OK ${OBJ}"
OK_LIST+=( "${OBJ}" )
else
rc=$?
echo "[$(date +%T)] FAIL ${OBJ} (exit ${rc})"
FAIL_LIST+=( "${OBJ}" )
fi
done
echo
echo "── Summary ────────────────────────────────────────────────"
echo "OK (${#OK_LIST[@]}): ${OK_LIST[*]:-—}"
echo "FAIL (${#FAIL_LIST[@]}): ${FAIL_LIST[*]:-—}"
if [ "${#FAIL_LIST[@]}" -gt 0 ]; then
exit 1
fi
|