From 672423d13e36d7a63fa89aa186dc35fc91475ab1 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Mon, 27 Jul 2026 14:18:29 -0500 Subject: baseline: add matched Transformer Forward-Forward --- experiments/transformer_feedback_smoke.py | 77 ++++++++++++++++++ sdil/transformer.py | 127 +++++++++++++++++++++++++++++- 2 files changed, 202 insertions(+), 2 deletions(-) diff --git a/experiments/transformer_feedback_smoke.py b/experiments/transformer_feedback_smoke.py index 1e27b38..2cd22d9 100644 --- a/experiments/transformer_feedback_smoke.py +++ b/experiments/transformer_feedback_smoke.py @@ -319,6 +319,82 @@ def audit_pepita(): } +def audit_forward_forward(): + config = LocalTransformerConfig( + vocab_size=7, context_length=5, depth=2, width=8, heads=2, + mlp_ratio=2, seed=151) + bp = LocalDecoderTransformer(config, "bp", dtype=torch.float64) + ff = LocalDecoderTransformer(config, "ff", dtype=torch.float64) + generator = torch.Generator().manual_seed(152) + tokens = torch.randint(0, 7, (3, 5), generator=generator) + positive = torch.randint(0, 7, (3, 5), generator=generator) + negative = (positive + torch.randint( + 1, 7, positive.shape, generator=generator)) % 7 + with torch.no_grad(): + forward_error = float(torch.max(torch.abs( + bp(tokens)["logits"] - ff(tokens)["logits"]))) + all_parameters = tuple(ff.parameters()) + audited = [] + for layer_index in range(ff.ff_num_layers): + for parameter in all_parameters: + parameter.grad = None + target_ids = { + id(parameter) + for parameter in ff.ff_layer_parameters(layer_index)} + loss, metrics = ff.ff_local_loss( + layer_index, tokens, positive, negative) + loss.backward() + unexpected = [ + index for index, parameter in enumerate(all_parameters) + if parameter.grad is not None and id(parameter) not in target_ids] + missing = [ + index for index, parameter in enumerate(all_parameters) + if parameter.grad is None and id(parameter) in target_ids] + explicit_positive = ff._ff_local_outputs( + layer_index, tokens, positive, + ff.ff_forward(tokens, positive)) + explicit_negative = ff._ff_local_outputs( + layer_index, tokens, negative, + ff.ff_forward(tokens, negative)) + explicit = ( + torch.nn.functional.softplus( + -explicit_positive.square().mean(dim=-1) + 2.0) + + torch.nn.functional.softplus( + explicit_negative.square().mean(dim=-1) - 2.0) + ).mean() + loss_error = abs(float(loss.detach() - explicit.detach())) + if unexpected or missing or loss_error >= 1e-12: + raise AssertionError({ + "layer": layer_index, + "unexpected": unexpected, + "missing": missing, + "loss_error": loss_error, + }) + audited.append({ + "layer": layer_index, + "loss": float(loss.detach()), + "loss_error": loss_error, + **metrics, + }) + scores = ff.ff_candidate_scores(tokens) + if (forward_error >= 1e-12 or scores.shape != (3, 5, 7) + or not torch.isfinite(scores).all()): + raise AssertionError({ + "forward_error": forward_error, + "score_shape": list(scores.shape), + }) + return { + "matched_forward_max_absolute_error": forward_error, + "audited_greedy_layers": len(audited), + "max_local_objective_absolute_error": max( + row["loss_error"] for row in audited), + "non_target_parameter_gradients": 0, + "candidate_score_shape": list(scores.shape), + "candidate_presentations_per_evaluation": + config.vocab_size, + } + + def main(): torch.set_num_threads(1) result = { @@ -329,6 +405,7 @@ def main(): "sdil_innovation_identity": audit_sdil_innovation_identity(), "strict_dfa": audit_strict_dfa(), "pepita": audit_pepita(), + "forward_forward": audit_forward_forward(), } print(json.dumps(result, indent=2, sort_keys=True)) diff --git a/sdil/transformer.py b/sdil/transformer.py index ee1e69f..d47fd18 100644 --- a/sdil/transformer.py +++ b/sdil/transformer.py @@ -15,7 +15,7 @@ import torch.nn.functional as F _FEEDBACK_METHODS = {"fa", "clean_kp", "sdil"} -_SUPPORTED_METHODS = {"bp", "dfa", "pepita"} | _FEEDBACK_METHODS +_SUPPORTED_METHODS = {"bp", "dfa", "pepita", "ff"} | _FEEDBACK_METHODS class _FeedbackLinearFunction(torch.autograd.Function): @@ -126,7 +126,7 @@ class FeedbackLinear(nn.Module): self.register_buffer("last_traffic_rms", torch.zeros((), dtype=dtype)) def forward(self, x): - if self.method in {"bp", "dfa", "pepita"}: + if self.method in {"bp", "dfa", "pepita", "ff"}: return F.linear(x, self.weight, self.bias) method_code = {"fa": 0, "clean_kp": 1, "sdil": 2}[self.method] return _FeedbackLinearFunction.apply( @@ -573,3 +573,126 @@ class LocalDecoderTransformer(nn.Module): if module.feedback is None: raise ValueError("BP modules do not contain feedback tensors") module.feedback.copy_(module.weight) + + @staticmethod + def _ff_normalize(value): + return value / ( + value.square().sum(dim=-1, keepdim=True).sqrt() + 1e-8) + + def _ff_overlay(self, embedded, candidate_labels): + if candidate_labels.shape != embedded.shape[:2]: + raise ValueError("candidate labels must match token positions") + if self.config.width < self.config.vocab_size: + raise ValueError( + "Forward-Forward overlay requires width >= vocabulary") + overlaid = embedded.clone() + overlaid[..., :self.config.vocab_size] = 0.0 + magnitude = embedded.detach().abs().amax( + dim=-1, keepdim=True).clamp_min(1e-4) + overlaid.scatter_( + -1, candidate_labels.unsqueeze(-1), magnitude) + return overlaid + + def ff_forward(self, tokens, candidate_labels): + """Run candidate-labelled, normalized data through the FF graph.""" + if self.method != "ff": + raise ValueError("ff_forward is only valid for method='ff'") + positions = self.position_embedding[:tokens.shape[1]] + base = self.token_embedding(tokens) + positions + embedded = self._ff_overlay(base, candidate_labels) + hidden = embedded + block_inputs = [] + block_outputs = [] + for block in self.blocks: + local_input = self._ff_normalize(hidden) + block_inputs.append(local_input) + hidden = block(local_input) + block_outputs.append(hidden) + final_input = self._ff_normalize(hidden) + logits = self.head(self.final_norm(final_input)) + return { + "embedded": embedded, + "block_inputs": block_inputs, + "block_outputs": block_outputs, + "final_input": final_input, + "logits": logits, + } + + @property + def ff_num_layers(self): + return self.config.depth + 2 + + def ff_layer_parameters(self, layer_index): + if layer_index == 0: + return (self.token_embedding.weight, self.position_embedding) + if 1 <= layer_index <= self.config.depth: + return tuple(self.blocks[layer_index - 1].parameters()) + if layer_index == self.config.depth + 1: + return ( + *tuple(self.final_norm.parameters()), + *tuple(self.head.parameters())) + raise ValueError("invalid Forward-Forward layer index") + + def _ff_local_outputs(self, layer_index, tokens, labels, cached): + if layer_index == 0: + base = ( + self.token_embedding(tokens) + + self.position_embedding[:tokens.shape[1]]) + return self._ff_overlay(base, labels) + if 1 <= layer_index <= self.config.depth: + local_input = cached["block_inputs"][ + layer_index - 1].detach() + return self.blocks[layer_index - 1](local_input) + if layer_index == self.config.depth + 1: + local_input = cached["final_input"].detach() + return self.head(self.final_norm(local_input)) + raise ValueError("invalid Forward-Forward layer index") + + def ff_local_loss( + self, layer_index, tokens, positive_labels, negative_labels, + threshold=2.0): + """Return one greedy FF objective with every prefix detached.""" + if self.method != "ff": + raise ValueError("ff_local_loss is only valid for method='ff'") + if not 0 <= layer_index < self.ff_num_layers: + raise ValueError("invalid Forward-Forward layer index") + with torch.no_grad(): + positive_cache = self.ff_forward(tokens, positive_labels) + negative_cache = self.ff_forward(tokens, negative_labels) + positive_output = self._ff_local_outputs( + layer_index, tokens, positive_labels, positive_cache) + negative_output = self._ff_local_outputs( + layer_index, tokens, negative_labels, negative_cache) + positive_goodness = positive_output.square().mean(dim=-1) + negative_goodness = negative_output.square().mean(dim=-1) + loss = ( + F.softplus(-positive_goodness + threshold) + + F.softplus(negative_goodness - threshold) + ).mean() + return loss, { + "positive_goodness": float( + positive_goodness.mean().detach()), + "negative_goodness": float( + negative_goodness.mean().detach()), + "pair_accuracy": float( + (positive_goodness > negative_goodness).float().mean()), + } + + @torch.no_grad() + def ff_candidate_scores(self, tokens, score_from_layer=1): + """Score every next-token candidate; cost is charged as V passes.""" + if not 0 <= score_from_layer < self.ff_num_layers: + raise ValueError("invalid Forward-Forward score range") + scores = [] + for candidate in range(self.config.vocab_size): + labels = torch.full_like(tokens, candidate) + candidate_forward = self.ff_forward(tokens, labels) + layer_outputs = ( + [candidate_forward["embedded"]] + + candidate_forward["block_outputs"] + + [candidate_forward["logits"]]) + goodness = [ + output.square().mean(dim=-1) + for output in layer_outputs] + scores.append(sum(goodness[score_from_layer:])) + return torch.stack(scores, dim=-1) -- cgit v1.2.3