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 /experiments | |
| parent | 7eb40a516849cbc3e93540d4fdf40366f6a8df9e (diff) | |
baseline: add unamortized node perturbation
Diffstat (limited to 'experiments')
| -rw-r--r-- | experiments/run.py | 26 | ||||
| -rw-r--r-- | experiments/smoke.py | 34 |
2 files changed, 58 insertions, 2 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)) |
