summaryrefslogtreecommitdiff
path: root/external/dualprop_patches/0007-crossover-add-reciprocal-clean-KP-training.patch
diff options
context:
space:
mode:
Diffstat (limited to 'external/dualprop_patches/0007-crossover-add-reciprocal-clean-KP-training.patch')
-rw-r--r--external/dualprop_patches/0007-crossover-add-reciprocal-clean-KP-training.patch321
1 files changed, 321 insertions, 0 deletions
diff --git a/external/dualprop_patches/0007-crossover-add-reciprocal-clean-KP-training.patch b/external/dualprop_patches/0007-crossover-add-reciprocal-clean-KP-training.patch
new file mode 100644
index 0000000..c869770
--- /dev/null
+++ b/external/dualprop_patches/0007-crossover-add-reciprocal-clean-KP-training.patch
@@ -0,0 +1,321 @@
+From 8b8dfd0fd0a0ba01e66bfca454ca521f68910bd5 Mon Sep 17 00:00:00 2001
+From: YurenHao0426 <Blackhao0426@gmail.com>
+Date: Mon, 27 Jul 2026 13:19:37 -0500
+Subject: [PATCH 07/19] crossover: add reciprocal clean KP training
+
+---
+ config/cli_config.py | 4 +-
+ src/__init__.py | 2 +-
+ src/models.py | 33 ++++++++++++++++
+ src/training_utils.py | 79 ++++++++++++++++++++++++++++++++++++++
+ tests/local_rules_smoke.py | 37 ++++++++++++++++++
+ train.py | 39 +++++++++++++++----
+ 6 files changed, 183 insertions(+), 11 deletions(-)
+
+diff --git a/config/cli_config.py b/config/cli_config.py
+index 795e6a3..56fdb9f 100644
+--- a/config/cli_config.py
++++ b/config/cli_config.py
+@@ -42,7 +42,7 @@ parser.add_argument('--experiment-name', default='test', help='A string denoting
+
+ parser.add_argument('--model', default='VGG16', choices=['VGG16', 'VGGlike', 'CNN', 'miniCNN', 'MLP'], help='')
+
+-parser.add_argument('--learning-algorithm', default='dualprop-lagr-ff', choices=['backprop', 'fa', 'dfa', 'pepita', 'ff', 'ep', 'dualprop-lagr-ff', 'dualprop-raovr-ff', 'dualprop-raovr-dampened-ff'])
++parser.add_argument('--learning-algorithm', default='dualprop-lagr-ff', choices=['backprop', 'fa', 'dfa', 'pepita', 'ff', 'ep', 'clean-kp', 'dualprop-lagr-ff', 'dualprop-raovr-ff', 'dualprop-raovr-dampened-ff'])
+
+ parser.add_argument(
+ '--feedback-seed', default=1729, type=int,
+@@ -156,7 +156,7 @@ elif config.model == "MLP":
+ dense_features = [1024, 1024, config.num_classes]
+
+ # Load model
+-modeltype = {"backprop":cnn_abstract, "fa": cnn_abstract, "dfa": cnn_abstract, "pepita": cnn_abstract, "ff": cnn_abstract, "ep": cnn_abstract, "dualprop-lagr-ff": cnn_dualprop_Lagr_ff, "dualprop-raovr-ff": cnn_dualprop_RAOVR_ff, "dualprop-raovr-dampened-ff": cnn_dualprop_RAOVR_dampened_ff}
++modeltype = {"backprop":cnn_abstract, "fa": cnn_abstract, "dfa": cnn_abstract, "pepita": cnn_abstract, "ff": cnn_abstract, "ep": cnn_abstract, "clean-kp": cnn_abstract, "dualprop-lagr-ff": cnn_dualprop_Lagr_ff, "dualprop-raovr-ff": cnn_dualprop_RAOVR_ff, "dualprop-raovr-dampened-ff": cnn_dualprop_RAOVR_dampened_ff}
+ activation={"relu": relu, "hs": hs, "sigmoid": sigmoid, "tanh": tanh}
+ config.model = modeltype[config.learning_algorithm](loss_func, Conv, Dense, activation[config.activation], config.num_classes, config.beta, config.alpha, config.dtype, config.param_dtype,
+ kernels=kernels, strides=strides, features=features, mp = mp,
+diff --git a/src/__init__.py b/src/__init__.py
+index 58c9adc..7ee7740 100644
+--- a/src/__init__.py
++++ b/src/__init__.py
+@@ -1,2 +1,2 @@
+ from .models import cnn_dualprop_Lagr_ff, cnn_dualprop_RAOVR_ff, cnn_dualprop_RAOVR_dampened_ff, cnn_abstract
+-from .training_utils import create_train_state, create_ff_train_state, create_local_feedback, train_epoch, train_ff_epoch, eval_model, eval_ep_model, eval_ff_model, get_mnist, get_svhn, get_fashionmnist, get_cifar10, get_cifar100, get_imagenet_32x32, heatmap_grads_batches, heatmap_grads_epochs, plot_L_or_gamma
++from .training_utils import create_train_state, create_ff_train_state, create_local_feedback, train_epoch, train_ff_epoch, train_kp_epoch, eval_model, eval_ep_model, eval_ff_model, get_mnist, get_svhn, get_fashionmnist, get_cifar10, get_cifar100, get_imagenet_32x32, heatmap_grads_batches, heatmap_grads_epochs, plot_L_or_gamma
+diff --git a/src/models.py b/src/models.py
+index 68e9804..36b2643 100644
+--- a/src/models.py
++++ b/src/models.py
+@@ -195,6 +195,39 @@ class cnn_abstract(nn.Module, ABC):
+ fields[i] = linear_pullback(linear_field)[0]
+ return fields
+
++ def kp_reciprocal_objective(self, s, forward_linear, teaching_fields):
++ """Recompute KP feedback correlations from local activities only.
++
++ This method is evaluated with Q parameters, but activation gates and
++ max-pool switches come from the cached W forward pass. Its derivative
++ therefore equals the corresponding W local correlation for any values
++ of W and Q, without reading either W or its update.
++ """
++ objective = 0.0
++ batch_size = s[0].shape[0]
++ for i, layer in enumerate(self.layers):
++ child_field = jax.lax.stop_gradient(teaching_fields[i + 1])
++ if i == self.num_layers - 1:
++ linear_field = child_field
++ elif i < self.num_convlayers:
++ _, post_pullback = jax.vjp(
++ lambda z: self.act(self.layers[i].pooling(z)),
++ jax.lax.stop_gradient(forward_linear[i]))
++ linear_field = post_pullback(child_field)[0]
++ else:
++ _, post_pullback = jax.vjp(
++ self.act, jax.lax.stop_gradient(forward_linear[i]))
++ linear_field = post_pullback(child_field)[0]
++
++ x = jax.lax.stop_gradient(s[i])
++ if i < self.num_convlayers:
++ prediction = layer.call_without_pooling(x)
++ else:
++ prediction = layer(x)
++ objective += jnp.sum(
++ prediction * jax.lax.stop_gradient(linear_field))
++ return objective / batch_size
++
+ def pepita_correlation_objective(self, modulated_states, original_linear,
+ modulated_linear, one_hot):
+ """Architecture-compatible PEPITA/ERIN two-presentation update.
+diff --git a/src/training_utils.py b/src/training_utils.py
+index f767b80..e573bb5 100644
+--- a/src/training_utils.py
++++ b/src/training_utils.py
+@@ -347,6 +347,35 @@ def train_epoch(state, train_ds, batch_size, rng, augmentation_on,
+ return state, epoch_metrics_np, runtime
+
+
++def train_kp_epoch(state, feedback_state, train_ds, batch_size, rng,
++ augmentation_on, num_classes, gradient_diagnostics=True):
++ """Train W and reciprocal Q once over a shared minibatch order."""
++ t0 = time.time()
++ size = len(train_ds["image"])
++ steps = size // batch_size
++ perms = jax.random.permutation(rng, size)[:steps * batch_size]
++ perms = perms.reshape((steps, batch_size))
++ metrics = []
++ for permutation in perms:
++ image = train_ds["image"][permutation]
++ labels = train_ds["label"][permutation]
++ one_hot = jax.nn.one_hot(labels, num_classes=num_classes)
++ rng, inf_rng, batch_rng = jax.random.split(rng, 3)
++ per_example_rng = jax.random.split(batch_rng, image.shape[0])
++ state, feedback_state, batch_metrics = train_step_kp(
++ state, feedback_state, image, one_hot, labels, per_example_rng,
++ inf_rng, augmentation_on, gradient_diagnostics)
++ metrics.append(batch_metrics)
++ host = jax.device_get(metrics)
++ summary = {}
++ for key in host[0]:
++ if key == "cosine_sim":
++ summary[key] = [record[key] for record in host]
++ else:
++ summary[key] = np.mean([record[key] for record in host], axis=0)
++ return state, feedback_state, summary, time.time() - t0
++
++
+ def direct_feedback_fields(states, output_field, direct_feedback):
+ """Apply per-hidden fixed output maps without a layerwise reverse chain."""
+ fields = [jnp.zeros_like(states[0])]
+@@ -357,6 +386,56 @@ def direct_feedback_fields(states, output_field, direct_feedback):
+ return fields
+
+
++@jax.jit
++def train_step_kp(state, feedback_state, image, labels_onehot, labels,
++ batch_rng, inf_rng, augmentation_on,
++ gradient_diagnostics):
++ """One simultaneous modified-KP W/Q step."""
++ image = jax.lax.cond(
++ augmentation_on, vmap_augment_train, no_aug, image, batch_rng)
++ states, linear = state.apply_fn(
++ {"params": state.params}, image, method="ff_with_local_cache")
++ states = tree_map(jax.lax.stop_gradient, states)
++ linear = tree_map(jax.lax.stop_gradient, linear)
++
++ def output_loss(logits):
++ return state.apply_fn(
++ {"params": state.params}, logits, labels_onehot,
++ method="output_loss")
++
++ output_field = jax.lax.stop_gradient(jax.grad(output_loss)(states[-1]))
++ teaching_fields = feedback_state.apply_fn(
++ {"params": feedback_state.params}, states, linear, output_field,
++ method="fa_teaching_fields")
++ teaching_fields = tree_map(jax.lax.stop_gradient, teaching_fields)
++
++ forward_grads = jax.grad(
++ lambda params: state.apply_fn(
++ {"params": params}, states, teaching_fields,
++ method="local_correlation_objective")
++ )(state.params)
++ reciprocal_grads = jax.grad(
++ lambda params: feedback_state.apply_fn(
++ {"params": params}, states, linear, teaching_fields,
++ method="kp_reciprocal_objective")
++ )(feedback_state.params)
++ metrics = compute_metrics(
++ image=image, labels_onehot=labels_onehot, labels=labels, state=state)
++ metrics["feedback_forward_cosine"] = cosine_sim_tree(
++ feedback_state.params, state.params)
++ metrics["reciprocal_gradient_error"] = jnp.linalg.norm(
++ jax.flatten_util.ravel_pytree(forward_grads)[0]
++ - jax.flatten_util.ravel_pytree(reciprocal_grads)[0])
++ metrics = jax.lax.cond(
++ gradient_diagnostics, ref_grad_and_angle, no_ref_grad_and_angle,
++ state, forward_grads, image, labels_onehot, metrics)
++
++ # Form both independent local correlations before changing either path.
++ feedback_state = feedback_state.apply_gradients(grads=reciprocal_grads)
++ state = state.apply_gradients(grads=forward_grads)
++ return state, feedback_state, metrics
++
++
+ @partial(
+ jax.jit,
+ static_argnames=("free_steps", "nudge_steps"),
+diff --git a/tests/local_rules_smoke.py b/tests/local_rules_smoke.py
+index 51d4c9f..f40f0ab 100644
+--- a/tests/local_rules_smoke.py
++++ b/tests/local_rules_smoke.py
+@@ -81,6 +81,40 @@ def main():
+ )(params)
+ symmetric_error = relative_error(symmetric_grads, bp_grads)
+ assert symmetric_error < 2e-12, symmetric_error
++ reciprocal_grads = jax.grad(
++ lambda candidate: model.apply(
++ {"params": candidate}, states, linear, symmetric_fields,
++ method="kp_reciprocal_objective")
++ )(params)
++ reciprocal_error = relative_error(reciprocal_grads, symmetric_grads)
++ assert reciprocal_error < 2e-12, reciprocal_error
++ changed_feedback = replace_layer(params, "c01", -0.5)
++ changed_reciprocal_grads = jax.grad(
++ lambda candidate: model.apply(
++ {"params": candidate}, states, linear, symmetric_fields,
++ method="kp_reciprocal_objective")
++ )(changed_feedback)
++ reciprocal_independence_error = relative_error(
++ changed_reciprocal_grads, reciprocal_grads)
++ assert reciprocal_independence_error == 0.0
++
++ kp_optimizer = optax.chain(
++ optax.add_decayed_weights(1e-2),
++ optax.sgd(0.1, momentum=0.9),
++ )
++ kp_w = params
++ kp_q = params
++ kp_w_state = kp_optimizer.init(kp_w)
++ kp_q_state = kp_optimizer.init(kp_q)
++ for _ in range(2):
++ w_updates, kp_w_state = kp_optimizer.update(
++ symmetric_grads, kp_w_state, kp_w)
++ q_updates, kp_q_state = kp_optimizer.update(
++ reciprocal_grads, kp_q_state, kp_q)
++ kp_w = optax.apply_updates(kp_w, w_updates)
++ kp_q = optax.apply_updates(kp_q, q_updates)
++ kp_symmetric_tracking_error = relative_error(kp_q, kp_w)
++ assert kp_symmetric_tracking_error == 0.0
+
+ feedback = create_local_feedback(
+ jax.random.PRNGKey(1729), model, params, (8, 8, 2), "fa", 3)
+@@ -224,6 +258,9 @@ def main():
+
+ report = {
+ "symmetric_fa_bp_relative_error": symmetric_error,
++ "kp_reciprocal_gradient_relative_error": reciprocal_error,
++ "kp_reciprocal_independence_error": reciprocal_independence_error,
++ "kp_symmetric_two_step_tracking_error": kp_symmetric_tracking_error,
+ "independent_feedback_forward_cosine": feedback_cosine,
+ "feedback_independence_error": feedback_independence_error,
+ "detached_local_boundary_error": local_boundary_error,
+diff --git a/train.py b/train.py
+index 08f42c9..6b2d8d8 100644
+--- a/train.py
++++ b/train.py
+@@ -8,7 +8,7 @@ from absl import logging # for logging
+ from matplotlib.colors import LogNorm
+
+ # Training utils
+-from src import create_train_state, create_local_feedback, train_epoch, eval_model, eval_ep_model, heatmap_grads_batches, heatmap_grads_epochs, plot_L_or_gamma
++from src import create_train_state, create_local_feedback, train_epoch, train_kp_epoch, eval_model, eval_ep_model, heatmap_grads_batches, heatmap_grads_epochs, plot_L_or_gamma
+
+ # Import configurations
+ # import config # Use this for the old method
+@@ -54,6 +54,14 @@ for experiment_index, seed in enumerate(config.seeds):
+ jax.random.PRNGKey(config.feedback_seed), config.model, state.params,
+ config.image_dims, config.learning_algorithm, config.num_classes,
+ pepita_projection_scale=config.pepita_projection_scale)
++ feedback_state = None
++ if config.learning_algorithm == "clean-kp":
++ feedback_state = create_train_state(
++ jax.random.PRNGKey(config.feedback_seed), config.model,
++ config.image_dims, config.learning_rate,
++ config.warmup_learning_rate, config.learning_rate_final,
++ config.momentum, config.weight_decay, config.num_epochs,
++ config.warmup_epochs, config.decay_epochs, steps_per_epoch)
+
+
+ del init_rng # Must not be used anymore.
+@@ -64,6 +72,8 @@ for experiment_index, seed in enumerate(config.seeds):
+ 'test_loss': np.nan, 'test_accuracy': np.nan, 'test_top5accuracy': np.nan, 'test_time': np.nan,
+ 'grad_cos_sim_batches': np.zeros((len(state.params), steps_per_epoch*config.num_epochs)),
+ 'grad_cos_sim_epochs': np.zeros((len(state.params), config.num_epochs)),
++ 'feedback_forward_cosine': np.zeros((len(state.params), config.num_epochs)),
++ 'reciprocal_gradient_error': np.zeros(config.num_epochs),
+ 'L10': np.zeros((len(state.params), config.num_epochs)),
+ 'L20': np.zeros((len(state.params), config.num_epochs)),
+ 'gamma10': np.zeros((len(state.params), config.num_epochs)),
+@@ -79,13 +89,21 @@ 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,
+- local_feedback=local_feedback,
+- ep_beta=config.ep_beta, ep_free_steps=config.ep_free_steps,
+- ep_nudge_steps=config.ep_nudge_steps, ep_dt=config.ep_dt,
+- gradient_diagnostics=(config.gradient_diagnostics == "full"))
++ if config.learning_algorithm == "clean-kp":
++ state, feedback_state, epoch_metrics, train_time = train_kp_epoch(
++ state, feedback_state, config.train_ds, config.batch_size,
++ input_rng, augmentation_on, config.num_classes,
++ gradient_diagnostics=(
++ config.gradient_diagnostics == "full"))
++ else:
++ state, epoch_metrics, train_time = train_epoch(
++ state, config.train_ds, config.batch_size, input_rng,
++ augmentation_on, config.learning_algorithm,
++ config.num_classes, local_feedback=local_feedback,
++ ep_beta=config.ep_beta, ep_free_steps=config.ep_free_steps,
++ ep_nudge_steps=config.ep_nudge_steps, ep_dt=config.ep_dt,
++ 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
+
+@@ -93,6 +111,11 @@ for experiment_index, seed in enumerate(config.seeds):
+ grad_cos_sim = np.stack(epoch_metrics["cosine_sim"]).T
+ hist["grad_cos_sim_batches"][:,(epoch-1)*steps_per_epoch:(epoch)*steps_per_epoch] = grad_cos_sim
+ hist["grad_cos_sim_epochs"][:,epoch-1] = grad_cos_sim.mean(axis=1)
++ if config.learning_algorithm == "clean-kp":
++ hist["feedback_forward_cosine"][:, epoch - 1] = (
++ epoch_metrics["feedback_forward_cosine"])
++ hist["reciprocal_gradient_error"][epoch - 1] = (
++ epoch_metrics["reciprocal_gradient_error"])
+
+
+ # Evaluate on the validation set after each training epoch
+--
+2.54.0
+