File size: 4,063 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 | #!/usr/bin/env bash
# =============================================================================
# Continue/warm-start probe training from an existing probes_gen_{relation}.pt
# =============================================================================
# Usage:
# PROBE_RESUME=/abs/path/to/probes_gen_kitchen_oven.pt \
# RELATION=kitchen_oven PROBE_EPOCHS=5 \
# bash experiment/scripts/probe/run_continue_probe_gen.sh
#
# Default behaviour: 5 additional epochs, cosine LR probe_lr/100 → probe_lr/10000
# (no hold phase). Checkpoints written to ./probe_outputs/probe_run_resume_<ts>/
# as probes_gen_{relation}_resume_epoch{NN}.pt + a final probes_gen_{relation}.pt.
# =============================================================================
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:-kitchen_oven}"
OUTPUT_DIR="${OUTPUT_DIR:-${PROJECT_ROOT}/probe_outputs}"
ADV_CONFIG="${ADV_CONFIG:-${PROJECT_ROOT}/experiment/adv_config.json}"
RUN_ID="${RUN_ID:-$(date +%Y%m%d_%H%M%S)}"
export RUN_ID
PROBE_EPOCHS="${PROBE_EPOCHS:-5}"
MAX_NEW_TOKENS_TRAIN="${MAX_NEW_TOKENS_TRAIN:-128}"
GEN_DO_SAMPLE="${GEN_DO_SAMPLE:-false}"
SAE_CHECKPOINT="${SAE_CHECKPOINT:-}"
PROBE_OUTPUT="${PROBE_OUTPUT:-}"
POOL="${POOL:-max}"
PROBE_RESUME="${PROBE_RESUME:-}"
if [ -z "$PROBE_RESUME" ]; then
echo "ERROR: PROBE_RESUME must be set to a path of a previous probes_gen_{relation}.pt" >&2
exit 1
fi
if [ ! -f "$PROBE_RESUME" ]; then
echo "ERROR: PROBE_RESUME file not found: $PROBE_RESUME" >&2
exit 1
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
mkdir -p "$OUTPUT_DIR"
echo "=========================================="
echo "Probe continuation (warm-start)"
echo "=========================================="
echo " Project: $PROJECT_ROOT"
echo " Config: $ADV_CONFIG"
echo " Relation: $RELATION"
echo " Resume: $PROBE_RESUME"
echo " Output: $OUTPUT_DIR"
echo " +Epochs: $PROBE_EPOCHS"
echo " MaxTok: $MAX_NEW_TOKENS_TRAIN"
echo " Pool: $POOL"
echo " GPUs: $NGPUS (visible=$VISIBLE_GPUS)"
echo "=========================================="
EXTRA_ARGS=(--config "$ADV_CONFIG" --relation "$RELATION" --output_dir "$OUTPUT_DIR")
EXTRA_ARGS+=(--probe_epochs "$PROBE_EPOCHS")
EXTRA_ARGS+=(--probe_resume "$PROBE_RESUME")
EXTRA_ARGS+=(--pool "$POOL")
EXTRA_ARGS+=(--adv.max_new_tokens_train "$MAX_NEW_TOKENS_TRAIN")
EXTRA_ARGS+=(--adv.gen_do_sample "$GEN_DO_SAMPLE")
[ -n "$SAE_CHECKPOINT" ] && EXTRA_ARGS+=(--adv.sae_checkpoint "$SAE_CHECKPOINT")
[ -n "$PROBE_OUTPUT" ] && EXTRA_ARGS+=(--probe_output "$PROBE_OUTPUT")
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.continue_probe_gen "${EXTRA_ARGS[@]}" "$@"
else
"$PYTHON_BIN" -m experiment.training.continue_probe_gen "${EXTRA_ARGS[@]}" "$@"
fi
|