diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 04:01:23 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 04:01:23 -0500 |
| commit | 3942eda789a64dbc4e213cb9ca2f376c1f334501 (patch) | |
| tree | 9f3fc556b3536c18668a037d27d8c6f7c99bfbcd | |
| parent | 7eb40a516849cbc3e93540d4fdf40366f6a8df9e (diff) | |
baseline: add unamortized node perturbation
| -rw-r--r-- | experiments/run.py | 26 | ||||
| -rw-r--r-- | experiments/smoke.py | 34 | ||||
| -rw-r--r-- | sdil/core.py | 13 | ||||
| -rw-r--r-- | sdil/probes.py | 20 |
4 files changed, 87 insertions, 6 deletions
diff --git a/experiments/run.py b/experiments/run.py index 3b0f961..ea19366 100644 --- a/experiments/run.py +++ b/experiments/run.py @@ -222,6 +222,15 @@ def build(args, device): p_update_on_neutral=bool(args.p_neutral), normalize_delta=bool(args.normalize_delta), raw_scale_control=args.raw_scale_control) + elif args.mode == "nodepert": + if args.pert_every != 1: + raise ValueError("direct node perturbation requires --pert_every 1") + cfg = SDILConfig( + eta=args.eta, use_residual=False, learn_A=False, learn_P=False, + pert_sigma=args.pert_sigma, pert_every=args.pert_every, + pert_ndirs=args.pert_ndirs, + pert_mode=args.pert_mode, momentum=args.momentum, + normalize_delta=bool(args.normalize_delta), direct_node_pert=True) else: raise ValueError(args.mode) return net, cfg @@ -366,7 +375,7 @@ def train(args): else: loss, aux = sdil_step(net, x, y, yoh, cfg, step, prev_error=prev_error) prev_error = aux["error"] - if args.mode == "sdil" and aux["did_pert"]: + if args.mode in ("sdil", "nodepert") and aux["did_pert"]: event = calibration_work_per_event(net, cfg) perturbation_events += 1 calibration_batch_loss_evaluations += event["batch_loss_evaluations"] @@ -387,6 +396,8 @@ def train(args): diagnostics_t0 = time.time() if args.mode == "fa" and inline_diagnostics: rec.update(probes.fa_alignment_report(net, px, py, poh)) + elif args.mode == "nodepert" and inline_diagnostics: + rec.update(probes.nodepert_alignment_report(net, px, py, cfg)) elif args.mode != "bp" and inline_diagnostics: al = probes.alignment_report(net, px, py, poh, cfg) rec["cos_r_negg"] = al["cos_r_negg"] @@ -442,6 +453,10 @@ def train(args): al = probes.fa_alignment_report(net, px, py, poh) meancos = sum(al["cos_fa_negg"]) / len(al["cos_fa_negg"]) msg += f" mean_cos(fa,-g) {meancos:+.3f} per-layer {['%.2f'%v for v in al['cos_fa_negg']]}" + elif args.mode == "nodepert": + al = probes.nodepert_alignment_report(net, px, py, cfg) + meancos = sum(al["cos_q_negg"]) / len(al["cos_q_negg"]) + msg += f" mean_cos(q,-g) {meancos:+.3f} per-layer {['%.2f'%v for v in al['cos_q_negg']]}" elif args.mode != "bp": al = probes.alignment_report(net, px, py, poh, cfg) meancos = sum(al["cos_r_negg"]) / len(al["cos_r_negg"]) @@ -468,6 +483,12 @@ def train(args): log["final"].update(probes.fa_alignment_report(net, px, py, poh)) device_sync(device) diagnostics_wall_s += time.time() - diagnostics_t0 + elif args.mode == "nodepert" and args.diagnostics != "none": + device_sync(device) + diagnostics_t0 = time.time() + log["final"].update(probes.nodepert_alignment_report(net, px, py, cfg)) + device_sync(device) + diagnostics_wall_s += time.time() - diagnostics_t0 elif args.mode != "bp" and args.diagnostics != "none": device_sync(device) diagnostics_t0 = time.time() @@ -534,7 +555,8 @@ def train(args): def get_args(): p = argparse.ArgumentParser() - p.add_argument("--mode", default="sdil", choices=["bp", "fa", "dfa", "sdil"]) + p.add_argument("--mode", default="sdil", + choices=["bp", "fa", "dfa", "sdil", "nodepert"]) p.add_argument("--dataset", default="mnist", choices=list(REAL_DATASETS + SYNTHETIC_DATASETS)) p.add_argument("--depth", type=int, default=3) # hidden layers diff --git a/experiments/smoke.py b/experiments/smoke.py index fcc0e9c..a49e4c6 100644 --- a/experiments/smoke.py +++ b/experiments/smoke.py @@ -155,6 +155,39 @@ def check_state_conditioned_vectorizers(): print("CHECK0e state-conditioned vectorizers: zero-init and local calibration passed") +def check_direct_node_perturbation(): + """Unamortized q must update hidden weights without using A.""" + torch.manual_seed(19) + x = torch.randn(32, 5) + y = torch.randint(0, 3, (32,)) + yoh = onehot(y, 3) + net = SDILNet([5, 7, 7, 3], act="tanh", device="cpu", seed=10, + residual=True) + for weight in net.A: + weight.zero_() + weights_before = [weight.clone() for weight in net.W] + apical_before = [weight.clone() for weight in net.A] + rng_state = torch.get_rng_state() + fwd = net.forward(x) + targets = simultaneous_node_perturbation_targets( + net, x, y, sigma=0.01, n_dirs=2) + expected = [] + for layer, target in enumerate(targets): + delta = target * net.act_prime(fwd["u"][layer]) + if layer >= 1: + delta = net.res_alpha * delta + expected.append(delta.t() @ fwd["h"][layer] / x.shape[0]) + torch.set_rng_state(rng_state) + cfg = SDILConfig(eta=0.01, learn_A=False, learn_P=False, pert_every=1, + pert_ndirs=2, pert_mode="simultaneous", direct_node_pert=True) + sdil_step(net, x, y, yoh, cfg, step=0) + for layer in range(net.L - 1): + observed = net.W[layer] - weights_before[layer] + assert torch.allclose(observed, cfg.eta * expected[layer], atol=1e-6, rtol=1e-5) + assert all(torch.equal(before, after) for before, after in zip(apical_before, net.A)) + print("CHECK0f direct node perturbation: exact q update; A untouched") + + def main(): torch.manual_seed(0) dev = "cpu" @@ -163,6 +196,7 @@ def main(): check_topdown_predictor() check_traffic_seed_isolation() check_state_conditioned_vectorizers() + check_direct_node_perturbation() 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 64c7c79..6536c17 100644 --- a/sdil/core.py +++ b/sdil/core.py @@ -452,7 +452,7 @@ 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"): + raw_scale_control="none", direct_node_pert=False): self.eta = eta self.eta_A = eta_A self.eta_P = eta_P @@ -479,6 +479,7 @@ class SDILConfig: # Normalising each layer's delta to unit RMS decouples step size (set by # eta) from that noise -- like normalized SGD. self.normalize_delta = normalize_delta + self.direct_node_pert = direct_node_pert def sdil_step(net, x, y, y_onehot, cfg, step, prev_error=None): @@ -514,7 +515,8 @@ def sdil_step(net, x, y, y_onehot, cfg, step, prev_error=None): # Measuring q after mutating W would pair a post-update causal target # with a pre-update prediction, introducing an avoidable stale-target # error (especially at large learning rates). - did_pert = cfg.learn_A and (step % cfg.pert_every == 0) + did_pert = ((cfg.learn_A or cfg.direct_node_pert) + and step % cfg.pert_every == 0) qs = None if did_pert: estimator = (simultaneous_node_perturbation_targets @@ -524,9 +526,12 @@ def sdil_step(net, x, y, y_onehot, cfg, step, prev_error=None): # ================= forward weight updates ================= # hidden layers: three-factor local rule + teaching_list = qs if cfg.direct_node_pert else r_list + if cfg.direct_node_pert and qs is None: + raise RuntimeError("direct node perturbation requires a target on every update step") for l in range(net.L - 1): gain = net.act_prime(u[l]) # phi'(u_l) (B, n_l) - delta = r_list[l] * gain # (B, n_l) + delta = teaching_list[l] * gain # (B, n_l) if cfg.normalize_delta: delta = delta / (delta.pow(2).mean().sqrt() + 1e-8) # For an interior residual block h'=h+alpha*phi(Wh), the local @@ -546,7 +551,7 @@ def sdil_step(net, x, y, y_onehot, cfg, step, prev_error=None): _apply(net, net.L - 1, dWL, dbL, cfg) # ================= apical vectorizer A via node perturbation ======= - if did_pert: + if did_pert and cfg.learn_A: for l in range(net.L - 1): calibration_error = qs[l] - r_list[l] dA = calibration_error.t() @ c / B # (n_l, n_classes) diff --git a/sdil/probes.py b/sdil/probes.py index 0075bca..efa52e6 100644 --- a/sdil/probes.py +++ b/sdil/probes.py @@ -116,6 +116,26 @@ def fa_alignment_report(net, x, y, y_onehot): @torch.no_grad() +def nodepert_alignment_report(net, x, y, cfg): + """Alignment of the unamortized perturbation signal actually used.""" + grads, loss = true_hidden_grads(net, x, y) + estimator = (core.simultaneous_node_perturbation_targets + if cfg.pert_mode == "simultaneous" + else core.node_perturbation_targets) + targets = estimator(net, x, y, sigma=cfg.pert_sigma, n_dirs=cfg.pert_ndirs) + fwd = net.forward(x) + cosines = [] + q_norm = [] + g_norm = [] + for layer, (target, grad) in enumerate(zip(targets, grads)): + gain = net.act_prime(fwd["u"][layer]) + cosines.append(_row_cos(target * gain, -grad * gain)) + q_norm.append(target.norm(dim=1).mean().item()) + g_norm.append(grad.norm(dim=1).mean().item()) + return {"cos_q_negg": cosines, "q_norm": q_norm, "g_norm": g_norm, "loss": loss} + + +@torch.no_grad() def loss_decrease_ratio(net, x, y, y_onehot, cfg, step): """Single-step descent quality: apply one SDIL update to a scratch copy and measure the actual loss drop on the same batch, compared to one plain-SGD |
