From 79169abac635715e2167f1bb8a089f239d06c431 Mon Sep 17 00:00:00 2001 From: SDIL replication runner Date: Thu, 6 Aug 2026 12:09:10 -0500 Subject: [PATCH] experiment: add neuron-specific bias to Dual Prop --- config/cli_config.py | 22 +++++ src/__init__.py | 2 +- src/models.py | 13 ++- src/training_utils.py | 194 +++++++++++++++++++++++++++++++++++++ tests/local_rules_smoke.py | 94 +++++++++++++++++- train.py | 49 +++++++++- 6 files changed, 368 insertions(+), 6 deletions(-) diff --git a/config/cli_config.py b/config/cli_config.py index b38ff7b..e3d44bd 100644 --- a/config/cli_config.py +++ b/config/cli_config.py @@ -93,6 +93,28 @@ parser.add_argument( '--sdil-calibration-examples', default=64, type=int, help='Neutral examples for the frozen slow affine predictor fit.') +parser.add_argument( + '--dp-bias-rule', default='none', + choices=['none', 'raw', 'innovation', 'oracle'], + help='Teaching-difference rule for the contrastive state-bias screen.') + +parser.add_argument( + '--dp-bias-kind', default='none', + choices=['none', 'common', 'fixed', 'activity'], + help='Neuron-specific bias injected into the DP compartment contrast.') + +parser.add_argument( + '--dp-bias-ratio', default=0.0, type=float, + help='Initialization-calibrated bias/clean-difference RMS ratio.') + +parser.add_argument( + '--dp-bias-seed', default=6100, type=int, + help='Fixed per-cell contrastive-bias pattern seed.') + +parser.add_argument( + '--dp-bias-calibration-examples', default=64, type=int, + help='Instruction-free observations used to calibrate the bias pattern.') + parser.add_argument( '--gradient-diagnostics', default='full', choices=['none', 'full'], help=('Compute the exact BP reference gradient and layerwise cosine on ' diff --git a/src/__init__.py b/src/__init__.py index 5acc8cd..3025ea8 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, 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 +from .training_utils import create_train_state, create_ff_train_state, create_local_feedback, create_sdil_auxiliary, create_dp_bias_auxiliary, dp_bias_differences, train_epoch, train_ff_epoch, train_kp_epoch, train_sdil_epoch, train_dp_bias_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 36b2643..76a4f55 100644 --- a/src/models.py +++ b/src/models.py @@ -445,11 +445,22 @@ class cnn_dualprop_abstract(cnn_abstract): return splus, sminus def get_J(self, splus, sminus): + deltas = [positive - negative + for positive, negative in zip(splus, sminus)] + return self.get_J_from_deltas(splus, sminus, deltas) + + def get_J_from_deltas(self, splus, sminus, deltas): + """Evaluate the DP local objective with explicit teaching differences. + + The ordinary DP rule supplies ``splus-sminus``. Keeping the inferred + activity states fixed while accepting an explicit difference lets the + bias experiment change only the neuron-local teaching variable. + """ J = 0.0 batchsize = splus[-1].shape[0] for i in range(1,len(splus)): sbar_previous = self.alpha*splus[i-1] + (1-self.alpha)*sminus[i-1] - delta = splus[i] - sminus[i] + delta = deltas[i] J += self.get_phi(-delta, sbar_previous, self.layers[i-1])/self.beta return J/batchsize diff --git a/src/training_utils.py b/src/training_utils.py index 271a81e..8a8ad64 100644 --- a/src/training_utils.py +++ b/src/training_utils.py @@ -294,6 +294,136 @@ def no_aug(image, batch_rng): return image +def create_dp_bias_auxiliary(rng, state, image, labels, num_classes, + alpha, bias_kind, bias_ratio): + """Freeze per-cell contrastive-bias coefficients and initial gains.""" + if bias_kind not in ("common", "fixed", "activity"): + raise ValueError("invalid DP bias kind") + if bias_ratio <= 0: + raise ValueError("DP bias ratio must be positive") + one_hot = jax.nn.one_hot(labels, num_classes=num_classes) + _, inference_rng = jax.random.split(rng) + plus, minus = state.apply_fn( + {"params": state.params}, image, one_hot, inference_rng, + method="infer_states_train") + keys = jax.random.split(rng, len(plus) - 1) + coefficients = [] + gains = [] + realized = [] + for key, positive, negative in zip(keys, plus[1:], minus[1:]): + activity = _dp_activity(alpha, positive, negative) + if bias_kind in ("activity", "common"): + coefficient = jnp.exp( + 0.25 * jax.random.normal( + key, activity.shape[1:], dtype=activity.dtype)) + source = coefficient * activity + else: + coefficient = jax.random.normal( + key, activity.shape[1:], dtype=activity.dtype) + source = jnp.broadcast_to(coefficient, activity.shape) + difference = positive - negative + difference_rms = jnp.sqrt(jnp.mean(jnp.square(difference))) + source_rms = jnp.sqrt(jnp.mean(jnp.square(source))) + gain = bias_ratio * difference_rms / jnp.maximum(source_rms, 1e-30) + coefficients.append(jax.lax.stop_gradient(coefficient)) + gains.append(jax.lax.stop_gradient(gain)) + realized.append( + jnp.sqrt(jnp.mean(jnp.square(gain * source))) + / jnp.maximum(difference_rms, 1e-30)) + auxiliary = { + "coefficients": tuple(coefficients), + "gains": tuple(gains), + "alpha": jnp.asarray(alpha, dtype=plus[0].dtype), + } + report = { + "bias_kind": bias_kind, + "bias_ratio_target": float(bias_ratio), + "realized_bias_difference_rms_ratio": [ + float(value) for value in jax.device_get(realized)], + "calibration_examples": int(image.shape[0]), + "instruction_observations_for_predictor": 0, + } + return auxiliary, report + + +def _dp_activity(alpha, positive, negative): + return alpha * positive + (1.0 - alpha) * negative + + +def _neutral_affine_prediction(activity, neutral): + """Per-cell affine fit using only the minibatch observation axis.""" + activity_mean = jnp.mean(activity, axis=0) + neutral_mean = jnp.mean(neutral, axis=0) + centered_activity = activity - activity_mean + centered_neutral = neutral - neutral_mean + variance = jnp.mean(jnp.square(centered_activity), axis=0) + covariance = jnp.mean(centered_activity * centered_neutral, axis=0) + slope = jnp.where( + variance > 1e-12, + covariance / jnp.maximum(variance, 1e-12), + jnp.zeros_like(variance)) + intercept = neutral_mean - slope * activity_mean + return slope * activity + intercept + + +@partial(jax.jit, static_argnames=("bias_kind", "bias_rule")) +def dp_bias_differences(plus, minus, auxiliary, bias_kind, bias_rule): + """Return clean or bias-corrected DP state differences and diagnostics.""" + clean = [positive - negative + for positive, negative in zip(plus, minus)] + used = [clean[0]] + raw_bias_power = jnp.asarray(0.0, dtype=plus[0].dtype) + post_bias_power = jnp.asarray(0.0, dtype=plus[0].dtype) + clean_power = jnp.asarray(0.0, dtype=plus[0].dtype) + maximum_relative_error = jnp.asarray(0.0, dtype=plus[0].dtype) + for positive, negative, difference, coefficient, gain in zip( + plus[1:], minus[1:], clean[1:], auxiliary["coefficients"], + auxiliary["gains"]): + activity = _dp_activity( + auxiliary["alpha"], positive, negative) + if bias_kind in ("activity", "common"): + generated = gain * coefficient * activity + else: + generated = gain * jnp.broadcast_to(coefficient, activity.shape) + # Identical compartment bias cancels before a contrast is formed. + differential = ( + jnp.zeros_like(generated) if bias_kind == "common" + else generated) + observed = difference + differential + if bias_rule == "raw": + corrected = observed + elif bias_rule == "oracle": + corrected = observed - differential + elif bias_rule == "innovation": + prediction = _neutral_affine_prediction(activity, differential) + corrected = observed - prediction + else: + raise ValueError("invalid DP bias rule") + residual = corrected - difference + used.append(corrected) + raw_bias_power += jnp.sum(jnp.square(differential)) + post_bias_power += jnp.sum(jnp.square(residual)) + clean_power += jnp.sum(jnp.square(difference)) + relative_error = ( + jnp.linalg.norm(residual.reshape(-1)) + / jnp.maximum(jnp.linalg.norm(difference.reshape(-1)), 1e-30)) + maximum_relative_error = jnp.maximum( + maximum_relative_error, relative_error) + report = { + "raw_bias_clean_difference_rms_ratio": jnp.sqrt( + raw_bias_power / jnp.maximum(clean_power, 1e-30)), + "post_bias_raw_bias_rms_ratio": jnp.sqrt( + post_bias_power / jnp.maximum(raw_bias_power, 1e-30)), + "used_clean_difference_rms_ratio": jnp.sqrt( + post_bias_power / jnp.maximum(clean_power, 1e-30)), + "maximum_used_clean_difference_relative_error": + maximum_relative_error, + "neutral_observations": jnp.asarray(plus[0].shape[0], jnp.int32), + "instruction_observations_for_predictor": jnp.asarray(0, jnp.int32), + } + return used, report + + def to_float16(ptree): return tree_map(lambda x: x.astype(jnp.float16), ptree) @@ -870,6 +1000,70 @@ def train_step_local(state, local_feedback, learning_algorithm, image, state = state.apply_gradients(grads=grads) return state, metrics +@partial(jax.jit, static_argnames=("bias_kind", "bias_rule")) +def train_step_dp_bias(state, auxiliary, image, labels_onehot, labels, + batch_rng, inf_rng, augmentation_on, bias_kind, + bias_rule, gradient_diagnostics): + """One DP update with an explicit biased or corrected teaching contrast.""" + image = jax.lax.cond( + augmentation_on, vmap_augment_train, no_aug, image, batch_rng) + plus, minus = state.apply_fn( + {"params": state.params}, image, labels_onehot, inf_rng, + method="infer_states_train") + plus = tree_map(jax.lax.stop_gradient, plus) + minus = tree_map(jax.lax.stop_gradient, minus) + differences, bias_metrics = dp_bias_differences( + plus, minus, auxiliary, bias_kind, bias_rule) + differences = tree_map(jax.lax.stop_gradient, differences) + + def loss_fn(params): + return state.apply_fn( + {"params": params}, plus, minus, differences, + method="get_J_from_deltas") + + loss, grads = jax.value_and_grad(loss_fn)(state.params) + metrics = compute_metrics( + image=image, labels_onehot=labels_onehot, labels=labels, state=state) + metrics.update(bias_metrics) + metrics["contrastive_objective"] = loss + 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 train_dp_bias_epoch(state, auxiliary, train_ds, batch_size, rng, + augmentation_on, num_classes, bias_kind, bias_rule, + gradient_diagnostics=True): + """Train one complete author-order epoch with a frozen bias condition.""" + t0 = time.time() + size = len(train_ds["image"]) + steps = size // batch_size + permutations = jax.random.permutation(rng, size)[:steps * batch_size] + permutations = permutations.reshape((steps, batch_size)) + metrics = [] + for permutation in permutations: + image = train_ds["image"][permutation] + labels = train_ds["label"][permutation] + one_hot = jax.nn.one_hot(labels, num_classes=num_classes) + rng, inference_rng, batch_rng = jax.random.split(rng, 3) + per_example_rng = jax.random.split(batch_rng, image.shape[0]) + state, batch_metrics = train_step_dp_bias( + state, auxiliary, image, one_hot, labels, per_example_rng, + inference_rng, augmentation_on, bias_kind, bias_rule, + 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, summary, time.time() - t0 + + @jax.jit def train_step(state, image, labels_onehot, labels, batch_rng, inf_rng, augmentation_on, gradient_diagnostics): diff --git a/tests/local_rules_smoke.py b/tests/local_rules_smoke.py index 3763ee8..3ff1bec 100644 --- a/tests/local_rules_smoke.py +++ b/tests/local_rules_smoke.py @@ -14,9 +14,10 @@ from flax.core.frozen_dict import unfreeze ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) sys.path.insert(0, ROOT) -from src.models import cnn_abstract +from src.models import cnn_abstract, cnn_dualprop_Lagr_ff from src.training_utils import ( create_local_feedback, + dp_bias_differences, direct_feedback_fields, ff_overlay, sdil_innovation_fields, @@ -139,6 +140,88 @@ def main(): ) < 2e-12 assert int(sdil_projection["instruction_observations"]) == 0 + dp_model = cnn_dualprop_Lagr_ff( + loss_func, nn.Conv, nn.Dense, nn.relu, 3, 0.1, 0.0, + jnp.float64, jnp.float64, + kernels=[(3, 3), (3, 3)], strides=[(1, 1), (1, 1)], + features=[4, 5], mp=[True, True], dense_features=[3], + inference_sequence="fwK", inference_passes_nudged=1) + dp_params = unfreeze(dp_model.init(jax.random.PRNGKey(12), x)["params"]) + dp_plus, dp_minus = dp_model.apply( + {"params": dp_params}, x, one_hot, jax.random.PRNGKey(13), + method="infer_states_train") + dp_clean_differences = [ + positive - negative + for positive, negative in zip(dp_plus, dp_minus)] + dp_hidden = dp_plus[1:] + dp_auxiliary = { + "coefficients": tuple( + jnp.ones(value.shape[1:], dtype=value.dtype) + for value in dp_hidden), + "gains": tuple(jnp.asarray(0.5, value.dtype) for value in dp_hidden), + "alpha": jnp.asarray(0.0, x.dtype), + } + zero_auxiliary = { + **dp_auxiliary, + "gains": tuple(jnp.asarray(0.0, value.dtype) for value in dp_hidden), + } + zero_raw, _ = dp_bias_differences( + dp_plus, dp_minus, zero_auxiliary, "activity", "raw") + zero_innovation, _ = dp_bias_differences( + dp_plus, dp_minus, zero_auxiliary, "activity", "innovation") + zero_oracle, _ = dp_bias_differences( + dp_plus, dp_minus, zero_auxiliary, "activity", "oracle") + dp_zero_bias_error = max( + relative_error(value[1:], dp_clean_differences[1:]) + for value in (zero_raw, zero_innovation, zero_oracle)) + assert dp_zero_bias_error < 2e-12, dp_zero_bias_error + + common_differences, common_report = dp_bias_differences( + dp_plus, dp_minus, dp_auxiliary, "common", "raw") + dp_common_bias_error = relative_error( + common_differences[1:], dp_clean_differences[1:]) + assert dp_common_bias_error < 2e-12, dp_common_bias_error + + raw_differences, raw_report = dp_bias_differences( + dp_plus, dp_minus, dp_auxiliary, "activity", "raw") + innovation_differences, innovation_report = dp_bias_differences( + dp_plus, dp_minus, dp_auxiliary, "activity", "innovation") + oracle_differences, oracle_report = dp_bias_differences( + dp_plus, dp_minus, dp_auxiliary, "activity", "oracle") + fixed_differences, fixed_report = dp_bias_differences( + dp_plus, dp_minus, dp_auxiliary, "fixed", "innovation") + dp_innovation_difference_error = relative_error( + innovation_differences[1:], dp_clean_differences[1:]) + dp_oracle_difference_error = relative_error( + oracle_differences[1:], dp_clean_differences[1:]) + dp_fixed_difference_error = relative_error( + fixed_differences[1:], dp_clean_differences[1:]) + assert dp_innovation_difference_error < 2e-12 + assert dp_oracle_difference_error < 2e-12 + assert dp_fixed_difference_error < 2e-12 + assert int( + innovation_report["instruction_observations_for_predictor"]) == 0 + + def dp_objective(candidate, differences): + return dp_model.apply( + {"params": candidate}, dp_plus, dp_minus, differences, + method="get_J_from_deltas") + + dp_clean_grads = jax.grad(dp_objective)( + dp_params, dp_clean_differences) + dp_raw_grads = jax.grad(dp_objective)(dp_params, raw_differences) + dp_innovation_grads = jax.grad(dp_objective)( + dp_params, innovation_differences) + dp_oracle_grads = jax.grad(dp_objective)( + dp_params, oracle_differences) + dp_raw_update_error = relative_error(dp_raw_grads, dp_clean_grads) + dp_innovation_update_error = relative_error( + dp_innovation_grads, dp_clean_grads) + dp_oracle_update_error = relative_error(dp_oracle_grads, dp_clean_grads) + assert dp_raw_update_error > 1e-3, dp_raw_update_error + assert dp_innovation_update_error < 2e-12, dp_innovation_update_error + assert dp_oracle_update_error < 2e-12, dp_oracle_update_error + feedback = create_local_feedback( jax.random.PRNGKey(1729), model, params, (8, 8, 2), "fa", 3) feedback_cosine = float( @@ -286,6 +369,15 @@ def main(): "kp_symmetric_two_step_tracking_error": kp_symmetric_tracking_error, "sdil_projection_instruction_identity_error": ( sdil_projection_identity_error), + "dp_zero_bias_difference_error": dp_zero_bias_error, + "dp_common_bias_difference_error": dp_common_bias_error, + "dp_activity_innovation_difference_error": ( + dp_innovation_difference_error), + "dp_fixed_innovation_difference_error": dp_fixed_difference_error, + "dp_oracle_difference_error": dp_oracle_difference_error, + "dp_raw_update_relative_error": dp_raw_update_error, + "dp_innovation_update_relative_error": dp_innovation_update_error, + "dp_oracle_update_relative_error": dp_oracle_update_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 2a03ada..def4da5 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, 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 +from src import create_train_state, create_local_feedback, create_sdil_auxiliary, create_dp_bias_auxiliary, train_epoch, train_kp_epoch, train_sdil_epoch, train_dp_bias_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 @@ -17,6 +17,12 @@ from config.cli_config import config if config.learning_algorithm == "ff": raise ValueError( "Forward-Forward uses greedy layerwise training; run train_ff.py") +dp_bias_enabled = config.dp_bias_rule != "none" +if dp_bias_enabled and config.learning_algorithm != "dualprop-lagr-ff": + raise ValueError("DP bias rules require --learning-algorithm dualprop-lagr-ff") +if not dp_bias_enabled and ( + config.dp_bias_kind != "none" or config.dp_bias_ratio != 0.0): + raise ValueError("DP bias kind/ratio require a non-none --dp-bias-rule") experiment_dir = "./runs/" + config.experiment_name + "/" if experiment_dir == "./runs/debug-test/" and os.path.isdir(experiment_dir): @@ -81,6 +87,19 @@ for experiment_index, seed in enumerate(config.seeds): feedback_state, config.train_ds["image"][:count], config.train_ds["label"][:count], config.num_classes, config.sdil_traffic_ratio) + dp_bias_auxiliary = None + dp_bias_initialization = None + if dp_bias_enabled: + count = config.dp_bias_calibration_examples + if count < 2 or count > len(config.train_ds["image"]): + raise ValueError("invalid --dp-bias-calibration-examples") + if config.dp_bias_kind == "none" or config.dp_bias_ratio <= 0: + raise ValueError("DP bias screen requires a positive bias condition") + dp_bias_auxiliary, dp_bias_initialization = create_dp_bias_auxiliary( + jax.random.PRNGKey(config.dp_bias_seed), state, + config.train_ds["image"][:count], + config.train_ds["label"][:count], config.num_classes, + config.alpha, config.dp_bias_kind, config.dp_bias_ratio) del init_rng # Must not be used anymore. @@ -96,11 +115,18 @@ for experiment_index, seed in enumerate(config.seeds): '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), + 'raw_bias_clean_difference_rms_ratio': np.zeros(config.num_epochs), + 'post_bias_raw_bias_rms_ratio': np.zeros(config.num_epochs), + 'used_clean_difference_rms_ratio': np.zeros(config.num_epochs), + 'maximum_used_clean_difference_relative_error': np.zeros(config.num_epochs), + 'neutral_observations': np.zeros(config.num_epochs), + 'instruction_observations_for_predictor': 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)), - 'sdil_initialization': sdil_initialization} + 'sdil_initialization': sdil_initialization, + 'dp_bias_initialization': dp_bias_initialization} best_accuracy, best_epoch = 0, 0 epoch = 0 @@ -112,7 +138,15 @@ for experiment_index, seed in enumerate(config.seeds): # Run an optimization step over a training batch # last augument turns off data augmentation for mnist augmentation_on = (config.dataset!="mnist") and (config.dataset!="fashionmnist") - if config.learning_algorithm == "clean-kp": + if dp_bias_enabled: + state, epoch_metrics, train_time = train_dp_bias_epoch( + state, dp_bias_auxiliary, config.train_ds, + config.batch_size, input_rng, augmentation_on, + config.num_classes, config.dp_bias_kind, + config.dp_bias_rule, + gradient_diagnostics=( + config.gradient_diagnostics == "full")) + elif config.learning_algorithm == "clean-kp": state, feedback_state, epoch_metrics, train_time = train_kp_epoch( state, feedback_state, config.train_ds, config.batch_size, input_rng, augmentation_on, config.num_classes, @@ -154,6 +188,15 @@ for experiment_index, seed in enumerate(config.seeds): hist["max_absolute_post_projection_soma_slope"][epoch - 1] = ( epoch_metrics[ "max_absolute_post_projection_soma_slope"]) + if dp_bias_enabled: + for key in ( + "raw_bias_clean_difference_rms_ratio", + "post_bias_raw_bias_rms_ratio", + "used_clean_difference_rms_ratio", + "maximum_used_clean_difference_relative_error", + "neutral_observations", + "instruction_observations_for_predictor"): + hist[key][epoch - 1] = epoch_metrics[key] # Evaluate on the validation set after each training epoch -- 2.54.0