summaryrefslogtreecommitdiff
path: root/external/dualprop_patches/0005-crossover-add-greedy-Forward-Forward-runner.patch
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-30 16:34:30 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-30 16:34:30 -0500
commit716205eba862076c1d0fbbff016601272e4614f2 (patch)
treedc8c31dcfefa49d53c98d305fe35419da76e58c9 /external/dualprop_patches/0005-crossover-add-greedy-Forward-Forward-runner.patch
parent906d17e896329e4d2dd8d7f728e4d5a234afe0e4 (diff)
experiment: package portable A6000 crossover
Diffstat (limited to 'external/dualprop_patches/0005-crossover-add-greedy-Forward-Forward-runner.patch')
-rw-r--r--external/dualprop_patches/0005-crossover-add-greedy-Forward-Forward-runner.patch422
1 files changed, 422 insertions, 0 deletions
diff --git a/external/dualprop_patches/0005-crossover-add-greedy-Forward-Forward-runner.patch b/external/dualprop_patches/0005-crossover-add-greedy-Forward-Forward-runner.patch
new file mode 100644
index 0000000..1b5f46f
--- /dev/null
+++ b/external/dualprop_patches/0005-crossover-add-greedy-Forward-Forward-runner.patch
@@ -0,0 +1,422 @@
+From fe79d6281891ac52977416bb9d65faac3c32e8e4 Mon Sep 17 00:00:00 2001
+From: YurenHao0426 <Blackhao0426@gmail.com>
+Date: Mon, 27 Jul 2026 13:12:55 -0500
+Subject: [PATCH 05/19] crossover: add greedy Forward-Forward runner
+
+---
+ config/cli_config.py | 12 ++++-
+ src/__init__.py | 2 +-
+ src/models.py | 43 +++++++++++++++
+ src/training_utils.py | 108 +++++++++++++++++++++++++++++++++++++
+ tests/local_rules_smoke.py | 40 +++++++++++++-
+ train.py | 4 ++
+ train_ff.py | 92 +++++++++++++++++++++++++++++++
+ 7 files changed, 297 insertions(+), 4 deletions(-)
+ create mode 100644 train_ff.py
+
+diff --git a/config/cli_config.py b/config/cli_config.py
+index 3088436..159e9bc 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', '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', 'dualprop-lagr-ff', 'dualprop-raovr-ff', 'dualprop-raovr-dampened-ff'])
+
+ parser.add_argument(
+ '--feedback-seed', default=1729, type=int,
+@@ -52,6 +52,14 @@ parser.add_argument(
+ '--pepita-projection-scale', default=0.05, type=float,
+ help='Multiplier on PEPITA He-uniform output-error-to-input projection.')
+
++parser.add_argument(
++ '--ff-threshold', default=2.0, type=float,
++ help='Forward-Forward positive/negative goodness threshold.')
++
++parser.add_argument(
++ '--ff-score-from-layer', default=1, type=int,
++ help='First zero-indexed FF layer included in candidate-label goodness.')
++
+ parser.add_argument(
+ '--gradient-diagnostics', default='full', choices=['none', 'full'],
+ help=('Compute the exact BP reference gradient and layerwise cosine on '
+@@ -132,7 +140,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, "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, "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 7a08c87..752dcc2 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_local_feedback, train_epoch, eval_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, train_epoch, train_ff_epoch, eval_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 127268a..0f5af4f 100644
+--- a/src/models.py
++++ b/src/models.py
+@@ -231,6 +231,49 @@ class cnn_abstract(nn.Module, ABC):
+ prediction = layer(x)
+ objective += jnp.sum(prediction * field) / batch_size
+ return objective
++
++ @staticmethod
++ def _ff_normalize(x):
++ axes = tuple(range(1, x.ndim))
++ norm = jnp.sqrt(jnp.sum(jnp.square(x), axis=axes, keepdims=True))
++ return x / (norm + 1e-8)
++
++ def _ff_layer(self, x, layer_index):
++ x = self._ff_normalize(x)
++ return self.act(self.layers[layer_index](x))
++
++ def ff_prefix(self, x, stop_layer):
++ """Detached input producer for greedy Forward-Forward training."""
++ for i in range(stop_layer):
++ x = self._ff_layer(x, i)
++ return x
++
++ def ff_layer_objective(self, positive_input, negative_input, layer_index,
++ threshold):
++ """Reference Forward-Forward local softplus goodness objective."""
++ positive = self._ff_layer(
++ jax.lax.stop_gradient(positive_input), layer_index)
++ negative = self._ff_layer(
++ jax.lax.stop_gradient(negative_input), layer_index)
++ axes = tuple(range(1, positive.ndim))
++ positive_goodness = jnp.mean(jnp.square(positive), axis=axes)
++ negative_goodness = jnp.mean(jnp.square(negative), axis=axes)
++ loss = jnp.mean(
++ jax.nn.softplus(-positive_goodness + threshold)
++ + jax.nn.softplus(negative_goodness - threshold))
++ pair_accuracy = jnp.mean(
++ positive_goodness > negative_goodness, dtype=jnp.float32)
++ return loss, (positive_goodness, negative_goodness, pair_accuracy)
++
++ def ff_goodness(self, x, score_from_layer=1):
++ """Candidate-label goodness used by supervised FF inference."""
++ score = jnp.zeros((x.shape[0],), dtype=x.dtype)
++ for i in range(self.num_layers):
++ x = self._ff_layer(x, i)
++ if i >= score_from_layer:
++ axes = tuple(range(1, x.ndim))
++ score += jnp.mean(jnp.square(x), axis=axes)
++ return score
+
+ def init_states_to_zero(self, x0):
+ s = [x0]
+diff --git a/src/training_utils.py b/src/training_utils.py
+index f5fae62..734f46d 100644
+--- a/src/training_utils.py
++++ b/src/training_utils.py
+@@ -218,6 +218,15 @@ def create_train_state(rng, model, image_dims, lr, wlr, lrf, momentum, weight_de
+ return train_state.TrainState.create(apply_fn=model.apply, params=unfreeze(params), tx=tx)
+
+
++def create_ff_train_state(rng, model, image_dims, learning_rate):
++ """Reference-style Adam state for greedy Forward-Forward layers."""
++ w, h, channels = image_dims
++ dummy = jnp.ones([1, w, h, channels])
++ params = unfreeze(model.init(rng, dummy)["params"])
++ return train_state.TrainState.create(
++ apply_fn=model.apply, params=params, tx=optax.adam(learning_rate))
++
++
+ def create_local_feedback(rng, model, params, image_dims, learning_algorithm,
+ num_classes, pepita_projection_scale=0.05):
+ """Create fixed feedback without reading a forward parameter value.
+@@ -342,6 +351,105 @@ def direct_feedback_fields(states, output_field, direct_feedback):
+ return fields
+
+
++def ff_overlay(image, labels, num_classes):
++ """Overlay a candidate label on the first input coordinates."""
++ flat = image.reshape((image.shape[0], -1))
++ flat = flat.at[:, :num_classes].set(0.0)
++ flat = flat.at[jnp.arange(image.shape[0]), labels].set(jnp.max(image))
++ return flat.reshape(image.shape)
++
++
++@partial(jax.jit, static_argnames=("layer_index", "num_classes"))
++def train_step_ff(state, image, labels, batch_rng, augmentation_on,
++ layer_index, num_classes, threshold):
++ """One greedy Forward-Forward update of exactly one layer."""
++ augmentation_rng, negative_rng = jax.random.split(batch_rng)
++ per_example_rng = jax.random.split(augmentation_rng, image.shape[0])
++ image = jax.lax.cond(
++ augmentation_on, vmap_augment_train, no_aug, image, per_example_rng)
++ offsets = jax.random.randint(
++ negative_rng, labels.shape, 1, num_classes)
++ negative_labels = (labels + offsets) % num_classes
++ positive = ff_overlay(image, labels, num_classes)
++ negative = ff_overlay(image, negative_labels, num_classes)
++ positive_input = state.apply_fn(
++ {"params": state.params}, positive, layer_index, method="ff_prefix")
++ negative_input = state.apply_fn(
++ {"params": state.params}, negative, layer_index, method="ff_prefix")
++ positive_input = jax.lax.stop_gradient(positive_input)
++ negative_input = jax.lax.stop_gradient(negative_input)
++
++ def objective(params):
++ return state.apply_fn(
++ {"params": params}, positive_input, negative_input, layer_index,
++ threshold, method="ff_layer_objective")
++
++ (loss, auxiliary), grads = jax.value_and_grad(
++ objective, has_aux=True)(state.params)
++ state = state.apply_gradients(grads=grads)
++ return state, {
++ "loss": loss,
++ "positive_goodness": jnp.mean(auxiliary[0]),
++ "negative_goodness": jnp.mean(auxiliary[1]),
++ "pair_accuracy": auxiliary[2],
++ }
++
++
++def train_ff_epoch(state, train_ds, batch_size, rng, augmentation_on,
++ layer_index, num_classes, threshold):
++ """Train one FF layer for one full data epoch."""
++ t0 = time.time()
++ train_ds_size = len(train_ds["image"])
++ steps_per_epoch = train_ds_size // batch_size
++ perms = jax.random.permutation(rng, train_ds_size)
++ perms = perms[:steps_per_epoch * batch_size]
++ perms = perms.reshape((steps_per_epoch, batch_size))
++ metrics = []
++ for permutation in perms:
++ image = train_ds["image"][permutation]
++ labels = train_ds["label"][permutation]
++ rng, batch_rng = jax.random.split(rng)
++ state, batch_metrics = train_step_ff(
++ state, image, labels, batch_rng, augmentation_on, layer_index,
++ num_classes, threshold)
++ metrics.append(batch_metrics)
++ host_metrics = jax.device_get(metrics)
++ summary = {
++ key: float(np.mean([record[key] for record in host_metrics]))
++ for key in host_metrics[0]
++ }
++ return state, summary, time.time() - t0
++
++
++@partial(jax.jit, static_argnames=("num_classes", "score_from_layer"))
++def predict_ff(state, image, num_classes, score_from_layer):
++ scores = []
++ for candidate in range(num_classes):
++ labels = jnp.full((image.shape[0],), candidate, dtype=jnp.int32)
++ overlaid = ff_overlay(image, labels, num_classes)
++ scores.append(state.apply_fn(
++ {"params": state.params}, overlaid, score_from_layer,
++ method="ff_goodness"))
++ return jnp.stack(scores, axis=-1)
++
++
++def eval_ff_model(state, dataset, batch_size, num_classes, score_from_layer):
++ """Evaluate all candidate-label overlays; no classifier head is assumed."""
++ t0 = time.time()
++ size = len(dataset["image"])
++ steps = size // batch_size
++ indices = jnp.arange(steps * batch_size).reshape((steps, batch_size))
++ correct = 0
++ total = 0
++ for index in indices:
++ image = dataset["image"][index]
++ labels = dataset["label"][index]
++ scores = predict_ff(state, image, num_classes, score_from_layer)
++ correct += int(jax.device_get(jnp.sum(jnp.argmax(scores, -1) == labels)))
++ total += labels.shape[0]
++ return 100.0 * correct / total, time.time() - t0
++
++
+ @jax.jit
+ def train_step_pepita(state, input_feedback, image, labels_onehot, labels,
+ batch_rng, augmentation_on, gradient_diagnostics):
+diff --git a/tests/local_rules_smoke.py b/tests/local_rules_smoke.py
+index e7da52e..8ce1e02 100644
+--- a/tests/local_rules_smoke.py
++++ b/tests/local_rules_smoke.py
+@@ -15,7 +15,11 @@ ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
+ sys.path.insert(0, ROOT)
+
+ from src.models import cnn_abstract
+-from src.training_utils import create_local_feedback, direct_feedback_fields
++from src.training_utils import (
++ create_local_feedback,
++ direct_feedback_fields,
++ ff_overlay,
++)
+
+
+ def loss_func(logits, one_hot):
+@@ -153,6 +157,38 @@ def main():
+ assert float(jnp.max(jnp.abs(pepita_projection))) <= projection_limit
+ assert bool(jnp.all(jnp.isfinite(flat(pepita_grads))))
+
++ positive = ff_overlay(x, labels, 3)
++ negative = ff_overlay(x, (labels + 1) % 3, 3)
++ assert bool(jnp.all(positive.reshape((x.shape[0], -1))[
++ jnp.arange(x.shape[0]), labels] == jnp.max(x)))
++ positive_prefix = model.apply(
++ {"params": params}, positive, 1, method="ff_prefix")
++ negative_prefix = model.apply(
++ {"params": params}, negative, 1, method="ff_prefix")
++ (ff_loss, ff_aux), ff_grads = jax.value_and_grad(
++ lambda candidate: model.apply(
++ {"params": candidate}, positive_prefix, negative_prefix, 1, 2.0,
++ method="ff_layer_objective"),
++ has_aux=True,
++ )(params)
++ ff_manual_loss = jnp.mean(
++ jax.nn.softplus(-ff_aux[0] + 2.0)
++ + jax.nn.softplus(ff_aux[1] - 2.0))
++ ff_objective_error = float(jnp.abs(ff_loss - ff_manual_loss))
++ assert ff_objective_error < 2e-12
++ assert float(jnp.linalg.norm(flat(ff_grads["c01"]))) > 0.0
++ assert float(jnp.linalg.norm(flat(ff_grads["c00"]))) == 0.0
++ assert float(jnp.linalg.norm(flat(ff_grads["d00"]))) == 0.0
++ ff_scores = jnp.stack([
++ model.apply(
++ {"params": params},
++ ff_overlay(x, jnp.full(labels.shape, candidate), 3), 1,
++ method="ff_goodness")
++ for candidate in range(3)
++ ], axis=-1)
++ assert ff_scores.shape == (x.shape[0], 3)
++ assert bool(jnp.all(jnp.isfinite(ff_scores)))
++
+ report = {
+ "symmetric_fa_bp_relative_error": symmetric_error,
+ "independent_feedback_forward_cosine": feedback_cosine,
+@@ -161,6 +197,8 @@ def main():
+ "dfa_hidden_maps": len(direct),
+ "pepita_readout_equation_max_error": pepita_readout_error,
+ "pepita_projection_shape": list(pepita_projection.shape),
++ "ff_local_objective_error": ff_objective_error,
++ "ff_score_shape": list(ff_scores.shape),
+ "status": "passed",
+ }
+ print(json.dumps(report, indent=2, sort_keys=True))
+diff --git a/train.py b/train.py
+index ba5f6ac..d3f0de6 100644
+--- a/train.py
++++ b/train.py
+@@ -14,6 +14,10 @@ from src import create_train_state, create_local_feedback, train_epoch, eval_mod
+ # import config # Use this for the old method
+ from config.cli_config import config
+
++if config.learning_algorithm == "ff":
++ raise ValueError(
++ "Forward-Forward uses greedy layerwise training; run train_ff.py")
++
+ experiment_dir = "./runs/" + config.experiment_name + "/"
+ if experiment_dir == "./runs/debug-test/" and os.path.isdir(experiment_dir):
+ shutil.rmtree(experiment_dir)
+diff --git a/train_ff.py b/train_ff.py
+new file mode 100644
+index 0000000..3839dd3
+--- /dev/null
++++ b/train_ff.py
+@@ -0,0 +1,92 @@
++"""Greedy supervised Forward-Forward on the author plain-CNN topologies."""
++import datetime
++import os
++import time
++
++import jax
++import numpy as np
++
++from config.cli_config import config
++from src import create_ff_train_state, eval_ff_model, train_ff_epoch
++
++
++if config.learning_algorithm != "ff":
++ raise ValueError("train_ff.py requires --learning-algorithm ff")
++if config.ff_score_from_layer < 0:
++ raise ValueError("--ff-score-from-layer must be nonnegative")
++
++experiment_dir = os.path.join("runs", config.experiment_name)
++if os.path.isdir(experiment_dir):
++ raise FileExistsError(
++ "experiment directory exists; refusing to overwrite "
++ + experiment_dir)
++os.makedirs(experiment_dir)
++
++for experiment_index, seed in enumerate(config.seeds):
++ timestamp = datetime.datetime.fromtimestamp(time.time())
++ outpath = os.path.join(
++ experiment_dir, timestamp.strftime("%Y_%m_%d_%H_%M_%S"))
++ os.makedirs(outpath)
++ print(
++ f"Starting FF experiment {experiment_index + 1}/"
++ f"{len(config.seeds)} seed={seed}",
++ flush=True)
++ rng = jax.random.PRNGKey(seed)
++ rng, init_rng = jax.random.split(rng)
++ state = create_ff_train_state(
++ init_rng, config.model, config.image_dims, config.learning_rate)
++ num_layers = len(state.params)
++ if config.ff_score_from_layer >= num_layers:
++ raise ValueError("--ff-score-from-layer excludes every layer")
++ augmentation_on = config.dataset not in ("mnist", "fashionmnist")
++ history = {
++ "method": "ff",
++ "epochs_per_layer": config.num_epochs,
++ "num_layers": num_layers,
++ "threshold": config.ff_threshold,
++ "score_from_layer": config.ff_score_from_layer,
++ "learning_rate": config.learning_rate,
++ "layers": [],
++ }
++ started = time.time()
++ for layer_index in range(num_layers):
++ layer_record = {"layer": layer_index, "epochs": []}
++ for epoch in range(config.num_epochs):
++ rng, epoch_rng = jax.random.split(rng)
++ state, metrics, runtime = train_ff_epoch(
++ state, config.train_ds, config.batch_size, epoch_rng,
++ augmentation_on, layer_index, config.num_classes,
++ config.ff_threshold)
++ record = {
++ "epoch": epoch + 1,
++ "runtime": runtime,
++ **metrics,
++ }
++ layer_record["epochs"].append(record)
++ print(
++ f"layer={layer_index + 1}/{num_layers} "
++ f"epoch={epoch + 1}/{config.num_epochs} "
++ f"loss={metrics['loss']:.5f} "
++ f"pair_acc={100 * metrics['pair_accuracy']:.2f}% "
++ f"runtime={runtime:.3f}s",
++ flush=True)
++ history["layers"].append(layer_record)
++
++ validation_accuracy, validation_time = eval_ff_model(
++ state, config.val_ds, config.batch_size, config.num_classes,
++ config.ff_score_from_layer)
++ test_accuracy, test_time = eval_ff_model(
++ state, config.test_ds, config.batch_size, config.num_classes,
++ config.ff_score_from_layer)
++ history["final"] = {
++ "validation_accuracy": validation_accuracy,
++ "validation_time": validation_time,
++ "test_accuracy": test_accuracy,
++ "test_time": test_time,
++ "train_and_eval_wall": time.time() - started,
++ }
++ print(
++ f"final val_accuracy={validation_accuracy:.3f}% "
++ f"test_accuracy={test_accuracy:.3f}%",
++ flush=True)
++ np.save(os.path.join(outpath, "hist.npy"), history)
+--
+2.54.0
+