summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-27 14:20:53 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-27 14:20:53 -0500
commite1e2c2eb9011eb75ff2108b06c9ac3fefa596767 (patch)
treefd3eec6a1bb5d8d6206a311476fb47e1d0067a98 /experiments
parent672423d13e36d7a63fa89aa186dc35fc91475ab1 (diff)
baseline: add Transformer Dual Prop and EP
Diffstat (limited to 'experiments')
-rw-r--r--experiments/transformer_feedback_smoke.py116
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))