From 4e00afb3c6151bede84908c1682b1e5649419fb0 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Wed, 22 Jul 2026 06:30:48 -0500 Subject: oral-a: release antithetic hidden states eagerly --- sdil/conv.py | 21 +++++++++------------ 1 file changed, 9 insertions(+), 12 deletions(-) diff --git a/sdil/conv.py b/sdil/conv.py index 0436525..878b5fc 100644 --- a/sdil/conv.py +++ b/sdil/conv.py @@ -596,22 +596,19 @@ def simultaneous_conv_node_perturbation(net, x, y, clean_forward, sigma=1e-2, diagnostic_derivatives = [] for _ in range(n_directions): directions = [] - plus_perturbations = [] - minus_perturbations = [] for hidden in clean_forward["hiddens"]: direction = torch.empty_like(hidden).bernoulli_( 0.5, generator=generator).mul_(2).sub_(1) directions.append(direction) - plus_perturbations.append(sigma * direction) - minus_perturbations.append(-sigma * direction) - plus_forward = net.forward( - x, perturbations=plus_perturbations, - training=True, update_stats=False) - minus_forward = net.forward( - x, perturbations=minus_perturbations, - training=True, update_stats=False) - plus = F.cross_entropy(plus_forward["logits"], y, reduction="none") - minus = F.cross_entropy(minus_forward["logits"], y, reduction="none") + # Build one signed intervention at a time and retain only its scalar + # losses. Keeping both complete hidden dictionaries would needlessly + # double peak memory at ResNet-56. + plus = F.cross_entropy(net.forward( + x, perturbations=[sigma * direction for direction in directions], + training=True, update_stats=False)["logits"], y, reduction="none") + minus = F.cross_entropy(net.forward( + x, perturbations=[-sigma * direction for direction in directions], + training=True, update_stats=False)["logits"], y, reduction="none") if net.normalization == "batchnorm": # BN couples examples. The per-example loss difference is not a # valid node-perturbation target because ell_i also responds to -- cgit v1.2.3