From 7ca9658999c8f882b3d5d295a4201cc9a1ce9cde Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Wed, 22 Jul 2026 13:10:14 -0500 Subject: baseline: add convolutional hierarchical feedback alignment --- experiments/conv_run.py | 35 +++++++++++++++++++++++++++++------ 1 file changed, 29 insertions(+), 6 deletions(-) (limited to 'experiments/conv_run.py') diff --git a/experiments/conv_run.py b/experiments/conv_run.py index 3c0de00..162a8be 100644 --- a/experiments/conv_run.py +++ b/experiments/conv_run.py @@ -12,9 +12,11 @@ import time import torch sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) -from sdil.conv import (CIFARLocalResNet, CIFARSDILResNet, ConvSDILConfig, +from sdil.conv import (CIFARHierarchicalFAResNet, CIFARLocalResNet, + CIFARSDILResNet, ConvSDILConfig, conv_alignment_report, conv_apical_calibration_step, - conv_local_step, evaluate_conv) + conv_hierarchical_alignment_report, + conv_hierarchical_step, conv_local_step, evaluate_conv) from sdil.data import DATA_DIR, get_cifar_image_splits @@ -101,6 +103,16 @@ def build(args): bn_momentum=args.bn_momentum, bn_eps=args.bn_eps) if args.mode == "bp": return CIFARLocalResNet(**common), None + if args.mode == "hfa": + net = CIFARHierarchicalFAResNet( + **common, feedback_seed=args.apical_seed, + feedback_scale=args.a_scale) + config = ConvSDILConfig( + eta=args.lr, eta_output=args.output_lr, eta_A=0.0, + momentum=args.momentum, weight_decay=args.weight_decay, + learn_A=False, learn_P=False) + config.validate() + return net, config net = CIFARSDILResNet( **common, a_scale=args.a_scale, apical_seed=args.apical_seed, vectorizer_mode=args.vectorizer_mode) @@ -213,7 +225,7 @@ def run(args): "schema_version": 1, "protocol_family": "oral_a_cifar_local_resnet_development", "calibration_metric_space": ( - None if config is None else { + None if config is None or args.mode == "hfa" else { "unit_targets": "full_hidden_field", "channel_subspace": "channel_basis_moments", "vectorizer_subspace": "vectorizer_parameter_gradients", @@ -238,6 +250,8 @@ def run(args): "vectorizer_mode": getattr(net, "vectorizer_mode", None), "fixed_traffic_coefficients": getattr( net, "n_fixed_traffic_coefficients", 0), + "fixed_feedback_parameters": getattr( + net, "n_fixed_feedback_parameters", 0), }, "epochs": [], } @@ -328,6 +342,10 @@ def run(args): x, y, lr, momentum=args.momentum, weight_decay=args.weight_decay) did_perturb = False + elif args.mode == "hfa": + result = conv_hierarchical_step(net, x, y, config) + loss = result["loss"] + did_perturb = False else: result = conv_local_step( net, x, y, config, step, generator=perturb_generator) @@ -404,8 +422,11 @@ def run(args): probe = min(args.alignment_probe, train.x.shape[0]) sync(args.device) diagnostic_start = time.time() - diagnostics = conv_alignment_report( - net, train.x[:probe], train.y[:probe], config) + diagnostics = (conv_hierarchical_alignment_report( + net, train.x[:probe], train.y[:probe]) + if args.mode == "hfa" else + conv_alignment_report( + net, train.x[:probe], train.y[:probe], config)) sync(args.device) diagnostics["wall_s"] = time.time() - diagnostic_start values = diagnostics["teaching_negative_gradient_cosine"] @@ -449,7 +470,9 @@ def run(args): def parse_args(): parser = argparse.ArgumentParser() - parser.add_argument("--mode", choices=("bp", "dfa", "sdil", "nodepert"), required=True) + parser.add_argument( + "--mode", choices=("bp", "dfa", "hfa", "sdil", "nodepert"), + required=True) parser.add_argument("--out", required=True) parser.add_argument("--device", default="cpu") parser.add_argument("--data_dir", default=DATA_DIR) -- cgit v1.2.3