From 2f4e0aca2e4f9064d6548eca876cab433ad769fc Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Mon, 27 Jul 2026 14:15:42 -0500 Subject: baseline: add strict Transformer DFA --- experiments/transformer_feedback_smoke.py | 92 +++++++++++++++++++ sdil/transformer.py | 142 ++++++++++++++++++++++++++++-- 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) - -- cgit v1.2.3