summaryrefslogtreecommitdiff
path: root/sdil
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-27 14:15:42 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-27 14:15:42 -0500
commit2f4e0aca2e4f9064d6548eca876cab433ad769fc (patch)
tree9e9e19e5e54aef58a98750c5cdda1eb03c64d385 /sdil
parent6a65b8b26e5cf9b7c6b4c429d5bc02ac0475b0b0 (diff)
baseline: add strict Transformer DFA
Diffstat (limited to 'sdil')
-rw-r--r--sdil/transformer.py142
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)
-