File size: 7,186 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 151 | # =============================================================================
# Train_Probe_SAE.sh — Train binary probe(s) on SAE feature activations
#
# Two-phase pipeline
# ──────────────────
# Phase 1 (one VLM pass + SAE): for each image generate caption, run forced
# forward pass, extract hook activations, SAE encode + max-pool, keep only
# pooled feature vectors in memory (no disk writes).
#
# Phase 2: gather pooled vectors from all ranks → train binary linear probes
# per hook point/layer. All probes share the same positive/negative split.
#
# Multi-GPU: set NUM_GPUS>1 to shard VLM inference with torchrun.
# Single-GPU: NUM_GPUS=1 (default).
#
# Override variables inline, e.g.:
# NUM_GPUS=4 PROBE_TYPE=bathroom bash training/scripts/Train_Probe_SAE.sh
# =============================================================================
wandb login --relogin "${WANDB_API_KEY}"
export HF_HOME="/data/caotue/hf_cache"
export HF_DATASETS_CACHE="/data/caotue/hf_cache/datasets"
export TORCH_HOME="/data/caotue/torch_cache"
export TMPDIR="/data/caotue/tmp"
export UV_CACHE_DIR="/data/caotue/uv"
export PIP_CACHE_DIR="/data/caotue/pip"
# ── GPU count ─────────────────────────────────────────────────────────────────
NUM_GPUS="${NUM_GPUS:-1}"
# ── Model / SAE ───────────────────────────────────────────────────────────────
SAE_CKPT="/data/caotue/SAE_checkpoints/16_d_model/last.ckpt"
MODEL_NAME="${MODEL_NAME:-llava-hf/llava-1.5-7b-hf}"
DEVICE_ID="${DEVICE_ID:-4}"
DTYPE="${DTYPE:-bfloat16}"
# ── Data ──────────────────────────────────────────────────────────────────────
IMAGE_FOLDER="/data/caotue/CC3M-Dataset/cc3m_images"
# HF dataset with binary relation labels (must have 'train' and 'validation' splits).
# ID_COL : column holding the image filename stem.
# PROBE_TYPE: binary label column used as positives (label==1 → B-only + A+B).
# Examples:
# HF_DATASET="pbcong/bathroom-toilet" PROBE_TYPE="toilet"
# HF_DATASET="pbcong/bathroom-toilet" PROBE_TYPE="bathroom"
# HF_DATASET="pbcong/cat-dog" PROBE_TYPE="cat"
HF_DATASET="${HF_DATASET:-pbcong/bathroom-toilet}"
ID_COL="${ID_COL:-image_id}"
PROBE_TYPE="${PROBE_TYPE:-toilet}"
# Contrastive object column (optional).
# When set: HF negatives = other_object=1 & probe_type=0 (from train/val splits).
# Example: OTHER_OBJECT=toilet when PROBE_TYPE=bathroom
OTHER_OBJECT="${OTHER_OBJECT:-bathroom}"
# JSON file with negative image IDs, pre-split:
# {"train": [...image_ids], "validation": [...image_ids]}
# Must be a valid JSON file.
NEG_JSONL="${NEG_JSONL:-mechanistic_interp/neg_cc3m_5k.json}"
# ── Hook specification ────────────────────────────────────────────────────────
# LAYERS : space-separated layer indices; one probe is trained per layer
# HOOK_TYPE : hook suffix — pre | mid | post |
LAYERS="${LAYERS:-}" # e.g. "10 15 20" for three probes, or empty for all layers
HOOK_TYPE="${HOOK_TYPE:-post}"
# ── SAE ───────────────────────────────────────────────────────────────────────
SAE_BATCH="${SAE_BATCH:-2048}"
# ── Generation ────────────────────────────────────────────────────────────────
MAX_NEW_TOKENS="${MAX_NEW_TOKENS:-256}"
QUESTION="${QUESTION:-Describe this image.}"
BATCH_SIZE="${BATCH_SIZE:-64}"
# ── Probe training ────────────────────────────────────────────────────────────
PROBE_BATCH_SIZE="${PROBE_BATCH_SIZE:-256}"
PROBE_EPOCHS="${PROBE_EPOCHS:-10}"
PROBE_LR="${PROBE_LR:-0.01}"
# ── Output ────────────────────────────────────────────────────────────────────
SAVE_DIR="/data/caotue/SAE_checkpoints/16_d_model/probes_${PROBE_TYPE}"
# ── Dry run ───────────────────────────────────────────────────────────────────
# Set DRY_RUN=1 to smoke-test the pipeline on a tiny sample before a real run.
# DRY_RUN=1 bash training/scripts/Train_Probe_SAE.sh
# DRY_RUN_IMAGES controls how many images are processed (default 8).
DRY_RUN="${DRY_RUN:-0}"
DRY_RUN_IMAGES="${DRY_RUN_IMAGES:-4}"
# =============================================================================
# Validation
# =============================================================================
cd "$(dirname "$0")/../.."
if [ ! -f "${SAE_CKPT}" ]; then
echo "Error: SAE checkpoint not found: ${SAE_CKPT}" >&2
exit 1
fi
if [ ! -d "${IMAGE_FOLDER}" ]; then
echo "Error: image_folder not found: ${IMAGE_FOLDER}" >&2
exit 1
fi
# =============================================================================
# Run
# =============================================================================
COMMON_ARGS=(
-m training.Train_Probe_SAE
--sae_ckpt "${SAE_CKPT}"
--model_name "${MODEL_NAME}"
--device_id "${DEVICE_ID}"
--dtype "${DTYPE}"
--image_folder "${IMAGE_FOLDER}"
--hf_dataset "${HF_DATASET}"
--id_col "${ID_COL}"
--probe_type "${PROBE_TYPE}"
--neg_jsonl "${NEG_JSONL}"
--hook_type "${HOOK_TYPE}"
--sae_batch "${SAE_BATCH}"
--max_new_tokens "${MAX_NEW_TOKENS}"
--question "${QUESTION}"
--probe_batch_size "${PROBE_BATCH_SIZE}"
--probe_epochs "${PROBE_EPOCHS}"
--probe_lr "${PROBE_LR}"
--batch_size "${BATCH_SIZE}"
--save_dir "${SAVE_DIR}"
)
# Only pass --layers when LAYERS is non-empty; otherwise Train_Probe_SAE.py
# defaults to every layer in the model (0 .. num_hidden_layers-1).
if [ -n "${LAYERS}" ]; then
COMMON_ARGS+=(--layers ${LAYERS})
fi
if [ -n "${OTHER_OBJECT}" ]; then
COMMON_ARGS+=(--other_object "${OTHER_OBJECT}")
fi
if [ "${DRY_RUN}" = "1" ] || [ "${DRY_RUN}" = "true" ]; then
COMMON_ARGS+=(--dry_run --dry_run_images "${DRY_RUN_IMAGES}")
fi
if [ "${NUM_GPUS}" -gt 1 ]; then
echo "Launching with torchrun on ${NUM_GPUS} GPUs"
torchrun --nproc_per_node="${NUM_GPUS}" "${COMMON_ARGS[@]}"
else
echo "Launching single-GPU mode"
python3 "${COMMON_ARGS[@]}"
fi
|