summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--experiments/conv_local_smoke.py39
-rw-r--r--experiments/conv_run.py9
-rw-r--r--sdil/conv.py81
3 files changed, 106 insertions, 23 deletions
diff --git a/experiments/conv_local_smoke.py b/experiments/conv_local_smoke.py
index d2428eb..99ca3b6 100644
--- a/experiments/conv_local_smoke.py
+++ b/experiments/conv_local_smoke.py
@@ -279,13 +279,46 @@ def apical_learning_checks():
targets = [torch.randn_like(value) * 0.01 for value in prediction]
before = sum(float((target - value).square().sum())
for target, value in zip(targets, prediction))
- net.calibrate_apical(output_signal, prediction, targets, eta=0.1)
+ net.calibrate_apical(
+ output_signal, clean["hiddens"], prediction, targets, eta=0.1)
after_prediction, _, _ = net.apical_components(
output_signal, clean["hiddens"], use_residual=True)
after = sum(float((target - value).square().sum())
for target, value in zip(targets, after_prediction))
assert after < before
+ gated = CIFARSDILResNet(
+ depth=8, base_width=2, seed=6, vectorizer_mode="channel_gated")
+ gated_clean = gated.forward(x)
+ gated_signal = (torch.softmax(gated_clean["logits"], dim=1)
+ - F.one_hot(y, 10))
+ gated_prediction, _, _ = gated.apical_components(
+ gated_signal, gated_clean["hiddens"], use_residual=True)
+ gated_targets = [torch.randn_like(value) * 0.01 for value in gated_prediction]
+ gated_before = sum(float((target - value).square().sum())
+ for target, value in zip(gated_targets, gated_prediction))
+ gated.calibrate_apical(
+ gated_signal, gated_clean["hiddens"], gated_prediction,
+ gated_targets, eta=0.1)
+ gated_after_prediction, _, _ = gated.apical_components(
+ gated_signal, gated_clean["hiddens"], use_residual=True)
+ gated_after = sum(float((target - value).square().sum())
+ for target, value in zip(gated_targets, gated_after_prediction))
+ assert gated_after < gated_before
+ shifted_hidden = [torch.roll(value, shifts=(3, -2), dims=(2, 3))
+ for value in gated_clean["hiddens"]]
+ shifted_instruction, _, _ = gated.apical_components(
+ gated_signal, shifted_hidden, use_residual=True)
+ original_instruction, _, _ = gated.apical_components(
+ gated_signal, gated_clean["hiddens"], use_residual=True)
+ assert all(torch.allclose(
+ shifted, torch.roll(original, shifts=(3, -2), dims=(2, 3)))
+ for shifted, original in zip(shifted_instruction, original_instruction))
+ spatial_56 = CIFARSDILResNet(depth=56, vectorizer_mode="spatial_template")
+ gated_56 = CIFARSDILResNet(depth=56, vectorizer_mode="channel_gated")
+ assert spatial_56.n_vectorizer_parameters == 5_324_800
+ assert gated_56.n_vectorizer_parameters == 40_640
+
predictor_net = CIFARSDILResNet(depth=8, base_width=2, seed=8)
hiddens = [torch.randn(64, *shape) for shape in predictor_net.hidden_shapes]
initial = predictor_net.predictor_step(hiddens, eta=0.1, nuisance_scale=0.5)
@@ -310,6 +343,10 @@ def apical_learning_checks():
assert all(not parameter.requires_grad
for parameter in net.W + [net.W_out, net.b_out])
return {"apical_mse_ratio": after / before,
+ "gated_apical_mse_ratio": gated_after / gated_before,
+ "gated_vectorizer_parameter_reduction": (
+ spatial_56.n_vectorizer_parameters
+ / gated_56.n_vectorizer_parameters),
"predictor_mse_ratio": final / initial}
diff --git a/experiments/conv_run.py b/experiments/conv_run.py
index dd3f6c5..9b41dd2 100644
--- a/experiments/conv_run.py
+++ b/experiments/conv_run.py
@@ -102,7 +102,8 @@ def build(args):
if args.mode == "bp":
return CIFARLocalResNet(**common), None
net = CIFARSDILResNet(
- **common, a_scale=args.a_scale, apical_seed=args.apical_seed)
+ **common, a_scale=args.a_scale, apical_seed=args.apical_seed,
+ vectorizer_mode=args.vectorizer_mode)
config = ConvSDILConfig(
eta=args.lr, eta_output=args.output_lr, eta_A=args.eta_A,
eta_P=args.eta_P, momentum=args.momentum, weight_decay=args.weight_decay,
@@ -225,6 +226,9 @@ def run(args):
"hidden_shapes": net.hidden_shapes,
"forward_parameters": net.n_forward_parameters,
"adaptive_apical_parameters": getattr(net, "n_apical_parameters", 0),
+ "vectorizer_parameters": getattr(net, "n_vectorizer_parameters", 0),
+ "predictor_parameters": getattr(net, "n_predictor_parameters", 0),
+ "vectorizer_mode": getattr(net, "vectorizer_mode", None),
"fixed_traffic_coefficients": getattr(
net, "n_fixed_traffic_coefficients", 0),
},
@@ -454,6 +458,9 @@ def parse_args():
parser.add_argument("--bn_momentum", type=float, default=0.1)
parser.add_argument("--bn_eps", type=float, default=1e-5)
parser.add_argument("--a_scale", type=float, default=1.0)
+ parser.add_argument("--vectorizer_mode",
+ choices=("spatial_template", "channel_gated"),
+ default="spatial_template")
parser.add_argument("--eta_A", type=float, default=0.01)
parser.add_argument("--eta_P", type=float, default=0.01)
parser.add_argument("--learn_P", type=int, choices=(0, 1), default=0)
diff --git a/sdil/conv.py b/sdil/conv.py
index 3b8fde1..40def21 100644
--- a/sdil/conv.py
+++ b/sdil/conv.py
@@ -414,21 +414,27 @@ class CIFARLocalResNet:
class CIFARSDILResNet(CIFARLocalResNet):
"""CIFAR local ResNet with per-unit apical vectorizers and predictors.
- Each spatial feature unit has a class-error vectorizer ``A_l``. A full
- spatial template is used rather than broadcasting one coefficient over a
- channel, because the true credit assigned to an early convolution is
- strongly position dependent. The predictor remains Harnett-faithful and
- diagonal: each unit fits its own affine soma--apical relation.
+ ``spatial_template`` gives every feature unit a class-error vectorizer.
+ ``channel_gated`` instead shares class coefficients across position and
+ obtains spatially heterogeneous credit through a local ``tanh(h)`` gate.
+ The latter respects convolutional translation sharing and sharply reduces
+ feedback parameters. The predictor remains Harnett-faithful and diagonal:
+ each unit fits its own affine soma--apical relation.
"""
- def __init__(self, *args, a_scale=1.0, apical_seed=None, **kwargs):
+ def __init__(self, *args, a_scale=1.0, apical_seed=None,
+ vectorizer_mode="spatial_template", **kwargs):
model_seed = kwargs.get("seed", 0)
super().__init__(*args, **kwargs)
generator = torch.Generator(device="cpu").manual_seed(
model_seed + 10007 if apical_seed is None else apical_seed)
nuisance_generator = torch.Generator(device="cpu").manual_seed(
model_seed + 20011 if apical_seed is None else apical_seed + 1)
+ if vectorizer_mode not in ("spatial_template", "channel_gated"):
+ raise ValueError(f"unknown convolutional vectorizer: {vectorizer_mode}")
+ self.vectorizer_mode = vectorizer_mode
self.A = []
+ self.A_gate = []
self.P = []
self.P_bias = []
self.Bnuis = []
@@ -438,8 +444,13 @@ class CIFARSDILResNet(CIFARLocalResNet):
# approximately as 1/(H*W). This scale keeps fixed-DFA controls
# finite while learned A remains free to change its gain.
std = a_scale / (height * width * math.sqrt(self.n_classes))
- self.A.append((torch.randn(units, self.n_classes, generator=generator)
- * std).to(device=self.device, dtype=self.dtype))
+ vectorizer_units = units if vectorizer_mode == "spatial_template" else channels
+ self.A.append((torch.randn(
+ vectorizer_units, self.n_classes, generator=generator)
+ * std).to(device=self.device, dtype=self.dtype))
+ if vectorizer_mode == "channel_gated":
+ self.A_gate.append(torch.zeros(
+ channels, self.n_classes, device=self.device, dtype=self.dtype))
shape = (channels, height, width)
self.P.append(torch.zeros(shape, device=self.device, dtype=self.dtype))
self.P_bias.append(torch.zeros(shape, device=self.device, dtype=self.dtype))
@@ -448,24 +459,40 @@ class CIFARSDILResNet(CIFARLocalResNet):
).to(device=self.device, dtype=self.dtype))
@property
- def n_apical_parameters(self):
+ def n_vectorizer_parameters(self):
return (sum(value.numel() for value in self.A)
- + sum(value.numel() for value in self.P)
+ + sum(value.numel() for value in self.A_gate))
+
+ @property
+ def n_predictor_parameters(self):
+ return (sum(value.numel() for value in self.P)
+ sum(value.numel() for value in self.P_bias))
@property
+ def n_apical_parameters(self):
+ return self.n_vectorizer_parameters + self.n_predictor_parameters
+
+ @property
def n_fixed_traffic_coefficients(self):
return sum(value.numel() for value in self.Bnuis)
@property
def apical_macs_per_example(self):
"""MACs for projecting one class-error vector to all hidden units."""
- return sum(value.numel() for value in self.A)
+ if self.vectorizer_mode == "spatial_template":
+ return sum(value.numel() for value in self.A)
+ projection = sum(value.numel() for value in self.A + self.A_gate)
+ gating = sum(math.prod(shape) for shape in self.hidden_shapes)
+ return projection + gating
- def instruction(self, index, output_signal):
+ def instruction(self, index, output_signal, hidden):
shape = self.hidden_shapes[index]
- return (output_signal @ self.A[index].t()).reshape(
- output_signal.shape[0], *shape)
+ if self.vectorizer_mode == "spatial_template":
+ return (output_signal @ self.A[index].t()).reshape(
+ output_signal.shape[0], *shape)
+ base = (output_signal @ self.A[index].t())[:, :, None, None]
+ gate = (output_signal @ self.A_gate[index].t())[:, :, None, None]
+ return base + torch.tanh(hidden) * gate
def apical_components(self, output_signal, hiddens, nuisance_scale=0.0,
use_residual=True):
@@ -476,7 +503,7 @@ class CIFARSDILResNet(CIFARLocalResNet):
raw_apical = []
innovations = []
for index, hidden in enumerate(hiddens):
- instruction = self.instruction(index, output_signal)
+ instruction = self.instruction(index, output_signal, hidden)
traffic = nuisance_scale * self.Bnuis[index] * hidden
raw = instruction + traffic
baseline = self.P[index] * hidden + self.P_bias[index]
@@ -506,19 +533,31 @@ class CIFARSDILResNet(CIFARLocalResNet):
return squared_error / units
@torch.no_grad()
- def calibrate_apical(self, output_signal, predicted_teaching, targets, eta):
+ def calibrate_apical(self, output_signal, hiddens, predicted_teaching,
+ targets, eta):
"""Local delta rule fitting innovation to causal perturbation targets."""
- if not (len(predicted_teaching) == len(targets) == self.n_hidden):
+ if not (len(hiddens) == len(predicted_teaching)
+ == len(targets) == self.n_hidden):
raise ValueError("calibration lists must cover every hidden population")
batch = output_signal.shape[0]
before_error = 0.0
target_power = 0.0
dot = 0.0
prediction_power = 0.0
- for index, (prediction, target) in enumerate(zip(predicted_teaching, targets)):
+ for index, (hidden, prediction, target) in enumerate(zip(
+ hiddens, predicted_teaching, targets)):
error = target - prediction
flat_error = error.flatten(1)
- self.A[index].add_(flat_error.t() @ output_signal / batch, alpha=eta)
+ if self.vectorizer_mode == "spatial_template":
+ self.A[index].add_(
+ flat_error.t() @ output_signal / batch, alpha=eta)
+ else:
+ spatial_error = error.mean(dim=(2, 3))
+ gated_error = (error * torch.tanh(hidden)).mean(dim=(2, 3))
+ self.A[index].add_(
+ spatial_error.t() @ output_signal / batch, alpha=eta)
+ self.A_gate[index].add_(
+ gated_error.t() @ output_signal / batch, alpha=eta)
before_error += float(error.square().sum())
target_power += float(target.square().sum())
prediction_power += float(prediction.square().sum())
@@ -667,7 +706,7 @@ def conv_local_step(net, x, y, config, step, generator=None):
calibration = None
if did_perturb and config.learn_A:
calibration = net.calibrate_apical(
- output_error, teaching, targets, config.eta_A)
+ output_error, forward["hiddens"], teaching, targets, config.eta_A)
predictor_mse = None
if config.learn_P:
predictor_mse = net.predictor_step(
@@ -705,7 +744,7 @@ def conv_apical_calibration_step(net, x, y, config, generator=None):
net, x, y, forward, sigma=config.pert_sigma,
n_directions=config.pert_directions, generator=generator)
calibration = net.calibrate_apical(
- output_error, teaching, targets, config.eta_A)
+ output_error, forward["hiddens"], teaching, targets, config.eta_A)
return float(F.cross_entropy(logits, y)), calibration