summaryrefslogtreecommitdiff
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
parent2f4e0aca2e4f9064d6548eca876cab433ad769fc (diff)
baseline: add matched Transformer PEPITA
-rw-r--r--experiments/transformer_feedback_smoke.py75
-rw-r--r--sdil/transformer.py122
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())