diff options
Diffstat (limited to 'experiments/run.py')
| -rw-r--r-- | experiments/run.py | 27 |
1 files changed, 23 insertions, 4 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) |
