diff options
Diffstat (limited to 'sdil/transformer.py')
| -rw-r--r-- | sdil/transformer.py | 335 |
1 files changed, 335 insertions, 0 deletions
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) + |
