diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-27 14:20:53 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-27 14:20:53 -0500 |
| commit | e1e2c2eb9011eb75ff2108b06c9ac3fefa596767 (patch) | |
| tree | fd3eec6a1bb5d8d6206a311476fb47e1d0067a98 /experiments | |
| parent | 672423d13e36d7a63fa89aa186dc35fc91475ab1 (diff) | |
baseline: add Transformer Dual Prop and EP
Diffstat (limited to 'experiments')
| -rw-r--r-- | experiments/transformer_feedback_smoke.py | 116 |
1 files changed, 116 insertions, 0 deletions
diff --git a/experiments/transformer_feedback_smoke.py b/experiments/transformer_feedback_smoke.py index 2cd22d9..ff34369 100644 --- a/experiments/transformer_feedback_smoke.py +++ b/experiments/transformer_feedback_smoke.py @@ -395,6 +395,120 @@ def audit_forward_forward(): } +def audit_dualprop(): + config = LocalTransformerConfig( + vocab_size=7, context_length=4, depth=2, width=8, heads=2, + mlp_ratio=2, seed=161) + bp = LocalDecoderTransformer(config, "bp", dtype=torch.float64) + dp = LocalDecoderTransformer(config, "dualprop", dtype=torch.float64) + generator = torch.Generator().manual_seed(162) + tokens = torch.randint(0, 7, (2, 4), generator=generator) + targets = torch.randint(0, 7, (2, 4), generator=generator) + with torch.no_grad(): + forward_error = float(torch.max(torch.abs( + bp(tokens)["logits"] - dp(tokens)["logits"]))) + plus, minus = dp.dualprop_states( + tokens, targets, alpha=0.0, beta=0.1, inference_passes=2) + metrics = dp.dualprop_gradients( + tokens, targets, alpha=0.0, beta=0.1, inference_passes=2) + actual = [ + parameter.grad.detach().clone() + for parameter in dp.blocks[0].parameters()] + missing = [ + name for name, parameter in dp.named_parameters() + if parameter.grad is None] + nonfinite = [ + name for name, parameter in dp.named_parameters() + if parameter.grad is not None + and not torch.isfinite(parameter.grad).all()] + states_finite = all( + torch.isfinite(state).all() for state in plus + minus) + + # Audit one complete block direction against its explicit detached local + # contrastive objective. + states = [negative.detach() for negative in minus[:-1]] + delta = ((plus[1] - minus[1]) / 0.1).detach() + parameters = tuple(dp.blocks[0].parameters()) + objective = -torch.sum( + dp.blocks[0](states[0]) * delta) / targets.numel() + expected = torch.autograd.grad(objective, parameters) + equation_error = max( + _relative_error(a, e) for a, e in zip(actual, expected)) + if (forward_error >= 1e-12 or equation_error >= 1e-12 + or missing or nonfinite or not states_finite): + raise AssertionError({ + "forward_error": forward_error, + "equation_error": equation_error, + "missing": missing, + "nonfinite": nonfinite, + "states_finite": states_finite, + }) + return { + "matched_forward_max_absolute_error": forward_error, + "local_contrastive_equation_max_relative_error": equation_error, + "inferred_state_tensors": len(plus) + len(minus), + "finite_inferred_states": states_finite, + "local_vjp_evaluations": metrics["local_vjp_evaluations"], + "uses_symmetric_block_vjps": True, + "uses_task_loss_backward": metrics["uses_task_loss_backward"], + } + + +def audit_equilibrium_propagation(): + config = LocalTransformerConfig( + vocab_size=7, context_length=4, depth=2, width=8, heads=2, + mlp_ratio=2, seed=171) + bp = LocalDecoderTransformer(config, "bp", dtype=torch.float64) + ep = LocalDecoderTransformer(config, "ep", dtype=torch.float64) + generator = torch.Generator().manual_seed(172) + tokens = torch.randint(0, 7, (2, 4), generator=generator) + targets = torch.randint(0, 7, (2, 4), generator=generator) + with torch.no_grad(): + forward_error = float(torch.max(torch.abs( + bp(tokens)["logits"] - ep(tokens)["logits"]))) + metrics, free, nudged = ep.ep_gradients( + tokens, targets, ep_beta=0.5, dt=0.5, + free_steps=3, nudge_steps=2, beta_sign=1.0) + missing = [ + name for name, parameter in ep.named_parameters() + if parameter.grad is None] + nonfinite = [ + name for name, parameter in ep.named_parameters() + if parameter.grad is not None + and not torch.isfinite(parameter.grad).all()] + states = free + nudged + lower = min(float(state.min()) for state in states) + upper = max(float(state.max()) for state in states) + + free_groups = ep._ep_phase_gradients(tokens, free) + nudged_groups = ep._ep_phase_gradients(tokens, nudged) + expected = [ + -(nudged_gradient - free_gradient) / 0.5 + for nudged_gradient, free_gradient in zip( + nudged_groups[1][1], free_groups[1][1])] + actual = [ + parameter.grad for parameter in free_groups[1][0]] + equation_error = max( + _relative_error(a, e) for a, e in zip(actual, expected)) + if (forward_error >= 1e-12 or equation_error >= 1e-12 + or missing or nonfinite or lower < 0.0 or upper > 1.0): + raise AssertionError({ + "forward_error": forward_error, + "equation_error": equation_error, + "missing": missing, + "nonfinite": nonfinite, + "state_range": [lower, upper], + }) + return { + "matched_forward_max_absolute_error": forward_error, + "phase_contrast_equation_max_relative_error": equation_error, + "state_range": [lower, upper], + "local_vjp_evaluations": metrics["local_vjp_evaluations"], + "uses_symmetric_block_vjps": True, + "uses_task_loss_backward": metrics["uses_task_loss_backward"], + } + + def main(): torch.set_num_threads(1) result = { @@ -406,6 +520,8 @@ def main(): "strict_dfa": audit_strict_dfa(), "pepita": audit_pepita(), "forward_forward": audit_forward_forward(), + "dual_propagation": audit_dualprop(), + "equilibrium_propagation": audit_equilibrium_propagation(), } print(json.dumps(result, indent=2, sort_keys=True)) |
