summaryrefslogtreecommitdiff
path: root/experiments/conv_local_smoke.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 16:43:39 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 16:43:39 -0500
commit72f6b758540e8c1d7a44fed63bafe71784f6dfa2 (patch)
treede271bb7948e31972e76070518677c12c4f79058 /experiments/conv_local_smoke.py
parent575530e7124c540e58571b50ef147a5f20e6cb54 (diff)
experiment: add closed-form neutral innovation fit
Diffstat (limited to 'experiments/conv_local_smoke.py')
-rw-r--r--experiments/conv_local_smoke.py38
1 files changed, 38 insertions, 0 deletions
diff --git a/experiments/conv_local_smoke.py b/experiments/conv_local_smoke.py
index 887674e..3267417 100644
--- a/experiments/conv_local_smoke.py
+++ b/experiments/conv_local_smoke.py
@@ -984,6 +984,19 @@ def kp_mixed_traffic_checks():
for slope, bias in zip(net.P_traffic, net.P_traffic_bias):
slope.zero_()
bias.zero_()
+ closed_form = net.predictor_closed_form_fit(forward["hiddens"])
+ fitted = net.mixed_apical_components(
+ instruction, forward["hiddens"], "innovation")
+ closed_form_error = max(float((left - right).abs().max())
+ for left, right in zip(
+ fitted["innovation"], instruction))
+ assert closed_form_error < 1e-14
+ assert closed_form["residual_traffic_rms_ratio"] < 1e-14
+ assert closed_form["max_absolute_residual_soma_slope"] < 1e-14
+
+ for slope, bias in zip(net.P_traffic, net.P_traffic_bias):
+ slope.zero_()
+ bias.zero_()
components = net.mixed_apical_components(
instruction, forward["hiddens"], "matched")
norm_errors = []
@@ -1037,6 +1050,24 @@ def kp_mixed_traffic_checks():
right.P_traffic + right.P_traffic_bias))
assert predictor_independence_error == 0.0
+ closed_left = CIFARKPMixedTrafficResNet(**common)
+ closed_right = CIFARKPMixedTrafficResNet(**common)
+ for left_gain, right_gain, source in zip(
+ closed_left.traffic_gain, closed_right.traffic_gain,
+ net.traffic_gain):
+ left_gain.copy_(source)
+ right_gain.copy_(source)
+ for value in (closed_right.W + closed_right.Q
+ + [closed_right.W_out, closed_right.R_out]):
+ value.add_(torch.randn_like(value))
+ closed_left.predictor_closed_form_fit(fixed_hiddens)
+ closed_right.predictor_closed_form_fit(fixed_hiddens)
+ closed_form_independence_error = max(float((a - b).abs().max())
+ for a, b in zip(
+ closed_left.P_traffic + closed_left.P_traffic_bias,
+ closed_right.P_traffic + closed_right.P_traffic_bias))
+ assert closed_form_independence_error == 0.0
+
net.traffic_rule = "innovation"
result = conv_kp_mixed_traffic_step(
net, x, y, ConvSDILConfig(
@@ -1053,11 +1084,18 @@ def kp_mixed_traffic_checks():
"kp_traffic_zero_limit_error": max(zero_errors),
"kp_traffic_ratio_error": ratio_error,
"kp_traffic_exact_predictor_error": exact_predictor_error,
+ "kp_traffic_closed_form_predictor_error": closed_form_error,
+ "kp_traffic_closed_form_residual_ratio": closed_form[
+ "residual_traffic_rms_ratio"],
+ "kp_traffic_closed_form_residual_slope": closed_form[
+ "max_absolute_residual_soma_slope"],
"kp_traffic_matched_norm_error": max(norm_errors),
"kp_traffic_matched_direction_error": max(direction_errors),
"kp_traffic_reciprocal_correlation_error": max(correlation_errors),
"kp_traffic_predictor_parameter_independence_error": (
predictor_independence_error),
+ "kp_traffic_closed_form_parameter_independence_error": (
+ closed_form_independence_error),
}