summaryrefslogtreecommitdiff
path: root/external/dualprop_patches/0003-fix-norm-preserve-direct-feedback-maps.patch
diff options
context:
space:
mode:
Diffstat (limited to 'external/dualprop_patches/0003-fix-norm-preserve-direct-feedback-maps.patch')
-rw-r--r--external/dualprop_patches/0003-fix-norm-preserve-direct-feedback-maps.patch43
1 files changed, 43 insertions, 0 deletions
diff --git a/external/dualprop_patches/0003-fix-norm-preserve-direct-feedback-maps.patch b/external/dualprop_patches/0003-fix-norm-preserve-direct-feedback-maps.patch
new file mode 100644
index 0000000..d2af3e1
--- /dev/null
+++ b/external/dualprop_patches/0003-fix-norm-preserve-direct-feedback-maps.patch
@@ -0,0 +1,43 @@
+From ce106aa21a77ee92cc98ea543ea0fc4fe211078a Mon Sep 17 00:00:00 2001
+From: YurenHao0426 <Blackhao0426@gmail.com>
+Date: Mon, 27 Jul 2026 13:05:38 -0500
+Subject: [PATCH 03/19] fix: norm-preserve direct feedback maps
+
+---
+ src/training_utils.py | 20 +++++++++++++-------
+ 1 file changed, 13 insertions(+), 7 deletions(-)
+
+diff --git a/src/training_utils.py b/src/training_utils.py
+index b6e1e3c..2b6b18a 100644
+--- a/src/training_utils.py
++++ b/src/training_utils.py
+@@ -236,13 +236,19 @@ def create_local_feedback(rng, model, params, image_dims, learning_algorithm,
+ states, _ = model.apply(
+ {"params": params}, dummy, method="ff_with_local_cache")
+ keys = jax.random.split(rng, len(states) - 2)
+- scale = jnp.asarray(num_classes ** -0.5, dtype=dummy.dtype)
+- return tuple(
+- scale * jax.random.normal(
+- key, (num_classes,) + tuple(state.shape[1:]),
+- dtype=state.dtype)
+- for key, state in zip(keys, states[1:-1])
+- )
++ feedback = []
++ for key, state in zip(keys, states[1:-1]):
++ hidden_units = int(np.prod(state.shape[1:]))
++ # Preserve the norm of an output teaching vector in expectation. A
++ # 1/sqrt(num_classes) per-coordinate scale would make the *whole*
++ # hidden teaching field grow as sqrt(num_hidden_units), which is
++ # catastrophic for early convolutional feature maps.
++ scale = jnp.asarray(hidden_units ** -0.5, dtype=state.dtype)
++ feedback.append(
++ scale * jax.random.normal(
++ key, (num_classes,) + tuple(state.shape[1:]),
++ dtype=state.dtype))
++ return tuple(feedback)
+
+ def augment_train(image, batch_rng):
+ w, h, c = image.shape
+--
+2.54.0
+