diff options
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.patch | 321 |
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 + |
