summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
Diffstat (limited to 'experiments')
-rw-r--r--experiments/conv_local_smoke.py39
-rw-r--r--experiments/conv_run.py9
2 files changed, 46 insertions, 2 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)