summaryrefslogtreecommitdiff
path: root/experiments/run.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/run.py')
-rw-r--r--experiments/run.py26
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