From ad557695cccdcdc679194f3ec88c0e30ff0b8596 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Mon, 27 Jul 2026 14:17:11 -0500 Subject: baseline: add matched Transformer PEPITA --- experiments/transformer_feedback_smoke.py | 75 ++++++++++++++++++ sdil/transformer.py | 122 +++++++++++++++++++++++++++++- 2 files changed, 193 insertions(+), 4 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)) diff --git a/sdil/transformer.py b/sdil/transformer.py index 608f43f..ee1e69f 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"} | _FEEDBACK_METHODS +_SUPPORTED_METHODS = {"bp", "dfa", "pepita"} | _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"}: + if self.method in {"bp", "dfa", "pepita"}: return F.linear(x, self.weight, self.bias) method_code = {"fa": 0, "clean_kp": 1, "sdil": 2}[self.method] return _FeedbackLinearFunction.apply( @@ -292,23 +292,44 @@ class LocalDecoderTransformer(nn.Module): self.register_buffer("dfa_block_feedback", None) self.register_buffer("dfa_embedding_feedback", None) self.register_buffer("dfa_final_norm_feedback", None) + if method == "pepita": + pepita_generator = torch.Generator().manual_seed( + config.seed + 3) + limit = math.sqrt(6.0 / config.width) * 0.05 + projection = ( + 2.0 * torch.rand( + config.vocab_size, config.width, + generator=pepita_generator, dtype=dtype) - 1.0 + ) * limit + self.register_buffer("pepita_input_feedback", projection) + else: + self.register_buffer("pepita_input_feedback", None) def forward( self, tokens, targets: Optional[torch.Tensor] = None, - return_cache: bool = False): + return_cache: bool = False, + embedding_offset: Optional[torch.Tensor] = None): if tokens.ndim != 2: raise ValueError("tokens must have shape (batch, time)") if tokens.shape[1] > self.config.context_length: raise ValueError("sequence exceeds configured context length") positions = self.position_embedding[:tokens.shape[1]] hidden = self.token_embedding(tokens) + positions + if embedding_offset is not None: + if embedding_offset.shape != hidden.shape: + raise ValueError("embedding_offset shape does not match tokens") + hidden = hidden + embedding_offset + embedded = hidden hidden = F.dropout( hidden, p=self.config.dropout, training=self.training) block_inputs = [] + block_outputs = [] for block in self.blocks: if return_cache: block_inputs.append(hidden.detach()) hidden = block(hidden) + if return_cache: + block_outputs.append(hidden.detach()) final_input = hidden normalized = self.final_norm(hidden) logits = self.head(normalized) @@ -320,6 +341,8 @@ class LocalDecoderTransformer(nn.Module): result = {"logits": logits, "loss": loss, "hidden": hidden} if return_cache: result["block_inputs"] = block_inputs + result["block_outputs"] = block_outputs + result["embedded"] = embedded.detach() result["final_input"] = final_input.detach() result["normalized"] = normalized.detach() return result @@ -410,6 +433,94 @@ class LocalDecoderTransformer(nn.Module): "detached_block_boundaries": len(self.blocks), } + def pepita_gradients(self, tokens, targets): + """Populate two-presentation PEPITA/ERIN local gradients. + + The analytical output error is projected into the continuous token + embedding stream. Each block then uses its first-minus-second output + difference and the second-presentation input. Block boundaries are + detached. Since discrete token IDs cannot themselves be perturbed, + the embedding table and positional code use the directly observable + embedding difference as their local field. + """ + if self.method != "pepita": + raise ValueError( + "pepita_gradients is only valid for method='pepita'") + with torch.no_grad(): + clean = self.forward(tokens, return_cache=True) + one_hot = F.one_hot( + targets, self.config.vocab_size).to(clean["logits"].dtype) + clean_error = torch.softmax(clean["logits"], dim=-1) - one_hot + embedding_offset = clean_error @ self.pepita_input_feedback + modulated = self.forward( + tokens, return_cache=True, + embedding_offset=embedding_offset) + modulated_error = ( + torch.softmax(modulated["logits"], dim=-1) - one_hot) + for parameter in self.parameters(): + parameter.grad = None + observations = targets.numel() + + error_flat = modulated_error.reshape( + -1, self.config.vocab_size) / observations + normalized_flat = modulated["normalized"].reshape( + -1, self.config.width) + self.head.weight.grad = error_flat.t() @ normalized_flat + if self.head.bias is not None: + self.head.bias.grad = error_flat.sum(dim=0) + + final_input = modulated["final_input"].detach() + normalized = self.final_norm(final_input) + final_field = ( + clean["normalized"] - modulated["normalized"]).detach() + final_parameters = tuple(self.final_norm.parameters()) + final_objective = torch.sum( + normalized * final_field) / observations + final_gradients = torch.autograd.grad( + final_objective, final_parameters) + self._assign_gradients(final_parameters, final_gradients) + + for block, local_input, clean_output, modulated_output in zip( + self.blocks, modulated["block_inputs"], + clean["block_outputs"], modulated["block_outputs"]): + local_output = block(local_input.detach()) + local_field = (clean_output - modulated_output).detach() + local_parameters = tuple(block.parameters()) + local_objective = torch.sum( + local_output * local_field) / observations + local_gradients = torch.autograd.grad( + local_objective, local_parameters) + self._assign_gradients(local_parameters, local_gradients) + + base_embedding = ( + self.token_embedding(tokens) + + self.position_embedding[:tokens.shape[1]]) + embedding_field = ( + clean["embedded"] - modulated["embedded"]).detach() + embedding_objective = torch.sum( + base_embedding * embedding_field) / observations + embedding_parameters = ( + self.token_embedding.weight, self.position_embedding) + embedding_gradients = torch.autograd.grad( + embedding_objective, embedding_parameters) + self._assign_gradients( + embedding_parameters, embedding_gradients) + return { + "clean_loss": float(F.cross_entropy( + clean["logits"].reshape(-1, self.config.vocab_size), + targets.reshape(-1))), + "embedding_offset_rms": float( + embedding_offset.square().mean().sqrt()), + "mean_block_field_rms": float(torch.stack([ + (first - second).square().mean().sqrt() + for first, second in zip( + clean["block_outputs"], + modulated["block_outputs"])]).mean()), + "training_presentations": 2, + "uses_task_loss_backward": False, + "detached_block_boundaries": len(self.blocks), + } + def feedback_linears(self) -> Iterable[FeedbackLinear]: return ( module for module in self.modules() @@ -439,7 +550,10 @@ class LocalDecoderTransformer(nn.Module): self.dfa_block_feedback, self.dfa_embedding_feedback, self.dfa_final_norm_feedback) if tensor is not None) - return affine_feedback + dfa_feedback + pepita_feedback = ( + self.pepita_input_feedback.numel() + if self.pepita_input_feedback is not None else 0) + return affine_feedback + dfa_feedback + pepita_feedback def teaching_statistics(self) -> Dict[str, float]: modules = list(self.feedback_linears()) -- cgit v1.2.3