summaryrefslogtreecommitdiff
path: root/experiments/conv_run.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 13:47:35 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 13:47:35 -0500
commit1e4fbaf8e773509798387d19b806f6daf38f8b5c (patch)
tree43d5fe00aa7e1252e884c35cf3ae5a102799f073 /experiments/conv_run.py
parent203e75987d02cdb8c317e3302e1f05c9a4beb0e3 (diff)
baseline: add residual response mirroring
Diffstat (limited to 'experiments/conv_run.py')
-rw-r--r--experiments/conv_run.py39
1 files changed, 30 insertions, 9 deletions
diff --git a/experiments/conv_run.py b/experiments/conv_run.py
index fcb6690..3b91da9 100644
--- a/experiments/conv_run.py
+++ b/experiments/conv_run.py
@@ -19,7 +19,8 @@ from sdil.conv import (CIFARHierarchicalFAResNet, CIFARLocalResNet,
conv_hierarchical_step,
conv_learned_hierarchical_step, conv_local_step,
evaluate_conv, hierarchical_parameter_subspace_calibration)
-from sdil.conv import normalized_response_mirror_step
+from sdil.conv import (normalized_residual_mirror_step,
+ normalized_response_mirror_step)
from sdil.data import DATA_DIR, get_cifar_image_splits
@@ -106,7 +107,7 @@ 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"):
+ if args.mode in ("hfa", "lhfa", "wm", "rrm"):
net = CIFARHierarchicalFAResNet(
**common, feedback_seed=args.apical_seed,
feedback_scale=args.a_scale)
@@ -161,6 +162,7 @@ def work_report(net, mode, counters):
mirror_readout_macs = getattr(net, "R_out", torch.empty(0)).numel()
mirror_forward = (counters["mirror_conv_examples"] * mirror_conv_macs
+ counters["mirror_readout_examples"] * mirror_readout_macs)
+ mirror_prediction = mirror_forward if mode == "rrm" else 0
mirror_correlation = mirror_forward
components = {
"ordinary_forward_macs": normal_forward,
@@ -171,6 +173,7 @@ def work_report(net, mode, counters):
"apical_projection_macs": apical_inference,
"apical_regression_macs": apical_regression,
"mirror_response_macs": mirror_forward,
+ "mirror_feedback_prediction_macs": mirror_prediction,
"mirror_local_correlation_macs": mirror_correlation,
}
return {
@@ -206,7 +209,7 @@ def run(args):
raise ValueError("apical/predictor warmup is restricted to SDIL")
if args.mode == "lhfa" and args.learn_P:
raise ValueError("predictor learning is not defined for learned HFA")
- if args.mode != "wm" and args.mirror_warmup_steps:
+ if args.mode not in ("wm", "rrm") and args.mirror_warmup_steps:
raise ValueError("mirror warmup is restricted to weight mirror mode")
if args.mirror_every < 1 or args.mirror_batch_size < 1:
raise ValueError("invalid mirror cadence or batch size")
@@ -261,10 +264,11 @@ def run(args):
None if config is None or args.mode == "hfa" else
"hierarchical_feedback_parameters" if args.mode == "lhfa" else {
"wm": "local_parent_child_response",
+ "rrm": "local_parent_child_response_residual",
"unit_targets": "full_hidden_field",
"channel_subspace": "channel_basis_moments",
"vectorizer_subspace": "vectorizer_parameter_gradients",
- }[args.mode if args.mode == "wm"
+ }[args.mode if args.mode in ("wm", "rrm")
else config.apical_calibration_mode]),
"args": vars(args),
"provenance": provenance(),
@@ -291,7 +295,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") else 0),
+ if args.mode in ("lhfa", "wm", "rrm") else 0),
},
"epochs": [],
}
@@ -302,12 +306,15 @@ def run(args):
apical_warmup_wall = 0.0
mirror_warmup_wall = 0.0
loader_state = train.g.get_state().clone()
- if args.mode == "wm" and args.mirror_warmup_steps:
+ if args.mode in ("wm", "rrm") and args.mirror_warmup_steps:
sync(args.device)
mirror_start = time.time()
mirror_metrics = []
for _ in range(args.mirror_warmup_steps):
- metric = normalized_response_mirror_step(
+ mirror_function = (normalized_residual_mirror_step
+ if args.mode == "rrm"
+ else normalized_response_mirror_step)
+ metric = mirror_function(
net, batch_size=args.mirror_batch_size,
noise_std=args.mirror_noise_std, eta=args.mirror_eta,
generator=mirror_generator)
@@ -444,6 +451,20 @@ def run(args):
result = conv_hierarchical_step(net, x, y, config)
loss = result["loss"]
did_perturb = False
+ elif args.mode == "rrm":
+ if step % args.mirror_every == 0:
+ mirror_metric = normalized_residual_mirror_step(
+ net, batch_size=args.mirror_batch_size,
+ noise_std=args.mirror_noise_std, eta=args.mirror_eta,
+ generator=mirror_generator)
+ mirror_metrics.append(mirror_metric)
+ counters["mirror_conv_examples"] += args.mirror_batch_size
+ counters["mirror_readout_examples"] += mirror_metric[
+ "readout_batch_size"]
+ counters["mirror_events"] += 1
+ 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)
@@ -527,7 +548,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") else
+ if args.mode in ("hfa", "lhfa", "wm", "rrm") else
conv_alignment_report(
net, train.x[:probe], train.y[:probe], config))
sync(args.device)
@@ -576,7 +597,7 @@ def parse_args():
parser = argparse.ArgumentParser()
parser.add_argument(
"--mode", choices=(
- "bp", "dfa", "hfa", "lhfa", "wm", "sdil", "nodepert"),
+ "bp", "dfa", "hfa", "lhfa", "wm", "rrm", "sdil", "nodepert"),
required=True)
parser.add_argument("--out", required=True)
parser.add_argument("--device", default="cpu")