File size: 6,253 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
#!/bin/bash
#SBATCH --job-name=visualize_multilayer        # Job name
#SBATCH --output=./log_slurm/result/visualize_multilayer.txt      # Output file
#SBATCH --error=./log_slurm/error/visualize_multilayer.txt       # Error file
#SBATCH --ntasks=1               # Number of tasks (processes)
#SBATCH --gpus=1                 # Number of GPUs per node
#SBATCH --nodes=1               # Số node yêu cầu
#SBATCH --cpus-per-task=20       # Số CPU cho mỗi task

# =============================================================================
# Visualize Multilayer SAE Features (single-pass)
#
# Runs training/visualize_multilayer_features.py to produce HTML reports
# for each specified feature, including:
#   - Layer/hook activation distribution chart
#   - Top activating image-patch crops
#   - Top activating text token contexts (green-highlighted)
#
# Two data modes — uncomment the one you want:
#   A) HuggingFace dataset with captions (COCO / CC3M)  ← recommended
#   B) Plain image folder (no captions)
#
# Usage:
#   bash scripts/visualize_multilayer_features.sh
#
# Override any variable inline:
#   FEATURE_IDS="0 1 42" DEVICE_ID=1 bash scripts/visualize_multilayer_features.sh
# =============================================================================

# ── GPU ──────────────────────────────────────────────────────────────────────
NUM_GPUS="${NUM_GPUS:-1}"
DEVICE_ID="${DEVICE_ID:-0}"

# ── Model & SAE ──────────────────────────────────────────────────────────────
SAE_CKPT="training/multilayer_sae_ckpt/last.ckpt"
MODEL_NAME="llava-hf/llava-1.5-7b-hf"
DTYPE="${DTYPE:-float16}"

# ── Data — Mode A: HF dataset with captions (COCO / CC3M) ───────────────────
HF_DATASET="yerevann/coco-karpathy"                          # e.g. lmms-lab/COCO-Caption
LOCAL_VAL_PATH="COCO-Dataset/val"                  # e.g. /path/to/coco/val2017
SPLIT="validation"
NUM_WORKERS="8"

# ── Data — Mode B: plain image folder (no captions) ─────────────────────────
DATA_DIR="${DATA_DIR:-}"                              # e.g. COCO-Dataset/filtered_val/hallucinated

# ── Feature selection (REQUIRED) ────────────────────────────────────────────
FEATURE_IDS="${FEATURE_IDS:-0 1}"                    # space-separated feature IDs

# ── Hook point (optional) ───────────────────────────────────────────────────
HOOK_POINT="model.language_model.model.layers.20.hook_resid_post"

# ── Processing ───────────────────────────────────────────────────────────────
BATCH_SIZE="${BATCH_SIZE:-4}"                           # keep small (4-8) for LLaVA-7B
SAE_BATCH="${SAE_BATCH:-4096}"
THRESHOLD="${THRESHOLD:-1e-3}"
MAX_BATCHES="${MAX_BATCHES:-}"                        # leave empty = all batches

# ── Visualisation ────────────────────────────────────────────────────────────
OUTPUT_DIR="training/visualize"
TOP_IMAGES="${TOP_IMAGES:-10}"
TOP_TEXTS="${TOP_TEXTS:-20}"
BUFFER="${BUFFER:-10}"

# =============================================================================
# Validation
# =============================================================================
cd "$(dirname "$0")/.."

if [ ! -f "${SAE_CKPT}" ]; then
    echo "Error: SAE checkpoint not found: ${SAE_CKPT}" >&2
    echo "Set SAE_CKPT= to a valid .ckpt file." >&2
    exit 1
fi

if [ -z "${HF_DATASET}" ] && [ -z "${DATA_DIR}" ]; then
    echo "Error: provide either HF_DATASET + LOCAL_VAL_PATH or DATA_DIR." >&2
    exit 1
fi

if [ -n "${HF_DATASET}" ] && [ -z "${LOCAL_VAL_PATH}" ]; then
    echo "Error: HF_DATASET is set but LOCAL_VAL_PATH is empty." >&2
    exit 1
fi

if [ -z "${FEATURE_IDS}" ]; then
    echo "Error: FEATURE_IDS must be set (e.g. FEATURE_IDS=\"0 1 42\")." >&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=(
    --sae_ckpt   "${SAE_CKPT}"
    --model_name "${MODEL_NAME}"
    --device_id  "${DEVICE_ID}"
    --dtype      "${DTYPE}"
    --output_dir "${OUTPUT_DIR}"
    --feature_ids ${FEATURE_IDS}
    --batch_size    "${BATCH_SIZE}"
    --hook_point     "${HOOK_POINT}"
    --sae_batch     "${SAE_BATCH}"
    --threshold     "${THRESHOLD}"
    --top_images    "${TOP_IMAGES}"
    --top_texts     "${TOP_TEXTS}"
    --buffer        "${BUFFER}"
    --num_workers   "${NUM_WORKERS}"
)

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

# Optional: limit batches
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_multilayer_features "${ARGS[@]}"
else
    echo "Launching single-GPU mode..."
    python training/visualize_multilayer_features.py "${ARGS[@]}"
fi