From b323ddb1bd0c11e313c452a63707c01b4254b71a Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Wed, 22 Jul 2026 06:04:12 -0500 Subject: oral-a: add convolutional apical calibration --- experiments/conv_local_smoke.py | 86 ++++++++++++++++++++++++++++++++++++++++- 1 file changed, 85 insertions(+), 1 deletion(-) (limited to 'experiments/conv_local_smoke.py') 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") -- cgit v1.2.3