diff options
Diffstat (limited to 'external/dualprop_patches/0001-crossover-separate-author-diagnostics-from-timed-tra.patch')
| -rw-r--r-- | external/dualprop_patches/0001-crossover-separate-author-diagnostics-from-timed-tra.patch | 198 |
1 files changed, 198 insertions, 0 deletions
diff --git a/external/dualprop_patches/0001-crossover-separate-author-diagnostics-from-timed-tra.patch b/external/dualprop_patches/0001-crossover-separate-author-diagnostics-from-timed-tra.patch new file mode 100644 index 0000000..6340b61 --- /dev/null +++ b/external/dualprop_patches/0001-crossover-separate-author-diagnostics-from-timed-tra.patch @@ -0,0 +1,198 @@ +From cd932728011d5fef91109dd24df8a5f2dfb9e6e7 Mon Sep 17 00:00:00 2001 +From: YurenHao0426 <Blackhao0426@gmail.com> +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 + |
