summaryrefslogtreecommitdiff
path: root/experiments/conv_local_smoke.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 14:15:09 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 14:33:50 -0500
commitd761f655301e3535368a87e9d4ffc26ab1ab8510 (patch)
tree247439cc0f270cc78555ff72dcd15d35d9c754f6 /experiments/conv_local_smoke.py
parent8a5c1045a43b1f91fd46b2fcc044b032cb14718c (diff)
baseline: add local reciprocal Kolen-Pollack
Diffstat (limited to 'experiments/conv_local_smoke.py')
-rw-r--r--experiments/conv_local_smoke.py101
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")