summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--experiments/conv_local_smoke.py38
-rw-r--r--experiments/diagnose_kp_traffic_nonfinite.py30
-rw-r--r--sdil/conv.py55
3 files changed, 114 insertions, 9 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),
}
diff --git a/experiments/diagnose_kp_traffic_nonfinite.py b/experiments/diagnose_kp_traffic_nonfinite.py
index a045a75..bdf79c3 100644
--- a/experiments/diagnose_kp_traffic_nonfinite.py
+++ b/experiments/diagnose_kp_traffic_nonfinite.py
@@ -86,6 +86,9 @@ def main():
parser = argparse.ArgumentParser()
parser.add_argument("--rule", choices=("raw", "matched", "innovation"),
default="innovation")
+ parser.add_argument("--predictor_mode",
+ choices=("nlms20", "closed_form"),
+ default="nlms20")
parser.add_argument("--device", default="cuda")
parser.add_argument("--max_steps", type=int, default=352)
parser.add_argument("--out", required=True)
@@ -127,21 +130,28 @@ def main():
calibration_error, calibration_forward)
calibration = net.calibrate_traffic_gain(
calibration_instruction, calibration_forward["hiddens"], 4.0)
- del calibration_forward, calibration_error, calibration_instruction
-
- iterator = iter(train)
+ closed_form_fit = None
warmup_mse = []
- for _ in range(20):
- x, _ = next(iterator)
- forward = net.forward(x, training=True, update_stats=False)
- warmup_mse.append(net.predictor_step(forward["hiddens"], 0.1))
+ if args.predictor_mode == "closed_form":
+ closed_form_fit = net.predictor_closed_form_fit(
+ calibration_forward["hiddens"])
+ warmup_mse.append(closed_form_fit["mse"])
+ else:
+ iterator = iter(train)
+ for _ in range(20):
+ x, _ = next(iterator)
+ forward = net.forward(x, training=True, update_stats=False)
+ warmup_mse.append(net.predictor_step(forward["hiddens"], 0.1))
+ del calibration_forward, calibration_error, calibration_instruction
train.g.set_state(loader_state)
loader_state_restored = torch.equal(train.g.get_state(), loader_state)
audit_forward = net.forward(
calibration_x, training=True, update_stats=False)
post_warmup_ratio = net.predictor_traffic_residual_rms_ratio(
audit_forward["hiddens"])
- del forward, audit_forward
+ if args.predictor_mode == "nlms20":
+ del forward
+ del audit_forward
initial_state = network_state(net)
if nonfinite_groups(initial_state):
@@ -212,15 +222,17 @@ def main():
"scope": "training_only_no_validation_or_test_evaluation",
"provenance": provenance(),
"rule": args.rule,
+ "predictor_mode": args.predictor_mode,
"max_steps": args.max_steps,
"split": split,
"traffic_calibration": calibration,
"predictor_warmup": {
- "steps": 20,
+ "steps": len(warmup_mse),
"first_mse": warmup_mse[0],
"last_mse": warmup_mse[-1],
"post_warmup_traffic_residual_rms_ratio": post_warmup_ratio,
"task_loader_state_restored": loader_state_restored,
+ "closed_form_fit": closed_form_fit,
},
"initial_parameter_state": initial_state,
"trajectory": trajectory,
diff --git a/sdil/conv.py b/sdil/conv.py
index ea62e95..d6bae4f 100644
--- a/sdil/conv.py
+++ b/sdil/conv.py
@@ -769,6 +769,61 @@ class CIFARKPMixedTrafficResNet(CIFARKPResNet):
return squared_error / units
@torch.no_grad()
+ def predictor_closed_form_fit(self, hiddens, min_variance=1e-12):
+ """Fit the local affine neutral relation by per-cell least squares.
+
+ Each spatial cell uses only its own soma and instruction-off apical
+ observations across the supplied batch. Cells with no observed soma
+ variance receive a zero slope and their local target mean as bias.
+ """
+ if min_variance < 0:
+ raise ValueError("minimum predictor variance must be nonnegative")
+ traffic = self.traffic_fields(hiddens)
+ residual_power = 0.0
+ traffic_power = 0.0
+ maximum_residual_slope = 0.0
+ squared_error = 0.0
+ units = 0
+ for hidden, target, slope, bias in zip(
+ hiddens, traffic, self.P_traffic, self.P_traffic_bias):
+ hidden_mean = hidden.mean(dim=0)
+ target_mean = target.mean(dim=0)
+ centered_h = hidden - hidden_mean
+ centered_target = target - target_mean
+ variance = centered_h.square().mean(dim=0)
+ covariance = (centered_h * centered_target).mean(dim=0)
+ active = variance > min_variance
+ fitted_slope = torch.where(
+ active, covariance / variance.clamp_min(min_variance),
+ torch.zeros_like(variance))
+ fitted_bias = target_mean - fitted_slope * hidden_mean
+ slope.copy_(fitted_slope)
+ bias.copy_(fitted_bias)
+ residual = target - slope * hidden - bias
+ residual_covariance = (centered_h * (
+ residual - residual.mean(dim=0))).mean(dim=0)
+ residual_slope = torch.where(
+ active,
+ residual_covariance / variance.clamp_min(min_variance),
+ torch.zeros_like(variance))
+ maximum_residual_slope = max(
+ maximum_residual_slope, float(residual_slope.abs().max()))
+ power = float(residual.square().sum())
+ residual_power += power
+ traffic_power += float(target.square().sum())
+ squared_error += power
+ units += residual.numel()
+ if traffic_power <= 0:
+ raise ValueError("closed-form predictor fit requires nonzero traffic")
+ return {
+ "mse": squared_error / units,
+ "residual_traffic_rms_ratio": math.sqrt(
+ residual_power / traffic_power),
+ "max_absolute_residual_soma_slope": maximum_residual_slope,
+ "observations": int(hiddens[0].shape[0]),
+ }
+
+ @torch.no_grad()
def predictor_traffic_residual_rms_ratio(self, hiddens):
traffic = self.traffic_fields(hiddens)
residual_power = 0.0