File size: 15,747 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 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 | # coding=utf-8
# Copyright 2024 The Google Research Authors.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Main file for Sudoku GPT experiments."""
import os
import sys
from absl import app
from absl import flags
from absl import logging
from clu import platform
import jax
import tensorflow as tf
import wandb
import ml_collections
from ml_collections import config_flags
from train import train_and_evaluate
from train import train_backtrack
import pdb
sys.dont_write_bytecode = True
logging.set_verbosity(logging.INFO)
FLAGS = flags.FLAGS
_WORKDIR = flags.DEFINE_string(
'workdir',
None,
'Directory to store model data.')
_EXP_NAME = flags.DEFINE_string(
'exp_name',
None,
'Experiment name.')
_CKPT_LOC = flags.DEFINE_string(
'ckpt_loc',
None,
'Directory to restore model.')
config_flags.DEFINE_config_file(
'config',
None,
'File path to the training hyperparameter configuration.',
lock_config=True)
flags.mark_flags_as_required(['workdir', 'exp_name'])
def get_config():
"""Get the default hyperparameter configuration.
Returns:
A ConfigDict object.
"""
# Common configuration for all experiments.
config = ml_collections.ConfigDict()
# Dataset choice
config.dataset = 'sudoku'
# Sequence order
config.seq_order = "solver-order" ## Choices = ["fixed", "solver-order", "random"]
# Training related parameters
config.max_steps = int(os.environ.get("SUDOKU_MAX_STEPS", 100000))
config.dtype = jax.numpy.bfloat16
config.minibatch_size = int(os.environ.get("SUDOKU_MINIBATCH", 64))
# Continuous latent thoughts (ATC / Coconut style) + curriculum.
# Set SUDOKU_LATENT_SLOTS=0 for the no-latent control baseline (plain
# transformer + difficulty curriculum, no recurrence, no candidate head).
config.num_latent_slots = int(os.environ.get("SUDOKU_LATENT_SLOTS", 6))
config.latent_token_id = 10 # placeholder id for inactive slots
# Recurrent latent feedback. 1 = full ATC/Coconut recurrence (slot k built
# from slot k-1's hidden). 0 = NO recurrence: the K stage slots + shared
# candidate head + curriculum + BCE supervision are all kept, but the latent
# hidden state is never fed back, so each stage is an independent parallel
# readout with no carried state (the "stagewise, no-recurrence" control).
config.recurrent_latent = int(os.environ.get("SUDOKU_RECURRENT", 1))
# Backtracking (stage-replay) training: when 1, use the train_backtrack loop
# which interleaves frontier-stage batches with replays of earlier stages at
# their MATCHED recurrence depth (num_passes=t for stage t) to mitigate
# forgetting. backtrack_prob = fraction of steps that are replay steps.
config.backtrack = int(os.environ.get("SUDOKU_BACKTRACK", 0))
# Replay strategy: "prob" = fixed-probability replay of a random earlier stage
# (backtrack_prob). "adaptive" = deficit-driven controller that only replays a
# stage once its held-out accuracy has regressed below its graduation value by
# backtrack_margin, repairs the most-deficient stage first with a FULL
# frontier-style step, and returns once it recovers or hits the step cap.
config.backtrack_mode = os.environ.get("SUDOKU_BACKTRACK_MODE", "prob")
config.backtrack_prob = float(os.environ.get("SUDOKU_BACKTRACK_PROB", 0.5))
config.backtrack_margin = float(os.environ.get("SUDOKU_BACKTRACK_MARGIN", 0.03))
config.backtrack_max_repair_steps = int(
os.environ.get("SUDOKU_BACKTRACK_MAX_REPAIR_STEPS", 6000))
# After leaving a repair episode, require this many frontier steps before
# another repair can trigger (prevents thrashing / starves-the-frontier).
config.backtrack_min_frontier_steps = int(
os.environ.get("SUDOKU_BACKTRACK_MIN_FRONTIER_STEPS", 0))
# While the frontier stage's own level accuracy is below this target, use a
# stricter (larger) repair trigger so stage-6 training is prioritized until
# it is itself "good". 0 disables.
config.backtrack_frontier_target_acc = float(
os.environ.get("SUDOKU_BACKTRACK_FRONTIER_TARGET_ACC", 0.0))
# Optional explicit graduation refs for inherited stages when resuming at
# stage > 1. Comma-separated floats for stages 1..start-1, e.g.
# "0.462,0.487,0.534,0.597,0.636". Empty = seed from first eval (old behavior).
config.backtrack_grad_acc_seed = os.environ.get("SUDOKU_GRAD_ACC_SEED", "")
# Max fraction of steps (since resume) that may be spent in repair mode.
# 0 disables the cap. Typical fair setting: 0.25.
config.backtrack_max_repair_fraction = float(
os.environ.get("SUDOKU_BACKTRACK_MAX_REPAIR_FRACTION", 0.0))
# Soft graduation refs: each eval, if current < grad, decay
# grad <- (1-d)*grad + d*current
# so chronic mild regression is forgiven; sharp drops still trigger.
# 0 disables (hard bars). Typical: 0.05–0.1.
config.backtrack_grad_decay = float(
os.environ.get("SUDOKU_BACKTRACK_GRAD_DECAY", 0.0))
# After this absolute step, force frontier-only (no new repairs).
# 0 disables. Use for "BT early, then freeze".
config.backtrack_freeze_after_step = int(
os.environ.get("SUDOKU_BACKTRACK_FREEZE_AFTER", 0))
# Frontier steps sample uniformly over ALL unlocked levels (standard-loop
# data mix) instead of only the frontier level. 1 = fair "standard+repairs"
# recipe (default); 0 = legacy frontier-level-only behavior.
config.backtrack_frontier_mix = int(
os.environ.get("SUDOKU_BACKTRACK_FRONTIER_MIX", 1))
config.curriculum_start_stage = int(os.environ.get("SUDOKU_START_STAGE", 1))
# Stage s = reasoning DEPTH s: s latent slots active, wave snapshots 1..s
# supervised. Difficulty is NOT gated -- every puzzle is available from
# step 0. Cap max_stage (e.g. =start_stage) to hold a single depth.
config.curriculum_max_stage = int(os.environ.get("SUDOKU_MAX_STAGE", 12))
# Promotion is gated on the candidate-set accuracy of the DEEPEST active
# slot, restricted to cells that changed from the previous snapshot.
config.promote_acc_threshold = float(os.environ.get("SUDOKU_PROMOTE_ACC", 0.85))
# Instance arm: inset is only scored on location-correct cells. Without a
# location floor, a 0.90 inset on 2% loc_acc would promote from a handful
# of cells. Stage-0 chance inset is ~3.7/9 ≈ 0.41, so 0.85 is the set bar
# and this is the "model actually emits the solver-order cells" bar.
config.promote_loc_threshold = float(os.environ.get("SUDOKU_PROMOTE_LOC", 0.70))
config.promote_patience_steps = int(os.environ.get("SUDOKU_PATIENCE", 8000))
config.min_stage_steps = int(os.environ.get("SUDOKU_MIN_STAGE_STEPS", 2000))
# Plateau promotion: advance once the frontier depth stops improving, rather
# than on a fixed timer. A stage is "done" when it has not gained
# plateau_delta over its best accuracy for plateau_steps steps. This is the
# primary rule; the accuracy threshold is a fast path for mastery and
# promote_patience_steps is a hard cap so a stuck stage cannot stall
# training forever. Set plateau_steps=0 to disable.
config.plateau_steps = int(os.environ.get("SUDOKU_PLATEAU_STEPS", 20000))
config.plateau_delta = float(os.environ.get("SUDOKU_PLATEAU_DELTA", 0.005))
# Train-time difficulty balancing. 0 (default) = draw puzzles uniformly from
# the corpus, so the difficulty tag selects nothing. 1 = uniform over the 6
# levels, which upsamples level 8 from 1.8% to 16.7% of batches. Eval is
# always level-balanced so per-level accuracy stays measurable.
config.level_balanced_sampling = int(
os.environ.get("SUDOKU_LEVEL_BALANCED", 0))
# Data curriculum over the puzzle POOL (as opposed to the latent-depth
# curriculum over the supervision).
# "none" = every puzzle available from step 0.
# "rounds" = stage t admits only puzzles whose solver round count falls in
# the first t of max_stage equal-count bins (rounds span 5..38,
# so 12 bins give a genuinely smooth 12-step ladder). This is
# the same propagation-depth axis the latent arm supervises,
# which makes the two arms directly comparable. Needs
# SUDOKU_TRAIN_META / SUDOKU_TEST_META.
config.data_curriculum = os.environ.get("SUDOKU_DATA_CURRICULUM", "none")
config.train_meta_path = os.environ.get("SUDOKU_TRAIN_META", "") or None
config.test_meta_path = os.environ.get("SUDOKU_TEST_META", "") or None
# Model related parameters
config.block_size = 81
config.seq_len = 3 * config.block_size + config.num_latent_slots
config.vocab_size = 11
# Model architecture
config.num_heads = 8
config.num_layers = 8
config.emb_dim = 576
config.qkv_dim = 576
config.mlp_dim = 6 * config.emb_dim
config.dropout_rate = float(os.environ.get("SUDOKU_DROPOUT", 0.2))
config.attention_dropout_rate = float(
os.environ.get("SUDOKU_ATTN_DROPOUT",
os.environ.get("SUDOKU_DROPOUT", 0.2)))
# Training hyperparameters
config.learning_rate = float(os.environ.get("SUDOKU_LR", 0.0002)) # Base learning rate.
config.end_lr_factor = float(os.environ.get("SUDOKU_END_LR_FACTOR", 0.2))
config.warmup_tokens = int(os.environ.get("SUDOKU_WARMUP", 10000))
config.weight_decay = float(os.environ.get("SUDOKU_WD", 0.005))
# Resume from a checkpoint (set SUDOKU_RESUME=1 and pass --ckpt_loc=<path>).
config.resume_training = os.environ.get("SUDOKU_RESUME", "0") == "1"
# Other hyperparameters
config.seed = 7
config.save_checkpoint = os.environ.get("SUDOKU_SAVE_CKPT", "1") == "1"
config.save_every_steps = int(os.environ.get("SUDOKU_SAVE_EVERY", 10000))
# How many checkpoints to retain. Large default so per-stage checkpoints are
# never rolled off (disk is plentiful; ~0.5GB each).
config.ckpt_keep = int(os.environ.get("SUDOKU_CKPT_KEEP", 100))
config.use_wandb = False
config.wandb_project_name = 'sudoku'
# Evaluation related parameters
config.eval_every_steps = int(os.environ.get("SUDOKU_EVAL_EVERY", 2000))
config.eval_epochs = int(os.environ.get("SUDOKU_EVAL_EPOCHS", 5))
# Path to dataset
config.train_puzzle_path = os.environ.get(
"SUDOKU_TRAIN_PATH", "datasets/train_sudoku_puzzles.npy")
config.train_candidate_path = "datasets/train_sudoku_puzzles_candidate.npy"
config.test_puzzle_path = os.environ.get(
"SUDOKU_TEST_PATH", "datasets/test_sudoku_puzzles.npy")
config.test_candidate_path = "datasets/test_sudoku_puzzles_candidate.npy"
# Staged multi-candidate supervision (per-latent-slot BCE targets).
# Empty string disables cand-mask loading (useful for K=0 baselines).
config.train_cand_masks_path = os.environ.get(
"SUDOKU_TRAIN_CAND", "datasets_multicandidate/train_cand_masks.npy") or None
config.test_cand_masks_path = os.environ.get(
"SUDOKU_TEST_CAND", "datasets_multicandidate/test_cand_masks.npy") or None
# Superposition-instance targets. When set, the output prompt's value tokens
# come from one sampled stage-k assignment instead of the unique solution:
# the input prompt (clues) is fixed and the same puzzle recurs with different
# legal completions, so the candidate set is represented across the batch
# rather than supervised as a multi-hot set. Setting this should go with
# SUDOKU_AUX_WEIGHT=0 (candidate head off) -- the masks are then read only
# for the in-set metric. Empty string = classic single-solution targets.
config.instance_dir = os.environ.get("SUDOKU_INSTANCE_DIR", "") or None
# Weight of the auxiliary candidate-set BCE loss relative to the LM CE loss.
config.aux_cand_weight = float(os.environ.get("SUDOKU_AUX_WEIGHT", 1.0))
# Positive-class weight inside the candidate BCE (counters the sparsity of
# the multi-hot masks so the head doesn't collapse to predicting all-zeros).
config.aux_pos_weight = float(os.environ.get("SUDOKU_CAND_POS_WEIGHT", 5.0))
# How many latent slots each example activates.
# "level" = k = clip(level-2, 1, K). Difficulty-matched, but puzzle level
# explains only ~19% of the variance in solver round count, so
# most examples leave the majority of the K slots inert (a
# level-3 puzzle activates ONE slot for ~21 rounds of work).
# "depth" = k = num_passes, uniform over the batch. Every slot the
# recurrence actually fills is active and supervised, and the
# curriculum advances reasoning depth rather than puzzle level.
config.cand_slot_mode = os.environ.get("SUDOKU_CAND_SLOT_MODE", "depth")
# Latent passes granted per curriculum stage: num_passes = min(pps*stage, K).
# pps=2 with K=12 reaches all 12 slots by stage 6.
config.passes_per_stage = int(os.environ.get("SUDOKU_PASSES_PER_STAGE", 1))
# Loss weight for candidate cells that did NOT change from the previous
# stage. Consecutive stages are highly redundant (at 12 stages ~97% of the
# target bits are copies of the previous slot), so plain BCE is dominated by
# echoing the previous slot. <1.0 down-weights the copied cells and puts the
# gradient on the digits actually eliminated at this stage. 1.0 = off.
config.aux_delta_bg = float(os.environ.get("SUDOKU_CAND_DELTA_BG", 1.0))
return config
def main(argv):
if len(argv) > 1:
raise app.UsageError('Too many command-line arguments.')
# # Hide any GPUs from TensorFlow. Otherwise TF might reserve memory and make
# # it unavailable to JAX.
tf.config.experimental.set_visible_devices([], 'GPU')
cfgs = get_config()
if cfgs.resume_training:
assert _CKPT_LOC.value is not None
if cfgs.use_wandb:
wandb.init(project=cfgs.wandb_project_name, name=_EXP_NAME.value, config=cfgs)
logging.info('JAX process: %d / %d', jax.process_index(), jax.process_count())
logging.info('JAX local devices: %r', jax.local_devices())
# Add a note so that we can tell which task is which JAX host.
# (Depending on the platform task 0 is not guaranteed to be host 0)
platform.work_unit().set_task_status(f'process_index: {jax.process_index()}, '
f'process_count: {jax.process_count()}')
platform.work_unit().create_artifact(platform.ArtifactType.DIRECTORY,
_WORKDIR.value, 'workdir')
logging.info(cfgs)
cfgs.workdir = _WORKDIR.value
cfgs.ckpt_loc = _CKPT_LOC.value
if int(getattr(cfgs, "backtrack", 0)):
if str(getattr(cfgs, "backtrack_mode", "prob")) == "adaptive":
train_backtrack.train_and_evaluate_backtrack_adaptive(cfgs, _WORKDIR.value)
else:
train_backtrack.train_and_evaluate_backtrack(cfgs, _WORKDIR.value)
else:
train_and_evaluate.train_and_evaluate(cfgs, _WORKDIR.value)
if cfgs.use_wandb:
wandb.finish()
if __name__ == '__main__':
jax.config.config_with_absl()
app.run(main) |