From 8b8dfd0fd0a0ba01e66bfca454ca521f68910bd5 Mon Sep 17 00:00:00 2001 From: YurenHao0426 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