summaryrefslogtreecommitdiff
path: root/experiments/conv_run.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 14:15:09 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 14:33:50 -0500
commitd761f655301e3535368a87e9d4ffc26ab1ab8510 (patch)
tree247439cc0f270cc78555ff72dcd15d35d9c754f6 /experiments/conv_run.py
parent8a5c1045a43b1f91fd46b2fcc044b032cb14718c (diff)
baseline: add local reciprocal Kolen-Pollack
Diffstat (limited to 'experiments/conv_run.py')
-rw-r--r--experiments/conv_run.py44
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")