summaryrefslogtreecommitdiff
path: root/external/dualprop_patches/0001-crossover-separate-author-diagnostics-from-timed-tra.patch
diff options
context:
space:
mode:
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.patch198
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
+