diff options
Diffstat (limited to 'external/dualprop_patches/0006-crossover-add-two-phase-equilibrium-propagation.patch')
| -rw-r--r-- | external/dualprop_patches/0006-crossover-add-two-phase-equilibrium-propagation.patch | 377 |
1 files changed, 377 insertions, 0 deletions
diff --git a/external/dualprop_patches/0006-crossover-add-two-phase-equilibrium-propagation.patch b/external/dualprop_patches/0006-crossover-add-two-phase-equilibrium-propagation.patch new file mode 100644 index 0000000..1d30ef2 --- /dev/null +++ b/external/dualprop_patches/0006-crossover-add-two-phase-equilibrium-propagation.patch @@ -0,0 +1,377 @@ +From b925ba7f07389054cba7b90d2b0db3d7e1699b94 Mon Sep 17 00:00:00 2001 +From: YurenHao0426 <Blackhao0426@gmail.com> +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 + |
