diff options
Diffstat (limited to 'experiments/conv_local_smoke.py')
| -rw-r--r-- | experiments/conv_local_smoke.py | 101 |
1 files changed, 98 insertions, 3 deletions
diff --git a/experiments/conv_local_smoke.py b/experiments/conv_local_smoke.py index 4e388a0..e8cf294 100644 --- a/experiments/conv_local_smoke.py +++ b/experiments/conv_local_smoke.py @@ -8,10 +8,11 @@ import torch import torch.nn.functional as F sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) -from sdil.conv import (CIFARHierarchicalFAResNet, CIFARLocalResNet, - CIFARSDILResNet, ConvSDILConfig, +from sdil.conv import (CIFARHierarchicalFAResNet, CIFARKPResNet, + CIFARLocalResNet, CIFARSDILResNet, ConvSDILConfig, channel_subspace_apical_calibration, - conv_hierarchical_step, conv_local_step, + conv_hierarchical_step, conv_kolen_pollack_step, + conv_local_step, hierarchical_mirror_observations, hierarchical_parameter_subspace_calibration, normalized_residual_mirror_update, @@ -843,6 +844,99 @@ def normalized_response_mirror_checks(): } +def kolen_pollack_checks(): + """KP's reciprocal correlations are local and preserve exact symmetry.""" + torch.manual_seed(127) + common = dict( + depth=8, base_width=2, seed=53, dtype=torch.float64, + normalization="batchnorm", residual_scale=1.0) + net = CIFARKPResNet(**common) + for index in range(1, len(net.Q)): + net.Q[index].copy_(net.W[index]) + net.R_out.copy_(-net.W_out.t()) + x = torch.randn(3, 3, 32, 32, dtype=torch.float64) + y = torch.tensor([1, 4, 7]) + forward = net.forward( + x, return_cache=True, training=True, update_stats=False) + output_error = (torch.softmax(forward["logits"], dim=1) + - F.one_hot(y, 10).to(torch.float64)) + teaching = net.hierarchical_teaching(output_error, forward) + (forward_directions, gamma_directions, beta_directions, + output_weight, output_bias) = net.local_ascent_directions( + teaching, output_error, forward) + reciprocal_directions, reciprocal_readout = ( + net.reciprocal_feedback_directions( + teaching, output_error, forward)) + direction_error = max([ + float((left - right).abs().max()) + for left, right in zip(forward_directions[1:], reciprocal_directions[1:]) + ] + [float((reciprocal_readout + output_weight.t()).abs().max())]) + assert direction_error < 1e-14 + + # Once local activities have been observed, neither forward nor feedback + # parameter values may alter the independently formed reciprocal update. + before = [value.clone() for value in reciprocal_directions[1:]] + before_readout = reciprocal_readout.clone() + for value in net.W + net.Q: + value.add_(torch.randn_like(value)) + net.W_out.add_(torch.randn_like(net.W_out)) + net.R_out.add_(torch.randn_like(net.R_out)) + independent_directions, independent_readout = ( + net.reciprocal_feedback_directions( + teaching, output_error, forward)) + independence_error = max([ + float((left - right).abs().max()) + for left, right in zip(before, independent_directions[1:]) + ] + [float((before_readout - independent_readout).abs().max())]) + assert independence_error == 0.0 + + # A fresh symmetric state must remain symmetric under two momentum steps. + net = CIFARKPResNet(**common) + for index in range(1, len(net.Q)): + net.Q[index].copy_(net.W[index]) + net.R_out.copy_(-net.W_out.t()) + for _ in range(2): + forward = net.forward( + x, return_cache=True, training=True, update_stats=False) + output_error = (torch.softmax(forward["logits"], dim=1) + - F.one_hot(y, 10).to(torch.float64)) + teaching = net.hierarchical_teaching(output_error, forward) + (forward_directions, gamma_directions, beta_directions, + output_weight, output_bias) = net.local_ascent_directions( + teaching, output_error, forward) + reciprocal_directions, reciprocal_readout = ( + net.reciprocal_feedback_directions( + teaching, output_error, forward)) + net.apply_reciprocal_ascent( + reciprocal_directions, reciprocal_readout, + eta_hidden=0.013, eta_output=0.017, momentum=0.9, + weight_decay=1e-4) + net.apply_ascent( + forward_directions, output_weight, output_bias, + eta_hidden=0.013, eta_output=0.017, momentum=0.9, + weight_decay=1e-4, gamma_directions=gamma_directions, + beta_directions=beta_directions) + symmetry_error = max([ + float((net.Q[index] - net.W[index]).abs().max()) + for index in range(1, len(net.Q)) + ] + [float((net.R_out + net.W_out.t()).abs().max())]) + assert symmetry_error < 1e-14 + + # Exercise the public training step and ensure it remains graph-free. + result = conv_kolen_pollack_step( + net, x, y, ConvSDILConfig( + eta=1e-3, eta_output=1e-3, momentum=0.9, + weight_decay=1e-4, learn_A=False, learn_P=False)) + assert math.isfinite(result["loss"]) + assert all(not value.requires_grad for value in + net.W + net.Q + [net.W_out, net.R_out, net.b_out]) + return { + "kp_local_direction_absolute_error": direction_error, + "kp_forward_parameter_independence_error": independence_error, + "kp_symmetric_update_absolute_error": symmetry_error, + } + + def apical_learning_checks(): torch.manual_seed(11) net = CIFARSDILResNet(depth=8, base_width=2, seed=6) @@ -951,6 +1045,7 @@ def main(): report.update(hierarchical_feedback_checks()) report.update(hierarchical_parameter_calibration_checks()) report.update(normalized_response_mirror_checks()) + report.update(kolen_pollack_checks()) report.update(apical_learning_checks()) print(report) print("ALL CONVOLUTIONAL LOCAL-ELIGIBILITY CHECKS PASSED") |
