From bf9ecb6d5470dde5e610250fb1ac3a2ecc896c90 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Wed, 22 Jul 2026 06:31:37 -0500 Subject: oral-a: reduce local-step signal residency --- sdil/conv.py | 21 +++++++++++++-------- 1 file 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) -- cgit v1.2.3