summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--experiments/transformer_feedback_smoke.py92
-rw-r--r--sdil/transformer.py142
2 files changed, 226 insertions, 8 deletions
diff --git a/experiments/transformer_feedback_smoke.py b/experiments/transformer_feedback_smoke.py
index 99bca48..c3cd4da 100644
--- a/experiments/transformer_feedback_smoke.py
+++ b/experiments/transformer_feedback_smoke.py
@@ -154,6 +154,97 @@ def audit_sdil_innovation_identity():
}
+def audit_strict_dfa():
+ config = LocalTransformerConfig(
+ vocab_size=13, context_length=6, depth=2, width=8, heads=2,
+ mlp_ratio=2, seed=131)
+ bp = LocalDecoderTransformer(config, "bp", dtype=torch.float64)
+ dfa = LocalDecoderTransformer(config, "dfa", dtype=torch.float64)
+ generator = torch.Generator().manual_seed(132)
+ tokens = torch.randint(0, 13, (3, 6), generator=generator)
+ targets = torch.randint(0, 13, (3, 6), generator=generator)
+ with torch.no_grad():
+ bp_output = bp(tokens)["logits"]
+ cache = dfa(tokens, return_cache=True)
+ output_error = (
+ torch.softmax(cache["logits"], dim=-1)
+ - torch.nn.functional.one_hot(targets, 13).to(torch.float64)
+ ) / targets.numel()
+ forward_error = float(torch.max(torch.abs(
+ bp_output - cache["logits"])))
+ feedback_before = (
+ dfa.dfa_block_feedback.detach().clone(),
+ dfa.dfa_embedding_feedback.detach().clone(),
+ dfa.dfa_final_norm_feedback.detach().clone(),
+ )
+ metrics = dfa.dfa_gradients(
+ tokens, cache=cache, output_error=output_error)
+ gradients = {
+ name: parameter.grad.detach().clone()
+ for name, parameter in dfa.named_parameters()}
+ missing = [
+ name for name, parameter in dfa.named_parameters()
+ if parameter.grad is None]
+ nonfinite = [
+ name for name, gradient in gradients.items()
+ if not torch.isfinite(gradient).all()]
+
+ expected_head = (
+ output_error.reshape(-1, 13).t()
+ @ cache["normalized"].reshape(-1, 8))
+ head_error = _relative_error(dfa.head.weight.grad, expected_head)
+ expected_position = (
+ output_error @ dfa.dfa_embedding_feedback).sum(dim=0)
+ position_error = _relative_error(
+ dfa.position_embedding.grad[:tokens.shape[1]], expected_position)
+
+ # With the cached boundary activities and output error fixed, a later
+ # block cannot influence an earlier block's DFA update.
+ first_before = {
+ name: parameter.grad.detach().clone()
+ for name, parameter in dfa.blocks[0].named_parameters()}
+ with torch.no_grad():
+ for parameter in dfa.blocks[1].parameters():
+ parameter.add_(torch.randn(
+ parameter.shape, generator=generator,
+ dtype=parameter.dtype))
+ dfa.dfa_gradients(tokens, cache=cache, output_error=output_error)
+ first_after = {
+ name: parameter.grad.detach()
+ for name, parameter in dfa.blocks[0].named_parameters()}
+ cross_block_error = max(
+ _relative_error(first_after[name], first_before[name])
+ for name in first_before)
+ feedback_change = max(
+ float(torch.max(torch.abs(current - old)))
+ for current, old in zip((
+ dfa.dfa_block_feedback, dfa.dfa_embedding_feedback,
+ dfa.dfa_final_norm_feedback), feedback_before))
+ if (forward_error >= 1e-12 or head_error >= 1e-12
+ or position_error >= 1e-12 or cross_block_error >= 1e-12
+ or feedback_change != 0.0 or missing or nonfinite):
+ raise AssertionError({
+ "forward_error": forward_error,
+ "head_error": head_error,
+ "position_error": position_error,
+ "cross_block_error": cross_block_error,
+ "feedback_change": feedback_change,
+ "missing": missing,
+ "nonfinite": nonfinite,
+ })
+ return {
+ "matched_forward_max_absolute_error": forward_error,
+ "head_local_equation_relative_error": head_error,
+ "position_local_equation_relative_error": position_error,
+ "earlier_gradient_change_after_later_weight_mutation":
+ cross_block_error,
+ "fixed_feedback_change": feedback_change,
+ "forward_parameter_tensors_with_gradients": len(gradients),
+ "detached_block_boundaries": metrics["detached_block_boundaries"],
+ "uses_task_loss_backward": metrics["uses_task_loss_backward"],
+ }
+
+
def main():
torch.set_num_threads(1)
result = {
@@ -162,6 +253,7 @@ def main():
"fa_independence_and_kp_locality":
audit_fa_independence_and_kp_locality(),
"sdil_innovation_identity": audit_sdil_innovation_identity(),
+ "strict_dfa": audit_strict_dfa(),
}
print(json.dumps(result, indent=2, sort_keys=True))
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)
-