diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-27 14:41:52 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-27 14:41:52 -0500 |
| commit | 711fc1485f5cae3fa0d0bfad2dfaf129b214b1f8 (patch) | |
| tree | 29879c7fe79fe225efa0488adfca8e95666880a6 | |
| parent | 08554657946e54d8d908a2866e82504e9b194617 (diff) | |
test: audit all Transformer crossover cells
| -rw-r--r-- | experiments/transformer_feedback_smoke.py | 56 |
1 files changed, 56 insertions, 0 deletions
diff --git a/experiments/transformer_feedback_smoke.py b/experiments/transformer_feedback_smoke.py index ff34369..546fca4 100644 --- a/experiments/transformer_feedback_smoke.py +++ b/experiments/transformer_feedback_smoke.py @@ -15,6 +15,12 @@ from sdil.transformer import ( # noqa: E402 ) +METHODS = ( + "bp", "fa", "dfa", "pepita", "ff", "ep", "dualprop", + "clean_kp", "sdil", +) + + def _named_forward_gradients(model): feedback_ids = { id(module.feedback) @@ -69,6 +75,54 @@ def audit_matched_forward_and_symmetric_limit(): } +def audit_formal_registry_forward_identity(): + expected_counts = { + 4: 813568, + 8: 1602048, + 12: 2390528, + } + tokens = torch.tensor([[0, 1, 2]], dtype=torch.long) + rows = {} + for depth, expected_count in expected_counts.items(): + config = LocalTransformerConfig(depth=depth) + reference = LocalDecoderTransformer(config, "bp") + with torch.no_grad(): + reference_logits = reference(tokens)["logits"] + method_errors = {} + method_counts = {} + for method in METHODS: + model = LocalDecoderTransformer(config, method) + with torch.no_grad(): + logits = model(tokens)["logits"] + method_errors[method] = float(torch.max(torch.abs( + logits - reference_logits))) + method_counts[method] = model.n_forward_parameters + if model.n_forward_parameters != expected_count: + raise AssertionError({ + "depth": depth, + "method": method, + "expected_count": expected_count, + "observed_count": model.n_forward_parameters, + }) + del model + maximum_error = max(method_errors.values()) + if maximum_error != 0.0: + raise AssertionError({ + "depth": depth, + "method_errors": method_errors, + }) + rows[str(depth)] = { + "forward_parameter_count": expected_count, + "methods": len(method_counts), + "max_initial_logit_absolute_error": maximum_error, + } + return { + "registered_architecture_method_cells": + len(expected_counts) * len(METHODS), + "depths": rows, + } + + def audit_fa_independence_and_kp_locality(): forward_generator = torch.Generator().manual_seed(111) feedback_generator = torch.Generator().manual_seed(112) @@ -512,6 +566,8 @@ def audit_equilibrium_propagation(): def main(): torch.set_num_threads(1) result = { + "formal_registry_forward_identity": + audit_formal_registry_forward_identity(), "matched_forward_and_symmetric_limit": audit_matched_forward_and_symmetric_limit(), "fa_independence_and_kp_locality": |
