diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 14:46:41 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 14:46:41 -0500 |
| commit | a4858a04acc5adca65a437295b63b426b6358831 (patch) | |
| tree | 1918ee8c07821687a187a856dfd79c149856e802 /sdil | |
| parent | bebbf6d34075bba089bffc39b18c33075a611deb (diff) | |
experiment: implement KP mixed-traffic innovation controls
Diffstat (limited to 'sdil')
| -rw-r--r-- | sdil/conv.py | 296 |
1 files changed, 296 insertions, 0 deletions
diff --git a/sdil/conv.py b/sdil/conv.py index 8331433..0beb249 100644 --- a/sdil/conv.py +++ b/sdil/conv.py @@ -598,6 +598,177 @@ class CIFARKPResNet(CIFARHierarchicalFAResNet): self.R_out.add_(update, alpha=eta_output) +class CIFARKPMixedTrafficResNet(CIFARKPResNet): + """KP credit with a per-unit mixed apical compartment and neutral predictor. + + The reciprocal KP pathway supplies the instructional field. Fixed + soma-coupled traffic is added locally at every hidden population, while a + diagonal affine predictor is fitted only to instruction-off observations. + Which of raw, norm-matched raw, or innovation drives plasticity is selected + by the runner; every condition still computes and trains the predictor. + """ + + def __init__(self, *args, traffic_seed=4000, **kwargs): + super().__init__(*args, **kwargs) + generator = torch.Generator(device="cpu").manual_seed(traffic_seed) + self.B_traffic = [] + self.traffic_gain = [] + self.P_traffic = [] + self.P_traffic_bias = [] + for shape in self.hidden_shapes: + coefficient = torch.exp( + 0.25 * torch.randn(shape, generator=generator)) + self.B_traffic.append(coefficient.to( + device=self.device, dtype=self.dtype)) + self.traffic_gain.append(torch.zeros( + (), device=self.device, dtype=self.dtype)) + self.P_traffic.append(torch.zeros( + shape, device=self.device, dtype=self.dtype)) + self.P_traffic_bias.append(torch.zeros( + shape, device=self.device, dtype=self.dtype)) + self.traffic_rule = None + + @property + def n_predictor_parameters(self): + return (sum(value.numel() for value in self.P_traffic) + + sum(value.numel() for value in self.P_traffic_bias)) + + @property + def n_fixed_traffic_coefficients(self): + return sum(value.numel() for value in self.B_traffic) + + @property + def n_apical_parameters(self): + # Reciprocal Q/R parameters are logged separately as adaptive feedback. + return self.n_predictor_parameters + + @property + def mixed_units_per_example(self): + return sum(math.prod(shape) for shape in self.hidden_shapes) + + def mixed_elementwise_ops_per_example(self, rule=None): + """Transparent arithmetic count beyond convolution/linear MACs. + + Five operations per unit form traffic, affine prediction, raw, and + innovation. Norm matching additionally charges squares/reductions, + rescaling, and scalar norm arithmetic conservatively as five per unit. + """ + rule = self.traffic_rule if rule is None else rule + if rule not in ("raw", "matched", "innovation"): + raise ValueError(f"unknown mixed-traffic rule: {rule}") + multiplier = 10 if rule == "matched" else 5 + return multiplier * self.mixed_units_per_example + + @property + def predictor_elementwise_ops_per_example(self): + # Residual, centering, moments, normalization, and affine updates. + return 12 * self.mixed_units_per_example + + @torch.no_grad() + def traffic_fields(self, hiddens): + if len(hiddens) != self.n_hidden: + raise ValueError("traffic requires every somatic population") + return [gain * coefficient * hidden for gain, coefficient, hidden in zip( + self.traffic_gain, self.B_traffic, hiddens)] + + @torch.no_grad() + def calibrate_traffic_gain(self, instruction, hiddens, target_ratio): + """Fix one gain per layer from an initialization-only training prefix.""" + if target_ratio <= 0: + raise ValueError("traffic ratio must be positive") + if not (len(instruction) == len(hiddens) == self.n_hidden): + raise ValueError("traffic calibration must cover every population") + instruction_rms = [] + unscaled_traffic_rms = [] + gains = [] + realized = [] + for signal, hidden, coefficient, gain in zip( + instruction, hiddens, self.B_traffic, self.traffic_gain): + signal_scale = signal.square().mean().sqrt() + traffic_scale = (coefficient * hidden).square().mean().sqrt() + if float(signal_scale) <= 0 or float(traffic_scale) <= 0: + raise ValueError("traffic calibration encountered a zero RMS") + value = target_ratio * signal_scale / traffic_scale + gain.copy_(value) + achieved = (gain * coefficient * hidden).square().mean().sqrt() + instruction_rms.append(float(signal_scale)) + unscaled_traffic_rms.append(float(traffic_scale)) + gains.append(float(gain)) + realized.append(float(achieved / signal_scale)) + return { + "target_ratio": float(target_ratio), + "instruction_rms": instruction_rms, + "unscaled_traffic_rms": unscaled_traffic_rms, + "traffic_gain": gains, + "realized_traffic_instruction_rms_ratio": realized, + } + + @torch.no_grad() + def mixed_apical_components(self, instruction, hiddens, rule): + """Return used, raw, innovation, matched, and traffic fields.""" + if rule not in ("raw", "matched", "innovation"): + raise ValueError(f"unknown mixed-traffic rule: {rule}") + if not (len(instruction) == len(hiddens) == self.n_hidden): + raise ValueError("mixed apical inputs must cover every population") + traffic = self.traffic_fields(hiddens) + raw = [] + innovation = [] + matched = [] + for signal, hidden, ordinary, slope, bias in zip( + instruction, hiddens, traffic, + self.P_traffic, self.P_traffic_bias): + apical = signal + ordinary + residual = apical - (slope * hidden + bias) + raw_norm = apical.flatten(1).norm(dim=1).clamp_min(1e-30) + residual_norm = residual.flatten(1).norm(dim=1) + scale = (residual_norm / raw_norm).reshape( + residual.shape[0], *([1] * (residual.ndim - 1))) + raw.append(apical) + innovation.append(residual) + matched.append(scale * apical) + choices = {"raw": raw, "matched": matched, "innovation": innovation} + return { + "used": choices[rule], "raw": raw, "innovation": innovation, + "matched": matched, "traffic": traffic, + "instruction": instruction, + } + + @torch.no_grad() + def predictor_step(self, hiddens, eta): + """Instruction-off normalized-LMS update from local soma/traffic pairs.""" + if not 0.0 < eta <= 1.0: + raise ValueError("predictor learning rate must lie in (0, 1]") + traffic = self.traffic_fields(hiddens) + squared_error = 0.0 + units = 0 + for hidden, target, slope, bias in zip( + hiddens, traffic, self.P_traffic, self.P_traffic_bias): + residual = target - slope * hidden - bias + centered_h = hidden - hidden.mean(dim=0) + centered_r = residual - residual.mean(dim=0) + variance = centered_h.square().mean(dim=0) + slope.add_((centered_r * centered_h).mean(dim=0) + / (variance + 1e-12), alpha=eta) + bias.add_(residual.mean(dim=0), alpha=eta) + squared_error += float(residual.square().sum()) + units += residual.numel() + return squared_error / units + + @torch.no_grad() + def predictor_traffic_residual_rms_ratio(self, hiddens): + traffic = self.traffic_fields(hiddens) + residual_power = 0.0 + traffic_power = 0.0 + for hidden, target, slope, bias in zip( + hiddens, traffic, self.P_traffic, self.P_traffic_bias): + residual = target - slope * hidden - bias + residual_power += float(residual.square().sum()) + traffic_power += float(target.square().sum()) + if traffic_power <= 0: + raise ValueError("predictor audit requires nonzero traffic") + return math.sqrt(residual_power / traffic_power) + + @torch.no_grad() def hierarchical_feedback_tracking_report(net): """Cheap parameter-space tracking diagnostics; never used for learning.""" @@ -1022,6 +1193,64 @@ def conv_kolen_pollack_step(net, x, y, config): } +def conv_kp_mixed_traffic_step(net, x, y, config, step, rule, + predictor_every): + """One reciprocal-KP update using a selected mixed-apical signal.""" + if not isinstance(net, CIFARKPMixedTrafficResNet): + raise TypeError("mixed-traffic KP step requires CIFARKPMixedTrafficResNet") + if predictor_every < 1: + raise ValueError("predictor cadence must be positive") + config.validate() + with torch.no_grad(): + forward = net.forward( + x, return_cache=True, training=True, update_stats=True) + logits = forward["logits"] + loss = F.cross_entropy(logits, y) + output_error = (torch.softmax(logits, dim=1) + - F.one_hot(y, net.n_classes).to(logits.dtype)) + instruction = net.hierarchical_teaching(output_error, forward) + components = net.mixed_apical_components( + instruction, forward["hiddens"], rule) + used = components["used"] + total_units = sum(value.numel() for value in used) + + def rms(values): + return math.sqrt( + sum(float(value.square().sum()) for value in values) / total_units) + + (directions, gamma_directions, beta_directions, + output_weight, output_bias) = net.local_ascent_directions( + used, output_error, forward) + reciprocal_directions, reciprocal_readout = ( + net.reciprocal_feedback_directions(used, output_error, forward)) + # Form both local correlations before either parameter path changes. + net.apply_reciprocal_ascent( + reciprocal_directions, reciprocal_readout, + eta_hidden=config.eta, eta_output=config.eta_output, + momentum=config.momentum, weight_decay=config.weight_decay) + net.apply_ascent( + directions, output_weight, output_bias, + eta_hidden=config.eta, eta_output=config.eta_output, + momentum=config.momentum, weight_decay=config.weight_decay, + gamma_directions=gamma_directions, + beta_directions=beta_directions) + did_predictor_update = step % predictor_every == 0 + predictor_mse = None + if did_predictor_update: + predictor_mse = net.predictor_step( + forward["hiddens"], config.eta_P) + return { + "loss": float(loss), "did_perturb": False, "calibration": None, + "did_predictor_update": did_predictor_update, + "predictor_mse": predictor_mse, + "teaching_rms": rms(used), + "instruction_rms": rms(components["instruction"]), + "raw_apical_rms": rms(components["raw"]), + "innovation_rms": rms(components["innovation"]), + "traffic_rms": rms(components["traffic"]), + } + + def conv_learned_hierarchical_step(net, x, y, config, step, generator=None): """One task update with optional causal calibration of hierarchical Q/R.""" config.validate() @@ -1089,6 +1318,73 @@ def conv_hierarchical_alignment_report(net, x, y): return report +def conv_kp_mixed_traffic_alignment_report(net, x, y, rule): + """Same-state audit of instruction/raw/innovation/matched directions.""" + if not isinstance(net, CIFARKPMixedTrafficResNet): + raise TypeError("mixed-traffic audit requires CIFARKPMixedTrafficResNet") + parameters = net.W + net.gamma + net.beta + [net.W_out, net.b_out] + for parameter in parameters: + parameter.requires_grad_(True) + forward = net.forward(x, return_cache=True, training=True, update_stats=False) + gradients = torch.autograd.grad( + F.cross_entropy(forward["logits"], y), forward["hiddens"]) + batch = x.shape[0] + negative_gradients = [-batch * value.detach() for value in gradients] + with torch.no_grad(): + output_error = (torch.softmax(forward["logits"], dim=1) + - F.one_hot(y, net.n_classes).to(forward["logits"].dtype)) + instruction = net.hierarchical_teaching(output_error, forward) + components = net.mixed_apical_components( + instruction, forward["hiddens"], rule) + + def align(values): + return [float(F.cosine_similarity( + left.flatten(1), right.flatten(1), dim=1).mean()) + for left, right in zip(values, negative_gradients)] + + norm_errors = [] + direction_errors = [] + for raw, innovation, matched in zip( + components["raw"], components["innovation"], + components["matched"]): + raw_flat = raw.flatten(1) + innovation_flat = innovation.flatten(1) + matched_flat = matched.flatten(1) + target_norm = innovation_flat.norm(dim=1) + norm_errors.append(float(( + (matched_flat.norm(dim=1) - target_norm).abs() + / target_norm.clamp_min(1e-30)).max())) + raw_match_cosine = F.cosine_similarity( + raw_flat, matched_flat, dim=1) + direction_errors.append(float((raw_match_cosine - 1.0).abs().max())) + + instruction_power = sum(float(value.square().sum()) + for value in components["instruction"]) + traffic_power = sum(float(value.square().sum()) + for value in components["traffic"]) + report = { + "normalization_state": "training_batch_stats_without_running_update", + "teaching_negative_gradient_cosine": align(components["used"]), + "used_negative_gradient_cosine": align(components["used"]), + "instruction_negative_gradient_cosine": align( + components["instruction"]), + "raw_negative_gradient_cosine": align(components["raw"]), + "innovation_negative_gradient_cosine": align( + components["innovation"]), + "matched_negative_gradient_cosine": align(components["matched"]), + "max_norm_match_relative_error": max(norm_errors), + "max_norm_match_direction_error": max(direction_errors), + "traffic_instruction_rms_ratio": math.sqrt( + traffic_power / instruction_power), + "predictor_traffic_residual_rms_ratio": ( + net.predictor_traffic_residual_rms_ratio(forward["hiddens"])), + } + report.update(hierarchical_feedback_tracking_report(net)) + for parameter in parameters: + parameter.requires_grad_(False) + return report + + class CIFARSDILResNet(CIFARLocalResNet): """CIFAR local ResNet with per-unit apical vectorizers and predictors. |
