diff options
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.patch | 430 |
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 + |
