summaryrefslogtreecommitdiff
path: root/external/dualprop_patches/0002-crossover-add-audited-ordinary-FA-and-DFA-rules.patch
diff options
context:
space:
mode:
Diffstat (limited to 'external/dualprop_patches/0002-crossover-add-audited-ordinary-FA-and-DFA-rules.patch')
-rw-r--r--external/dualprop_patches/0002-crossover-add-audited-ordinary-FA-and-DFA-rules.patch453
1 files changed, 453 insertions, 0 deletions
diff --git a/external/dualprop_patches/0002-crossover-add-audited-ordinary-FA-and-DFA-rules.patch b/external/dualprop_patches/0002-crossover-add-audited-ordinary-FA-and-DFA-rules.patch
new file mode 100644
index 0000000..3138039
--- /dev/null
+++ b/external/dualprop_patches/0002-crossover-add-audited-ordinary-FA-and-DFA-rules.patch
@@ -0,0 +1,453 @@
+From ff5d9524a56d8b440c8524e6e571cf1488ce6aaa Mon Sep 17 00:00:00 2001
+From: YurenHao0426 <Blackhao0426@gmail.com>
+Date: Mon, 27 Jul 2026 13:03:33 -0500
+Subject: [PATCH 02/19] crossover: add audited ordinary FA and DFA rules
+
+---
+ config/cli_config.py | 8 ++-
+ src/__init__.py | 2 +-
+ src/models.py | 80 ++++++++++++++++++++-
+ src/training_utils.py | 90 +++++++++++++++++++++++-
+ tests/local_rules_smoke.py | 138 +++++++++++++++++++++++++++++++++++++
+ train.py | 6 +-
+ 6 files changed, 317 insertions(+), 7 deletions(-)
+ create mode 100644 tests/local_rules_smoke.py
+
+diff --git a/config/cli_config.py b/config/cli_config.py
+index c1b23ba..e1756e1 100644
+--- a/config/cli_config.py
++++ b/config/cli_config.py
+@@ -42,7 +42,11 @@ 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', 'dualprop-lagr-ff', 'dualprop-raovr-ff', 'dualprop-raovr-dampened-ff'])
++parser.add_argument('--learning-algorithm', default='dualprop-lagr-ff', choices=['backprop', 'fa', 'dfa', 'dualprop-lagr-ff', 'dualprop-raovr-ff', 'dualprop-raovr-dampened-ff'])
++
++parser.add_argument(
++ '--feedback-seed', default=1729, type=int,
++ help='Independent fixed-feedback initialization seed for FA/DFA.')
+
+ parser.add_argument(
+ '--gradient-diagnostics', default='full', choices=['none', 'full'],
+@@ -124,7 +128,7 @@ elif config.model == "MLP":
+ dense_features = [1024, 1024, config.num_classes]
+
+ # Load model
+-modeltype = {"backprop":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, "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 f8ce731..7a08c87 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, train_epoch, eval_model, get_mnist, get_svhn, get_fashionmnist, get_cifar10, get_cifar100, get_imagenet_32x32, heatmap_grads_batches, heatmap_grads_epochs, plot_L_or_gamma
+\ No newline at end of file
++from .training_utils import create_train_state, create_local_feedback, train_epoch, eval_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 c162209..edc432c 100644
+--- a/src/models.py
++++ b/src/models.py
+@@ -116,6 +116,84 @@ class cnn_abstract(nn.Module, ABC):
+ s.append(self.act(layer(s[-1])))
+ s.append(self.layers[-1](s[-1]))
+ return s
++
++ def ff_with_local_cache(self, x0):
++ """Forward pass with the linear outputs needed by audited local rules.
++
++ ``s[i]`` is the post-nonlinearity input to layer ``i`` and
++ ``linear[i]`` is that layer's output before its parameter-free
++ pooling/nonlinearity. Keeping the two pieces separate lets feedback
++ alignment use the *forward* ReLU mask and max-pool switches while
++ replacing only the learned linear transpose by an independent random
++ operator.
++ """
++ s = [x0]
++ linear = []
++ for i, layer in enumerate(self.layers):
++ if i < self.num_convlayers:
++ z = layer.call_without_pooling(s[-1])
++ h = self.act(layer.pooling(z))
++ else:
++ z = layer(s[-1])
++ h = z if i == self.num_layers - 1 else self.act(z)
++ linear.append(z)
++ s.append(h)
++ return s, linear
++
++ def local_correlation_objective(self, s, teaching_fields):
++ """Sum independent per-layer correlations for a local weight update.
++
++ Both the layer inputs and teaching fields are detached. Differentiating
++ this scalar with respect to the model parameters therefore evaluates
++ only each layer's eligibility/Jacobian and cannot create a reverse
++ graph through another forward layer.
++ """
++ objective = 0.0
++ batch_size = s[0].shape[0]
++ for i, layer in enumerate(self.layers):
++ x = jax.lax.stop_gradient(s[i])
++ field = jax.lax.stop_gradient(teaching_fields[i + 1])
++ if i < self.num_convlayers:
++ z = layer.call_without_pooling(x)
++ h = self.act(layer.pooling(z))
++ else:
++ z = layer(x)
++ h = z if i == self.num_layers - 1 else self.act(z)
++ objective += jnp.sum(h * field)
++ return objective / batch_size
++
++ def fa_teaching_fields(self, s, linear, output_field):
++ """Transport a post-output field through fixed random linear maps.
++
++ This method is evaluated with an independent parameter tree. The
++ activation and max-pool pullbacks are evaluated at the cached forward
++ linear outputs, while only the convolution/dense transpose comes from
++ this method's parameters. In the audit-only symmetric limit where the
++ parameter tree equals the forward tree, the resulting fields are exact
++ reverse-mode fields.
++ """
++ fields = [jnp.zeros_like(value) for value in s]
++ fields[-1] = output_field
++ for i in range(self.num_layers - 1, 0, -1):
++ child_field = 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)),
++ linear[i])
++ linear_field = post_pullback(child_field)[0]
++ else:
++ _, post_pullback = jax.vjp(self.act, linear[i])
++ linear_field = post_pullback(child_field)[0]
++
++ if i < self.num_convlayers:
++ linear_map = lambda x: self.layers[i].call_without_pooling(x)
++ else:
++ linear_map = self.layers[i]
++ _, linear_pullback = jax.vjp(linear_map, s[i])
++ fields[i] = linear_pullback(linear_field)[0]
++ return fields
+
+ def init_states_to_zero(self, x0):
+ s = [x0]
+@@ -327,4 +405,4 @@ class cnn_dualprop_RAOVR_dampened_ff(cnn_dualprop_abstract):
+ # delta = s[i+1] - fa[i+1]
+ # s[i], fa[i] = self.infer_hidden(s[i], s[i-1], delta, self.layers[i-1], self.layers[i], L[i])
+
+-# return s, fa
+\ No newline at end of file
++# return s, fa
+diff --git a/src/training_utils.py b/src/training_utils.py
+index 2fbfccd..b6e1e3c 100644
+--- a/src/training_utils.py
++++ b/src/training_utils.py
+@@ -15,6 +15,7 @@ import time, datetime # measuring runetime and generating timestamps for experim
+ import os # for os.makdirs() function
+ import seaborn as sns
+ import matplotlib.pylab as plt
++from functools import partial
+
+ class SumGreaterThan100Error(Exception):
+ pass
+@@ -216,6 +217,33 @@ def create_train_state(rng, model, image_dims, lr, wlr, lrf, momentum, weight_de
+
+ return train_state.TrainState.create(apply_fn=model.apply, params=unfreeze(params), tx=tx)
+
++
++def create_local_feedback(rng, model, params, image_dims, learning_algorithm,
++ num_classes):
++ """Create fixed feedback without reading a forward parameter value.
++
++ FA uses an independently initialized model-shaped linear parameter tree.
++ DFA uses one independent output-to-hidden tensor per hidden population.
++ The forward tree is used only to infer activation shapes for DFA.
++ """
++ if learning_algorithm not in ("fa", "dfa"):
++ return None
++ w, h, channels = image_dims
++ dummy = jnp.ones([1, w, h, channels])
++ if learning_algorithm == "fa":
++ return unfreeze(model.init(rng, dummy)["params"])
++
++ states, _ = model.apply(
++ {"params": params}, dummy, method="ff_with_local_cache")
++ keys = jax.random.split(rng, len(states) - 2)
++ scale = jnp.asarray(num_classes ** -0.5, dtype=dummy.dtype)
++ return tuple(
++ scale * jax.random.normal(
++ key, (num_classes,) + tuple(state.shape[1:]),
++ dtype=state.dtype)
++ for key, state in zip(keys, states[1:-1])
++ )
++
+ def augment_train(image, batch_rng):
+ w, h, c = image.shape
+
+@@ -241,7 +269,8 @@ 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, gradient_diagnostics=True):
++ learning_algorithm, num_classes, local_feedback=None,
++ gradient_diagnostics=True):
+ """Train for a single epoch."""
+ t0 = time.time()
+ train_ds_size = len(train_ds['image'])
+@@ -261,7 +290,12 @@ def train_epoch(state, train_ds, batch_size, rng, augmentation_on,
+ batch_rng = jax.random.split(batch_rng, image.shape[0])
+ # image = vmap_augment_train_imagenet(image, batch_rng)
+
+- if learning_algorithm != "backprop":
++ if learning_algorithm in ("fa", "dfa"):
++ state, metrics = train_step_local(
++ state, local_feedback, learning_algorithm, image,
++ labels_onehot, labels, batch_rng, augmentation_on,
++ gradient_diagnostics)
++ elif learning_algorithm != "backprop":
+ state, metrics = train_step(
+ state, image, labels_onehot, labels, batch_rng, inf_rng,
+ augmentation_on, gradient_diagnostics)
+@@ -281,6 +315,58 @@ def train_epoch(state, train_ds, batch_size, rng, augmentation_on,
+
+ return state, epoch_metrics_np, runtime
+
++
++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])]
++ for feedback in direct_feedback:
++ fields.append(jnp.tensordot(
++ output_field, feedback, axes=((-1,), (0,))))
++ fields.append(output_field)
++ return fields
++
++
++@partial(jax.jit, static_argnames=("learning_algorithm",))
++def train_step_local(state, local_feedback, learning_algorithm, image,
++ labels_onehot, labels, batch_rng, augmentation_on,
++ gradient_diagnostics):
++ """One ordinary-FA or DFA step with detached local eligibilities."""
++ 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]))
++ if learning_algorithm == "fa":
++ teaching_fields = state.apply_fn(
++ {"params": local_feedback}, states, linear, output_field,
++ method="fa_teaching_fields")
++ else:
++ teaching_fields = direct_feedback_fields(
++ states, output_field, local_feedback)
++ teaching_fields = tree_map(jax.lax.stop_gradient, teaching_fields)
++
++ def local_objective(params):
++ return state.apply_fn(
++ {"params": params}, states, teaching_fields,
++ method="local_correlation_objective")
++
++ grads = jax.grad(local_objective)(state.params)
++ metrics = compute_metrics(
++ image=image, labels_onehot=labels_onehot, labels=labels, state=state)
++ metrics = jax.lax.cond(
++ gradient_diagnostics, ref_grad_and_angle, no_ref_grad_and_angle,
++ state, grads, image, labels_onehot, metrics)
++ state = state.apply_gradients(grads=grads)
++ return state, metrics
++
+ @jax.jit
+ def train_step(state, image, labels_onehot, labels, batch_rng, inf_rng,
+ augmentation_on, gradient_diagnostics):
+diff --git a/tests/local_rules_smoke.py b/tests/local_rules_smoke.py
+new file mode 100644
+index 0000000..2366e61
+--- /dev/null
++++ b/tests/local_rules_smoke.py
+@@ -0,0 +1,138 @@
++#!/usr/bin/env python3
++"""Deterministic mechanics checks for matched ordinary FA and DFA."""
++import json
++import os
++import sys
++
++import jax
++import jax.numpy as jnp
++import optax
++from flax import linen as nn
++from flax.core.frozen_dict import unfreeze
++
++
++ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
++sys.path.insert(0, ROOT)
++
++from src.models import cnn_abstract
++from src.training_utils import create_local_feedback, direct_feedback_fields
++
++
++def loss_func(logits, one_hot):
++ return jnp.sum(optax.softmax_cross_entropy(logits, one_hot))
++
++
++def flat(tree):
++ return jax.flatten_util.ravel_pytree(tree)[0].astype(jnp.float64)
++
++
++def relative_error(left, right):
++ return float(
++ jnp.linalg.norm(flat(left) - flat(right))
++ / jnp.maximum(jnp.linalg.norm(flat(right)), 1e-30)
++ )
++
++
++def replace_layer(tree, layer_name, delta):
++ changed = unfreeze(tree)
++ changed[layer_name] = jax.tree_util.tree_map(
++ lambda value: value + delta, changed[layer_name])
++ return changed
++
++
++def main():
++ jax.config.update("jax_enable_x64", True)
++ model = cnn_abstract(
++ loss_func, nn.Conv, nn.Dense, nn.relu, 3, 0.1, 0.0,
++ jnp.float64, jnp.float64,
++ kernels=[(3, 3), (3, 3)], strides=[(1, 1), (1, 1)],
++ features=[4, 5], mp=[True, True], dense_features=[3],
++ inference_sequence="fwK", inference_passes_nudged=1)
++ x = jax.random.normal(jax.random.PRNGKey(1), (3, 8, 8, 2),
++ dtype=jnp.float64)
++ labels = jnp.asarray([0, 2, 1])
++ one_hot = jax.nn.one_hot(labels, 3, dtype=jnp.float64)
++ params = unfreeze(model.init(jax.random.PRNGKey(2), x)["params"])
++ states, linear = model.apply(
++ {"params": params}, x, method="ff_with_local_cache")
++
++ def batch_loss(candidate):
++ logits = model.apply({"params": candidate}, x)
++ return model.apply(
++ {"params": candidate}, logits, one_hot,
++ method="output_loss") / x.shape[0]
++
++ bp_grads = jax.grad(batch_loss)(params)
++ output_field = jax.grad(
++ lambda logits: model.apply(
++ {"params": params}, logits, one_hot, method="output_loss")
++ )(states[-1])
++ symmetric_fields = model.apply(
++ {"params": params}, states, linear, output_field,
++ method="fa_teaching_fields")
++ symmetric_grads = jax.grad(
++ lambda candidate: model.apply(
++ {"params": candidate}, states, symmetric_fields,
++ method="local_correlation_objective")
++ )(params)
++ symmetric_error = relative_error(symmetric_grads, bp_grads)
++ assert symmetric_error < 2e-12, symmetric_error
++
++ feedback = create_local_feedback(
++ jax.random.PRNGKey(1729), model, params, (8, 8, 2), "fa", 3)
++ feedback_cosine = float(
++ jnp.vdot(flat(feedback), flat(params))
++ / (jnp.linalg.norm(flat(feedback)) * jnp.linalg.norm(flat(params)))
++ )
++ assert abs(feedback_cosine) < 0.25, feedback_cosine
++ random_fields = model.apply(
++ {"params": feedback}, states, linear, output_field,
++ method="fa_teaching_fields")
++ changed_forward = replace_layer(params, "d00", 0.75)
++ random_fields_again = model.apply(
++ {"params": feedback}, states, linear, output_field,
++ method="fa_teaching_fields")
++ feedback_independence_error = relative_error(
++ random_fields[1:-1], random_fields_again[1:-1])
++ assert feedback_independence_error == 0.0
++
++ random_fields = jax.tree_util.tree_map(
++ jax.lax.stop_gradient, random_fields)
++ local_grads = jax.grad(
++ lambda candidate: model.apply(
++ {"params": candidate}, states, random_fields,
++ method="local_correlation_objective")
++ )(params)
++ changed_local_grads = jax.grad(
++ lambda candidate: model.apply(
++ {"params": candidate}, states, random_fields,
++ method="local_correlation_objective")
++ )(changed_forward)
++ local_boundary_error = relative_error(
++ local_grads["c00"], changed_local_grads["c00"])
++ assert local_boundary_error == 0.0
++
++ direct = create_local_feedback(
++ jax.random.PRNGKey(1730), model, params, (8, 8, 2), "dfa", 3)
++ assert len(direct) == len(states) - 2
++ dfa_fields = direct_feedback_fields(states, output_field, direct)
++ assert len(dfa_fields) == len(states)
++ assert all(
++ field.shape == state.shape
++ for field, state in zip(dfa_fields, states)
++ )
++ assert all(bool(jnp.all(jnp.isfinite(field))) for field in dfa_fields)
++
++ report = {
++ "symmetric_fa_bp_relative_error": symmetric_error,
++ "independent_feedback_forward_cosine": feedback_cosine,
++ "feedback_independence_error": feedback_independence_error,
++ "detached_local_boundary_error": local_boundary_error,
++ "dfa_hidden_maps": len(direct),
++ "status": "passed",
++ }
++ print(json.dumps(report, indent=2, sort_keys=True))
++
++
++if __name__ == "__main__":
++ main()
+diff --git a/train.py b/train.py
+index 139fd6c..c030158 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, train_epoch, eval_model, heatmap_grads_batches, heatmap_grads_epochs, plot_L_or_gamma
++from src import create_train_state, create_local_feedback, train_epoch, eval_model, heatmap_grads_batches, heatmap_grads_epochs, plot_L_or_gamma
+
+ # Import configurations
+ # import config # Use this for the old method
+@@ -46,6 +46,9 @@ for experiment_index, seed in enumerate(config.seeds):
+ steps_per_epoch = len(config.train_ds['image']) // config.batch_size
+
+ state = create_train_state(init_rng, 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)
++ local_feedback = create_local_feedback(
++ jax.random.PRNGKey(config.feedback_seed), config.model, state.params,
++ config.image_dims, config.learning_algorithm, config.num_classes)
+
+
+ del init_rng # Must not be used anymore.
+@@ -74,6 +77,7 @@ for experiment_index, seed in enumerate(config.seeds):
+ 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,
+ 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
+--
+2.54.0
+