diff options
Diffstat (limited to 'experiments/conv_local_smoke.py')
| -rw-r--r-- | experiments/conv_local_smoke.py | 84 |
1 files changed, 84 insertions, 0 deletions
diff --git a/experiments/conv_local_smoke.py b/experiments/conv_local_smoke.py index 9207e04..894f9d3 100644 --- a/experiments/conv_local_smoke.py +++ b/experiments/conv_local_smoke.py @@ -1,5 +1,6 @@ #!/usr/bin/env python3 """Prove the convolutional local eligibility matches exact BP when instructed.""" +import math import os import sys @@ -11,6 +12,7 @@ from sdil.conv import (CIFARHierarchicalFAResNet, CIFARLocalResNet, CIFARSDILResNet, ConvSDILConfig, channel_subspace_apical_calibration, conv_hierarchical_step, conv_local_step, + hierarchical_parameter_subspace_calibration, simultaneous_conv_node_perturbation, vectorizer_subspace_apical_calibration) @@ -673,6 +675,87 @@ def hierarchical_feedback_checks(): } +def hierarchical_parameter_calibration_checks(): + """Audit the causal JVP and the exact local Q/R delta-rule moments.""" + torch.manual_seed(109) + net = CIFARHierarchicalFAResNet( + depth=8, base_width=2, seed=110, dtype=torch.float64, + normalization="batchnorm", residual_scale=1.0) + x = torch.randn(3, 3, 32, 32, dtype=torch.float64) + y = torch.tensor([0, 4, 7]) + parameters = net.W + net.gamma + net.beta + [net.W_out, net.b_out] + for parameter in parameters: + parameter.requires_grad_(True) + forward = net.forward( + x, return_cache=True, training=True, update_stats=False) + loss = F.cross_entropy(forward["logits"], y) + hidden_gradients = torch.autograd.grad(loss, forward["hiddens"]) + output_signal = (torch.softmax(forward["logits"].detach(), dim=1) + - F.one_hot(y, 10).to(torch.float64)) + _, diagnostic = hierarchical_parameter_subspace_calibration( + net, x, y, forward, output_signal, sigma=1e-5, + n_directions=1, eta=0.0, + generator=torch.Generator().manual_seed(111), + return_diagnostics=True) + directions = diagnostic["directions"][0]["hidden"] + exact_directional = x.shape[0] * sum( + (gradient * direction).sum() + for gradient, direction in zip(hidden_gradients, directions)) + estimated_directional = diagnostic[ + "directional_derivatives"][0]["scaled_directional"] + jvp_relative = float((estimated_directional - exact_directional).abs() + / exact_directional.abs().clamp_min(1e-30)) + assert jvp_relative < 3e-7 + for parameter in parameters: + parameter.requires_grad_(False) + + # Under an audit-only symmetric copy, the hierarchical field is the exact + # negative gradient. Consequently every local Q/R predicted moment equals + # its exact causal regression target, including option-A shortcut terms. + exact = CIFARHierarchicalFAResNet( + depth=8, base_width=2, seed=112, dtype=torch.float64, + normalization="batchnorm", residual_scale=1.0) + exact.Q = [value.clone() for value in exact.W] + exact.R_out.copy_(-exact.W_out.t()) + for parameter in exact.W + exact.gamma + exact.beta + [ + exact.W_out, exact.b_out]: + parameter.requires_grad_(True) + clean = exact.forward( + x, return_cache=True, training=True, update_stats=False) + gradients = torch.autograd.grad( + F.cross_entropy(clean["logits"], y), clean["hiddens"]) + negative = [-x.shape[0] * value.detach() for value in gradients] + signal = (torch.softmax(clean["logits"].detach(), dim=1) + - F.one_hot(y, 10).to(torch.float64)) + teaching, contexts, recipients = exact.hierarchical_teaching( + signal, clean, return_edge_contexts=True) + numerator = 0.0 + denominator = 0.0 + for index in range(1, len(exact.Q)): + recipient = recipients[index] + spec = exact.layer_specs[index] + spatial = (negative[recipient].shape[2] + * negative[recipient].shape[3]) + target = torch.nn.grad.conv2d_weight( + negative[recipient], exact.Q[index].shape, contexts[index], + stride=spec.stride, padding=spec.padding) / (x.shape[0] * spatial) + prediction = torch.nn.grad.conv2d_weight( + teaching[recipient], exact.Q[index].shape, contexts[index], + stride=spec.stride, padding=spec.padding) / (x.shape[0] * spatial) + numerator += float((target - prediction).square().sum()) + denominator += float(target.square().sum()) + target_r = negative[-1].mean(dim=(2, 3)).t() @ signal / x.shape[0] + prediction_r = teaching[-1].mean(dim=(2, 3)).t() @ signal / x.shape[0] + numerator += float((target_r - prediction_r).square().sum()) + denominator += float(target_r.square().sum()) + delta_rule_relative = math.sqrt(numerator / max(denominator, 1e-300)) + assert delta_rule_relative < 2e-12 + return { + "hierarchical_parameter_subspace_jvp_relative_error": jvp_relative, + "hierarchical_parameter_delta_rule_relative_error": delta_rule_relative, + } + + def apical_learning_checks(): torch.manual_seed(11) net = CIFARSDILResNet(depth=8, base_width=2, seed=6) @@ -779,6 +862,7 @@ def main(): report.update(channel_subspace_estimator_check()) report.update(vectorizer_subspace_estimator_check()) report.update(hierarchical_feedback_checks()) + report.update(hierarchical_parameter_calibration_checks()) report.update(apical_learning_checks()) print(report) print("ALL CONVOLUTIONAL LOCAL-ELIGIBILITY CHECKS PASSED") |
