summaryrefslogtreecommitdiff
path: root/sdil
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 06:30:48 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 06:30:48 -0500
commit4e00afb3c6151bede84908c1682b1e5649419fb0 (patch)
tree884f24ef2a6e86abd1ce04b37be8d2835531d420 /sdil
parentfe52db7f66db34045a178a793fcc7daf3413aea0 (diff)
oral-a: release antithetic hidden states eagerly
Diffstat (limited to 'sdil')
-rw-r--r--sdil/conv.py21
1 files 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