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