summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-27 14:18:29 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-27 14:18:29 -0500
commit672423d13e36d7a63fa89aa186dc35fc91475ab1 (patch)
treecd1f7d548938750f2117ba50bca0c8e55c8f3c19 /experiments
parentad557695cccdcdc679194f3ec88c0e30ff0b8596 (diff)
baseline: add matched Transformer Forward-Forward
Diffstat (limited to 'experiments')
-rw-r--r--experiments/transformer_feedback_smoke.py77
1 files changed, 77 insertions, 0 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))