diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-27 14:17:11 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-27 14:17:11 -0500 |
| commit | ad557695cccdcdc679194f3ec88c0e30ff0b8596 (patch) | |
| tree | 48f906457f56a29e25581a6de420069f96ccfd7d /experiments | |
| parent | 2f4e0aca2e4f9064d6548eca876cab433ad769fc (diff) | |
baseline: add matched Transformer PEPITA
Diffstat (limited to 'experiments')
| -rw-r--r-- | experiments/transformer_feedback_smoke.py | 75 |
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)) |
