summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-27 14:17:11 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-27 14:17:11 -0500
commitad557695cccdcdc679194f3ec88c0e30ff0b8596 (patch)
tree48f906457f56a29e25581a6de420069f96ccfd7d /experiments
parent2f4e0aca2e4f9064d6548eca876cab433ad769fc (diff)
baseline: add matched Transformer PEPITA
Diffstat (limited to 'experiments')
-rw-r--r--experiments/transformer_feedback_smoke.py75
1 files changed, 75 insertions, 0 deletions
diff --git a/experiments/transformer_feedback_smoke.py b/experiments/transformer_feedback_smoke.py
index c3cd4da..1e27b38 100644
--- a/experiments/transformer_feedback_smoke.py
+++ b/experiments/transformer_feedback_smoke.py
@@ -1,6 +1,7 @@
#!/usr/bin/env python3
"""Deterministic audits for Transformer BP/FA/KP/SDIL feedback transport."""
import json
+import math
import os
import sys
@@ -245,6 +246,79 @@ def audit_strict_dfa():
}
+def audit_pepita():
+ config = LocalTransformerConfig(
+ vocab_size=13, context_length=6, depth=2, width=8, heads=2,
+ mlp_ratio=2, seed=141)
+ bp = LocalDecoderTransformer(config, "bp", dtype=torch.float64)
+ pepita = LocalDecoderTransformer(config, "pepita", dtype=torch.float64)
+ generator = torch.Generator().manual_seed(142)
+ tokens = torch.randint(0, 13, (3, 6), generator=generator)
+ targets = torch.randint(0, 13, (3, 6), generator=generator)
+ with torch.no_grad():
+ bp_output = bp(tokens)["logits"]
+ clean = pepita(tokens, return_cache=True)
+ one_hot = torch.nn.functional.one_hot(
+ targets, 13).to(torch.float64)
+ clean_error = torch.softmax(clean["logits"], dim=-1) - one_hot
+ offset = clean_error @ pepita.pepita_input_feedback
+ modulated = pepita(
+ tokens, return_cache=True, embedding_offset=offset)
+ modulated_error = (
+ torch.softmax(modulated["logits"], dim=-1) - one_hot)
+ forward_error = float(torch.max(torch.abs(
+ bp_output - clean["logits"])))
+ feedback_before = pepita.pepita_input_feedback.detach().clone()
+ metrics = pepita.pepita_gradients(tokens, targets)
+ missing = [
+ name for name, parameter in pepita.named_parameters()
+ if parameter.grad is None]
+ nonfinite = [
+ name for name, parameter in pepita.named_parameters()
+ if parameter.grad is not None
+ and not torch.isfinite(parameter.grad).all()]
+ observations = targets.numel()
+ expected_head = (
+ (modulated_error / observations).reshape(-1, 13).t()
+ @ modulated["normalized"].reshape(-1, 8))
+ head_error = _relative_error(
+ pepita.head.weight.grad, expected_head)
+ embedding_field = clean["embedded"] - modulated["embedded"]
+ expected_position = embedding_field.sum(dim=0) / observations
+ position_error = _relative_error(
+ pepita.position_embedding.grad[:tokens.shape[1]],
+ expected_position)
+ feedback_change = float(torch.max(torch.abs(
+ pepita.pepita_input_feedback - feedback_before)))
+ projection_limit = math.sqrt(6.0 / config.width) * 0.05
+ observed_limit = float(torch.max(torch.abs(
+ pepita.pepita_input_feedback)))
+ if (forward_error >= 1e-12 or head_error >= 1e-12
+ or position_error >= 1e-12 or feedback_change != 0.0
+ or observed_limit > projection_limit or missing or nonfinite):
+ raise AssertionError({
+ "forward_error": forward_error,
+ "head_error": head_error,
+ "position_error": position_error,
+ "feedback_change": feedback_change,
+ "projection_limit": projection_limit,
+ "observed_limit": observed_limit,
+ "missing": missing,
+ "nonfinite": nonfinite,
+ })
+ return {
+ "matched_forward_max_absolute_error": forward_error,
+ "head_local_equation_relative_error": head_error,
+ "position_local_equation_relative_error": position_error,
+ "fixed_projection_change": feedback_change,
+ "projection_limit": projection_limit,
+ "observed_projection_max": observed_limit,
+ "training_presentations": metrics["training_presentations"],
+ "detached_block_boundaries": metrics["detached_block_boundaries"],
+ "uses_task_loss_backward": metrics["uses_task_loss_backward"],
+ }
+
+
def main():
torch.set_num_threads(1)
result = {
@@ -254,6 +328,7 @@ def main():
audit_fa_independence_and_kp_locality(),
"sdil_innovation_identity": audit_sdil_innovation_identity(),
"strict_dfa": audit_strict_dfa(),
+ "pepita": audit_pepita(),
}
print(json.dumps(result, indent=2, sort_keys=True))