diff options
Diffstat (limited to 'experiments/transformer_feedback_smoke.py')
| -rw-r--r-- | experiments/transformer_feedback_smoke.py | 92 |
1 files changed, 92 insertions, 0 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)) |
