diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 16:43:39 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 16:43:39 -0500 |
| commit | 72f6b758540e8c1d7a44fed63bafe71784f6dfa2 (patch) | |
| tree | de271bb7948e31972e76070518677c12c4f79058 /experiments/conv_local_smoke.py | |
| parent | 575530e7124c540e58571b50ef147a5f20e6cb54 (diff) | |
experiment: add closed-form neutral innovation fit
Diffstat (limited to 'experiments/conv_local_smoke.py')
| -rw-r--r-- | experiments/conv_local_smoke.py | 38 |
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), } |
