From dd705590a6210b6ada988cec0a392f7669c5cb52 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Wed, 22 Jul 2026 06:17:50 -0500 Subject: oral-a: add translation-shared apical vectorizer --- experiments/conv_local_smoke.py | 39 ++++++++++++++++++++++++++++++++++++++- experiments/conv_run.py | 9 ++++++++- 2 files changed, 46 insertions(+), 2 deletions(-) (limited to 'experiments') 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} diff --git a/experiments/conv_run.py b/experiments/conv_run.py index dd3f6c5..9b41dd2 100644 --- a/experiments/conv_run.py +++ b/experiments/conv_run.py @@ -102,7 +102,8 @@ def build(args): if args.mode == "bp": return CIFARLocalResNet(**common), None net = CIFARSDILResNet( - **common, a_scale=args.a_scale, apical_seed=args.apical_seed) + **common, a_scale=args.a_scale, apical_seed=args.apical_seed, + vectorizer_mode=args.vectorizer_mode) config = ConvSDILConfig( eta=args.lr, eta_output=args.output_lr, eta_A=args.eta_A, eta_P=args.eta_P, momentum=args.momentum, weight_decay=args.weight_decay, @@ -225,6 +226,9 @@ def run(args): "hidden_shapes": net.hidden_shapes, "forward_parameters": net.n_forward_parameters, "adaptive_apical_parameters": getattr(net, "n_apical_parameters", 0), + "vectorizer_parameters": getattr(net, "n_vectorizer_parameters", 0), + "predictor_parameters": getattr(net, "n_predictor_parameters", 0), + "vectorizer_mode": getattr(net, "vectorizer_mode", None), "fixed_traffic_coefficients": getattr( net, "n_fixed_traffic_coefficients", 0), }, @@ -454,6 +458,9 @@ def parse_args(): parser.add_argument("--bn_momentum", type=float, default=0.1) parser.add_argument("--bn_eps", type=float, default=1e-5) parser.add_argument("--a_scale", type=float, default=1.0) + parser.add_argument("--vectorizer_mode", + choices=("spatial_template", "channel_gated"), + default="spatial_template") parser.add_argument("--eta_A", type=float, default=0.01) parser.add_argument("--eta_P", type=float, default=0.01) parser.add_argument("--learn_P", type=int, choices=(0, 1), default=0) -- cgit v1.2.3