File size: 3,824 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
#!/usr/bin/env bash
# =============================================================================
# Eval a trained gen-time probe on a held-out split (default: val)
# =============================================================================
# Usage:
#   PROBE_CHECKPOINT=/path/to/probes_gen_<rel>.pt \
#     bash experiment/scripts/probe/run_eval_probe_gen.sh
# Common overrides:
#   RELATION=bathroom_toilet SPLIT=val NGPUS=4 \
#     PROBE_CHECKPOINT=... OUTPUT_JSON=eval.json \
#     bash experiment/scripts/probe/run_eval_probe_gen.sh
# =============================================================================

set -euo pipefail

PROJECT_ROOT="$(cd "$(dirname "$0")/../../.." && pwd)"
export PYTHONPATH="${PROJECT_ROOT}:${PYTHONPATH:-}"
export TOKENIZERS_PARALLELISM="${TOKENIZERS_PARALLELISM:-false}"
export TORCHINDUCTOR_CACHE_DIR="${TORCHINDUCTOR_CACHE_DIR:-${HOME}/scratch/.cache/torchinductor}"
export TRITON_CACHE_DIR="${TRITON_CACHE_DIR:-${HOME}/scratch/.cache/triton}"

RELATION="${RELATION:-bathroom_toilet}"
ADV_CONFIG="${ADV_CONFIG:-${PROJECT_ROOT}/experiment/adv_config.json}"
SPLIT="${SPLIT:-val}"
PROBE_CHECKPOINT="${PROBE_CHECKPOINT:-}"
SAE_CHECKPOINT="${SAE_CHECKPOINT:-}"
OUTPUT_JSON="${OUTPUT_JSON:-}"
MAX_EVAL_SAMPLES="${MAX_EVAL_SAMPLES:-}"
MAX_NEW_TOKENS_TRAIN="${MAX_NEW_TOKENS_TRAIN:-}"
GEN_DO_SAMPLE="${GEN_DO_SAMPLE:-}"
POOL="${POOL:-max}"

if [ -z "$PROBE_CHECKPOINT" ]; then
    echo "ERROR: PROBE_CHECKPOINT is required (path to probes_gen_<rel>.pt)" >&2
    exit 2
fi
if [ ! -f "$PROBE_CHECKPOINT" ]; then
    echo "ERROR: probe checkpoint not found: $PROBE_CHECKPOINT" >&2
    exit 2
fi

PYTHON_BIN="${PYTHON_BIN:-${PROJECT_ROOT}/.venv/bin/python}"
TORCHRUN_BIN="${TORCHRUN_BIN:-${PROJECT_ROOT}/.venv/bin/torchrun}"
[ -x "$PYTHON_BIN" ] || PYTHON_BIN="python3"
[ -x "$TORCHRUN_BIN" ] || TORCHRUN_BIN="torchrun"

count_visible_gpus() {
    local n
    n=$(nvidia-smi -L 2>/dev/null | grep -c '^GPU' || true)
    n=$(echo "$n" | tr -d '[:space:]')
    [ -z "$n" ] && n=1
    [ "$n" -lt 1 ] && n=1
    echo "$n"
}

VISIBLE_GPUS=$(count_visible_gpus)
NGPUS_RAW="${NGPUS:-auto}"
if [ -z "$NGPUS_RAW" ] || [ "$NGPUS_RAW" = "auto" ]; then
    NGPUS=$VISIBLE_GPUS
else
    NGPUS=$NGPUS_RAW
fi
[ "$NGPUS" -lt 1 ] && NGPUS=1
[ "$NGPUS" -gt "$VISIBLE_GPUS" ] && NGPUS=$VISIBLE_GPUS

echo "=========================================="
echo "Probe eval (gen-time SAE features)"
echo "=========================================="
echo "  Project:   $PROJECT_ROOT"
echo "  Config:    $ADV_CONFIG"
echo "  Relation:  $RELATION"
echo "  Split:     $SPLIT"
echo "  ProbeCkpt: $PROBE_CHECKPOINT"
echo "  Pool:      $POOL"
echo "  GPUs:      $NGPUS (visible=$VISIBLE_GPUS)"
echo "=========================================="

EXTRA_ARGS=(--config "$ADV_CONFIG" --relation "$RELATION")
EXTRA_ARGS+=(--probe_checkpoint "$PROBE_CHECKPOINT" --split "$SPLIT" --pool "$POOL")
[ -n "$SAE_CHECKPOINT" ]      && EXTRA_ARGS+=(--adv.sae_checkpoint "$SAE_CHECKPOINT")
[ -n "$MAX_NEW_TOKENS_TRAIN" ] && EXTRA_ARGS+=(--adv.max_new_tokens_train "$MAX_NEW_TOKENS_TRAIN")
[ -n "$GEN_DO_SAMPLE" ]        && EXTRA_ARGS+=(--adv.gen_do_sample "$GEN_DO_SAMPLE")
[ -n "$MAX_EVAL_SAMPLES" ]     && EXTRA_ARGS+=(--max_eval_samples "$MAX_EVAL_SAMPLES")
[ -n "$OUTPUT_JSON" ]          && EXTRA_ARGS+=(--output_json "$OUTPUT_JSON")

if [ "$NGPUS" -gt 1 ]; then
    export MASTER_ADDR="${MASTER_ADDR:-127.0.0.1}"
    [ -z "${MASTER_PORT:-}" ] && export MASTER_PORT=$((29500 + RANDOM % 1000))
    echo "  MASTER_ADDR: $MASTER_ADDR  MASTER_PORT: $MASTER_PORT"
    "$TORCHRUN_BIN" --nproc_per_node="$NGPUS" \
        --master_addr="$MASTER_ADDR" --master_port="$MASTER_PORT" \
        -m experiment.training.eval_probe_gen "${EXTRA_ARGS[@]}" "$@"
else
    "$PYTHON_BIN" -m experiment.training.eval_probe_gen "${EXTRA_ARGS[@]}" "$@"
fi