diff options
Diffstat (limited to 'sdil/transformer.py')
| -rw-r--r-- | sdil/transformer.py | 127 |
1 files changed, 125 insertions, 2 deletions
diff --git a/sdil/transformer.py b/sdil/transformer.py index ee1e69f..d47fd18 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", "pepita"} | _FEEDBACK_METHODS +_SUPPORTED_METHODS = {"bp", "dfa", "pepita", "ff"} | _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", "pepita"}: + if self.method in {"bp", "dfa", "pepita", "ff"}: return F.linear(x, self.weight, self.bias) method_code = {"fa": 0, "clean_kp": 1, "sdil": 2}[self.method] return _FeedbackLinearFunction.apply( @@ -573,3 +573,126 @@ class LocalDecoderTransformer(nn.Module): if module.feedback is None: raise ValueError("BP modules do not contain feedback tensors") module.feedback.copy_(module.weight) + + @staticmethod + def _ff_normalize(value): + return value / ( + value.square().sum(dim=-1, keepdim=True).sqrt() + 1e-8) + + def _ff_overlay(self, embedded, candidate_labels): + if candidate_labels.shape != embedded.shape[:2]: + raise ValueError("candidate labels must match token positions") + if self.config.width < self.config.vocab_size: + raise ValueError( + "Forward-Forward overlay requires width >= vocabulary") + overlaid = embedded.clone() + overlaid[..., :self.config.vocab_size] = 0.0 + magnitude = embedded.detach().abs().amax( + dim=-1, keepdim=True).clamp_min(1e-4) + overlaid.scatter_( + -1, candidate_labels.unsqueeze(-1), magnitude) + return overlaid + + def ff_forward(self, tokens, candidate_labels): + """Run candidate-labelled, normalized data through the FF graph.""" + if self.method != "ff": + raise ValueError("ff_forward is only valid for method='ff'") + positions = self.position_embedding[:tokens.shape[1]] + base = self.token_embedding(tokens) + positions + embedded = self._ff_overlay(base, candidate_labels) + hidden = embedded + block_inputs = [] + block_outputs = [] + for block in self.blocks: + local_input = self._ff_normalize(hidden) + block_inputs.append(local_input) + hidden = block(local_input) + block_outputs.append(hidden) + final_input = self._ff_normalize(hidden) + logits = self.head(self.final_norm(final_input)) + return { + "embedded": embedded, + "block_inputs": block_inputs, + "block_outputs": block_outputs, + "final_input": final_input, + "logits": logits, + } + + @property + def ff_num_layers(self): + return self.config.depth + 2 + + def ff_layer_parameters(self, layer_index): + if layer_index == 0: + return (self.token_embedding.weight, self.position_embedding) + if 1 <= layer_index <= self.config.depth: + return tuple(self.blocks[layer_index - 1].parameters()) + if layer_index == self.config.depth + 1: + return ( + *tuple(self.final_norm.parameters()), + *tuple(self.head.parameters())) + raise ValueError("invalid Forward-Forward layer index") + + def _ff_local_outputs(self, layer_index, tokens, labels, cached): + if layer_index == 0: + base = ( + self.token_embedding(tokens) + + self.position_embedding[:tokens.shape[1]]) + return self._ff_overlay(base, labels) + if 1 <= layer_index <= self.config.depth: + local_input = cached["block_inputs"][ + layer_index - 1].detach() + return self.blocks[layer_index - 1](local_input) + if layer_index == self.config.depth + 1: + local_input = cached["final_input"].detach() + return self.head(self.final_norm(local_input)) + raise ValueError("invalid Forward-Forward layer index") + + def ff_local_loss( + self, layer_index, tokens, positive_labels, negative_labels, + threshold=2.0): + """Return one greedy FF objective with every prefix detached.""" + if self.method != "ff": + raise ValueError("ff_local_loss is only valid for method='ff'") + if not 0 <= layer_index < self.ff_num_layers: + raise ValueError("invalid Forward-Forward layer index") + with torch.no_grad(): + positive_cache = self.ff_forward(tokens, positive_labels) + negative_cache = self.ff_forward(tokens, negative_labels) + positive_output = self._ff_local_outputs( + layer_index, tokens, positive_labels, positive_cache) + negative_output = self._ff_local_outputs( + layer_index, tokens, negative_labels, negative_cache) + positive_goodness = positive_output.square().mean(dim=-1) + negative_goodness = negative_output.square().mean(dim=-1) + loss = ( + F.softplus(-positive_goodness + threshold) + + F.softplus(negative_goodness - threshold) + ).mean() + return loss, { + "positive_goodness": float( + positive_goodness.mean().detach()), + "negative_goodness": float( + negative_goodness.mean().detach()), + "pair_accuracy": float( + (positive_goodness > negative_goodness).float().mean()), + } + + @torch.no_grad() + def ff_candidate_scores(self, tokens, score_from_layer=1): + """Score every next-token candidate; cost is charged as V passes.""" + if not 0 <= score_from_layer < self.ff_num_layers: + raise ValueError("invalid Forward-Forward score range") + scores = [] + for candidate in range(self.config.vocab_size): + labels = torch.full_like(tokens, candidate) + candidate_forward = self.ff_forward(tokens, labels) + layer_outputs = ( + [candidate_forward["embedded"]] + + candidate_forward["block_outputs"] + + [candidate_forward["logits"]]) + goodness = [ + output.square().mean(dim=-1) + for output in layer_outputs] + scores.append(sum(goodness[score_from_layer:])) + return torch.stack(scores, dim=-1) |
