summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-27 14:41:52 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-27 14:41:52 -0500
commit711fc1485f5cae3fa0d0bfad2dfaf129b214b1f8 (patch)
tree29879c7fe79fe225efa0488adfca8e95666880a6 /experiments
parent08554657946e54d8d908a2866e82504e9b194617 (diff)
test: audit all Transformer crossover cells
Diffstat (limited to 'experiments')
-rw-r--r--experiments/transformer_feedback_smoke.py56
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":