diff options
Diffstat (limited to 'experiments/conv_local_smoke.py')
| -rw-r--r-- | experiments/conv_local_smoke.py | 39 |
1 files changed, 38 insertions, 1 deletions
diff --git a/experiments/conv_local_smoke.py b/experiments/conv_local_smoke.py index d2428eb..99ca3b6 100644 --- a/experiments/conv_local_smoke.py +++ b/experiments/conv_local_smoke.py @@ -279,13 +279,46 @@ def apical_learning_checks(): targets = [torch.randn_like(value) * 0.01 for value in prediction] before = sum(float((target - value).square().sum()) for target, value in zip(targets, prediction)) - net.calibrate_apical(output_signal, prediction, targets, eta=0.1) + net.calibrate_apical( + output_signal, clean["hiddens"], prediction, targets, eta=0.1) after_prediction, _, _ = net.apical_components( output_signal, clean["hiddens"], use_residual=True) after = sum(float((target - value).square().sum()) for target, value in zip(targets, after_prediction)) assert after < before + gated = CIFARSDILResNet( + depth=8, base_width=2, seed=6, vectorizer_mode="channel_gated") + gated_clean = gated.forward(x) + gated_signal = (torch.softmax(gated_clean["logits"], dim=1) + - F.one_hot(y, 10)) + gated_prediction, _, _ = gated.apical_components( + gated_signal, gated_clean["hiddens"], use_residual=True) + gated_targets = [torch.randn_like(value) * 0.01 for value in gated_prediction] + gated_before = sum(float((target - value).square().sum()) + for target, value in zip(gated_targets, gated_prediction)) + gated.calibrate_apical( + gated_signal, gated_clean["hiddens"], gated_prediction, + gated_targets, eta=0.1) + gated_after_prediction, _, _ = gated.apical_components( + gated_signal, gated_clean["hiddens"], use_residual=True) + gated_after = sum(float((target - value).square().sum()) + for target, value in zip(gated_targets, gated_after_prediction)) + assert gated_after < gated_before + shifted_hidden = [torch.roll(value, shifts=(3, -2), dims=(2, 3)) + for value in gated_clean["hiddens"]] + shifted_instruction, _, _ = gated.apical_components( + gated_signal, shifted_hidden, use_residual=True) + original_instruction, _, _ = gated.apical_components( + gated_signal, gated_clean["hiddens"], use_residual=True) + assert all(torch.allclose( + shifted, torch.roll(original, shifts=(3, -2), dims=(2, 3))) + for shifted, original in zip(shifted_instruction, original_instruction)) + spatial_56 = CIFARSDILResNet(depth=56, vectorizer_mode="spatial_template") + gated_56 = CIFARSDILResNet(depth=56, vectorizer_mode="channel_gated") + assert spatial_56.n_vectorizer_parameters == 5_324_800 + assert gated_56.n_vectorizer_parameters == 40_640 + predictor_net = CIFARSDILResNet(depth=8, base_width=2, seed=8) hiddens = [torch.randn(64, *shape) for shape in predictor_net.hidden_shapes] initial = predictor_net.predictor_step(hiddens, eta=0.1, nuisance_scale=0.5) @@ -310,6 +343,10 @@ def apical_learning_checks(): assert all(not parameter.requires_grad for parameter in net.W + [net.W_out, net.b_out]) return {"apical_mse_ratio": after / before, + "gated_apical_mse_ratio": gated_after / gated_before, + "gated_vectorizer_parameter_reduction": ( + spatial_56.n_vectorizer_parameters + / gated_56.n_vectorizer_parameters), "predictor_mse_ratio": final / initial} |
