#!/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()