summaryrefslogtreecommitdiff
path: root/experiments/conv_run.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 06:17:50 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 06:17:50 -0500
commitdd705590a6210b6ada988cec0a392f7669c5cb52 (patch)
treef5787d098db0c2b32c4784e4a06fbe4e8ff17c9e /experiments/conv_run.py
parent200625df021c02e83aeecd496b4d4f5f3ffe8ad5 (diff)
oral-a: add translation-shared apical vectorizer
Diffstat (limited to 'experiments/conv_run.py')
-rw-r--r--experiments/conv_run.py9
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)