From b925ba7f07389054cba7b90d2b0db3d7e1699b94 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Mon, 27 Jul 2026 13:16:22 -0500 Subject: [PATCH 06/19] crossover: add two-phase equilibrium propagation --- config/cli_config.py | 20 +++++++- src/__init__.py | 2 +- src/models.py | 57 +++++++++++++++++++++++ src/training_utils.py | 93 ++++++++++++++++++++++++++++++++++++++ tests/local_rules_smoke.py | 34 ++++++++++++++ train.py | 30 ++++++++---- 6 files changed, 224 insertions(+), 12 deletions(-) diff --git a/config/cli_config.py b/config/cli_config.py index 159e9bc..795e6a3 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', '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', 'dualprop-lagr-ff', 'dualprop-raovr-ff', 'dualprop-raovr-dampened-ff']) parser.add_argument( '--feedback-seed', default=1729, type=int, @@ -60,6 +60,22 @@ parser.add_argument( '--ff-score-from-layer', default=1, type=int, help='First zero-indexed FF layer included in candidate-label goodness.') +parser.add_argument( + '--ep-beta', default=0.5, type=float, + help='Magnitude of the randomly signed EP output nudge.') + +parser.add_argument( + '--ep-dt', default=0.5, type=float, + help='Euler step size for EP state relaxation.') + +parser.add_argument( + '--ep-free-steps', default=20, type=int, + help='Number of free-phase EP relaxation steps.') + +parser.add_argument( + '--ep-nudge-steps', default=4, type=int, + help='Number of nudged-phase EP relaxation steps.') + parser.add_argument( '--gradient-diagnostics', default='full', choices=['none', 'full'], help=('Compute the exact BP reference gradient and layerwise cosine on ' @@ -140,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, "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, "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 752dcc2..58c9adc 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_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, 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 0f5af4f..68e9804 100644 --- a/src/models.py +++ b/src/models.py @@ -274,6 +274,63 @@ class cnn_abstract(nn.Module, ABC): axes = tuple(range(1, x.ndim)) score += jnp.mean(jnp.square(x), axis=axes) return score + + @staticmethod + def _ep_rho(state): + return jnp.clip(state, 0.0, 1.0) + + @staticmethod + def _ep_rhop(state): + return jnp.asarray( + (state >= 0.0) & (state <= 1.0), dtype=state.dtype) + + def ep_relax(self, x, one_hot, beta, steps, dt, initial_states=None): + """Canonical hard-sigmoid EP dynamics on the layered author graph. + + The update follows the original two-phase implementation: synchronous + leaky state dynamics, symmetric top-down interactions through the + forward weights, and a squared-error output nudge. Max-pool pullbacks + use the local switches induced by the current lower state. + """ + if initial_states is None: + states = self.init_states_to_zero(x) + else: + states = initial_states + for _ in range(steps): + updated = [x] + for i in range(1, self.num_layers): + below = self.layers[i - 1](self._ep_rho(states[i - 1])) + above = grad(self.get_phi, argnums=(1))( + self._ep_rho(states[i + 1]), + self._ep_rho(states[i]), + self.layers[i]) + drive = ( + -self._ep_rho(states[i]) + below + above) + next_state = states[i] + dt * self._ep_rhop(states[i]) * drive + updated.append(self._ep_rho(next_state)) + + below = self.layers[-1](self._ep_rho(states[-2])) + drive = -self._ep_rho(states[-1]) + below + drive += 2.0 * beta * (one_hot - states[-1]) + next_output = ( + states[-1] + dt * self._ep_rhop(states[-1]) * drive) + updated.append(self._ep_rho(next_output)) + states = updated + return states + + def ep_contrastive_objective(self, free_states, nudged_states, beta): + """Local EP correlation difference for optimizer-style descent.""" + free_phi = 0.0 + nudged_phi = 0.0 + for i, layer in enumerate(self.layers): + free_phi += self.get_phi( + self._ep_rho(free_states[i + 1]), + self._ep_rho(free_states[i]), layer) + nudged_phi += self.get_phi( + self._ep_rho(nudged_states[i + 1]), + self._ep_rho(nudged_states[i]), layer) + return (free_phi - nudged_phi) / ( + beta * free_states[0].shape[0]) def init_states_to_zero(self, x0): s = [x0] diff --git a/src/training_utils.py b/src/training_utils.py index 734f46d..f767b80 100644 --- a/src/training_utils.py +++ b/src/training_utils.py @@ -291,6 +291,7 @@ def to_float32(ptree): def train_epoch(state, train_ds, batch_size, rng, augmentation_on, learning_algorithm, num_classes, local_feedback=None, + ep_beta=0.5, ep_free_steps=20, ep_nudge_steps=4, ep_dt=0.5, gradient_diagnostics=True): """Train for a single epoch.""" t0 = time.time() @@ -320,6 +321,11 @@ def train_epoch(state, train_ds, batch_size, rng, augmentation_on, state, metrics = train_step_pepita( state, local_feedback, image, labels_onehot, labels, batch_rng, augmentation_on, gradient_diagnostics) + elif learning_algorithm == "ep": + state, metrics = train_step_ep( + state, image, labels_onehot, labels, batch_rng, inf_rng, + augmentation_on, ep_beta, ep_free_steps, ep_nudge_steps, + ep_dt, gradient_diagnostics) elif learning_algorithm != "backprop": state, metrics = train_step( state, image, labels_onehot, labels, batch_rng, inf_rng, @@ -351,6 +357,47 @@ def direct_feedback_fields(states, output_field, direct_feedback): return fields +@partial( + jax.jit, + static_argnames=("free_steps", "nudge_steps"), +) +def train_step_ep(state, image, labels_onehot, labels, batch_rng, inf_rng, + augmentation_on, beta, free_steps, nudge_steps, dt, + gradient_diagnostics): + """One canonical two-phase equilibrium-propagation update.""" + image = jax.lax.cond( + augmentation_on, vmap_augment_train, no_aug, image, batch_rng) + free_states = state.apply_fn( + {"params": state.params}, image, labels_onehot, 0.0, free_steps, dt, + method="ep_relax") + beta_sign = jnp.where( + jax.random.bernoulli(inf_rng), 1.0, -1.0) + signed_beta = beta * beta_sign + nudged_states = state.apply_fn( + {"params": state.params}, image, labels_onehot, signed_beta, + nudge_steps, dt, free_states, method="ep_relax") + free_states = tree_map(jax.lax.stop_gradient, free_states) + nudged_states = tree_map(jax.lax.stop_gradient, nudged_states) + + def contrastive_objective(params): + return state.apply_fn( + {"params": params}, free_states, nudged_states, signed_beta, + method="ep_contrastive_objective") + + grads = jax.grad(contrastive_objective)(state.params) + logits = free_states[-1] + metrics = { + "loss": jnp.mean(jnp.square(logits - labels_onehot)), + "accuracy": 100.0 * jnp.mean( + jnp.argmax(logits, -1) == labels, dtype=jnp.float32), + } + 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 + + def ff_overlay(image, labels, num_classes): """Overlay a candidate label on the first input coordinates.""" flat = image.reshape((image.shape[0], -1)) @@ -637,6 +684,52 @@ def eval_model(state, params, test_ds, batch_size, num_classes, eval_rng, return summary['loss'], summary['accuracy'], summary['top5accuracy'], runtime, L10, L20, gamma10, gamma20 + +@partial(jax.jit, static_argnames=("free_steps",)) +def eval_step_ep(state, params, image, labels_onehot, labels, free_steps, dt): + states = state.apply_fn( + {"params": params}, image, labels_onehot, 0.0, free_steps, dt, + method="ep_relax") + output = states[-1] + _, top5_indices = jax.lax.top_k(output, 5) + return { + "loss": jnp.mean(jnp.square(output - labels_onehot)), + "accuracy": 100.0 * jnp.mean( + jnp.argmax(output, -1) == labels, dtype=jnp.float32), + "top5accuracy": 100.0 * jnp.mean( + jnp.any(top5_indices == labels[:, None], axis=1), + dtype=jnp.float32), + } + + +def eval_ep_model(state, params, dataset, batch_size, num_classes, free_steps, + dt): + """Evaluate classification from the settled free-phase EP output.""" + t0 = time.time() + size = len(dataset["image"]) + steps = size // batch_size + indices = jnp.arange(steps * batch_size).reshape((steps, batch_size)) + metrics = [] + for index in indices: + labels = dataset["label"][index] + one_hot = jax.nn.one_hot(labels, num_classes=num_classes) + metrics.append(eval_step_ep( + state, params, dataset["image"][index], one_hot, labels, + free_steps, dt)) + host = jax.device_get(metrics) + summary = { + key: float(np.mean([record[key] for record in host])) + for key in host[0] + } + runtime = time.time() - t0 + nan_diagnostics = [ + jnp.asarray(jnp.nan, dtype=jnp.float32) for _ in range(len(params)) + ] + return ( + summary["loss"], summary["accuracy"], summary["top5accuracy"], + runtime, nan_diagnostics, nan_diagnostics, nan_diagnostics, + nan_diagnostics) + def ref_grad_and_angle(state, grads, image, labels_onehot, metrics): # Compute a reference backprop gradient (but don't use it). def loss_fn_ref(params, state): diff --git a/tests/local_rules_smoke.py b/tests/local_rules_smoke.py index 8ce1e02..51d4c9f 100644 --- a/tests/local_rules_smoke.py +++ b/tests/local_rules_smoke.py @@ -189,6 +189,39 @@ def main(): assert ff_scores.shape == (x.shape[0], 3) assert bool(jnp.all(jnp.isfinite(ff_scores))) + ep_free = model.apply( + {"params": params}, x, one_hot, 0.0, 2, 0.5, + method="ep_relax") + ep_nudged = model.apply( + {"params": params}, x, one_hot, 0.5, 1, 0.5, ep_free, + method="ep_relax") + assert all( + bool(jnp.all((state >= 0.0) & (state <= 1.0))) + for state in ep_free[1:] + ep_nudged[1:] + ) + ep_grads = jax.grad( + lambda candidate: model.apply( + {"params": candidate}, ep_free, ep_nudged, 0.5, + method="ep_contrastive_objective") + )(params) + free_readout_input = ep_free[-2].reshape((x.shape[0], -1)) + nudged_readout_input = ep_nudged[-2].reshape((x.shape[0], -1)) + ep_manual_kernel = ( + free_readout_input.T @ ep_free[-1] + - nudged_readout_input.T @ ep_nudged[-1] + ) / (0.5 * x.shape[0]) + ep_manual_bias = (ep_free[-1] - ep_nudged[-1]).mean(axis=0) / 0.5 + ep_readout_leaves = jax.tree_util.tree_leaves(ep_grads["d00"]) + ep_readout_kernel = next( + value for value in ep_readout_leaves if value.ndim == 2) + ep_readout_bias = next( + value for value in ep_readout_leaves if value.ndim == 1) + ep_contrastive_error = max( + float(jnp.max(jnp.abs(ep_readout_kernel - ep_manual_kernel))), + float(jnp.max(jnp.abs(ep_readout_bias - ep_manual_bias))), + ) + assert ep_contrastive_error < 2e-12, ep_contrastive_error + report = { "symmetric_fa_bp_relative_error": symmetric_error, "independent_feedback_forward_cosine": feedback_cosine, @@ -199,6 +232,7 @@ def main(): "pepita_projection_shape": list(pepita_projection.shape), "ff_local_objective_error": ff_objective_error, "ff_score_shape": list(ff_scores.shape), + "ep_contrastive_readout_max_error": ep_contrastive_error, "status": "passed", } print(json.dumps(report, indent=2, sort_keys=True)) diff --git a/train.py b/train.py index d3f0de6..08f42c9 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, heatmap_grads_batches, heatmap_grads_epochs, plot_L_or_gamma +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 # Import configurations # import config # Use this for the old method @@ -83,6 +83,8 @@ for experiment_index, seed in enumerate(config.seeds): 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 @@ -95,10 +97,15 @@ for experiment_index, seed in enumerate(config.seeds): # Evaluate on the validation set after each training epoch rng, input_rng = jax.random.split(rng) - val_loss, val_accuracy, val_top5_accuracy, val_time, L10, L20, gamma10, gamma20 = eval_model( - state, state.params, config.val_ds, config.batch_size, - config.num_classes, input_rng, - spectral_diagnostics=(config.spectral_diagnostics == "full")) + if config.learning_algorithm == "ep": + val_loss, val_accuracy, val_top5_accuracy, val_time, L10, L20, gamma10, gamma20 = eval_ep_model( + state, state.params, config.val_ds, config.batch_size, + config.num_classes, config.ep_free_steps, config.ep_dt) + else: + val_loss, val_accuracy, val_top5_accuracy, val_time, L10, L20, gamma10, gamma20 = eval_model( + state, state.params, config.val_ds, config.batch_size, + config.num_classes, input_rng, + spectral_diagnostics=(config.spectral_diagnostics == "full")) loginfo_and_print('val: \tloss: %.4f, \taccuracy: %.4f, \ttop5_accuracy: %.4f, \truntime: %.4f' % (val_loss, val_accuracy, val_top5_accuracy, val_time)) loginfo_and_print(f"L20: {[np.round(Li.item(), decimals=4) for Li in L20]}") loginfo_and_print(f"gamma20: {[np.round(gi.item(), decimals=4) for gi in gamma20]}") @@ -123,10 +130,15 @@ for experiment_index, seed in enumerate(config.seeds): loginfo_and_print(f"\n====Loading model with best validation accuracy (epoch {best_epoch})====") best_state = checkpoints.restore_checkpoint(ckpt_dir=CKPT_DIR, target=state) rng, input_rng = jax.random.split(rng) - test_loss, test_accuracy, test_top5accuracy, test_time, _, _, _, _ = eval_model( - best_state, best_state.params, config.test_ds, config.batch_size, - config.num_classes, input_rng, - spectral_diagnostics=(config.spectral_diagnostics == "full")) + if config.learning_algorithm == "ep": + test_loss, test_accuracy, test_top5accuracy, test_time, _, _, _, _ = eval_ep_model( + best_state, best_state.params, config.test_ds, config.batch_size, + config.num_classes, config.ep_free_steps, config.ep_dt) + else: + test_loss, test_accuracy, test_top5accuracy, test_time, _, _, _, _ = eval_model( + best_state, best_state.params, config.test_ds, config.batch_size, + config.num_classes, input_rng, + spectral_diagnostics=(config.spectral_diagnostics == "full")) hist['test_loss'], hist['test_accuracy'], hist['test_top5accuracy'], hist['test_time'] = test_loss, test_accuracy, test_top5accuracy, test_time loginfo_and_print('test: \tloss: %.4f, \taccuracy: %.4f, \ttop5_accuracy: %.4f, \truntime: %.4f' % (test_loss, test_accuracy, test_top5accuracy, test_time)) -- 2.54.0