summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--experiments/protocol_smoke.py4
-rw-r--r--experiments/run.py79
-rw-r--r--experiments/smoke.py30
-rw-r--r--sdil/core.py65
4 files changed, 161 insertions, 17 deletions
diff --git a/experiments/protocol_smoke.py b/experiments/protocol_smoke.py
index 18de4c9..7ee9b4b 100644
--- a/experiments/protocol_smoke.py
+++ b/experiments/protocol_smoke.py
@@ -59,6 +59,9 @@ def main():
net, SDILConfig(pert_ndirs=2, pert_mode="simultaneous"))
layerwise = calibration_work_per_event(
net, SDILConfig(pert_ndirs=2, pert_mode="layerwise"))
+ overridden = calibration_work_per_event(
+ net, SDILConfig(pert_ndirs=1, pert_mode="simultaneous"),
+ pert_mode="layerwise", pert_ndirs=2)
assert simultaneous == {
"batch_loss_evaluations": 4,
"forward_equivalent_batches": 5.0,
@@ -66,6 +69,7 @@ def main():
}
assert layerwise["batch_loss_evaluations"] == 9
assert abs(layerwise["forward_equivalent_batches"] - 11.0 / 3.0) < 1e-12
+ assert overridden == layerwise
lesion_net = SDILNet([1, 4, 4, 4, 4, 2], act="relu", residual=True, device="cpu")
lesion_x = torch.linspace(0, 1, 16).view(-1, 1)
diff --git a/experiments/run.py b/experiments/run.py
index ea19366..ced7585 100644
--- a/experiments/run.py
+++ b/experiments/run.py
@@ -22,7 +22,8 @@ import torch
import torch.nn.functional as F
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.core import (SDILNet, SDILConfig, apical_calibration_step, sdil_step,
+ neutral_p_update)
from sdil.baselines import BPNet, dfa_config, evaluate
from sdil.local_baselines import FANet
from sdil import probes
@@ -156,15 +157,16 @@ def residual_lesion_report(net, evaluation_loader, probe_x, fraction):
}
-def calibration_work_per_event(net, cfg):
+def calibration_work_per_event(net, cfg, pert_mode=None, pert_ndirs=None):
"""Hardware-independent causal-calibration work for one minibatch.
Forward equivalents use affine multiply-add work as the denominator. The
returned loss evaluations count calls producing per-example scalar losses;
the caller multiplies by the actual minibatch size for scalar observations.
"""
- directions = cfg.pert_ndirs
- if cfg.pert_mode == "simultaneous":
+ directions = cfg.pert_ndirs if pert_ndirs is None else pert_ndirs
+ mode = cfg.pert_mode if pert_mode is None else pert_mode
+ if mode == "simultaneous":
return {
"batch_loss_evaluations": 2 * directions,
# The implementation performs one duplicate clean forward plus one
@@ -341,6 +343,8 @@ def train(args):
calibration_example_loss_evaluations = 0
calibration_forward_equivalent_examples = 0.0
perturbation_batch_expansion = 0
+ feedback_warmup_examples = 0
+ feedback_warmup_events = 0
# predictor warmup on neutral-period (c=0) drive, so P cancels the apical
# nuisance before task plasticity relies on the residual (no-op when rho=0).
@@ -361,6 +365,63 @@ def train(args):
warmup_wall_s = time.time() - warmup_t0
else:
warmup_wall_s = 0.0
+
+ # Feedback-first timescale separation. Causal perturbations fit A on a
+ # stationary forward network before noisy predictions can move W into a
+ # bad basin. The output readout is frozen as well. This is supervised
+ # calibration work, not free initialization, so every forward and scalar
+ # loss observation is included in the same hardware-independent ledger.
+ if args.a_warmup_steps > 0:
+ if args.mode != "sdil" or not args.learn_A:
+ raise ValueError("--a_warmup_steps requires SDIL with --learn_A 1")
+ device_sync(device)
+ feedback_warmup_t0 = time.time()
+ loader_state = (train_loader.g.get_state().clone()
+ if hasattr(train_loader, "g") else None)
+ rng_devices = ([torch.cuda.current_device()]
+ if str(device).startswith("cuda") and torch.cuda.is_available()
+ else [])
+ # Warmup length must not silently change the minibatch order or random
+ # directions used by the subsequent joint phase. fork_rng and restoring
+ # the loader generator isolate those nuisance differences while keeping
+ # the learned A parameters.
+ try:
+ with torch.random.fork_rng(devices=rng_devices):
+ it = iter(train_loader)
+ for _ in range(args.a_warmup_steps):
+ try:
+ wx, wy = next(it)
+ except StopIteration:
+ it = iter(train_loader)
+ wx, wy = next(it)
+ wx, wy = wx.to(device), wy.to(device)
+ wyoh = onehot(wy, n_out, device=device)
+ apical_calibration_step(
+ net, wx, wy, wyoh, cfg,
+ pert_mode=args.a_warmup_mode,
+ pert_ndirs=args.a_warmup_ndirs)
+ feedback_warmup_examples += wx.shape[0]
+ feedback_warmup_events += 1
+ event = calibration_work_per_event(
+ net, cfg, pert_mode=args.a_warmup_mode,
+ pert_ndirs=args.a_warmup_ndirs)
+ perturbation_events += 1
+ calibration_batch_loss_evaluations += event["batch_loss_evaluations"]
+ calibration_example_loss_evaluations += (
+ event["batch_loss_evaluations"] * wx.shape[0])
+ calibration_forward_equivalent_examples += (
+ event["forward_equivalent_batches"] * wx.shape[0])
+ perturbation_batch_expansion = max(
+ perturbation_batch_expansion,
+ event["perturbation_batch_expansion"])
+ finally:
+ if loader_state is not None:
+ train_loader.g.set_state(loader_state)
+ device_sync(device)
+ feedback_warmup_wall_s = time.time() - feedback_warmup_t0
+ else:
+ feedback_warmup_wall_s = 0.0
+
for epoch in range(args.epochs):
device_sync(device)
train_t0 = time.time()
@@ -524,6 +585,7 @@ def train(args):
log["final"]["wall_s"] = time.time() - t0
log["timing"] = {
"warmup_wall_s": warmup_wall_s,
+ "feedback_warmup_wall_s": feedback_warmup_wall_s,
"training_loop_wall_s": train_wall_s,
"diagnostics_wall_s": diagnostics_wall_s,
"inline_diagnostics_wall_s": inline_diagnostics_wall_s,
@@ -535,12 +597,14 @@ def train(args):
"train_steps": step,
"ordinary_training_forward_examples": ordinary_forward_examples,
"predictor_warmup_forward_examples": warmup_examples,
+ "feedback_warmup_forward_examples": feedback_warmup_examples,
+ "feedback_warmup_perturbation_events": feedback_warmup_events,
"perturbation_events": perturbation_events,
"calibration_batch_loss_evaluations": calibration_batch_loss_evaluations,
"calibration_example_loss_evaluations": calibration_example_loss_evaluations,
"calibration_forward_equivalent_examples": calibration_forward_equivalent_examples,
"training_forward_equivalent_examples": (
- ordinary_forward_examples + warmup_examples
+ ordinary_forward_examples + warmup_examples + feedback_warmup_examples
+ calibration_forward_equivalent_examples),
"max_perturbation_batch_expansion": perturbation_batch_expansion,
}
@@ -602,6 +666,11 @@ def get_args():
p.add_argument("--p_neutral", type=int, default=1) # P update on neutral (c=0) drive
p.add_argument("--p_warmup_steps", type=int, default=200) # pre-task neutral P warmup
p.add_argument("--p_warmup_eta", type=float, default=0.05)
+ p.add_argument("--a_warmup_steps", type=int, default=0,
+ help="A-only causal-calibration prefix with all forward weights frozen")
+ p.add_argument("--a_warmup_ndirs", type=int, default=4)
+ p.add_argument("--a_warmup_mode", default="layerwise",
+ choices=["layerwise", "simultaneous"])
p.add_argument("--nuis_rho", type=float, default=0.0)
p.add_argument("--traffic_seed", type=int, default=1234,
help="ordinary apical-traffic projection seed; independent of model/data seeds")
diff --git a/experiments/smoke.py b/experiments/smoke.py
index a49e4c6..9971800 100644
--- a/experiments/smoke.py
+++ b/experiments/smoke.py
@@ -6,7 +6,7 @@ import sys
import torch
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
-from sdil.core import (SDILNet, SDILConfig, sdil_step,
+from sdil.core import (SDILNet, SDILConfig, apical_calibration_step, sdil_step,
node_perturbation_targets, simultaneous_node_perturbation_targets,
neutral_p_update, teaching_signal)
from sdil.baselines import dfa_config
@@ -188,6 +188,33 @@ def check_direct_node_perturbation():
print("CHECK0f direct node perturbation: exact q update; A untouched")
+def check_feedback_first_calibration():
+ """A-only calibration must leave the entire forward network unchanged."""
+ torch.manual_seed(23)
+ x = torch.randn(32, 5)
+ y = torch.randint(0, 3, (32,))
+ yoh = onehot(y, 3)
+ net = SDILNet([5, 7, 7, 3], act="relu", device="cpu", seed=11,
+ residual=True, vectorizer_mode="context_gated")
+ weights_before = [weight.clone() for weight in net.W]
+ biases_before = [bias.clone() for bias in net.b]
+ apical_before = [weight.clone() for weight in net.A]
+ gates_before = [weight.clone() for weight in net.A_gate]
+ cfg = SDILConfig(eta_A=0.02, learn_A=True, learn_P=True,
+ pert_ndirs=1, pert_mode="simultaneous")
+ apical_calibration_step(
+ net, x, y, yoh, cfg, pert_mode="simultaneous", pert_ndirs=2)
+ assert all(torch.equal(before, after)
+ for before, after in zip(weights_before, net.W))
+ assert all(torch.equal(before, after)
+ for before, after in zip(biases_before, net.b))
+ assert any(not torch.equal(before, after)
+ for before, after in zip(apical_before, net.A))
+ assert any(not torch.equal(before, after)
+ for before, after in zip(gates_before, net.A_gate))
+ print("CHECK0g feedback-first calibration: A changed; W/readout frozen")
+
+
def main():
torch.manual_seed(0)
dev = "cpu"
@@ -197,6 +224,7 @@ def main():
check_traffic_seed_isolation()
check_state_conditioned_vectorizers()
check_direct_node_perturbation()
+ check_feedback_first_calibration()
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 6536c17..20e94eb 100644
--- a/sdil/core.py
+++ b/sdil/core.py
@@ -482,6 +482,59 @@ class SDILConfig:
self.direct_node_pert = direct_node_pert
+def _update_apical_vectorizer(net, h, c, r_list, qs, eta_A):
+ """Apply the local causal-regression update to all apical pathways."""
+ B = c.shape[0]
+ context = h[-2]
+ for l in range(net.L - 1):
+ calibration_error = qs[l] - r_list[l]
+ dA = calibration_error.t() @ c / B
+ net.A[l] += eta_A * dA
+ if net.vectorizer_mode == "soma_gated":
+ dgate = (calibration_error * torch.tanh(h[l + 1])).t() @ c / B
+ net.A_gate[l] += eta_A * dgate
+ elif net.vectorizer_mode == "context_gated":
+ features = (c.unsqueeze(2) * torch.tanh(context).unsqueeze(1)).flatten(1)
+ dgate = calibration_error.t() @ features / B
+ net.A_gate[l] += eta_A * dgate
+
+
+def apical_calibration_step(net, x, y, y_onehot, cfg,
+ pert_mode=None, pert_ndirs=None):
+ """Calibrate the apical vectorizer while keeping all forward weights fixed.
+
+ This is the feedback-first half of a two-timescale protocol. It uses labels
+ and forward-only causal interventions exactly like an ordinary perturbation
+ event, but performs no hidden or output weight update. The caller controls
+ and accounts for the number of prefix steps and perturbation work.
+ """
+ if not cfg.learn_A:
+ raise ValueError("apical calibration requires learn_A=True")
+ if cfg.feedback != "error":
+ raise ValueError("feedback-first calibration currently requires error feedback")
+ mode = cfg.pert_mode if pert_mode is None else pert_mode
+ directions = cfg.pert_ndirs if pert_ndirs is None else pert_ndirs
+ if mode not in ("layerwise", "simultaneous"):
+ raise ValueError(f"unknown perturbation mode: {mode}")
+ if directions < 1:
+ raise ValueError("perturbation directions must be positive")
+
+ with torch.no_grad():
+ fwd = net.forward(x)
+ h = fwd["h"]
+ logits = h[-1]
+ loss = loss_ce(logits, y).mean().item()
+ c = net.output_error(logits, y_onehot)
+ context = h[-2]
+ r_list = [teaching_signal(net, l, c, h[l + 1], cfg, context)[0]
+ for l in range(net.L - 1)]
+ estimator = (simultaneous_node_perturbation_targets
+ if mode == "simultaneous" else node_perturbation_targets)
+ qs = estimator(net, x, y, sigma=cfg.pert_sigma, n_dirs=directions)
+ _update_apical_vectorizer(net, h, c, r_list, qs, cfg.eta_A)
+ return loss, {"error": c, "did_pert": True}
+
+
def sdil_step(net, x, y, y_onehot, cfg, step, prev_error=None):
"""One SDIL training step (minibatch). Returns (loss, aux dict). No autograd."""
with torch.no_grad():
@@ -552,17 +605,7 @@ def sdil_step(net, x, y, y_onehot, cfg, step, prev_error=None):
# ================= apical vectorizer A via node perturbation =======
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)
- net.A[l] += cfg.eta_A * dA
- if net.vectorizer_mode == "soma_gated":
- dgate = (calibration_error * torch.tanh(h[l + 1])).t() @ c / B
- net.A_gate[l] += cfg.eta_A * dgate
- elif net.vectorizer_mode == "context_gated":
- features = (c.unsqueeze(2) * torch.tanh(context).unsqueeze(1)).flatten(1)
- dgate = calibration_error.t() @ features / B
- net.A_gate[l] += cfg.eta_A * dgate
+ _update_apical_vectorizer(net, h, c, r_list, qs, cfg.eta_A)
# ================= predictor P (neutral) ==========================
# KEY identification condition. P must learn the soma->apical coupling