| #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 | |