summaryrefslogtreecommitdiff
path: root/experiments/conv_local_smoke.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 06:17:50 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 06:17:50 -0500
commitdd705590a6210b6ada988cec0a392f7669c5cb52 (patch)
treef5787d098db0c2b32c4784e4a06fbe4e8ff17c9e /experiments/conv_local_smoke.py
parent200625df021c02e83aeecd496b4d4f5f3ffe8ad5 (diff)
oral-a: add translation-shared apical vectorizer
Diffstat (limited to 'experiments/conv_local_smoke.py')
-rw-r--r--experiments/conv_local_smoke.py39
1 files changed, 38 insertions, 1 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}