diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 06:31:37 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 06:31:37 -0500 |
| commit | bf9ecb6d5470dde5e610250fb1ac3a2ecc896c90 (patch) | |
| tree | 55cd510cb0586e97c9c4bfc1cbde49a56794ad3a /sdil | |
| parent | 4e00afb3c6151bede84908c1682b1e5649419fb0 (diff) | |
oral-a: reduce local-step signal residency
Diffstat (limited to 'sdil')
| -rw-r--r-- | sdil/conv.py | 21 |
1 files changed, 13 insertions, 8 deletions
diff --git a/sdil/conv.py b/sdil/conv.py index 878b5fc..279acac 100644 --- a/sdil/conv.py +++ b/sdil/conv.py @@ -681,6 +681,14 @@ def conv_local_step(net, x, y, config, step, generator=None): teaching, raw, innovations = net.apical_components( output_error, forward["hiddens"], config.nuisance_scale, config.use_residual) + total_units = sum(value.numel() for value in teaching) + teaching_rms = math.sqrt( + sum(float(value.square().sum()) for value in teaching) / total_units) + raw_apical_rms = math.sqrt( + sum(float(value.square().sum()) for value in raw) / total_units) + innovation_rms = math.sqrt( + sum(float(value.square().sum()) for value in innovations) / total_units) + del raw, innovations did_perturb = ((config.learn_A or config.direct_node_perturbation) and step % config.pert_every == 0) targets = None @@ -713,13 +721,9 @@ def conv_local_step(net, x, y, config, step, generator=None): "did_perturb": did_perturb, "calibration": calibration, "predictor_mse": predictor_mse, - "teaching_rms": math.sqrt(sum(float(value.square().sum()) for value in teaching) - / sum(value.numel() for value in teaching)), - "raw_apical_rms": math.sqrt(sum(float(value.square().sum()) for value in raw) - / sum(value.numel() for value in raw)), - "innovation_rms": math.sqrt( - sum(float(value.square().sum()) for value in innovations) - / sum(value.numel() for value in innovations)), + "teaching_rms": teaching_rms, + "raw_apical_rms": raw_apical_rms, + "innovation_rms": innovation_rms, } @@ -734,9 +738,10 @@ def conv_apical_calibration_step(net, x, y, config, generator=None): logits = forward["logits"] output_error = (torch.softmax(logits, dim=1) - F.one_hot(y, net.n_classes).to(logits.dtype)) - teaching, _, _ = net.apical_components( + teaching, raw, innovations = net.apical_components( output_error, forward["hiddens"], config.nuisance_scale, config.use_residual) + del raw, innovations targets = simultaneous_conv_node_perturbation( net, x, y, forward, sigma=config.pert_sigma, n_directions=config.pert_directions, generator=generator) |
