#!/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}