From cd932728011d5fef91109dd24df8a5f2dfb9e6e7 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Mon, 27 Jul 2026 12:50:47 -0500 Subject: [PATCH 01/19] crossover: separate author diagnostics from timed training --- config/cli_config.py | 14 +++++++++++++- src/training_utils.py | 42 ++++++++++++++++++++++++++++++------------ train.py | 27 ++++++++++++++++++++------- 3 files changed, 63 insertions(+), 20 deletions(-) diff --git a/config/cli_config.py b/config/cli_config.py index 0aec5f6..c1b23ba 100644 --- a/config/cli_config.py +++ b/config/cli_config.py @@ -44,6 +44,18 @@ parser.add_argument('--model', default='VGG16', choices=['VGG16', 'VGGlike', 'CN parser.add_argument('--learning-algorithm', default='dualprop-lagr-ff', choices=['backprop', 'dualprop-lagr-ff', 'dualprop-raovr-ff', 'dualprop-raovr-dampened-ff']) +parser.add_argument( + '--gradient-diagnostics', default='full', choices=['none', 'full'], + help=('Compute the exact BP reference gradient and layerwise cosine on ' + 'every training minibatch. The author-compatible default is full; ' + 'use none for diagnostic-free timed crossover runs.')) + +parser.add_argument( + '--spectral-diagnostics', default='full', choices=['none', 'full'], + help=('Run the author power-iteration L/gamma probes after every ' + 'validation evaluation. The author-compatible default is full; ' + 'use none for diagnostic-free timed crossover runs.')) + dtypes = {'bfloat16': jnp.bfloat16, 'float16':jnp.float16, 'float32':jnp.float32} parser.add_argument('--dtype', default='float32', choices=dtypes.keys()) parser.add_argument('--param-dtype', default='float32', choices=['bfloat16', 'float16', 'float32']) @@ -120,4 +132,4 @@ config.model = modeltype[config.learning_algorithm](loss_func, Conv, Dense, acti ) # Load datasets -config.train_ds, config.val_ds, config.test_ds = datasets[config.dataset](config.dtype, config.percent_train, config.percent_val) \ No newline at end of file +config.train_ds, config.val_ds, config.test_ds = datasets[config.dataset](config.dtype, config.percent_train, config.percent_val) diff --git a/src/training_utils.py b/src/training_utils.py index 7152451..2fbfccd 100644 --- a/src/training_utils.py +++ b/src/training_utils.py @@ -240,7 +240,8 @@ def to_float16(ptree): def to_float32(ptree): return tree_map(lambda x: x.astype(jnp.float32), ptree) -def train_epoch(state, train_ds, batch_size, rng, augmentation_on, learning_algorithm, num_classes): +def train_epoch(state, train_ds, batch_size, rng, augmentation_on, + learning_algorithm, num_classes, gradient_diagnostics=True): """Train for a single epoch.""" t0 = time.time() train_ds_size = len(train_ds['image']) @@ -261,7 +262,9 @@ def train_epoch(state, train_ds, batch_size, rng, augmentation_on, learning_algo # image = vmap_augment_train_imagenet(image, batch_rng) if learning_algorithm != "backprop": - state, metrics = train_step(state, image, labels_onehot, labels, batch_rng, inf_rng, augmentation_on) + state, metrics = train_step( + state, image, labels_onehot, labels, batch_rng, inf_rng, + augmentation_on, gradient_diagnostics) elif learning_algorithm == "backprop": state, metrics = train_step_bp(state, image, labels_onehot, labels, batch_rng, inf_rng, augmentation_on) batch_metrics.append(metrics) @@ -279,7 +282,8 @@ def train_epoch(state, train_ds, batch_size, rng, augmentation_on, learning_algo return state, epoch_metrics_np, runtime @jax.jit -def train_step(state, image, labels_onehot, labels, batch_rng, inf_rng, augmentation_on): +def train_step(state, image, labels_onehot, labels, batch_rng, inf_rng, + augmentation_on, gradient_diagnostics): """Train for a single step.""" # batch_rng = jax.random.split(batch_rng, batch['image'].shape[0]) @@ -298,8 +302,9 @@ def train_step(state, image, labels_onehot, labels, batch_rng, inf_rng, augmenta inf_rng, _ = jax.random.split(inf_rng) metrics = compute_metrics(image=image, labels_onehot=labels_onehot, labels=labels, state=state) - get_ref_grad_angle = True - metrics = jax.lax.cond(get_ref_grad_angle, ref_grad_and_angle, no_ref_grad_and_angle, state, grads, image, labels_onehot, metrics) + metrics = jax.lax.cond( + gradient_diagnostics, ref_grad_and_angle, no_ref_grad_and_angle, + state, grads, image, labels_onehot, metrics) # The optimizer may modify grads, so we need to compare grads and ref_grads before performing the gradient step. state = state.apply_gradients(grads=grads) @@ -340,7 +345,8 @@ def eval_step(state, params, image, labels_onehot, labels, inf_rng): metrics = compute_metrics(image=image, labels_onehot=labels_onehot, labels=labels, state=state) return metrics -def eval_model(state, params, test_ds, batch_size, num_classes, eval_rng): +def eval_model(state, params, test_ds, batch_size, num_classes, eval_rng, + spectral_diagnostics=True): t0 = time.time() test_ds_size = len(test_ds['image']) steps = test_ds_size // batch_size @@ -365,11 +371,23 @@ def eval_model(state, params, test_ds, batch_size, num_classes, eval_rng): runtime = time.time() - t0 - # dummy states, used by get_L_and_gamma to infer correct array shape when generating random arrays - sdummy = state.apply_fn({'params': params}, batch["image"], method='make_predictions') - eval_rng, _ = jax.random.split(eval_rng) - L10, gamma10 = state.apply_fn({'params': state.params}, s=sdummy, rng_key=eval_rng, numiter=10, method='get_L_and_gamma') - L20, gamma20 = state.apply_fn({'params': state.params}, s=sdummy, rng_key=eval_rng, numiter=20, method='get_L_and_gamma') + if spectral_diagnostics: + # Dummy states infer the array shapes used by the author's power + # iterations. These diagnostics are expensive and are deliberately + # optional in timed crossover runs. + sdummy = state.apply_fn( + {'params': params}, batch["image"], method='make_predictions') + eval_rng, _ = jax.random.split(eval_rng) + L10, gamma10 = state.apply_fn( + {'params': state.params}, s=sdummy, rng_key=eval_rng, numiter=10, + method='get_L_and_gamma') + L20, gamma20 = state.apply_fn( + {'params': state.params}, s=sdummy, rng_key=eval_rng, numiter=20, + method='get_L_and_gamma') + else: + count = len(state.params) + L10 = L20 = gamma10 = gamma20 = [ + jnp.asarray(jnp.nan, dtype=jnp.float32) for _ in range(count)] return summary['loss'], summary['accuracy'], summary['top5accuracy'], runtime, L10, L20, gamma10, gamma20 @@ -470,4 +488,4 @@ def plot_L_or_gamma(L20, L10, ylabel, save_path): plt.tight_layout() plt.savefig(save_path) plt.close() - return \ No newline at end of file + return diff --git a/train.py b/train.py index bd7f6b7..139fd6c 100644 --- a/train.py +++ b/train.py @@ -71,7 +71,10 @@ for experiment_index, seed in enumerate(config.seeds): # Run an optimization step over a training batch # last augument turns off data augmentation for mnist augmentation_on = (config.dataset!="mnist") and (config.dataset!="fashionmnist") - state, epoch_metrics, train_time = train_epoch(state, config.train_ds, config.batch_size, input_rng, augmentation_on, config.learning_algorithm, config.num_classes) + state, epoch_metrics, train_time = train_epoch( + state, config.train_ds, config.batch_size, input_rng, + augmentation_on, config.learning_algorithm, config.num_classes, + gradient_diagnostics=(config.gradient_diagnostics == "full")) loginfo_and_print('train: \tloss: %.4f, \taccuracy: %.4f, \truntime: %.4f' % (epoch_metrics["loss"], epoch_metrics["accuracy"], train_time)) hist['train_loss'][epoch-1], hist['train_accuracy'][epoch-1], hist['train_time'][epoch-1] = epoch_metrics["loss"], epoch_metrics["accuracy"], train_time @@ -83,7 +86,10 @@ for experiment_index, seed in enumerate(config.seeds): # Evaluate on the validation set after each training epoch rng, input_rng = jax.random.split(rng) - val_loss, val_accuracy, val_top5_accuracy, val_time, L10, L20, gamma10, gamma20 = eval_model(state, state.params, config.val_ds, config.batch_size, config.num_classes, input_rng) + val_loss, val_accuracy, val_top5_accuracy, val_time, L10, L20, gamma10, gamma20 = eval_model( + state, state.params, config.val_ds, config.batch_size, + config.num_classes, input_rng, + spectral_diagnostics=(config.spectral_diagnostics == "full")) loginfo_and_print('val: \tloss: %.4f, \taccuracy: %.4f, \ttop5_accuracy: %.4f, \truntime: %.4f' % (val_loss, val_accuracy, val_top5_accuracy, val_time)) loginfo_and_print(f"L20: {[np.round(Li.item(), decimals=4) for Li in L20]}") loginfo_and_print(f"gamma20: {[np.round(gi.item(), decimals=4) for gi in gamma20]}") @@ -108,18 +114,25 @@ for experiment_index, seed in enumerate(config.seeds): loginfo_and_print(f"\n====Loading model with best validation accuracy (epoch {best_epoch})====") best_state = checkpoints.restore_checkpoint(ckpt_dir=CKPT_DIR, target=state) rng, input_rng = jax.random.split(rng) - test_loss, test_accuracy, test_top5accuracy, test_time, _, _, _, _ = eval_model(best_state, best_state.params, config.test_ds, config.batch_size, config.num_classes, input_rng) + test_loss, test_accuracy, test_top5accuracy, test_time, _, _, _, _ = eval_model( + best_state, best_state.params, config.test_ds, config.batch_size, + config.num_classes, input_rng, + spectral_diagnostics=(config.spectral_diagnostics == "full")) hist['test_loss'], hist['test_accuracy'], hist['test_top5accuracy'], hist['test_time'] = test_loss, test_accuracy, test_top5accuracy, test_time loginfo_and_print('test: \tloss: %.4f, \taccuracy: %.4f, \ttop5_accuracy: %.4f, \truntime: %.4f' % (test_loss, test_accuracy, test_top5accuracy, test_time)) - if config.learning_algorithm != "backprop": + if (config.learning_algorithm != "backprop" + and config.gradient_diagnostics == "full"): # Use color_norm=LogNorm(clip=True) for logscale plot heatmap_grads_epochs(hist["grad_cos_sim_epochs"], outpath+"grad_angle_epochs.pdf", True, color_norm=None) #Grad angle across batches in the first epoch first_N = 100 heatmap_grads_batches(hist["grad_cos_sim_batches"][:,0:first_N], outpath+f"grad_angle_first_{first_N}_batches.pdf", True, color_norm=None) - plot_L_or_gamma(hist["L20"], hist["L10"], "L", outpath+"L.pdf") - plot_L_or_gamma(hist["gamma20"], hist["gamma10"], r"$\gamma$", outpath+"gamma.pdf") + if config.spectral_diagnostics == "full": + plot_L_or_gamma(hist["L20"], hist["L10"], "L", outpath+"L.pdf") + plot_L_or_gamma( + hist["gamma20"], hist["gamma10"], r"$\gamma$", + outpath+"gamma.pdf") - np.save(outpath+"hist.npy", hist) \ No newline at end of file + np.save(outpath+"hist.npy", hist) -- 2.54.0