diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 13:21:20 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 13:21:20 -0500 |
| commit | 225459f7ddfcd3c8d89f0515a0d766bf2a15c848 (patch) | |
| tree | 329ec70d95000f00b358db70f8d6273cb6d9d8a9 /experiments | |
| parent | f37644eb452f0eb3d7f368740f8b39493f11b2aa (diff) | |
algorithm: calibrate hierarchical feedback parameter subspaces
Diffstat (limited to 'experiments')
| -rw-r--r-- | experiments/conv_local_smoke.py | 84 | ||||
| -rw-r--r-- | experiments/conv_run.py | 63 |
2 files changed, 134 insertions, 13 deletions
diff --git a/experiments/conv_local_smoke.py b/experiments/conv_local_smoke.py index 9207e04..894f9d3 100644 --- a/experiments/conv_local_smoke.py +++ b/experiments/conv_local_smoke.py @@ -1,5 +1,6 @@ #!/usr/bin/env python3 """Prove the convolutional local eligibility matches exact BP when instructed.""" +import math import os import sys @@ -11,6 +12,7 @@ from sdil.conv import (CIFARHierarchicalFAResNet, CIFARLocalResNet, CIFARSDILResNet, ConvSDILConfig, channel_subspace_apical_calibration, conv_hierarchical_step, conv_local_step, + hierarchical_parameter_subspace_calibration, simultaneous_conv_node_perturbation, vectorizer_subspace_apical_calibration) @@ -673,6 +675,87 @@ def hierarchical_feedback_checks(): } +def hierarchical_parameter_calibration_checks(): + """Audit the causal JVP and the exact local Q/R delta-rule moments.""" + torch.manual_seed(109) + net = CIFARHierarchicalFAResNet( + depth=8, base_width=2, seed=110, dtype=torch.float64, + normalization="batchnorm", residual_scale=1.0) + x = torch.randn(3, 3, 32, 32, dtype=torch.float64) + y = torch.tensor([0, 4, 7]) + parameters = net.W + net.gamma + net.beta + [net.W_out, net.b_out] + for parameter in parameters: + parameter.requires_grad_(True) + forward = net.forward( + x, return_cache=True, training=True, update_stats=False) + loss = F.cross_entropy(forward["logits"], y) + hidden_gradients = torch.autograd.grad(loss, forward["hiddens"]) + output_signal = (torch.softmax(forward["logits"].detach(), dim=1) + - F.one_hot(y, 10).to(torch.float64)) + _, diagnostic = hierarchical_parameter_subspace_calibration( + net, x, y, forward, output_signal, sigma=1e-5, + n_directions=1, eta=0.0, + generator=torch.Generator().manual_seed(111), + return_diagnostics=True) + directions = diagnostic["directions"][0]["hidden"] + exact_directional = x.shape[0] * sum( + (gradient * direction).sum() + for gradient, direction in zip(hidden_gradients, directions)) + estimated_directional = diagnostic[ + "directional_derivatives"][0]["scaled_directional"] + jvp_relative = float((estimated_directional - exact_directional).abs() + / exact_directional.abs().clamp_min(1e-30)) + assert jvp_relative < 3e-7 + for parameter in parameters: + parameter.requires_grad_(False) + + # Under an audit-only symmetric copy, the hierarchical field is the exact + # negative gradient. Consequently every local Q/R predicted moment equals + # its exact causal regression target, including option-A shortcut terms. + exact = CIFARHierarchicalFAResNet( + depth=8, base_width=2, seed=112, dtype=torch.float64, + normalization="batchnorm", residual_scale=1.0) + exact.Q = [value.clone() for value in exact.W] + exact.R_out.copy_(-exact.W_out.t()) + for parameter in exact.W + exact.gamma + exact.beta + [ + exact.W_out, exact.b_out]: + parameter.requires_grad_(True) + clean = exact.forward( + x, return_cache=True, training=True, update_stats=False) + gradients = torch.autograd.grad( + F.cross_entropy(clean["logits"], y), clean["hiddens"]) + negative = [-x.shape[0] * value.detach() for value in gradients] + signal = (torch.softmax(clean["logits"].detach(), dim=1) + - F.one_hot(y, 10).to(torch.float64)) + teaching, contexts, recipients = exact.hierarchical_teaching( + signal, clean, return_edge_contexts=True) + numerator = 0.0 + denominator = 0.0 + for index in range(1, len(exact.Q)): + recipient = recipients[index] + spec = exact.layer_specs[index] + spatial = (negative[recipient].shape[2] + * negative[recipient].shape[3]) + target = torch.nn.grad.conv2d_weight( + negative[recipient], exact.Q[index].shape, contexts[index], + stride=spec.stride, padding=spec.padding) / (x.shape[0] * spatial) + prediction = torch.nn.grad.conv2d_weight( + teaching[recipient], exact.Q[index].shape, contexts[index], + stride=spec.stride, padding=spec.padding) / (x.shape[0] * spatial) + numerator += float((target - prediction).square().sum()) + denominator += float(target.square().sum()) + target_r = negative[-1].mean(dim=(2, 3)).t() @ signal / x.shape[0] + prediction_r = teaching[-1].mean(dim=(2, 3)).t() @ signal / x.shape[0] + numerator += float((target_r - prediction_r).square().sum()) + denominator += float(target_r.square().sum()) + delta_rule_relative = math.sqrt(numerator / max(denominator, 1e-300)) + assert delta_rule_relative < 2e-12 + return { + "hierarchical_parameter_subspace_jvp_relative_error": jvp_relative, + "hierarchical_parameter_delta_rule_relative_error": delta_rule_relative, + } + + def apical_learning_checks(): torch.manual_seed(11) net = CIFARSDILResNet(depth=8, base_width=2, seed=6) @@ -779,6 +862,7 @@ def main(): report.update(channel_subspace_estimator_check()) report.update(vectorizer_subspace_estimator_check()) report.update(hierarchical_feedback_checks()) + report.update(hierarchical_parameter_calibration_checks()) report.update(apical_learning_checks()) print(report) print("ALL CONVOLUTIONAL LOCAL-ELIGIBILITY CHECKS PASSED") 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") |
