File size: 7,434 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 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 | #!/bin/bash
#SBATCH --job-name=visualize_probe_features
#SBATCH --output=./log_slurm/result/visualize_probe_features.txt
#SBATCH --error=./log_slurm/error/visualize_probe_features.txt
#SBATCH --ntasks=1
#SBATCH --gpus=1
#SBATCH --nodes=1
#SBATCH --cpus-per-task=20
# =============================================================================
# Visualize Probe Features — Multi-layer interactive HTML
#
# Runs training/visualize_probe_features.py which:
# 1. Loads linear probe checkpoints for each layer → top-k features by weight
# 2. Runs a single forward pass per layer to collect top activations
# 3. Generates ONE self-contained interactive HTML:
# Layers → Top Features (+ probe weight) → Image patches / Text tokens
#
# DATA MODES (set DATA_MODE below)
# ---------
# toilet — only HF "pbcong/bathroom-toilet" images matching TOILET_MODE.
# Set IMAGE_FOLDER to the local CC3M image directory.
# Supports CAPTION_MODE=generated|dataset.
#
# cc3m — full CC3M or COCO via HF dataset + local path, OR a plain image
# folder. Set HF_DATASET+LOCAL_VAL_PATH -or- DATA_DIR.
#
# Usage:
# bash training/scripts/visualize_probe_features.sh
#
# Override any variable inline:
# DATA_MODE=cc3m LAYERS="0 1 2 3 4 5 6" \
# bash training/scripts/visualize_probe_features.sh
# =============================================================================
# ── GPU ──────────────────────────────────────────────────────────────────────
NUM_GPUS="${NUM_GPUS:-1}"
DEVICE_ID="${DEVICE_ID:-0}"
# ── Model & SAE ──────────────────────────────────────────────────────────────
SAE_CKPT="${SAE_CKPT:-training/multilayer_sae_ckpt/last.ckpt}"
MODEL_NAME="${MODEL_NAME:-llava-hf/llava-1.5-7b-hf}"
DTYPE="${DTYPE:-float16}"
# ── Probe ────────────────────────────────────────────────────────────────────
PROBE_DIR="${PROBE_DIR:-training/multilayer_sae_ckpt}"
PROBE_INPUT_DIM="${PROBE_INPUT_DIM:-65536}"
LAYERS="${LAYERS:-0 1 2 3 4 5 6}"
TOP_PROBE_K="${TOP_PROBE_K:-10}"
# ── Data mode: toilet | cc3m ──────────────────────────────────────────────────
DATA_MODE="${DATA_MODE:-toilet}"
# ── Data — toilet mode ────────────────────────────────────────────────────────
IMAGE_FOLDER="${IMAGE_FOLDER:-CC3M-Dataset/cc3m_images/train}"
TOILET_MODE="${TOILET_MODE:-toilet}" # toilet | bathroom
CAPTION_MODE="${CAPTION_MODE:-generated}" # generated | dataset
MAX_NEW_TOKENS="${MAX_NEW_TOKENS:-128}"
# ── Data — cc3m mode (HF dataset + local path) ───────────────────────────────
HF_DATASET="${HF_DATASET:-pixparse/cc3m-wds}" # or yerevann/coco-karpathy
LOCAL_VAL_PATH="${LOCAL_VAL_PATH:-CC3M-Dataset/cc3m_images/train}"
SPLIT="${SPLIT:-train}"
# ── Data — cc3m mode (plain image folder, alternative to HF_DATASET) ─────────
DATA_DIR="${DATA_DIR:-}"
# ── Common data ───────────────────────────────────────────────────────────────
NUM_WORKERS="${NUM_WORKERS:-8}"
# ── Processing ────────────────────────────────────────────────────────────────
BATCH_SIZE="${BATCH_SIZE:-4}"
SAE_BATCH="${SAE_BATCH:-4096}"
THRESHOLD="${THRESHOLD:-1e-3}"
MAX_BATCHES="${MAX_BATCHES:-}"
# ── Visualisation ─────────────────────────────────────────────────────────────
OUTPUT_DIR="${OUTPUT_DIR:-outputs/probe_features}"
TOP_IMAGES="${TOP_IMAGES:-10}"
TOP_TEXTS="${TOP_TEXTS:-10}"
BUFFER="${BUFFER:-10}"
# =============================================================================
# Validation
# =============================================================================
cd "$(dirname "$0")/../.."
if [ ! -f "${SAE_CKPT}" ]; then
echo "Error: SAE checkpoint not found: ${SAE_CKPT}" >&2
exit 1
fi
if [ ! -d "${PROBE_DIR}" ]; then
echo "Error: PROBE_DIR not found: ${PROBE_DIR}" >&2
exit 1
fi
if [ "${DATA_MODE}" = "toilet" ] && [ ! -d "${IMAGE_FOLDER}" ]; then
echo "Error: IMAGE_FOLDER not found: ${IMAGE_FOLDER}" >&2
exit 1
fi
if [ "${DATA_MODE}" = "cc3m" ] && [ -z "${HF_DATASET}" ] && [ -z "${DATA_DIR}" ]; then
echo "Error: cc3m mode requires HF_DATASET+LOCAL_VAL_PATH or DATA_DIR." >&2
exit 1
fi
# =============================================================================
# Environment
# =============================================================================
export HF_HOME="${HF_HOME:-${HOME}/scratch/hf_home}"
export PYTHONPATH="$(pwd):${PYTHONPATH:-}"
if [ -f .env ]; then
set -a; source .env; set +a
fi
# =============================================================================
# Build argument list
# =============================================================================
ARGS=(
--data_mode "${DATA_MODE}"
--sae_ckpt "${SAE_CKPT}"
--model_name "${MODEL_NAME}"
--device_id "${DEVICE_ID}"
--dtype "${DTYPE}"
--probe_dir "${PROBE_DIR}"
--probe_input_dim "${PROBE_INPUT_DIM}"
--layers ${LAYERS}
--top_probe_k "${TOP_PROBE_K}"
--num_workers "${NUM_WORKERS}"
--batch_size "${BATCH_SIZE}"
--sae_batch "${SAE_BATCH}"
--threshold "${THRESHOLD}"
--output_dir "${OUTPUT_DIR}"
--top_images "${TOP_IMAGES}"
--top_texts "${TOP_TEXTS}"
--buffer "${BUFFER}"
)
# Data-mode-specific args
if [ "${DATA_MODE}" = "toilet" ]; then
ARGS+=(
--image_folder "${IMAGE_FOLDER}"
--toilet_mode "${TOILET_MODE}"
--caption_mode "${CAPTION_MODE}"
--max_new_tokens "${MAX_NEW_TOKENS}"
)
else
# cc3m mode: HF dataset or plain folder
if [ -n "${HF_DATASET}" ]; then
ARGS+=(--hf_dataset "${HF_DATASET}" --local_val_path "${LOCAL_VAL_PATH}" --split "${SPLIT}")
else
ARGS+=(--data_dir "${DATA_DIR}")
fi
fi
if [ -n "${MAX_BATCHES}" ]; then
ARGS+=(--max_batches "${MAX_BATCHES}")
fi
# =============================================================================
# Run
# =============================================================================
if [ "${NUM_GPUS}" -gt 1 ]; then
echo "Launching with torchrun on ${NUM_GPUS} GPUs..."
torchrun --nproc_per_node="${NUM_GPUS}" -m training.visualize_probe_features "${ARGS[@]}"
else
echo "Launching single-GPU mode (device ${DEVICE_ID})..."
python -m training.visualize_probe_features "${ARGS[@]}"
fi
|