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_run.py | 9 ++++++++- 1 file changed, 8 insertions(+), 1 deletion(-) (limited to 'experiments/conv_run.py') 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