summaryrefslogtreecommitdiff
path: root/sdil/transformer.py
diff options
context:
space:
mode:
Diffstat (limited to 'sdil/transformer.py')
-rw-r--r--sdil/transformer.py127
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)