#!/usr/bin/env python3 """Deterministic audits for Transformer BP/FA/KP/SDIL feedback transport.""" import json import math 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 audit_pepita(): config = LocalTransformerConfig( vocab_size=13, context_length=6, depth=2, width=8, heads=2, mlp_ratio=2, seed=141) bp = LocalDecoderTransformer(config, "bp", dtype=torch.float64) pepita = LocalDecoderTransformer(config, "pepita", dtype=torch.float64) generator = torch.Generator().manual_seed(142) 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"] clean = pepita(tokens, return_cache=True) one_hot = torch.nn.functional.one_hot( targets, 13).to(torch.float64) clean_error = torch.softmax(clean["logits"], dim=-1) - one_hot offset = clean_error @ pepita.pepita_input_feedback modulated = pepita( tokens, return_cache=True, embedding_offset=offset) modulated_error = ( torch.softmax(modulated["logits"], dim=-1) - one_hot) forward_error = float(torch.max(torch.abs( bp_output - clean["logits"]))) feedback_before = pepita.pepita_input_feedback.detach().clone() metrics = pepita.pepita_gradients(tokens, targets) missing = [ name for name, parameter in pepita.named_parameters() if parameter.grad is None] nonfinite = [ name for name, parameter in pepita.named_parameters() if parameter.grad is not None and not torch.isfinite(parameter.grad).all()] observations = targets.numel() expected_head = ( (modulated_error / observations).reshape(-1, 13).t() @ modulated["normalized"].reshape(-1, 8)) head_error = _relative_error( pepita.head.weight.grad, expected_head) embedding_field = clean["embedded"] - modulated["embedded"] expected_position = embedding_field.sum(dim=0) / observations position_error = _relative_error( pepita.position_embedding.grad[:tokens.shape[1]], expected_position) feedback_change = float(torch.max(torch.abs( pepita.pepita_input_feedback - feedback_before))) projection_limit = math.sqrt(6.0 / config.width) * 0.05 observed_limit = float(torch.max(torch.abs( pepita.pepita_input_feedback))) if (forward_error >= 1e-12 or head_error >= 1e-12 or position_error >= 1e-12 or feedback_change != 0.0 or observed_limit > projection_limit or missing or nonfinite): raise AssertionError({ "forward_error": forward_error, "head_error": head_error, "position_error": position_error, "feedback_change": feedback_change, "projection_limit": projection_limit, "observed_limit": observed_limit, "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, "fixed_projection_change": feedback_change, "projection_limit": projection_limit, "observed_projection_max": observed_limit, "training_presentations": metrics["training_presentations"], "detached_block_boundaries": metrics["detached_block_boundaries"], "uses_task_loss_backward": metrics["uses_task_loss_backward"], } def audit_forward_forward(): config = LocalTransformerConfig( vocab_size=7, context_length=5, depth=2, width=8, heads=2, mlp_ratio=2, seed=151) bp = LocalDecoderTransformer(config, "bp", dtype=torch.float64) ff = LocalDecoderTransformer(config, "ff", dtype=torch.float64) generator = torch.Generator().manual_seed(152) tokens = torch.randint(0, 7, (3, 5), generator=generator) positive = torch.randint(0, 7, (3, 5), generator=generator) negative = (positive + torch.randint( 1, 7, positive.shape, generator=generator)) % 7 with torch.no_grad(): forward_error = float(torch.max(torch.abs( bp(tokens)["logits"] - ff(tokens)["logits"]))) all_parameters = tuple(ff.parameters()) audited = [] for layer_index in range(ff.ff_num_layers): for parameter in all_parameters: parameter.grad = None target_ids = { id(parameter) for parameter in ff.ff_layer_parameters(layer_index)} loss, metrics = ff.ff_local_loss( layer_index, tokens, positive, negative) loss.backward() unexpected = [ index for index, parameter in enumerate(all_parameters) if parameter.grad is not None and id(parameter) not in target_ids] missing = [ index for index, parameter in enumerate(all_parameters) if parameter.grad is None and id(parameter) in target_ids] explicit_positive = ff._ff_local_outputs( layer_index, tokens, positive, ff.ff_forward(tokens, positive)) explicit_negative = ff._ff_local_outputs( layer_index, tokens, negative, ff.ff_forward(tokens, negative)) explicit = ( torch.nn.functional.softplus( -explicit_positive.square().mean(dim=-1) + 2.0) + torch.nn.functional.softplus( explicit_negative.square().mean(dim=-1) - 2.0) ).mean() loss_error = abs(float(loss.detach() - explicit.detach())) if unexpected or missing or loss_error >= 1e-12: raise AssertionError({ "layer": layer_index, "unexpected": unexpected, "missing": missing, "loss_error": loss_error, }) audited.append({ "layer": layer_index, "loss": float(loss.detach()), "loss_error": loss_error, **metrics, }) scores = ff.ff_candidate_scores(tokens) if (forward_error >= 1e-12 or scores.shape != (3, 5, 7) or not torch.isfinite(scores).all()): raise AssertionError({ "forward_error": forward_error, "score_shape": list(scores.shape), }) return { "matched_forward_max_absolute_error": forward_error, "audited_greedy_layers": len(audited), "max_local_objective_absolute_error": max( row["loss_error"] for row in audited), "non_target_parameter_gradients": 0, "candidate_score_shape": list(scores.shape), "candidate_presentations_per_evaluation": config.vocab_size, } 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(), "pepita": audit_pepita(), "forward_forward": audit_forward_forward(), } print(json.dumps(result, indent=2, sort_keys=True)) if __name__ == "__main__": main()