Avra98's picture
sync training code: stage-1 instance-epoch sampler, multi-stage run, superposition metrics
6a1771b verified
Raw
History Blame Contribute Delete
9.11 kB
#!/bin/bash
#SBATCH --partition=songmei
#SBATCH --nodelist=feanor
#SBATCH --gres=gpu:1
#SBATCH --cpus-per-task=8
#SBATCH --mem=64G
#SBATCH --time=72:00:00
#SBATCH --job-name=wave12
#SBATCH --output=/scratch/users/gatmiry/llm-reasoning-logic-puzzles/sudoku-code/wavecurriculum_run/logs/slurm_%x_%j.out
#SBATCH --error=/scratch/users/gatmiry/llm-reasoning-logic-puzzles/sudoku-code/wavecurriculum_run/logs/slurm_%x_%j.err
# Wave-depth curriculum, 12 stages.
#
# The curriculum axis is PROPAGATION DEPTH, not puzzle difficulty. Each puzzle's
# ~25 solver waves are subsampled to 12 evenly spaced candidate snapshots
# (staged_candidate_gen.py --stages 12, last snapshot = the unique solution).
# Curriculum stage t means: t recurrence passes, latent slots 1..t active, wave
# snapshots 1..t supervised. Promotion adds exactly one wave block.
#
# Difficulty is NOT gated: every puzzle is available from step 0, drawn
# uniformly from the corpus. The old level<=stage+2 filter is gone -- the
# difficulty tag has only 6 values (so it cannot express a 12-step ladder), is
# uncorrelated with puzzle size (r=-0.003 vs empty cells), and explains only
# ~19% of the variance in solver round count.
#
# Promotion is plateau-driven: a stage graduates when the candidate-set accuracy
# at its deepest slot stops improving, scored on cells that CHANGED from the
# previous snapshot. Unchanged cells are copies of slot j-1 and stay correct for
# a head that learned nothing, so they cannot signal acquisition. The accuracy
# threshold is a fast path and the patience cap prevents a stuck stage stalling.
#
# Submit one job per arm:
# sbatch --job-name=w12_flat --export=ALL,ARM=flat sbatch_wave12.sh
# sbatch --job-name=w12_datacur --export=ALL,ARM=datacur sbatch_wave12.sh
# sbatch --job-name=w12_latent --export=ALL,ARM=latent sbatch_wave12.sh
# sbatch --job-name=w12_latent_bt --export=ALL,ARM=latent_bt sbatch_wave12.sh
#
# ARM:
# flat K=0, no curriculum at all: pure control.
# datacur K=0 with the 12-stage round-count DATA curriculum. Stage t admits
# only puzzles in the first t round-count bins, so the propagation
# ladder is expressed in the puzzle POOL rather than in latent
# supervision. This is the no-latent curriculum arm: same axis as
# the latent arms, no latent tokens. Promotion is gated on accuracy
# over the newest round-bin.
# latent K=12 latent chain, 12-stage wave supervision, plateau promotion.
# latent_bt as latent, plus adaptive backtracking keyed on DEPTH: a repair
# replays num_passes=t when snapshot t's accuracy regresses below
# its graduation value.
#
# From scratch is forced, not chosen: pos_embeddings is exactly (3*81+K, emb_dim),
# so no 6-slot checkpoint can restore into a 12-slot model.
#
# The 12-stage masks live on feanor's node-local /tmp (sbatch_gen_s12_masks.sh)
# because the 20 G scratch quota has ~2 G free and the train masks are 3.5 G.
set -u
hostname
nvidia-smi -L
echo "[$(date)] ARM=${ARM:-unset}"
ARM="${ARM:?set ARM=flat|datacur|latent|latent_bt via --export=ALL,ARM=...}"
NAME="w12_${ARM}"
SCRATCH_ROOT=/scratch/users/gatmiry/llm-reasoning-logic-puzzles
RUN_DIR=${SCRATCH_ROOT}/sudoku-code/wavecurriculum_run
ENV_LOCAL=/tmp/logicpuzzles
CAND_DIR=/tmp/sudoku_s12
LOCAL_LOG=/tmp/sudoku_wave_runs/${NAME}
TARBALL_GANDALF=/tmp/logicpuzzles_env.tar.gz
TARBALL_LOCAL=/tmp/logicpuzzles_env_${SLURM_JOB_ID}.tar.gz
mkdir -p "${RUN_DIR}/logs" "${LOCAL_LOG}" /tmp/sudoku_wave_runs
if ${ENV_LOCAL}/bin/python -u -c "import jax; assert jax.default_backend()=='gpu' or 'cuda' in str(jax.devices()[0]).lower()" 2>/dev/null; then
echo "[$(date)] reusing ${ENV_LOCAL}"
else
echo "[$(date)] fetching env tarball"
rm -rf "${ENV_LOCAL}"
SCP_OPTS="-o IdentitiesOnly=yes -o StrictHostKeyChecking=accept-new"
[ -f "${HOME}/.ssh/id_ed25519_berkeley" ] && SCP_OPTS="${SCP_OPTS} -i ${HOME}/.ssh/id_ed25519_berkeley"
scp ${SCP_OPTS} "gandalf.berkeley.edu:${TARBALL_GANDALF}" "${TARBALL_LOCAL}"
tar xzf "${TARBALL_LOCAL}" -C /tmp
rm -f "${TARBALL_LOCAL}"
fi
export PY=${ENV_LOCAL}/bin/python
export LD_LIBRARY_PATH=\
${ENV_LOCAL}/lib/python3.9/site-packages/nvidia/cudnn/lib:\
${ENV_LOCAL}/lib/python3.9/site-packages/nvidia/cublas/lib:\
${ENV_LOCAL}/lib/python3.9/site-packages/nvidia/cuda_runtime/lib:\
${ENV_LOCAL}/lib/python3.9/site-packages/nvidia/cuda_nvrtc/lib:\
${ENV_LOCAL}/lib/python3.9/site-packages/nvidia/nccl/lib:\
${LD_LIBRARY_PATH:-}
${PY} -u -c "import jax; print(jax.__version__, jax.devices(), jax.default_backend())"
# ---- Shared recipe: identical across arms ----
export SUDOKU_RESUME=0
export SUDOKU_START_STAGE=1
export SUDOKU_MAX_STAGE="${SUDOKU_MAX_STAGE:-12}"
# Plateau promotion. A stage graduates when its frontier-depth accuracy has not
# gained 0.005 for 20k steps; PATIENCE is only a hard ceiling. This replaces the
# old fixed 15k timer, under which every promotion fired on patience at 0.32-0.46
# accuracy and the 0.70 threshold was dead code.
export SUDOKU_PLATEAU_STEPS=20000
export SUDOKU_PLATEAU_DELTA=0.005
export SUDOKU_PATIENCE=80000
export SUDOKU_MIN_STAGE_STEPS=8000
export SUDOKU_PROMOTE_ACC=0.90
# Matches the flat reference run (which reached level-3 acc 0.945 at ~800k), so
# the arms are comparable to it at equal steps rather than stopping at 250k.
export SUDOKU_MAX_STEPS="${SUDOKU_MAX_STEPS:-800000}"
# Difficulty tag selects nothing: uniform over the corpus (~68% level 3).
export SUDOKU_LEVEL_BALANCED=0
export SUDOKU_EVAL_EVERY=2000
export SUDOKU_SAVE_EVERY=10000
export SUDOKU_CKPT_KEEP=3
export SUDOKU_LR=0.0002
export SUDOKU_DROPOUT=0.2
export SUDOKU_WD=0.005
export SUDOKU_TRAIN_PATH="../datasets/train_sudoku_puzzles.npy"
export SUDOKU_TEST_PATH="../datasets/test_sudoku_puzzles.npy"
export XLA_PYTHON_CLIENT_MEM_FRACTION=0.9
# ---- Per-arm knobs ----
export SUDOKU_BACKTRACK=0
case "${ARM}" in
flat)
# No latents: no candidate head, no depth ladder. Masks left unset.
export SUDOKU_LATENT_SLOTS=0
export SUDOKU_RECURRENT=0
export SUDOKU_AUX_WEIGHT=0.0
export SUDOKU_TRAIN_CAND=""
export SUDOKU_TEST_CAND=""
;;
datacur)
# No latents, but a real 12-stage curriculum over the puzzle pool, ordered
# by solver round count. Candidate masks stay unset (nothing to supervise);
# only the meta files are needed, for the round counts.
export SUDOKU_LATENT_SLOTS=0
export SUDOKU_RECURRENT=0
export SUDOKU_AUX_WEIGHT=0.0
export SUDOKU_TRAIN_CAND=""
export SUDOKU_TEST_CAND=""
export SUDOKU_DATA_CURRICULUM=rounds
export SUDOKU_TRAIN_META="${CAND_DIR}/train_meta.npy"
export SUDOKU_TEST_META="${CAND_DIR}/test_meta.npy"
for f in "${SUDOKU_TRAIN_META}" "${SUDOKU_TEST_META}"; do
[ -s "${f}" ] || { echo "missing meta ${f}; run sbatch_gen_s12_masks.sh on this node" >&2; exit 1; }
done
;;
latent|latent_bt)
export SUDOKU_LATENT_SLOTS=12
export SUDOKU_RECURRENT=1
export SUDOKU_AUX_WEIGHT=1.0
export SUDOKU_CAND_SLOT_MODE=depth # k = num_passes, slot j -> snapshot j
export SUDOKU_PASSES_PER_STAGE=1 # stage t -> depth t, 12 stages -> 12 slots
export SUDOKU_CAND_DELTA_BG=0.25 # down-weight cells copied from slot j-1
export SUDOKU_TRAIN_CAND="${CAND_DIR}/train_cand_masks.npy"
export SUDOKU_TEST_CAND="${CAND_DIR}/test_cand_masks.npy"
for f in "${SUDOKU_TRAIN_CAND}" "${SUDOKU_TEST_CAND}"; do
[ -s "${f}" ] || { echo "missing masks ${f}; run sbatch_gen_s12_masks.sh on this node" >&2; exit 1; }
done
;;
*) echo "unknown ARM '${ARM}'" >&2; exit 1 ;;
esac
if [ "${ARM}" = "latent_bt" ]; then
# Adaptive replay keyed on depth: deficits are measured on snapshot t's
# candidate-set accuracy, and a repair trains at num_passes=t on the same
# full-corpus batch distribution the frontier uses.
export SUDOKU_BACKTRACK=1
export SUDOKU_BACKTRACK_MODE=adaptive
export SUDOKU_BACKTRACK_MARGIN=0.03
export SUDOKU_BACKTRACK_MAX_REPAIR_STEPS=4000
export SUDOKU_BACKTRACK_MIN_FRONTIER_STEPS=8000
export SUDOKU_BACKTRACK_MAX_REPAIR_FRACTION=0.25
export SUDOKU_BACKTRACK_GRAD_DECAY=0.05
export SUDOKU_BACKTRACK_FRONTIER_MIX=1
fi
cd "${RUN_DIR}"
# Log-only periodic sync: checkpoints stay on /tmp (20 G scratch quota).
(
while true; do sleep 900
rsync -a "${LOCAL_LOG}.log" "${RUN_DIR}/logs/${NAME}.log" 2>/dev/null || true
done
) &
SYNC_PID=$!
echo "[$(date)] starting ${NAME} from scratch"
echo " K=${SUDOKU_LATENT_SLOTS} recurrent=${SUDOKU_RECURRENT} bt=${SUDOKU_BACKTRACK}"
echo " max_stage=${SUDOKU_MAX_STAGE} plateau=${SUDOKU_PLATEAU_STEPS} patience=${SUDOKU_PATIENCE}"
echo " slot_mode=${SUDOKU_CAND_SLOT_MODE:-n/a} pps=${SUDOKU_PASSES_PER_STAGE:-n/a} delta_bg=${SUDOKU_CAND_DELTA_BG:-n/a}"
CUDA_VISIBLE_DEVICES=0 ${PY} -u -m train.main \
--workdir="${LOCAL_LOG}" --exp_name="${NAME}" \
> "${LOCAL_LOG}.log" 2>&1
EC=$?
kill ${SYNC_PID} 2>/dev/null || true
rsync -a "${LOCAL_LOG}.log" "${RUN_DIR}/logs/${NAME}.log" 2>/dev/null || true
echo "[$(date)] ${NAME} exit ${EC}"
exit ${EC}