summaryrefslogtreecommitdiff
path: root/experiments/conv_run.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 13:21:20 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 13:21:20 -0500
commit225459f7ddfcd3c8d89f0515a0d766bf2a15c848 (patch)
tree329ec70d95000f00b358db70f8d6273cb6d9d8a9 /experiments/conv_run.py
parentf37644eb452f0eb3d7f368740f8b39493f11b2aa (diff)
algorithm: calibrate hierarchical feedback parameter subspaces
Diffstat (limited to 'experiments/conv_run.py')
-rw-r--r--experiments/conv_run.py63
1 files changed, 50 insertions, 13 deletions
diff --git a/experiments/conv_run.py b/experiments/conv_run.py
index 162a8be..fb759b4 100644
--- a/experiments/conv_run.py
+++ b/experiments/conv_run.py
@@ -16,7 +16,9 @@ from sdil.conv import (CIFARHierarchicalFAResNet, CIFARLocalResNet,
CIFARSDILResNet, ConvSDILConfig,
conv_alignment_report, conv_apical_calibration_step,
conv_hierarchical_alignment_report,
- conv_hierarchical_step, conv_local_step, evaluate_conv)
+ conv_hierarchical_step,
+ conv_learned_hierarchical_step, conv_local_step,
+ evaluate_conv, hierarchical_parameter_subspace_calibration)
from sdil.data import DATA_DIR, get_cifar_image_splits
@@ -103,14 +105,18 @@ 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":
+ if args.mode in ("hfa", "lhfa"):
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,
+ eta=args.lr, eta_output=args.output_lr,
+ eta_A=(args.eta_A if args.mode == "lhfa" else 0.0),
momentum=args.momentum, weight_decay=args.weight_decay,
- learn_A=False, learn_P=False)
+ learn_A=args.mode == "lhfa", learn_P=False,
+ pert_sigma=args.pert_sigma, pert_every=args.pert_every,
+ pert_directions=args.pert_directions,
+ apical_calibration_mode="hierarchical_parameter_subspace")
config.validate()
return net, config
net = CIFARSDILResNet(
@@ -146,7 +152,10 @@ def work_report(net, mode, counters):
local_correlation = counters["ordinary_examples"] * forward_macs
apical_inference = ((counters["ordinary_examples"]
+ counters["apical_warmup_examples"]) * apical_macs)
- apical_regression = counters["calibration_event_examples"] * apical_macs
+ regression_multiplier = 2 if mode == "lhfa" else 1
+ apical_regression = (regression_multiplier
+ * counters["calibration_event_examples"]
+ * apical_macs)
components = {
"ordinary_forward_macs": normal_forward,
"warmup_clean_forward_macs": warmup_forward,
@@ -182,8 +191,11 @@ def work_report(net, mode, counters):
def run(args):
if args.eval_split == "test" and args.eval_every:
raise ValueError("test protocols must use --eval_every 0 (one final evaluation)")
- if args.mode != "sdil" and (args.a_warmup_steps or args.learn_P):
+ if args.mode not in ("sdil", "lhfa") and (
+ args.a_warmup_steps or args.learn_P):
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")
torch.manual_seed(args.seed)
if str(args.device).startswith("cuda"):
if not torch.cuda.is_available():
@@ -225,7 +237,8 @@ def run(args):
"schema_version": 1,
"protocol_family": "oral_a_cifar_local_resnet_development",
"calibration_metric_space": (
- None if config is None or args.mode == "hfa" else {
+ None if config is None or args.mode == "hfa" else
+ "hierarchical_feedback_parameters" if args.mode == "lhfa" else {
"unit_targets": "full_hidden_field",
"channel_subspace": "channel_basis_moments",
"vectorizer_subspace": "vectorizer_parameter_gradients",
@@ -250,8 +263,12 @@ 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),
+ "fixed_feedback_parameters": (
+ getattr(net, "n_fixed_feedback_parameters", 0)
+ if args.mode == "hfa" else 0),
+ "adaptive_feedback_parameters": (
+ getattr(net, "n_fixed_feedback_parameters", 0)
+ if args.mode == "lhfa" else 0),
},
"epochs": [],
}
@@ -290,8 +307,21 @@ def run(args):
except StopIteration:
iterator = iter(train)
x, y = next(iterator)
- _, metric = conv_apical_calibration_step(
- net, x, y, config, generator=warmup_generator)
+ if args.mode == "lhfa":
+ forward = net.forward(
+ x, return_cache=True, training=True, update_stats=False)
+ output_signal = (
+ torch.softmax(forward["logits"], dim=1)
+ - torch.nn.functional.one_hot(
+ y, net.n_classes).to(forward["logits"].dtype))
+ metric = hierarchical_parameter_subspace_calibration(
+ net, x, y, forward, output_signal,
+ sigma=config.pert_sigma,
+ n_directions=config.pert_directions, eta=config.eta_A,
+ generator=warmup_generator)
+ else:
+ _, metric = conv_apical_calibration_step(
+ net, x, y, config, generator=warmup_generator)
warmup_metrics.append(metric)
batch = x.shape[0]
counters["apical_warmup_examples"] += batch
@@ -346,6 +376,13 @@ def run(args):
result = conv_hierarchical_step(net, x, y, config)
loss = result["loss"]
did_perturb = False
+ elif args.mode == "lhfa":
+ result = conv_learned_hierarchical_step(
+ net, x, y, config, step, generator=perturb_generator)
+ loss = result["loss"]
+ did_perturb = result["did_perturb"]
+ if result["calibration"] is not None:
+ calibration_metrics.append(result["calibration"])
else:
result = conv_local_step(
net, x, y, config, step, generator=perturb_generator)
@@ -424,7 +461,7 @@ def run(args):
diagnostic_start = time.time()
diagnostics = (conv_hierarchical_alignment_report(
net, train.x[:probe], train.y[:probe])
- if args.mode == "hfa" else
+ if args.mode in ("hfa", "lhfa") else
conv_alignment_report(
net, train.x[:probe], train.y[:probe], config))
sync(args.device)
@@ -471,7 +508,7 @@ def run(args):
def parse_args():
parser = argparse.ArgumentParser()
parser.add_argument(
- "--mode", choices=("bp", "dfa", "hfa", "sdil", "nodepert"),
+ "--mode", choices=("bp", "dfa", "hfa", "lhfa", "sdil", "nodepert"),
required=True)
parser.add_argument("--out", required=True)
parser.add_argument("--device", default="cpu")