diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-27 14:14:25 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-27 14:14:25 -0500 |
| commit | 6a65b8b26e5cf9b7c6b4c429d5bc02ac0475b0b0 (patch) | |
| tree | 159bec05704a181766af49e05f1a17b095c71586 /experiments | |
| parent | 0b2f4d22bbb1fa934ae07f8e6d8f0470fff1f2fa (diff) | |
baseline: add matched Transformer feedback transport
Diffstat (limited to 'experiments')
| -rw-r--r-- | experiments/transformer_feedback_smoke.py | 170 |
1 files changed, 170 insertions, 0 deletions
diff --git a/experiments/transformer_feedback_smoke.py b/experiments/transformer_feedback_smoke.py new file mode 100644 index 0000000..99bca48 --- /dev/null +++ b/experiments/transformer_feedback_smoke.py @@ -0,0 +1,170 @@ +#!/usr/bin/env python3 +"""Deterministic audits for Transformer BP/FA/KP/SDIL feedback transport.""" +import json +import os +import sys + +import torch + +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +from sdil.transformer import ( # noqa: E402 + FeedbackLinear, + LocalDecoderTransformer, + LocalTransformerConfig, +) + + +def _named_forward_gradients(model): + feedback_ids = { + id(module.feedback) + for module in model.feedback_linears() + if isinstance(module.feedback, torch.nn.Parameter)} + return { + name: parameter.grad.detach().clone() + for name, parameter in model.named_parameters() + if id(parameter) not in feedback_ids} + + +def _relative_error(actual, expected): + return float( + torch.linalg.vector_norm(actual - expected) + / torch.linalg.vector_norm(expected).clamp_min(1e-30)) + + +def audit_matched_forward_and_symmetric_limit(): + config = LocalTransformerConfig( + vocab_size=11, context_length=7, depth=2, width=8, heads=2, + mlp_ratio=2, seed=101) + bp = LocalDecoderTransformer(config, "bp", dtype=torch.float64) + fa = LocalDecoderTransformer(config, "fa", dtype=torch.float64) + fa.set_feedback_equal_to_forward() + generator = torch.Generator().manual_seed(102) + tokens = torch.randint(0, 11, (3, 7), generator=generator) + targets = torch.randint(0, 11, (3, 7), generator=generator) + bp_output = bp(tokens, targets) + fa_output = fa(tokens, targets) + forward_error = float(torch.max(torch.abs( + bp_output["logits"] - fa_output["logits"])).detach()) + bp_output["loss"].backward() + fa_output["loss"].backward() + bp_gradients = _named_forward_gradients(bp) + fa_gradients = _named_forward_gradients(fa) + if bp_gradients.keys() != fa_gradients.keys(): + raise AssertionError("forward parameter names differ") + errors = { + name: _relative_error(fa_gradients[name], bp_gradients[name]) + for name in bp_gradients} + if forward_error >= 1e-12 or max(errors.values()) >= 2e-12: + raise AssertionError({ + "forward_error": forward_error, + "gradient_errors": errors, + }) + return { + "matched_forward_max_absolute_error": forward_error, + "symmetric_limit_max_relative_error": max(errors.values()), + "audited_forward_parameter_tensors": len(errors), + "matched_forward_parameter_count": bp.n_forward_parameters, + "matched_feedback_tensor_count": len(list(fa.feedback_linears())), + } + + +def audit_fa_independence_and_kp_locality(): + forward_generator = torch.Generator().manual_seed(111) + feedback_generator = torch.Generator().manual_seed(112) + fa = FeedbackLinear( + 5, 3, "fa", forward_generator, feedback_generator, + dtype=torch.float64) + feedback_before = fa.feedback.detach().clone() + with torch.no_grad(): + fa.weight.add_(torch.randn( + fa.weight.shape, generator=forward_generator, + dtype=fa.weight.dtype)) + fa_independence_error = float(torch.max(torch.abs( + fa.feedback - feedback_before))) + + forward_generator = torch.Generator().manual_seed(113) + feedback_generator = torch.Generator().manual_seed(114) + kp = FeedbackLinear( + 5, 3, "clean_kp", forward_generator, feedback_generator, + dtype=torch.float64) + x = torch.randn( + 2, 4, 5, generator=forward_generator, dtype=torch.float64, + requires_grad=True) + upstream = torch.randn( + 2, 4, 3, generator=forward_generator, dtype=torch.float64) + output = kp(x) + (output * upstream).sum().backward() + expected = upstream.reshape(-1, 3).t() @ x.detach().reshape(-1, 5) + forward_error = _relative_error(kp.weight.grad, expected) + feedback_error = _relative_error(kp.feedback.grad, expected) + if max(fa_independence_error, forward_error, feedback_error) >= 1e-12: + raise AssertionError({ + "fa_independence_error": fa_independence_error, + "forward_local_error": forward_error, + "feedback_local_error": feedback_error, + }) + return { + "fa_feedback_change_after_forward_weight_mutation": + fa_independence_error, + "kp_forward_local_correlation_relative_error": forward_error, + "kp_feedback_local_correlation_relative_error": feedback_error, + "kp_feedback_gradient_recomputed_locally": True, + } + + +def audit_sdil_innovation_identity(): + config = LocalTransformerConfig( + vocab_size=13, context_length=6, depth=2, width=8, heads=2, + mlp_ratio=2, traffic_ratio=4.0, seed=121) + kp = LocalDecoderTransformer(config, "clean_kp", dtype=torch.float64) + sdil = LocalDecoderTransformer(config, "sdil", dtype=torch.float64) + generator = torch.Generator().manual_seed(122) + tokens = torch.randint(0, 13, (3, 6), generator=generator) + targets = torch.randint(0, 13, (3, 6), generator=generator) + kp(tokens, targets)["loss"].backward() + sdil(tokens, targets)["loss"].backward() + kp_gradients = { + name: parameter.grad.detach() + for name, parameter in kp.named_parameters()} + sdil_gradients = { + name: parameter.grad.detach() + for name, parameter in sdil.named_parameters()} + errors = { + name: _relative_error(sdil_gradients[name], kp_gradients[name]) + for name in kp_gradients} + statistics = sdil.teaching_statistics() + traffic_ratio = ( + statistics["traffic_rms"] + / max(statistics["innovation_rms"], 1e-30)) + raw_is_contaminated = ( + statistics["raw_rms"] > statistics["innovation_rms"]) + if max(errors.values()) >= 2e-12 or not raw_is_contaminated: + raise AssertionError({ + "gradient_errors": errors, + "statistics": statistics, + "observed_traffic_ratio": traffic_ratio, + }) + return { + "kp_sdil_max_gradient_relative_error": max(errors.values()), + "mean_raw_rms": statistics["raw_rms"], + "mean_innovation_rms": statistics["innovation_rms"], + "mean_traffic_rms": statistics["traffic_rms"], + "raw_signal_contaminated": raw_is_contaminated, + "paired_neutral_subtraction": True, + } + + +def main(): + torch.set_num_threads(1) + result = { + "matched_forward_and_symmetric_limit": + audit_matched_forward_and_symmetric_limit(), + "fa_independence_and_kp_locality": + audit_fa_independence_and_kp_locality(), + "sdil_innovation_identity": audit_sdil_innovation_identity(), + } + print(json.dumps(result, indent=2, sort_keys=True)) + + +if __name__ == "__main__": + main() |
