summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
Diffstat (limited to 'experiments')
-rw-r--r--experiments/run.py27
-rw-r--r--experiments/synthetic_smoke.py2
2 files changed, 24 insertions, 5 deletions
diff --git a/experiments/run.py b/experiments/run.py
index cb5de4e..b412b04 100644
--- a/experiments/run.py
+++ b/experiments/run.py
@@ -22,6 +22,7 @@ import torch
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from sdil.core import SDILNet, SDILConfig, sdil_step, neutral_p_update
from sdil.baselines import BPNet, dfa_config, evaluate
+from sdil.local_baselines import FANet
from sdil import probes
from sdil.data import (get_dataset, onehot, make_hierarchical,
make_teacher_student, make_tentmap)
@@ -54,6 +55,13 @@ def build(args, device):
predictor_mode=args.predictor_mode)
cfg = SDILConfig(eta=args.eta, momentum=args.momentum)
return net, cfg
+ if args.mode == "fa":
+ net = FANet(sizes, act=args.act, device=device, seed=args.seed,
+ w_scale=args.w_scale, nuis_rho=0.0,
+ residual=bool(args.residual),
+ predictor_mode=args.predictor_mode,
+ b_scale=args.feedback_scale)
+ return net, SDILConfig(eta=args.eta, momentum=args.momentum)
net = SDILNet(sizes, act=args.act, device=device, seed=args.seed,
w_scale=args.w_scale, a_scale=args.a_scale,
nuis_rho=args.nuis_rho, feedback=args.feedback,
@@ -144,13 +152,17 @@ def train(args):
yoh = onehot(y, n_out, device=device)
if args.mode == "bp":
loss = net.bp_step(x, y, cfg.eta, momentum=cfg.momentum)
+ elif args.mode == "fa":
+ loss = net.fa_step(x, y, yoh, cfg.eta, momentum=cfg.momentum)
else:
loss, aux = sdil_step(net, x, y, yoh, cfg, step, prev_error=prev_error)
prev_error = aux["error"]
if step % args.log_every == 0:
rec = {"step": step, "epoch": epoch, "train_loss": float(loss)}
- if args.mode != "bp":
+ if args.mode == "fa":
+ rec.update(probes.fa_alignment_report(net, px, py, poh))
+ elif args.mode != "bp":
al = probes.alignment_report(net, px, py, poh, cfg)
rec["cos_r_negg"] = al["cos_r_negg"]
rec["cos_innovation_negg"] = al["cos_innovation_negg"]
@@ -168,7 +180,11 @@ def train(args):
acc, tloss = evaluate(net, test_loader)
msg = f"[{args.tag}] epoch {epoch} step {step} loss {loss:.4f} test_acc {acc:.4f}"
- if args.mode != "bp":
+ if args.mode == "fa":
+ 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 != "bp":
al = probes.alignment_report(net, px, py, poh, cfg)
meancos = sum(al["cos_r_negg"]) / len(al["cos_r_negg"])
msg += f" mean_cos(r,-g) {meancos:+.3f} per-layer {['%.2f'%v for v in al['cos_r_negg']]}"
@@ -177,7 +193,9 @@ def train(args):
acc, tloss = evaluate(net, test_loader)
log["final"] = {"test_acc": acc, "test_loss": tloss, "wall_s": time.time() - t0}
- if args.mode != "bp":
+ if args.mode == "fa":
+ log["final"].update(probes.fa_alignment_report(net, px, py, poh))
+ elif args.mode != "bp":
al = probes.alignment_report(net, px, py, poh, cfg)
log["final"]["cos_r_negg"] = al["cos_r_negg"]
log["final"]["cos_innovation_negg"] = al["cos_innovation_negg"]
@@ -193,7 +211,7 @@ def train(args):
def get_args():
p = argparse.ArgumentParser()
- p.add_argument("--mode", default="sdil", choices=["bp", "dfa", "sdil"])
+ p.add_argument("--mode", default="sdil", choices=["bp", "fa", "dfa", "sdil"])
p.add_argument("--dataset", default="mnist",
choices=list(REAL_DATASETS + SYNTHETIC_DATASETS))
p.add_argument("--depth", type=int, default=3) # hidden layers
@@ -220,6 +238,7 @@ def get_args():
p.add_argument("--momentum", type=float, default=0.9)
p.add_argument("--w_scale", type=float, default=1.0)
p.add_argument("--a_scale", type=float, default=1.0)
+ p.add_argument("--feedback_scale", type=float, default=1.0)
p.add_argument("--pert_sigma", type=float, default=1e-2)
p.add_argument("--pert_every", type=int, default=4)
p.add_argument("--pert_ndirs", type=int, default=4)
diff --git a/experiments/synthetic_smoke.py b/experiments/synthetic_smoke.py
index c118813..cd6e2af 100644
--- a/experiments/synthetic_smoke.py
+++ b/experiments/synthetic_smoke.py
@@ -26,7 +26,7 @@ def check_task(dataset, expected_in, expected_out, extra):
x, y = next(iter(train))
assert x.shape == (16, n_in)
assert y.min() >= 0 and y.max() < n_out
- for mode in ("bp", "dfa", "sdil"):
+ for mode in ("bp", "fa", "dfa", "sdil"):
args.mode = mode
net, _ = build(args, "cpu")
assert net.logits(x).shape == (16, n_out)