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