File size: 5,520 Bytes
bb23b91
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
146
#!/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=w12_instbt
#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

# Same instance-CE latent curriculum as w12_inst, plus adaptive backtrack:
# if a graduated stage's in-set rate falls more than the margin below its
# graduation value, train at that stage's depth until it recovers.

set -u
hostname
nvidia-smi -L
echo "[$(date)] w12_inst_latent_bt"

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
INST_DIR=/tmp/sudoku_superposition
LOCAL_LOG=/tmp/sudoku_wave_runs/w12_inst_latent_bt
TARBALL_GANDALF=/tmp/logicpuzzles_env.tar.gz
TARBALL_LOCAL=/tmp/logicpuzzles_env_${SLURM_JOB_ID}.tar.gz
GANDALF_INST=gandalf.berkeley.edu:/tmp/sudoku_superposition
GANDALF_CAND=gandalf.berkeley.edu:/tmp/sudoku_s12

mkdir -p "${RUN_DIR}/logs" "${LOCAL_LOG}" /tmp/sudoku_wave_runs

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"

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 ${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())"

need_inst=0
for f in train_assignments.npy train_starts.npy train_counts.npy \
         test_assignments.npy test_starts.npy test_counts.npy; do
  [ -s "${INST_DIR}/${f}" ] || need_inst=1
done
if [ "${need_inst}" = 1 ]; then
  echo "[$(date)] pulling instances from ${GANDALF_INST}"
  mkdir -p "${INST_DIR}"
  rsync -a --progress -e "ssh ${SCP_OPTS}" \
    "${GANDALF_INST}/" "${INST_DIR}/"
fi
for f in train_assignments.npy train_starts.npy train_counts.npy \
         test_assignments.npy test_starts.npy test_counts.npy; do
  [ -s "${INST_DIR}/${f}" ] || { echo "missing ${INST_DIR}/${f}" >&2; exit 1; }
done

need_cand=0
for f in train_cand_masks.npy test_cand_masks.npy; do
  [ -s "${CAND_DIR}/${f}" ] || need_cand=1
done
if [ "${need_cand}" = 1 ]; then
  echo "[$(date)] pulling s12 masks from ${GANDALF_CAND}"
  mkdir -p "${CAND_DIR}"
  rsync -a -e "ssh ${SCP_OPTS}" "${GANDALF_CAND}/" "${CAND_DIR}/" || true
fi
for f in train_cand_masks.npy test_cand_masks.npy; do
  [ -s "${CAND_DIR}/${f}" ] || { echo "missing ${CAND_DIR}/${f}" >&2; exit 1; }
done

export SUDOKU_RESUME=0
export SUDOKU_START_STAGE=1
export SUDOKU_MAX_STAGE=12
export SUDOKU_LATENT_SLOTS=12
export SUDOKU_RECURRENT=1
export SUDOKU_CAND_SLOT_MODE=depth
export SUDOKU_PASSES_PER_STAGE=1
export SUDOKU_AUX_WEIGHT=0.0
export SUDOKU_LEVEL_BALANCED=0
export SUDOKU_DATA_CURRICULUM=none
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.85
export SUDOKU_PROMOTE_LOC=0.70
export SUDOKU_MAX_STEPS="${SUDOKU_MAX_STEPS:-800000}"
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="${SCRATCH_ROOT}/sudoku-code/datasets/train_sudoku_puzzles.npy"
export SUDOKU_TEST_PATH="${SCRATCH_ROOT}/sudoku-code/datasets/test_sudoku_puzzles.npy"
export SUDOKU_TRAIN_CAND="${CAND_DIR}/train_cand_masks.npy"
export SUDOKU_TEST_CAND="${CAND_DIR}/test_cand_masks.npy"
export SUDOKU_INSTANCE_DIR="${INST_DIR}"
export XLA_PYTHON_CLIENT_MEM_FRACTION=0.9

# Adaptive repair: if stage t's in-set rate drops more than 0.03 below
# the value it graduated at, replay that depth until it recovers.
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

cd "${RUN_DIR}"
(
  while true; do sleep 900
    rsync -a "${LOCAL_LOG}.log" "${RUN_DIR}/logs/w12_inst_latent_bt.log" 2>/dev/null || true
  done
) &
SYNC_PID=$!

echo "[$(date)] starting w12_inst_latent_bt from scratch"
echo "  K=12 recurrent=1 bt=adaptive aux=0 instance_dir=${INST_DIR}"
CUDA_VISIBLE_DEVICES=0 ${PY} -u -m train.main \
  --workdir="${LOCAL_LOG}" --exp_name="w12_inst_latent_bt" \
  > "${LOCAL_LOG}.log" 2>&1
EC=$?
kill ${SYNC_PID} 2>/dev/null || true
rsync -a "${LOCAL_LOG}.log" "${RUN_DIR}/logs/w12_inst_latent_bt.log" 2>/dev/null || true
echo "[$(date)] w12_inst_latent_bt exit ${EC}"
exit ${EC}