diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 14:15:09 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 14:33:50 -0500 |
| commit | d761f655301e3535368a87e9d4ffc26ab1ab8510 (patch) | |
| tree | 247439cc0f270cc78555ff72dcd15d35d9c754f6 /experiments/conv_run.py | |
| parent | 8a5c1045a43b1f91fd46b2fcc044b032cb14718c (diff) | |
baseline: add local reciprocal Kolen-Pollack
Diffstat (limited to 'experiments/conv_run.py')
| -rw-r--r-- | experiments/conv_run.py | 44 |
1 files changed, 34 insertions, 10 deletions
diff --git a/experiments/conv_run.py b/experiments/conv_run.py index 3b91da9..738c86d 100644 --- a/experiments/conv_run.py +++ b/experiments/conv_run.py @@ -12,13 +12,14 @@ import time import torch sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) -from sdil.conv import (CIFARHierarchicalFAResNet, CIFARLocalResNet, - CIFARSDILResNet, ConvSDILConfig, +from sdil.conv import (CIFARHierarchicalFAResNet, CIFARKPResNet, + CIFARLocalResNet, CIFARSDILResNet, ConvSDILConfig, conv_alignment_report, conv_apical_calibration_step, conv_hierarchical_alignment_report, - conv_hierarchical_step, + conv_hierarchical_step, conv_kolen_pollack_step, conv_learned_hierarchical_step, conv_local_step, - evaluate_conv, hierarchical_parameter_subspace_calibration) + evaluate_conv, hierarchical_feedback_tracking_report, + hierarchical_parameter_subspace_calibration) from sdil.conv import (normalized_residual_mirror_step, normalized_response_mirror_step) from sdil.data import DATA_DIR, get_cifar_image_splits @@ -107,8 +108,10 @@ def build(args): bn_momentum=args.bn_momentum, bn_eps=args.bn_eps) if args.mode == "bp": return CIFARLocalResNet(**common), None - if args.mode in ("hfa", "lhfa", "wm", "rrm"): - net = CIFARHierarchicalFAResNet( + if args.mode in ("hfa", "lhfa", "wm", "rrm", "kp"): + network_class = ( + CIFARKPResNet if args.mode == "kp" else CIFARHierarchicalFAResNet) + net = network_class( **common, feedback_seed=args.apical_seed, feedback_scale=args.a_scale) config = ConvSDILConfig( @@ -164,6 +167,8 @@ def work_report(net, mode, counters): + counters["mirror_readout_examples"] * mirror_readout_macs) mirror_prediction = mirror_forward if mode == "rrm" else 0 mirror_correlation = mirror_forward + kp_feedback_correlation = ( + counters["ordinary_examples"] * apical_macs if mode == "kp" else 0) components = { "ordinary_forward_macs": normal_forward, "warmup_clean_forward_macs": warmup_forward, @@ -175,6 +180,7 @@ def work_report(net, mode, counters): "mirror_response_macs": mirror_forward, "mirror_feedback_prediction_macs": mirror_prediction, "mirror_local_correlation_macs": mirror_correlation, + "kp_reciprocal_correlation_macs": kp_feedback_correlation, } return { "forward_macs_per_example": forward_macs, @@ -262,7 +268,8 @@ def run(args): "protocol_family": "oral_a_cifar_local_resnet_development", "calibration_metric_space": ( None if config is None or args.mode == "hfa" else - "hierarchical_feedback_parameters" if args.mode == "lhfa" else { + "hierarchical_feedback_parameters" if args.mode == "lhfa" else + "reciprocal_local_activity_products" if args.mode == "kp" else { "wm": "local_parent_child_response", "rrm": "local_parent_child_response_residual", "unit_targets": "full_hidden_field", @@ -295,7 +302,7 @@ def run(args): if args.mode == "hfa" else 0), "adaptive_feedback_parameters": ( getattr(net, "n_fixed_feedback_parameters", 0) - if args.mode in ("lhfa", "wm", "rrm") else 0), + if args.mode in ("lhfa", "wm", "rrm", "kp") else 0), }, "epochs": [], } @@ -465,6 +472,10 @@ def run(args): result = conv_hierarchical_step(net, x, y, config) loss = result["loss"] did_perturb = False + elif args.mode == "kp": + result = conv_kolen_pollack_step(net, x, y, config) + loss = result["loss"] + did_perturb = False else: result = conv_local_step( net, x, y, config, step, generator=perturb_generator) @@ -510,6 +521,18 @@ def run(args): key: sum(value[key] for value in mirror_metrics) / len(mirror_metrics) for key in mirror_metrics[0] } + if args.mode == "kp": + tracking = hierarchical_feedback_tracking_report(net) + record["feedback_tracking"] = { + "mean_feedback_forward_cosine": tracking[ + "mean_feedback_forward_cosine"], + "mean_feedback_forward_relative_error": tracking[ + "mean_feedback_forward_relative_error"], + "min_feedback_forward_cosine": min( + tracking["feedback_forward_cosine"]), + "max_feedback_forward_relative_error": max( + tracking["feedback_forward_relative_error"]), + } if args.eval_every and (epoch + 1) % args.eval_every == 0: sync(args.device) eval_start = time.time() @@ -548,7 +571,7 @@ def run(args): diagnostic_start = time.time() diagnostics = (conv_hierarchical_alignment_report( net, train.x[:probe], train.y[:probe]) - if args.mode in ("hfa", "lhfa", "wm", "rrm") else + if args.mode in ("hfa", "lhfa", "wm", "rrm", "kp") else conv_alignment_report( net, train.x[:probe], train.y[:probe], config)) sync(args.device) @@ -597,7 +620,8 @@ def parse_args(): parser = argparse.ArgumentParser() parser.add_argument( "--mode", choices=( - "bp", "dfa", "hfa", "lhfa", "wm", "rrm", "sdil", "nodepert"), + "bp", "dfa", "hfa", "lhfa", "wm", "rrm", "kp", "sdil", + "nodepert"), required=True) parser.add_argument("--out", required=True) parser.add_argument("--device", default="cpu") |
