summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--experiments/transformer_feedback_smoke.py170
-rw-r--r--sdil/__init__.py2
-rw-r--r--sdil/transformer.py335
3 files changed, 506 insertions, 1 deletions
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)
+