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