diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 06:30:48 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 06:30:48 -0500 |
| commit | 4e00afb3c6151bede84908c1682b1e5649419fb0 (patch) | |
| tree | 884f24ef2a6e86abd1ce04b37be8d2835531d420 /sdil | |
| parent | fe52db7f66db34045a178a793fcc7daf3413aea0 (diff) | |
oral-a: release antithetic hidden states eagerly
Diffstat (limited to 'sdil')
| -rw-r--r-- | sdil/conv.py | 21 |
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 |
