summaryrefslogtreecommitdiff
path: root/external/dualprop_patches/0008-crossover-add-dynamic-innovation-on-reciprocal-KP.patch
diff options
context:
space:
mode:
Diffstat (limited to 'external/dualprop_patches/0008-crossover-add-dynamic-innovation-on-reciprocal-KP.patch')
-rw-r--r--external/dualprop_patches/0008-crossover-add-dynamic-innovation-on-reciprocal-KP.patch430
1 files changed, 430 insertions, 0 deletions
diff --git a/external/dualprop_patches/0008-crossover-add-dynamic-innovation-on-reciprocal-KP.patch b/external/dualprop_patches/0008-crossover-add-dynamic-innovation-on-reciprocal-KP.patch
new file mode 100644
index 0000000..d1da5b2
--- /dev/null
+++ b/external/dualprop_patches/0008-crossover-add-dynamic-innovation-on-reciprocal-KP.patch
@@ -0,0 +1,430 @@
+From d4b231c502fe158752e9294ece450554f42f7cae Mon Sep 17 00:00:00 2001
+From: YurenHao0426 <Blackhao0426@gmail.com>
+Date: Mon, 27 Jul 2026 13:22:51 -0500
+Subject: [PATCH 08/19] crossover: add dynamic innovation on reciprocal KP
+
+---
+ config/cli_config.py | 16 ++-
+ src/__init__.py | 2 +-
+ src/training_utils.py | 205 +++++++++++++++++++++++++++++++++++++
+ tests/local_rules_smoke.py | 25 +++++
+ train.py | 40 +++++++-
+ 5 files changed, 281 insertions(+), 7 deletions(-)
+
+diff --git a/config/cli_config.py b/config/cli_config.py
+index 56fdb9f..4979b37 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', 'clean-kp', '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', 'sdil', 'dualprop-lagr-ff', 'dualprop-raovr-ff', 'dualprop-raovr-dampened-ff'])
+
+ parser.add_argument(
+ '--feedback-seed', default=1729, type=int,
+@@ -76,6 +76,18 @@ parser.add_argument(
+ '--ep-nudge-steps', default=4, type=int,
+ help='Number of nudged-phase EP relaxation steps.')
+
++parser.add_argument(
++ '--sdil-traffic-ratio', default=4.0, type=float,
++ help='Initialization-calibrated traffic/instruction RMS ratio.')
++
++parser.add_argument(
++ '--sdil-traffic-seed', default=4000, type=int,
++ help='Fixed per-cell soma-predictable traffic seed.')
++
++parser.add_argument(
++ '--sdil-calibration-examples', default=64, type=int,
++ help='Neutral examples for the frozen slow affine predictor fit.')
++
+ parser.add_argument(
+ '--gradient-diagnostics', default='full', choices=['none', 'full'],
+ help=('Compute the exact BP reference gradient and layerwise cosine on '
+@@ -156,7 +168,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, "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}
++modeltype = {"backprop":cnn_abstract, "fa": cnn_abstract, "dfa": cnn_abstract, "pepita": cnn_abstract, "ff": cnn_abstract, "ep": cnn_abstract, "clean-kp": cnn_abstract, "sdil": 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 7ee7740..5acc8cd 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, 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
++from .training_utils import create_train_state, create_ff_train_state, create_local_feedback, create_sdil_auxiliary, train_epoch, train_ff_epoch, train_kp_epoch, train_sdil_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/training_utils.py b/src/training_utils.py
+index e573bb5..ac05857 100644
+--- a/src/training_utils.py
++++ b/src/training_utils.py
+@@ -376,6 +376,159 @@ def train_kp_epoch(state, feedback_state, train_ds, batch_size, rng,
+ return state, feedback_state, summary, time.time() - t0
+
+
++def create_sdil_auxiliary(rng, state, feedback_state, image, labels,
++ num_classes, traffic_ratio):
++ """Initialization-only traffic calibration and neutral predictor fit."""
++ one_hot = jax.nn.one_hot(labels, num_classes=num_classes)
++ states, linear = state.apply_fn(
++ {"params": state.params}, image, method="ff_with_local_cache")
++
++ def output_loss(logits):
++ return state.apply_fn(
++ {"params": state.params}, logits, one_hot,
++ method="output_loss")
++
++ output_field = jax.grad(output_loss)(states[-1])
++ instruction = feedback_state.apply_fn(
++ {"params": feedback_state.params}, states, linear, output_field,
++ method="fa_teaching_fields")
++ keys = jax.random.split(rng, len(states) - 2)
++ coefficients = []
++ gains = []
++ slopes = []
++ biases = []
++ realized = []
++ predictor_residual_ratios = []
++ for key, hidden, signal in zip(keys, states[1:-1], instruction[1:-1]):
++ coefficient = jnp.exp(
++ 0.25 * jax.random.normal(
++ key, hidden.shape[1:], dtype=hidden.dtype))
++ signal_rms = jnp.sqrt(jnp.mean(jnp.square(signal)))
++ unscaled_rms = jnp.sqrt(jnp.mean(jnp.square(coefficient * hidden)))
++ gain = traffic_ratio * signal_rms / jnp.maximum(unscaled_rms, 1e-30)
++ traffic = gain * coefficient * hidden
++ hidden_mean = jnp.mean(hidden, axis=0)
++ traffic_mean = jnp.mean(traffic, axis=0)
++ centered_hidden = hidden - hidden_mean
++ centered_traffic = traffic - traffic_mean
++ variance = jnp.mean(jnp.square(centered_hidden), axis=0)
++ covariance = jnp.mean(
++ centered_hidden * centered_traffic, axis=0)
++ slope = jnp.where(
++ variance > 1e-12,
++ covariance / jnp.maximum(variance, 1e-12),
++ jnp.zeros_like(variance))
++ bias = traffic_mean - slope * hidden_mean
++ residual = traffic - (slope * hidden + bias)
++ coefficients.append(jax.lax.stop_gradient(coefficient))
++ gains.append(jax.lax.stop_gradient(gain))
++ slopes.append(jax.lax.stop_gradient(slope))
++ biases.append(jax.lax.stop_gradient(bias))
++ realized.append(jnp.sqrt(jnp.mean(jnp.square(traffic)))
++ / jnp.maximum(signal_rms, 1e-30))
++ predictor_residual_ratios.append(
++ jnp.sqrt(jnp.mean(jnp.square(residual)))
++ / jnp.maximum(
++ jnp.sqrt(jnp.mean(jnp.square(traffic))), 1e-30))
++ auxiliary = {
++ "coefficients": tuple(coefficients),
++ "gains": tuple(gains),
++ "slopes": tuple(slopes),
++ "biases": tuple(biases),
++ }
++ report = {
++ "traffic_ratio_target": float(traffic_ratio),
++ "realized_traffic_instruction_rms_ratio": [
++ float(value) for value in jax.device_get(realized)],
++ "predictor_residual_traffic_rms_ratio": [
++ float(value)
++ for value in jax.device_get(predictor_residual_ratios)],
++ "observations": int(image.shape[0]),
++ "instruction_observations_for_predictor": 0,
++ }
++ return auxiliary, report
++
++
++def sdil_innovation_fields(states, instruction, auxiliary):
++ """Paired-neutral affine projection for every hidden population."""
++ fields = [instruction[0]]
++ pre_power = jnp.asarray(0.0, dtype=states[0].dtype)
++ post_power = jnp.asarray(0.0, dtype=states[0].dtype)
++ traffic_power = jnp.asarray(0.0, dtype=states[0].dtype)
++ maximum_post_slope = jnp.asarray(0.0, dtype=states[0].dtype)
++ for hidden, signal, coefficient, gain, slope, bias in zip(
++ states[1:-1], instruction[1:-1],
++ auxiliary["coefficients"], auxiliary["gains"],
++ auxiliary["slopes"], auxiliary["biases"]):
++ traffic = gain * coefficient * hidden
++ neutral = traffic - (slope * hidden + bias)
++ centered_hidden = hidden - jnp.mean(hidden, axis=0)
++ centered_neutral = neutral - jnp.mean(neutral, axis=0)
++ variance = jnp.mean(jnp.square(centered_hidden), axis=0)
++ covariance = jnp.mean(
++ centered_hidden * centered_neutral, axis=0)
++ correction = jnp.where(
++ variance > 1e-12,
++ covariance / jnp.maximum(variance, 1e-12),
++ jnp.zeros_like(variance))
++ remainder = centered_neutral - correction * centered_hidden
++ centered_remainder = remainder - jnp.mean(remainder, axis=0)
++ post_covariance = jnp.mean(
++ centered_hidden * centered_remainder, axis=0)
++ post_slope = jnp.where(
++ variance > 1e-12,
++ post_covariance / jnp.maximum(variance, 1e-12),
++ jnp.zeros_like(variance))
++ maximum_post_slope = jnp.maximum(
++ maximum_post_slope, jnp.max(jnp.abs(post_slope)))
++ pre_power += jnp.sum(jnp.square(neutral))
++ post_power += jnp.sum(jnp.square(remainder))
++ traffic_power += jnp.sum(jnp.square(traffic))
++ fields.append(signal + remainder)
++ fields.append(instruction[-1])
++ report = {
++ "pre_projection_traffic_rms_ratio": jnp.sqrt(
++ pre_power / jnp.maximum(traffic_power, 1e-30)),
++ "post_projection_traffic_rms_ratio": jnp.sqrt(
++ post_power / jnp.maximum(traffic_power, 1e-30)),
++ "max_absolute_post_projection_soma_slope": maximum_post_slope,
++ "instruction_observations": jnp.asarray(0, dtype=jnp.int32),
++ "observations": jnp.asarray(states[0].shape[0], dtype=jnp.int32),
++ }
++ return fields, report
++
++
++def train_sdil_epoch(state, feedback_state, auxiliary, train_ds, batch_size,
++ rng, augmentation_on, num_classes,
++ gradient_diagnostics=True):
++ """Train dynamic innovation W/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_sdil(
++ state, feedback_state, auxiliary, 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])]
+@@ -436,6 +589,58 @@ def train_step_kp(state, feedback_state, image, labels_onehot, labels,
+ return state, feedback_state, metrics
+
+
++@jax.jit
++def train_step_sdil(state, feedback_state, auxiliary, image, labels_onehot,
++ labels, batch_rng, inf_rng, augmentation_on,
++ gradient_diagnostics):
++ """One simultaneous dynamic-innovation 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]))
++ instruction = feedback_state.apply_fn(
++ {"params": feedback_state.params}, states, linear, output_field,
++ method="fa_teaching_fields")
++ instruction = tree_map(jax.lax.stop_gradient, instruction)
++ teaching_fields, projection = sdil_innovation_fields(
++ states, instruction, auxiliary)
++ 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.update(projection)
++ 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)
++ 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 f40f0ab..3763ee8 100644
+--- a/tests/local_rules_smoke.py
++++ b/tests/local_rules_smoke.py
+@@ -19,6 +19,7 @@ from src.training_utils import (
+ create_local_feedback,
+ direct_feedback_fields,
+ ff_overlay,
++ sdil_innovation_fields,
+ )
+
+
+@@ -116,6 +117,28 @@ def main():
+ kp_symmetric_tracking_error = relative_error(kp_q, kp_w)
+ assert kp_symmetric_tracking_error == 0.0
+
++ hidden_states = states[1:-1]
++ sdil_auxiliary = {
++ "coefficients": tuple(jnp.ones(hidden.shape[1:], hidden.dtype)
++ for hidden in hidden_states),
++ "gains": tuple(jnp.asarray(4.0, hidden.dtype)
++ for hidden in hidden_states),
++ # Leave a one-times-soma neutral residual for the fast projection.
++ "slopes": tuple(jnp.full(hidden.shape[1:], 3.0, hidden.dtype)
++ for hidden in hidden_states),
++ "biases": tuple(jnp.full(hidden.shape[1:], 0.2, hidden.dtype)
++ for hidden in hidden_states),
++ }
++ sdil_fields, sdil_projection = sdil_innovation_fields(
++ states, symmetric_fields, sdil_auxiliary)
++ sdil_projection_identity_error = relative_error(
++ sdil_fields[1:-1], symmetric_fields[1:-1])
++ assert sdil_projection_identity_error < 2e-12
++ assert float(
++ sdil_projection["max_absolute_post_projection_soma_slope"]
++ ) < 2e-12
++ assert int(sdil_projection["instruction_observations"]) == 0
++
+ feedback = create_local_feedback(
+ jax.random.PRNGKey(1729), model, params, (8, 8, 2), "fa", 3)
+ feedback_cosine = float(
+@@ -261,6 +284,8 @@ def main():
+ "kp_reciprocal_gradient_relative_error": reciprocal_error,
+ "kp_reciprocal_independence_error": reciprocal_independence_error,
+ "kp_symmetric_two_step_tracking_error": kp_symmetric_tracking_error,
++ "sdil_projection_instruction_identity_error": (
++ sdil_projection_identity_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 6b2d8d8..b1df1b0 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, train_kp_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, create_sdil_auxiliary, train_epoch, train_kp_epoch, train_sdil_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
+@@ -55,13 +55,26 @@ for experiment_index, seed in enumerate(config.seeds):
+ 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":
++ if config.learning_algorithm in ("clean-kp", "sdil"):
+ 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)
++ sdil_auxiliary = None
++ sdil_initialization = None
++ if config.learning_algorithm == "sdil":
++ count = config.sdil_calibration_examples
++ if count < 2 or count > len(config.train_ds["image"]):
++ raise ValueError("invalid --sdil-calibration-examples")
++ if config.sdil_traffic_ratio <= 0:
++ raise ValueError("--sdil-traffic-ratio must be positive")
++ sdil_auxiliary, sdil_initialization = create_sdil_auxiliary(
++ jax.random.PRNGKey(config.sdil_traffic_seed), state,
++ feedback_state, config.train_ds["image"][:count],
++ config.train_ds["label"][:count], config.num_classes,
++ config.sdil_traffic_ratio)
+
+
+ del init_rng # Must not be used anymore.
+@@ -74,10 +87,14 @@ for experiment_index, seed in enumerate(config.seeds):
+ '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),
++ 'pre_projection_traffic_rms_ratio': np.zeros(config.num_epochs),
++ 'post_projection_traffic_rms_ratio': np.zeros(config.num_epochs),
++ 'max_absolute_post_projection_soma_slope': 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)),
+- 'gamma20': np.zeros((len(state.params), config.num_epochs))}
++ 'gamma20': np.zeros((len(state.params), config.num_epochs)),
++ 'sdil_initialization': sdil_initialization}
+
+ best_accuracy, best_epoch = 0, 0
+ epoch = 0
+@@ -95,6 +112,13 @@ for experiment_index, seed in enumerate(config.seeds):
+ input_rng, augmentation_on, config.num_classes,
+ gradient_diagnostics=(
+ config.gradient_diagnostics == "full"))
++ elif config.learning_algorithm == "sdil":
++ state, feedback_state, epoch_metrics, train_time = train_sdil_epoch(
++ state, feedback_state, sdil_auxiliary, 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,
+@@ -111,11 +135,19 @@ 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":
++ if config.learning_algorithm in ("clean-kp", "sdil"):
+ hist["feedback_forward_cosine"][:, epoch - 1] = (
+ epoch_metrics["feedback_forward_cosine"])
+ hist["reciprocal_gradient_error"][epoch - 1] = (
+ epoch_metrics["reciprocal_gradient_error"])
++ if config.learning_algorithm == "sdil":
++ hist["pre_projection_traffic_rms_ratio"][epoch - 1] = (
++ epoch_metrics["pre_projection_traffic_rms_ratio"])
++ hist["post_projection_traffic_rms_ratio"][epoch - 1] = (
++ epoch_metrics["post_projection_traffic_rms_ratio"])
++ hist["max_absolute_post_projection_soma_slope"][epoch - 1] = (
++ epoch_metrics[
++ "max_absolute_post_projection_soma_slope"])
+
+
+ # Evaluate on the validation set after each training epoch
+--
+2.54.0
+