From 6a65b8b26e5cf9b7c6b4c429d5bc02ac0475b0b0 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Mon, 27 Jul 2026 14:14:25 -0500 Subject: baseline: add matched Transformer feedback transport --- experiments/transformer_feedback_smoke.py | 170 +++++++++++++++ sdil/__init__.py | 2 +- sdil/transformer.py | 335 ++++++++++++++++++++++++++++++ 3 files changed, 506 insertions(+), 1 deletion(-) create mode 100644 experiments/transformer_feedback_smoke.py create mode 100644 sdil/transformer.py diff --git a/experiments/transformer_feedback_smoke.py b/experiments/transformer_feedback_smoke.py new file mode 100644 index 0000000..99bca48 --- /dev/null +++ b/experiments/transformer_feedback_smoke.py @@ -0,0 +1,170 @@ +#!/usr/bin/env python3 +"""Deterministic audits for Transformer BP/FA/KP/SDIL feedback transport.""" +import json +import os +import sys + +import torch + +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +from sdil.transformer import ( # noqa: E402 + FeedbackLinear, + LocalDecoderTransformer, + LocalTransformerConfig, +) + + +def _named_forward_gradients(model): + feedback_ids = { + id(module.feedback) + for module in model.feedback_linears() + if isinstance(module.feedback, torch.nn.Parameter)} + return { + name: parameter.grad.detach().clone() + for name, parameter in model.named_parameters() + if id(parameter) not in feedback_ids} + + +def _relative_error(actual, expected): + return float( + torch.linalg.vector_norm(actual - expected) + / torch.linalg.vector_norm(expected).clamp_min(1e-30)) + + +def audit_matched_forward_and_symmetric_limit(): + config = LocalTransformerConfig( + vocab_size=11, context_length=7, depth=2, width=8, heads=2, + mlp_ratio=2, seed=101) + bp = LocalDecoderTransformer(config, "bp", dtype=torch.float64) + fa = LocalDecoderTransformer(config, "fa", dtype=torch.float64) + fa.set_feedback_equal_to_forward() + generator = torch.Generator().manual_seed(102) + tokens = torch.randint(0, 11, (3, 7), generator=generator) + targets = torch.randint(0, 11, (3, 7), generator=generator) + bp_output = bp(tokens, targets) + fa_output = fa(tokens, targets) + forward_error = float(torch.max(torch.abs( + bp_output["logits"] - fa_output["logits"])).detach()) + bp_output["loss"].backward() + fa_output["loss"].backward() + bp_gradients = _named_forward_gradients(bp) + fa_gradients = _named_forward_gradients(fa) + if bp_gradients.keys() != fa_gradients.keys(): + raise AssertionError("forward parameter names differ") + errors = { + name: _relative_error(fa_gradients[name], bp_gradients[name]) + for name in bp_gradients} + if forward_error >= 1e-12 or max(errors.values()) >= 2e-12: + raise AssertionError({ + "forward_error": forward_error, + "gradient_errors": errors, + }) + return { + "matched_forward_max_absolute_error": forward_error, + "symmetric_limit_max_relative_error": max(errors.values()), + "audited_forward_parameter_tensors": len(errors), + "matched_forward_parameter_count": bp.n_forward_parameters, + "matched_feedback_tensor_count": len(list(fa.feedback_linears())), + } + + +def audit_fa_independence_and_kp_locality(): + forward_generator = torch.Generator().manual_seed(111) + feedback_generator = torch.Generator().manual_seed(112) + fa = FeedbackLinear( + 5, 3, "fa", forward_generator, feedback_generator, + dtype=torch.float64) + feedback_before = fa.feedback.detach().clone() + with torch.no_grad(): + fa.weight.add_(torch.randn( + fa.weight.shape, generator=forward_generator, + dtype=fa.weight.dtype)) + fa_independence_error = float(torch.max(torch.abs( + fa.feedback - feedback_before))) + + forward_generator = torch.Generator().manual_seed(113) + feedback_generator = torch.Generator().manual_seed(114) + kp = FeedbackLinear( + 5, 3, "clean_kp", forward_generator, feedback_generator, + dtype=torch.float64) + x = torch.randn( + 2, 4, 5, generator=forward_generator, dtype=torch.float64, + requires_grad=True) + upstream = torch.randn( + 2, 4, 3, generator=forward_generator, dtype=torch.float64) + output = kp(x) + (output * upstream).sum().backward() + expected = upstream.reshape(-1, 3).t() @ x.detach().reshape(-1, 5) + forward_error = _relative_error(kp.weight.grad, expected) + feedback_error = _relative_error(kp.feedback.grad, expected) + if max(fa_independence_error, forward_error, feedback_error) >= 1e-12: + raise AssertionError({ + "fa_independence_error": fa_independence_error, + "forward_local_error": forward_error, + "feedback_local_error": feedback_error, + }) + return { + "fa_feedback_change_after_forward_weight_mutation": + fa_independence_error, + "kp_forward_local_correlation_relative_error": forward_error, + "kp_feedback_local_correlation_relative_error": feedback_error, + "kp_feedback_gradient_recomputed_locally": True, + } + + +def audit_sdil_innovation_identity(): + config = LocalTransformerConfig( + vocab_size=13, context_length=6, depth=2, width=8, heads=2, + mlp_ratio=2, traffic_ratio=4.0, seed=121) + kp = LocalDecoderTransformer(config, "clean_kp", dtype=torch.float64) + sdil = LocalDecoderTransformer(config, "sdil", dtype=torch.float64) + generator = torch.Generator().manual_seed(122) + tokens = torch.randint(0, 13, (3, 6), generator=generator) + targets = torch.randint(0, 13, (3, 6), generator=generator) + kp(tokens, targets)["loss"].backward() + sdil(tokens, targets)["loss"].backward() + kp_gradients = { + name: parameter.grad.detach() + for name, parameter in kp.named_parameters()} + sdil_gradients = { + name: parameter.grad.detach() + for name, parameter in sdil.named_parameters()} + errors = { + name: _relative_error(sdil_gradients[name], kp_gradients[name]) + for name in kp_gradients} + statistics = sdil.teaching_statistics() + traffic_ratio = ( + statistics["traffic_rms"] + / max(statistics["innovation_rms"], 1e-30)) + raw_is_contaminated = ( + statistics["raw_rms"] > statistics["innovation_rms"]) + if max(errors.values()) >= 2e-12 or not raw_is_contaminated: + raise AssertionError({ + "gradient_errors": errors, + "statistics": statistics, + "observed_traffic_ratio": traffic_ratio, + }) + return { + "kp_sdil_max_gradient_relative_error": max(errors.values()), + "mean_raw_rms": statistics["raw_rms"], + "mean_innovation_rms": statistics["innovation_rms"], + "mean_traffic_rms": statistics["traffic_rms"], + "raw_signal_contaminated": raw_is_contaminated, + "paired_neutral_subtraction": True, + } + + +def main(): + torch.set_num_threads(1) + result = { + "matched_forward_and_symmetric_limit": + audit_matched_forward_and_symmetric_limit(), + "fa_independence_and_kp_locality": + audit_fa_independence_and_kp_locality(), + "sdil_innovation_identity": audit_sdil_innovation_identity(), + } + print(json.dumps(result, indent=2, sort_keys=True)) + + +if __name__ == "__main__": + main() diff --git a/sdil/__init__.py b/sdil/__init__.py index 4a7cf44..09f751f 100644 --- a/sdil/__init__.py +++ b/sdil/__init__.py @@ -1 +1 @@ -from . import core, probes, data, baselines # noqa: F401 +from . import core, probes, data, baselines, transformer # noqa: F401 diff --git a/sdil/transformer.py b/sdil/transformer.py new file mode 100644 index 0000000..d6433c4 --- /dev/null +++ b/sdil/transformer.py @@ -0,0 +1,335 @@ +"""Matched decoder-Transformer components for local-learning crossovers. + +The forward graph is deliberately identical for BP, ordinary FA, clean KP, +and SDIL. Only the vector transported through parameterized affine maps is +changed. Parameter-free Jacobians (residual addition, LayerNorm, GELU, and +softmax attention) remain local and exact. +""" +from dataclasses import dataclass +import math +from typing import Dict, Iterable, Optional + +import torch +from torch import nn +import torch.nn.functional as F + + +_FEEDBACK_METHODS = {"fa", "clean_kp", "sdil"} + + +class _FeedbackLinearFunction(torch.autograd.Function): + """Linear map with fixed or locally plastic feedback. + + ``feedback`` has the same orientation as ``weight``. Consequently the + transported vector is ``delta @ feedback``, while both plastic matrices + can be updated from the locally available ``delta.T @ input`` correlation. + The feedback correlation is recomputed rather than copied from the + forward-weight gradient. + """ + + @staticmethod + def forward( + ctx, x, weight, feedback, bias, method_code, traffic_ratio, + raw_rms, innovation_rms, traffic_rms): + output = F.linear(x, weight, bias) + ctx.save_for_backward(x, feedback, output) + ctx.method_code = int(method_code) + ctx.traffic_ratio = float(traffic_ratio) + ctx.has_bias = bias is not None + ctx.raw_rms = raw_rms + ctx.innovation_rms = innovation_rms + ctx.traffic_rms = traffic_rms + return output + + @staticmethod + def backward(ctx, grad_output): + x, feedback, soma = ctx.saved_tensors + task_instruction = grad_output + traffic = torch.zeros_like(task_instruction) + if ctx.method_code == 2 and ctx.traffic_ratio: + # A paired neutral observation exposes the component predictable + # from the local somatic response. Match it to a frozen multiple + # of task-instruction RMS without changing its somatic direction. + task_rms = task_instruction.square().mean().sqrt() + centered_soma = soma - soma.mean(dim=-1, keepdim=True) + soma_rms = centered_soma.square().mean().sqrt().clamp_min(1e-30) + traffic = ( + centered_soma * task_rms * ctx.traffic_ratio / soma_rms) + raw_apical = task_instruction + traffic + neutral_prediction = traffic + innovation = raw_apical - neutral_prediction + + input_flat = x.reshape(-1, x.shape[-1]) + delta_flat = innovation.reshape(-1, innovation.shape[-1]) + grad_input = innovation @ feedback + grad_weight = delta_flat.t() @ input_flat + grad_feedback = None + if ctx.method_code in (1, 2): + # This is intentionally a second evaluation of the local + # correlation, not an assignment from grad_weight. + grad_feedback = delta_flat.t() @ input_flat + grad_bias = None + if ctx.has_bias: + grad_bias = delta_flat.sum(dim=0) + + with torch.no_grad(): + ctx.raw_rms.copy_(raw_apical.square().mean().sqrt()) + ctx.innovation_rms.copy_(innovation.square().mean().sqrt()) + ctx.traffic_rms.copy_(traffic.square().mean().sqrt()) + return ( + grad_input, grad_weight, grad_feedback, grad_bias, + None, None, None, None, None) + + +class FeedbackLinear(nn.Module): + """A forward-matched affine map for BP, FA, clean KP, or SDIL.""" + + def __init__( + self, in_features: int, out_features: int, method: str, + forward_generator: torch.Generator, + feedback_generator: torch.Generator, + bias: bool = False, init_std: float = 0.02, + traffic_ratio: float = 4.0, dtype=torch.float32): + super().__init__() + if method not in {"bp"} | _FEEDBACK_METHODS: + raise ValueError(f"unsupported feedback-linear method: {method}") + self.in_features = int(in_features) + self.out_features = int(out_features) + self.method = method + self.traffic_ratio = float(traffic_ratio if method == "sdil" else 0.0) + self.weight = nn.Parameter(torch.empty( + out_features, in_features, dtype=dtype)) + nn.init.normal_( + self.weight, mean=0.0, std=init_std, + generator=forward_generator) + if bias: + self.bias = nn.Parameter(torch.zeros(out_features, dtype=dtype)) + else: + self.register_parameter("bias", None) + + if method in _FEEDBACK_METHODS: + feedback = torch.empty( + out_features, in_features, dtype=dtype) + nn.init.normal_( + feedback, mean=0.0, std=init_std, + generator=feedback_generator) + if method == "fa": + self.register_buffer("feedback", feedback) + else: + self.feedback = nn.Parameter(feedback) + else: + self.register_buffer("feedback", None) + self.register_buffer("last_raw_rms", torch.zeros((), dtype=dtype)) + self.register_buffer( + "last_innovation_rms", torch.zeros((), dtype=dtype)) + self.register_buffer("last_traffic_rms", torch.zeros((), dtype=dtype)) + + def forward(self, x): + if self.method == "bp": + return F.linear(x, self.weight, self.bias) + method_code = {"fa": 0, "clean_kp": 1, "sdil": 2}[self.method] + return _FeedbackLinearFunction.apply( + x, self.weight, self.feedback, self.bias, method_code, + self.traffic_ratio, self.last_raw_rms, + self.last_innovation_rms, self.last_traffic_rms) + + def extra_repr(self): + return ( + f"in_features={self.in_features}, " + f"out_features={self.out_features}, method={self.method}") + + +@dataclass(frozen=True) +class LocalTransformerConfig: + vocab_size: int = 65 + context_length: int = 64 + depth: int = 4 + width: int = 128 + heads: int = 4 + mlp_ratio: int = 4 + dropout: float = 0.0 + bias: bool = False + init_std: float = 0.02 + traffic_ratio: float = 4.0 + seed: int = 2027 + + def __post_init__(self): + if self.width % self.heads: + raise ValueError("width must be divisible by heads") + if self.depth < 1 or self.context_length < 1: + raise ValueError("depth and context_length must be positive") + + +class LocalCausalSelfAttention(nn.Module): + + def __init__( + self, config: LocalTransformerConfig, method: str, + forward_generator: torch.Generator, + feedback_generator: torch.Generator, dtype=torch.float32): + super().__init__() + self.heads = config.heads + self.head_width = config.width // config.heads + self.width = config.width + common = dict( + method=method, forward_generator=forward_generator, + feedback_generator=feedback_generator, bias=config.bias, + init_std=config.init_std, traffic_ratio=config.traffic_ratio, + dtype=dtype) + self.q = FeedbackLinear(config.width, config.width, **common) + self.k = FeedbackLinear(config.width, config.width, **common) + self.v = FeedbackLinear(config.width, config.width, **common) + self.output = FeedbackLinear(config.width, config.width, **common) + causal = torch.tril(torch.ones( + config.context_length, config.context_length, dtype=torch.bool)) + self.register_buffer("causal_mask", causal, persistent=False) + self.dropout = float(config.dropout) + + def forward(self, x): + batch, time, width = x.shape + + def split_heads(value): + return value.view( + batch, time, self.heads, self.head_width).transpose(1, 2) + + query = split_heads(self.q(x)) + key = split_heads(self.k(x)) + value = split_heads(self.v(x)) + scores = query @ key.transpose(-2, -1) + scores = scores * self.head_width ** -0.5 + mask = self.causal_mask[:time, :time] + scores = scores.masked_fill(~mask, float("-inf")) + attention = F.softmax(scores, dim=-1) + attention = F.dropout( + attention, p=self.dropout, training=self.training) + mixed = attention @ value + mixed = mixed.transpose(1, 2).contiguous().view(batch, time, width) + return self.output(mixed) + + +class LocalTransformerBlock(nn.Module): + + def __init__( + self, config: LocalTransformerConfig, method: str, + forward_generator: torch.Generator, + feedback_generator: torch.Generator, dtype=torch.float32): + super().__init__() + self.ln_attention = nn.LayerNorm(config.width, dtype=dtype) + self.attention = LocalCausalSelfAttention( + config, method, forward_generator, feedback_generator, dtype) + self.ln_mlp = nn.LayerNorm(config.width, dtype=dtype) + hidden = config.mlp_ratio * config.width + common = dict( + method=method, forward_generator=forward_generator, + feedback_generator=feedback_generator, bias=config.bias, + init_std=config.init_std, traffic_ratio=config.traffic_ratio, + dtype=dtype) + self.mlp_in = FeedbackLinear(config.width, hidden, **common) + self.mlp_out = FeedbackLinear(hidden, config.width, **common) + self.dropout = float(config.dropout) + + def forward(self, x): + x = x + F.dropout( + self.attention(self.ln_attention(x)), + p=self.dropout, training=self.training) + x = x + F.dropout( + self.mlp_out(F.gelu(self.mlp_in(self.ln_mlp(x)))), + p=self.dropout, training=self.training) + return x + + +class LocalDecoderTransformer(nn.Module): + """Depth-scaled, forward-matched character decoder.""" + + def __init__( + self, config: LocalTransformerConfig, method: str = "bp", + dtype=torch.float32): + super().__init__() + if method not in {"bp"} | _FEEDBACK_METHODS: + raise ValueError(f"unsupported Transformer method: {method}") + self.config = config + self.method = method + forward_generator = torch.Generator().manual_seed(config.seed) + feedback_generator = torch.Generator().manual_seed(config.seed + 1) + self.token_embedding = nn.Embedding( + config.vocab_size, config.width, dtype=dtype) + self.position_embedding = nn.Parameter(torch.empty( + config.context_length, config.width, dtype=dtype)) + nn.init.normal_( + self.token_embedding.weight, mean=0.0, std=config.init_std, + generator=forward_generator) + nn.init.normal_( + self.position_embedding, mean=0.0, std=config.init_std, + generator=forward_generator) + self.blocks = nn.ModuleList([ + LocalTransformerBlock( + config, method, forward_generator, feedback_generator, dtype) + for _ in range(config.depth)]) + self.final_norm = nn.LayerNorm(config.width, dtype=dtype) + self.head = FeedbackLinear( + config.width, config.vocab_size, method, forward_generator, + feedback_generator, bias=False, init_std=config.init_std, + traffic_ratio=config.traffic_ratio, dtype=dtype) + + def forward(self, tokens, targets: 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 + hidden = F.dropout( + hidden, p=self.config.dropout, training=self.training) + for block in self.blocks: + hidden = block(hidden) + logits = self.head(self.final_norm(hidden)) + loss = None + if targets is not None: + loss = F.cross_entropy( + logits.reshape(-1, logits.shape[-1]), + targets.reshape(-1)) + return {"logits": logits, "loss": loss, "hidden": hidden} + + def feedback_linears(self) -> Iterable[FeedbackLinear]: + return ( + module for module in self.modules() + if isinstance(module, FeedbackLinear)) + + def forward_parameters(self) -> Iterable[nn.Parameter]: + feedback_ids = { + id(module.feedback) + for module in self.feedback_linears() + if isinstance(module.feedback, nn.Parameter)} + return ( + parameter for parameter in self.parameters() + if id(parameter) not in feedback_ids) + + @property + def n_forward_parameters(self): + return sum(parameter.numel() for parameter in self.forward_parameters()) + + @property + def n_feedback_parameters(self): + return sum( + module.feedback.numel() + for module in self.feedback_linears() + if module.feedback is not None) + + def teaching_statistics(self) -> Dict[str, float]: + modules = list(self.feedback_linears()) + if not modules: + return { + "raw_rms": 0.0, "innovation_rms": 0.0, + "traffic_rms": 0.0} + return { + name: float(torch.stack([ + getattr(module, f"last_{name}") + for module in modules]).mean()) + for name in ("raw_rms", "innovation_rms", "traffic_rms")} + + @torch.no_grad() + def set_feedback_equal_to_forward(self): + for module in self.feedback_linears(): + if module.feedback is None: + raise ValueError("BP modules do not contain feedback tensors") + module.feedback.copy_(module.weight) + -- cgit v1.2.3