summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--experiments/run.py26
-rw-r--r--experiments/smoke.py34
-rw-r--r--sdil/core.py13
-rw-r--r--sdil/probes.py20
4 files changed, 87 insertions, 6 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))
diff --git a/sdil/core.py b/sdil/core.py
index 64c7c79..6536c17 100644
--- a/sdil/core.py
+++ b/sdil/core.py
@@ -452,7 +452,7 @@ class SDILConfig:
pert_sigma=1e-2, pert_every=5, pert_ndirs=1, momentum=0.0, wd=0.0,
settle_steps=0, kappa=0.0, feedback="error", p_update_on_neutral=True,
normalize_delta=False, pert_mode="layerwise",
- raw_scale_control="none"):
+ raw_scale_control="none", direct_node_pert=False):
self.eta = eta
self.eta_A = eta_A
self.eta_P = eta_P
@@ -479,6 +479,7 @@ class SDILConfig:
# Normalising each layer's delta to unit RMS decouples step size (set by
# eta) from that noise -- like normalized SGD.
self.normalize_delta = normalize_delta
+ self.direct_node_pert = direct_node_pert
def sdil_step(net, x, y, y_onehot, cfg, step, prev_error=None):
@@ -514,7 +515,8 @@ def sdil_step(net, x, y, y_onehot, cfg, step, prev_error=None):
# Measuring q after mutating W would pair a post-update causal target
# with a pre-update prediction, introducing an avoidable stale-target
# error (especially at large learning rates).
- did_pert = cfg.learn_A and (step % cfg.pert_every == 0)
+ did_pert = ((cfg.learn_A or cfg.direct_node_pert)
+ and step % cfg.pert_every == 0)
qs = None
if did_pert:
estimator = (simultaneous_node_perturbation_targets
@@ -524,9 +526,12 @@ def sdil_step(net, x, y, y_onehot, cfg, step, prev_error=None):
# ================= forward weight updates =================
# hidden layers: three-factor local rule
+ teaching_list = qs if cfg.direct_node_pert else r_list
+ if cfg.direct_node_pert and qs is None:
+ raise RuntimeError("direct node perturbation requires a target on every update step")
for l in range(net.L - 1):
gain = net.act_prime(u[l]) # phi'(u_l) (B, n_l)
- delta = r_list[l] * gain # (B, n_l)
+ delta = teaching_list[l] * gain # (B, n_l)
if cfg.normalize_delta:
delta = delta / (delta.pow(2).mean().sqrt() + 1e-8)
# For an interior residual block h'=h+alpha*phi(Wh), the local
@@ -546,7 +551,7 @@ def sdil_step(net, x, y, y_onehot, cfg, step, prev_error=None):
_apply(net, net.L - 1, dWL, dbL, cfg)
# ================= apical vectorizer A via node perturbation =======
- if did_pert:
+ if did_pert and cfg.learn_A:
for l in range(net.L - 1):
calibration_error = qs[l] - r_list[l]
dA = calibration_error.t() @ c / B # (n_l, n_classes)
diff --git a/sdil/probes.py b/sdil/probes.py
index 0075bca..efa52e6 100644
--- a/sdil/probes.py
+++ b/sdil/probes.py
@@ -116,6 +116,26 @@ def fa_alignment_report(net, x, y, y_onehot):
@torch.no_grad()
+def nodepert_alignment_report(net, x, y, cfg):
+ """Alignment of the unamortized perturbation signal actually used."""
+ grads, loss = true_hidden_grads(net, x, y)
+ estimator = (core.simultaneous_node_perturbation_targets
+ if cfg.pert_mode == "simultaneous"
+ else core.node_perturbation_targets)
+ targets = estimator(net, x, y, sigma=cfg.pert_sigma, n_dirs=cfg.pert_ndirs)
+ fwd = net.forward(x)
+ cosines = []
+ q_norm = []
+ g_norm = []
+ for layer, (target, grad) in enumerate(zip(targets, grads)):
+ gain = net.act_prime(fwd["u"][layer])
+ cosines.append(_row_cos(target * gain, -grad * gain))
+ q_norm.append(target.norm(dim=1).mean().item())
+ g_norm.append(grad.norm(dim=1).mean().item())
+ return {"cos_q_negg": cosines, "q_norm": q_norm, "g_norm": g_norm, "loss": loss}
+
+
+@torch.no_grad()
def loss_decrease_ratio(net, x, y, y_onehot, cfg, step):
"""Single-step descent quality: apply one SDIL update to a scratch copy and
measure the actual loss drop on the same batch, compared to one plain-SGD