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