#!/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 audit_strict_dfa(): config = LocalTransformerConfig( vocab_size=13, context_length=6, depth=2, width=8, heads=2, mlp_ratio=2, seed=131) bp = LocalDecoderTransformer(config, "bp", dtype=torch.float64) dfa = LocalDecoderTransformer(config, "dfa", dtype=torch.float64) generator = torch.Generator().manual_seed(132) tokens = torch.randint(0, 13, (3, 6), generator=generator) targets = torch.randint(0, 13, (3, 6), generator=generator) with torch.no_grad(): bp_output = bp(tokens)["logits"] cache = dfa(tokens, return_cache=True) output_error = ( torch.softmax(cache["logits"], dim=-1) - torch.nn.functional.one_hot(targets, 13).to(torch.float64) ) / targets.numel() forward_error = float(torch.max(torch.abs( bp_output - cache["logits"]))) feedback_before = ( dfa.dfa_block_feedback.detach().clone(), dfa.dfa_embedding_feedback.detach().clone(), dfa.dfa_final_norm_feedback.detach().clone(), ) metrics = dfa.dfa_gradients( tokens, cache=cache, output_error=output_error) gradients = { name: parameter.grad.detach().clone() for name, parameter in dfa.named_parameters()} missing = [ name for name, parameter in dfa.named_parameters() if parameter.grad is None] nonfinite = [ name for name, gradient in gradients.items() if not torch.isfinite(gradient).all()] expected_head = ( output_error.reshape(-1, 13).t() @ cache["normalized"].reshape(-1, 8)) head_error = _relative_error(dfa.head.weight.grad, expected_head) expected_position = ( output_error @ dfa.dfa_embedding_feedback).sum(dim=0) position_error = _relative_error( dfa.position_embedding.grad[:tokens.shape[1]], expected_position) # With the cached boundary activities and output error fixed, a later # block cannot influence an earlier block's DFA update. first_before = { name: parameter.grad.detach().clone() for name, parameter in dfa.blocks[0].named_parameters()} with torch.no_grad(): for parameter in dfa.blocks[1].parameters(): parameter.add_(torch.randn( parameter.shape, generator=generator, dtype=parameter.dtype)) dfa.dfa_gradients(tokens, cache=cache, output_error=output_error) first_after = { name: parameter.grad.detach() for name, parameter in dfa.blocks[0].named_parameters()} cross_block_error = max( _relative_error(first_after[name], first_before[name]) for name in first_before) feedback_change = max( float(torch.max(torch.abs(current - old))) for current, old in zip(( dfa.dfa_block_feedback, dfa.dfa_embedding_feedback, dfa.dfa_final_norm_feedback), feedback_before)) if (forward_error >= 1e-12 or head_error >= 1e-12 or position_error >= 1e-12 or cross_block_error >= 1e-12 or feedback_change != 0.0 or missing or nonfinite): raise AssertionError({ "forward_error": forward_error, "head_error": head_error, "position_error": position_error, "cross_block_error": cross_block_error, "feedback_change": feedback_change, "missing": missing, "nonfinite": nonfinite, }) return { "matched_forward_max_absolute_error": forward_error, "head_local_equation_relative_error": head_error, "position_local_equation_relative_error": position_error, "earlier_gradient_change_after_later_weight_mutation": cross_block_error, "fixed_feedback_change": feedback_change, "forward_parameter_tensors_with_gradients": len(gradients), "detached_block_boundaries": metrics["detached_block_boundaries"], "uses_task_loss_backward": metrics["uses_task_loss_backward"], } 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(), "strict_dfa": audit_strict_dfa(), } print(json.dumps(result, indent=2, sort_keys=True)) if __name__ == "__main__": main()