diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 06:17:50 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 06:17:50 -0500 |
| commit | dd705590a6210b6ada988cec0a392f7669c5cb52 (patch) | |
| tree | f5787d098db0c2b32c4784e4a06fbe4e8ff17c9e /experiments/conv_run.py | |
| parent | 200625df021c02e83aeecd496b4d4f5f3ffe8ad5 (diff) | |
oral-a: add translation-shared apical vectorizer
Diffstat (limited to 'experiments/conv_run.py')
| -rw-r--r-- | experiments/conv_run.py | 9 |
1 files changed, 8 insertions, 1 deletions
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) |
