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