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