From f56e60c9bc095904e1ee721e70a991b77b801e53 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Wed, 22 Jul 2026 05:17:51 -0500 Subject: sdil: normalize causal vectorizer regression --- experiments/run.py | 6 +++++- experiments/smoke.py | 30 +++++++++++++++++++++++++++++- sdil/core.py | 42 +++++++++++++++++++++++++++++++----------- 3 files changed, 65 insertions(+), 13 deletions(-) diff --git a/experiments/run.py b/experiments/run.py index ced7585..0019c67 100644 --- a/experiments/run.py +++ b/experiments/run.py @@ -223,7 +223,9 @@ def build(args, device): kappa=args.kappa, feedback=args.feedback, p_update_on_neutral=bool(args.p_neutral), normalize_delta=bool(args.normalize_delta), - raw_scale_control=args.raw_scale_control) + raw_scale_control=args.raw_scale_control, + vectorizer_optimizer=args.vectorizer_optimizer, + vectorizer_eps=args.vectorizer_eps) elif args.mode == "nodepert": if args.pert_every != 1: raise ValueError("direct node perturbation requires --pert_every 1") @@ -679,6 +681,8 @@ def get_args(): p.add_argument("--predictor_mode", default="diagonal", choices=["diagonal", "full"]) p.add_argument("--vectorizer_mode", default="linear", choices=["linear", "soma_gated", "context_gated"]) + p.add_argument("--vectorizer_optimizer", default="sgd", choices=["sgd", "nlms"]) + p.add_argument("--vectorizer_eps", type=float, default=1e-6) p.add_argument("--normalize_delta", type=int, default=0) p.add_argument("--settle_steps", type=int, default=0) p.add_argument("--kappa", type=float, default=0.0) diff --git a/experiments/smoke.py b/experiments/smoke.py index 9971800..37c38e3 100644 --- a/experiments/smoke.py +++ b/experiments/smoke.py @@ -8,7 +8,7 @@ import torch sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) from sdil.core import (SDILNet, SDILConfig, apical_calibration_step, sdil_step, node_perturbation_targets, simultaneous_node_perturbation_targets, - neutral_p_update, teaching_signal) + neutral_p_update, teaching_signal, _update_apical_vectorizer) from sdil.baselines import dfa_config from sdil import probes from sdil.data import get_dataset, onehot @@ -215,6 +215,33 @@ def check_feedback_first_calibration(): print("CHECK0g feedback-first calibration: A changed; W/readout frozen") +def check_vectorizer_nlms(): + """NLMS must divide the joint linear/context update by feature power.""" + torch.manual_seed(29) + batch = 16 + net = SDILNet([5, 7, 7, 3], act="relu", device="cpu", seed=12, + residual=True, vectorizer_mode="context_gated") + for weight in net.A: + weight.zero_() + for weight in net.A_gate: + weight.zero_() + h = net.forward(torch.randn(batch, 5))["h"] + c = torch.randn(batch, 3) + qs = [torch.randn_like(hidden) for hidden in h[1:-1]] + residuals = [torch.zeros_like(target) for target in qs] + cfg = SDILConfig(eta_A=0.03, vectorizer_optimizer="nlms", vectorizer_eps=1e-6) + features = (c.unsqueeze(2) * torch.tanh(h[-2]).unsqueeze(1)).flatten(1) + denominator = (c.square().sum(1, keepdim=True) + + features.square().sum(1, keepdim=True)).clamp_min(cfg.vectorizer_eps) + expected_a0 = cfg.eta_A * ((qs[0] / denominator).t() @ c / batch) + expected_gate0 = cfg.eta_A * ((qs[0] / denominator).t() @ features / batch) + _update_apical_vectorizer(net, h, c, residuals, qs, cfg) + assert torch.allclose(net.A[0], expected_a0, atol=1e-7, rtol=1e-6) + assert torch.allclose(net.A_gate[0], expected_gate0, atol=1e-7, rtol=1e-6) + assert torch.isfinite(net.A[0]).all() and torch.isfinite(net.A_gate[0]).all() + print("CHECK0h vectorizer NLMS: exact joint-feature normalization") + + def main(): torch.manual_seed(0) dev = "cpu" @@ -225,6 +252,7 @@ def main(): check_state_conditioned_vectorizers() check_direct_node_perturbation() check_feedback_first_calibration() + check_vectorizer_nlms() print("loading MNIST subset...") tr, te, n_in, n_out = get_dataset("mnist", batch_size=128, device=dev) xb, yb = next(iter(tr)) diff --git a/sdil/core.py b/sdil/core.py index 20e94eb..d816549 100644 --- a/sdil/core.py +++ b/sdil/core.py @@ -452,7 +452,8 @@ class SDILConfig: pert_sigma=1e-2, pert_every=5, pert_ndirs=1, momentum=0.0, wd=0.0, settle_steps=0, kappa=0.0, feedback="error", p_update_on_neutral=True, normalize_delta=False, pert_mode="layerwise", - raw_scale_control="none", direct_node_pert=False): + raw_scale_control="none", direct_node_pert=False, + vectorizer_optimizer="sgd", vectorizer_eps=1e-6): self.eta = eta self.eta_A = eta_A self.eta_P = eta_P @@ -480,23 +481,42 @@ class SDILConfig: # eta) from that noise -- like normalized SGD. self.normalize_delta = normalize_delta self.direct_node_pert = direct_node_pert + if vectorizer_optimizer not in ("sgd", "nlms"): + raise ValueError(f"unknown vectorizer optimizer: {vectorizer_optimizer}") + self.vectorizer_optimizer = vectorizer_optimizer + self.vectorizer_eps = vectorizer_eps -def _update_apical_vectorizer(net, h, c, r_list, qs, eta_A): +def _update_apical_vectorizer(net, h, c, r_list, qs, cfg): """Apply the local causal-regression update to all apical pathways.""" B = c.shape[0] context = h[-2] + context_features = None + if net.vectorizer_mode == "context_gated": + context_features = (c.unsqueeze(2) * torch.tanh(context).unsqueeze(1)).flatten(1) for l in range(net.L - 1): calibration_error = qs[l] - r_list[l] - dA = calibration_error.t() @ c / B - net.A[l] += eta_A * dA + scaled_error = calibration_error + if cfg.vectorizer_optimizer == "nlms": + if net.vectorizer_mode == "linear": + feature_power = c.square().sum(1, keepdim=True) + elif net.vectorizer_mode == "soma_gated": + # Each cell i sees [c, tanh(h_i)c], so its effective feature + # power is ||c||^2 (1+tanh(h_i)^2). + feature_power = (c.square().sum(1, keepdim=True) + * (1.0 + torch.tanh(h[l + 1]).square())) + else: + feature_power = (c.square().sum(1, keepdim=True) + + context_features.square().sum(1, keepdim=True)) + scaled_error = calibration_error / feature_power.clamp_min(cfg.vectorizer_eps) + dA = scaled_error.t() @ c / B + net.A[l] += cfg.eta_A * dA if net.vectorizer_mode == "soma_gated": - dgate = (calibration_error * torch.tanh(h[l + 1])).t() @ c / B - net.A_gate[l] += eta_A * dgate + dgate = (scaled_error * torch.tanh(h[l + 1])).t() @ c / B + net.A_gate[l] += cfg.eta_A * dgate elif net.vectorizer_mode == "context_gated": - features = (c.unsqueeze(2) * torch.tanh(context).unsqueeze(1)).flatten(1) - dgate = calibration_error.t() @ features / B - net.A_gate[l] += eta_A * dgate + dgate = scaled_error.t() @ context_features / B + net.A_gate[l] += cfg.eta_A * dgate def apical_calibration_step(net, x, y, y_onehot, cfg, @@ -531,7 +551,7 @@ def apical_calibration_step(net, x, y, y_onehot, cfg, estimator = (simultaneous_node_perturbation_targets if mode == "simultaneous" else node_perturbation_targets) qs = estimator(net, x, y, sigma=cfg.pert_sigma, n_dirs=directions) - _update_apical_vectorizer(net, h, c, r_list, qs, cfg.eta_A) + _update_apical_vectorizer(net, h, c, r_list, qs, cfg) return loss, {"error": c, "did_pert": True} @@ -605,7 +625,7 @@ def sdil_step(net, x, y, y_onehot, cfg, step, prev_error=None): # ================= apical vectorizer A via node perturbation ======= if did_pert and cfg.learn_A: - _update_apical_vectorizer(net, h, c, r_list, qs, cfg.eta_A) + _update_apical_vectorizer(net, h, c, r_list, qs, cfg) # ================= predictor P (neutral) ========================== # KEY identification condition. P must learn the soma->apical coupling -- cgit v1.2.3