summaryrefslogtreecommitdiff
path: root/external/dualprop_patches
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-06 12:12:41 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-06 12:12:41 -0500
commitd91cfe4d806f4c1e09c6cb75829a8625ff6506ec (patch)
treea5cf10dddfbdd904877872e38c03b8e81dff0107 /external/dualprop_patches
parent051414af6f3b7016ce8ee4125a41dfacf0a01a3e (diff)
experiment: add contrastive state-bias screen
Diffstat (limited to 'external/dualprop_patches')
-rw-r--r--external/dualprop_patches/0020-experiment-add-neuron-specific-bias-to-Dual-Prop.patch518
1 files changed, 518 insertions, 0 deletions
diff --git a/external/dualprop_patches/0020-experiment-add-neuron-specific-bias-to-Dual-Prop.patch b/external/dualprop_patches/0020-experiment-add-neuron-specific-bias-to-Dual-Prop.patch
new file mode 100644
index 0000000..29a5be8
--- /dev/null
+++ b/external/dualprop_patches/0020-experiment-add-neuron-specific-bias-to-Dual-Prop.patch
@@ -0,0 +1,518 @@
+From 79169abac635715e2167f1bb8a089f239d06c431 Mon Sep 17 00:00:00 2001
+From: SDIL replication runner <sdil-replication@invalid.example>
+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
+