summaryrefslogtreecommitdiff
path: root/experiments/conv_local_smoke.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 06:04:12 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 06:04:12 -0500
commitb323ddb1bd0c11e313c452a63707c01b4254b71a (patch)
tree107b14af68925962d3af786072bc2d306428fad2 /experiments/conv_local_smoke.py
parentc0a12f4ad897f6f9f9e432d622cabdf46c40806f (diff)
oral-a: add convolutional apical calibration
Diffstat (limited to 'experiments/conv_local_smoke.py')
-rw-r--r--experiments/conv_local_smoke.py86
1 files changed, 85 insertions, 1 deletions
diff --git a/experiments/conv_local_smoke.py b/experiments/conv_local_smoke.py
index 35073a6..4a0f47d 100644
--- a/experiments/conv_local_smoke.py
+++ b/experiments/conv_local_smoke.py
@@ -7,7 +7,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.conv import CIFARLocalResNet
+from sdil.conv import (CIFARLocalResNet, CIFARSDILResNet, ConvSDILConfig,
+ conv_local_step, simultaneous_conv_node_perturbation)
def architecture_checks():
@@ -111,10 +112,93 @@ def perturbation_checks():
raise AssertionError("short perturbation list was accepted")
+def perturbation_estimator_check():
+ """Antithetic finite differences equal the simultaneous hidden JVP."""
+ torch.manual_seed(3)
+ batch = 2
+ net = CIFARSDILResNet(
+ depth=8, base_width=2, seed=4, dtype=torch.float64)
+ x = torch.randn(batch, 3, 32, 32, dtype=torch.float64)
+ y = torch.tensor([2, 7])
+ parameters = net.W + [net.W_out, net.b_out]
+ for parameter in parameters:
+ parameter.requires_grad_(True)
+ clean = net.forward(x, return_cache=True)
+ for hidden in clean["hiddens"]:
+ hidden.retain_grad()
+ F.cross_entropy(clean["logits"], y).backward()
+ generator = torch.Generator(device="cpu").manual_seed(99)
+ targets, diagnostics = simultaneous_conv_node_perturbation(
+ net, x, y, clean, sigma=1e-6, n_directions=1,
+ generator=generator, return_diagnostics=True)
+ directions = diagnostics["directions"][0]
+ finite_difference = diagnostics["directional_derivatives"][0]
+ exact = sum(
+ (batch * hidden.grad * direction).flatten(1).sum(dim=1)
+ for hidden, direction in zip(clean["hiddens"], directions))
+ relative = (finite_difference - exact).abs() / exact.abs().clamp_min(1e-12)
+ assert float(relative.max()) < 2e-6
+ for target, direction in zip(targets, directions):
+ expected = -finite_difference[:, None, None, None] * direction
+ assert torch.equal(target, expected)
+ for parameter in parameters:
+ parameter.requires_grad_(False)
+ return float(relative.max())
+
+
+def apical_learning_checks():
+ torch.manual_seed(11)
+ net = CIFARSDILResNet(depth=8, base_width=2, seed=6)
+ x = torch.randn(8, 3, 32, 32)
+ y = torch.arange(8) % 10
+ clean = net.forward(x, return_cache=True)
+ output_signal = (torch.softmax(clean["logits"], dim=1)
+ - F.one_hot(y, 10))
+ prediction, _, _ = net.apical_components(
+ output_signal, clean["hiddens"], use_residual=True)
+ 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)
+ 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
+
+ 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)
+ final = initial
+ for _ in range(60):
+ final = predictor_net.predictor_step(hiddens, eta=0.1, nuisance_scale=0.5)
+ assert final < initial * 1e-3
+
+ weights_before = [weight.clone() for weight in net.W]
+ result = conv_local_step(
+ net, x[:2], y[:2],
+ ConvSDILConfig(
+ eta=1e-3, eta_A=1e-3, momentum=0.0, weight_decay=0.0,
+ pert_every=1),
+ step=0, generator=torch.Generator(device="cpu").manual_seed(7))
+ assert result["did_perturb"] and result["calibration"] is not None
+ assert torch.isfinite(torch.tensor(list(
+ value for key, value in result.items()
+ if isinstance(value, float) and key != "predictor_mse"))).all()
+ assert any(not torch.equal(before_weight, after_weight)
+ for before_weight, after_weight in zip(weights_before, net.W))
+ assert all(not parameter.requires_grad
+ for parameter in net.W + [net.W_out, net.b_out])
+ return {"apical_mse_ratio": after / before,
+ "predictor_mse_ratio": final / initial}
+
+
def main():
architecture_checks()
perturbation_checks()
report = exact_local_gradient_check()
+ report["perturbation_jvp_max_relative_error"] = perturbation_estimator_check()
+ report.update(apical_learning_checks())
print(report)
print("ALL CONVOLUTIONAL LOCAL-ELIGIBILITY CHECKS PASSED")