summaryrefslogtreecommitdiff
path: root/experiments/transformer_feedback_smoke.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/transformer_feedback_smoke.py')
-rw-r--r--experiments/transformer_feedback_smoke.py170
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()