diff options
Diffstat (limited to 'sdil/transformer.py')
| -rw-r--r-- | sdil/transformer.py | 142 |
1 files changed, 134 insertions, 8 deletions
diff --git a/sdil/transformer.py b/sdil/transformer.py index d6433c4..608f43f 100644 --- a/sdil/transformer.py +++ b/sdil/transformer.py @@ -15,6 +15,7 @@ import torch.nn.functional as F _FEEDBACK_METHODS = {"fa", "clean_kp", "sdil"} +_SUPPORTED_METHODS = {"bp", "dfa"} | _FEEDBACK_METHODS class _FeedbackLinearFunction(torch.autograd.Function): @@ -91,7 +92,7 @@ class FeedbackLinear(nn.Module): 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: + if method not in _SUPPORTED_METHODS: raise ValueError(f"unsupported feedback-linear method: {method}") self.in_features = int(in_features) self.out_features = int(out_features) @@ -125,7 +126,7 @@ class FeedbackLinear(nn.Module): self.register_buffer("last_traffic_rms", torch.zeros((), dtype=dtype)) def forward(self, x): - if self.method == "bp": + if self.method in {"bp", "dfa"}: return F.linear(x, self.weight, self.bias) method_code = {"fa": 0, "clean_kp": 1, "sdil": 2}[self.method] return _FeedbackLinearFunction.apply( @@ -244,7 +245,7 @@ class LocalDecoderTransformer(nn.Module): self, config: LocalTransformerConfig, method: str = "bp", dtype=torch.float32): super().__init__() - if method not in {"bp"} | _FEEDBACK_METHODS: + if method not in _SUPPORTED_METHODS: raise ValueError(f"unsupported Transformer method: {method}") self.config = config self.method = method @@ -269,8 +270,32 @@ class LocalDecoderTransformer(nn.Module): config.width, config.vocab_size, method, forward_generator, feedback_generator, bias=False, init_std=config.init_std, traffic_ratio=config.traffic_ratio, dtype=dtype) + if method == "dfa": + dfa_generator = torch.Generator().manual_seed(config.seed + 2) + feedback_scale = config.init_std + self.register_buffer("dfa_block_feedback", torch.empty( + config.depth, config.vocab_size, config.width, dtype=dtype)) + self.register_buffer("dfa_embedding_feedback", torch.empty( + config.vocab_size, config.width, dtype=dtype)) + self.register_buffer("dfa_final_norm_feedback", torch.empty( + config.vocab_size, config.width, dtype=dtype)) + nn.init.normal_( + self.dfa_block_feedback, mean=0.0, std=feedback_scale, + generator=dfa_generator) + nn.init.normal_( + self.dfa_embedding_feedback, mean=0.0, std=feedback_scale, + generator=dfa_generator) + nn.init.normal_( + self.dfa_final_norm_feedback, mean=0.0, std=feedback_scale, + generator=dfa_generator) + else: + self.register_buffer("dfa_block_feedback", None) + self.register_buffer("dfa_embedding_feedback", None) + self.register_buffer("dfa_final_norm_feedback", None) - def forward(self, tokens, targets: Optional[torch.Tensor] = None): + def forward( + self, tokens, targets: Optional[torch.Tensor] = None, + return_cache: bool = False): if tokens.ndim != 2: raise ValueError("tokens must have shape (batch, time)") if tokens.shape[1] > self.config.context_length: @@ -279,15 +304,111 @@ class LocalDecoderTransformer(nn.Module): hidden = self.token_embedding(tokens) + positions hidden = F.dropout( hidden, p=self.config.dropout, training=self.training) + block_inputs = [] for block in self.blocks: + if return_cache: + block_inputs.append(hidden.detach()) hidden = block(hidden) - logits = self.head(self.final_norm(hidden)) + final_input = hidden + normalized = self.final_norm(hidden) + logits = self.head(normalized) 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} + result = {"logits": logits, "loss": loss, "hidden": hidden} + if return_cache: + result["block_inputs"] = block_inputs + result["final_input"] = final_input.detach() + result["normalized"] = normalized.detach() + return result + + @staticmethod + def _assign_gradients(parameters, gradients): + for parameter, gradient in zip(parameters, gradients): + if parameter.grad is None: + parameter.grad = gradient.detach().clone() + else: + parameter.grad.copy_(gradient.detach()) + + def dfa_gradients( + self, tokens, targets: Optional[torch.Tensor] = None, + cache: Optional[Dict[str, torch.Tensor]] = None, + output_error: Optional[torch.Tensor] = None): + """Populate strict block-DFA gradients without a task-loss backward. + + Every decoder block receives a separate fixed projection of the + analytical output error. Inputs are detached at block boundaries; + autograd is used only for each explicitly local block objective. + Embeddings and the final normalization receive their own fixed direct + projections, while the vocabulary head uses its exact local delta. + """ + if self.method != "dfa": + raise ValueError("dfa_gradients is only valid for method='dfa'") + if cache is None: + with torch.no_grad(): + cache = self.forward(tokens, return_cache=True) + if output_error is None: + if targets is None: + raise ValueError("targets or output_error must be provided") + probabilities = torch.softmax(cache["logits"], dim=-1) + one_hot = F.one_hot( + targets, self.config.vocab_size).to(probabilities.dtype) + output_error = ( + probabilities - one_hot) / targets.numel() + output_error = output_error.detach() + for parameter in self.parameters(): + parameter.grad = None + + error_flat = output_error.reshape(-1, self.config.vocab_size) + normalized_flat = cache["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 = cache["final_input"].detach() + normalized = self.final_norm(final_input) + final_field = output_error @ self.dfa_final_norm_feedback + final_parameters = tuple(self.final_norm.parameters()) + final_objective = torch.sum(normalized * final_field) + final_gradients = torch.autograd.grad( + final_objective, final_parameters) + self._assign_gradients(final_parameters, final_gradients) + + for index, (block, block_input) in enumerate(zip( + self.blocks, cache["block_inputs"])): + local_input = block_input.detach() + local_output = block(local_input) + local_field = output_error @ self.dfa_block_feedback[index] + local_parameters = tuple(block.parameters()) + local_objective = torch.sum(local_output * local_field) + local_gradients = torch.autograd.grad( + local_objective, local_parameters) + self._assign_gradients(local_parameters, local_gradients) + + embedding_field = output_error @ self.dfa_embedding_feedback + token_activity = self.token_embedding(tokens) + position_activity = self.position_embedding[:tokens.shape[1]] + embedding_objective = ( + torch.sum(token_activity * embedding_field) + + torch.sum(position_activity * embedding_field.sum(dim=0))) + 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 { + "output_error_rms": float( + output_error.square().mean().sqrt()), + "block_field_rms": [ + float((output_error @ feedback).square().mean().sqrt()) + for feedback in self.dfa_block_feedback], + "uses_task_loss_backward": False, + "detached_block_boundaries": len(self.blocks), + } def feedback_linears(self) -> Iterable[FeedbackLinear]: return ( @@ -309,10 +430,16 @@ class LocalDecoderTransformer(nn.Module): @property def n_feedback_parameters(self): - return sum( + affine_feedback = sum( module.feedback.numel() for module in self.feedback_linears() if module.feedback is not None) + dfa_feedback = sum( + tensor.numel() for tensor in ( + self.dfa_block_feedback, self.dfa_embedding_feedback, + self.dfa_final_norm_feedback) + if tensor is not None) + return affine_feedback + dfa_feedback def teaching_statistics(self) -> Dict[str, float]: modules = list(self.feedback_linears()) @@ -332,4 +459,3 @@ class LocalDecoderTransformer(nn.Module): if module.feedback is None: raise ValueError("BP modules do not contain feedback tensors") module.feedback.copy_(module.weight) - |
