From 0db651ac83608ea8f5b94950a391da22f076a142 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Wed, 22 Jul 2026 13:33:33 -0500 Subject: baseline: add normalized local response mirroring --- experiments/conv_local_smoke.py | 49 ++++++++++++++++++++++++ experiments/conv_run.py | 84 ++++++++++++++++++++++++++++++++++++++--- 2 files changed, 128 insertions(+), 5 deletions(-) (limited to 'experiments') diff --git a/experiments/conv_local_smoke.py b/experiments/conv_local_smoke.py index 894f9d3..7d24154 100644 --- a/experiments/conv_local_smoke.py +++ b/experiments/conv_local_smoke.py @@ -12,7 +12,9 @@ from sdil.conv import (CIFARHierarchicalFAResNet, CIFARLocalResNet, CIFARSDILResNet, ConvSDILConfig, channel_subspace_apical_calibration, conv_hierarchical_step, conv_local_step, + hierarchical_mirror_observations, hierarchical_parameter_subspace_calibration, + normalized_response_mirror_update, simultaneous_conv_node_perturbation, vectorizer_subspace_apical_calibration) @@ -756,6 +758,52 @@ def hierarchical_parameter_calibration_checks(): } +def normalized_response_mirror_checks(): + """Audit local response estimation and absence of W access in the update.""" + net = CIFARHierarchicalFAResNet( + depth=8, base_width=2, seed=121, dtype=torch.float64, + normalization="batchnorm") + observations = hierarchical_mirror_observations( + net, batch_size=16, noise_std=1.0, + generator=torch.Generator().manual_seed(122)) + metrics, _ = normalized_response_mirror_update( + net, observations, eta=1.0) + pairs = list(zip(net.Q[1:], net.W[1:])) + [ + (net.R_out, -net.W_out.t())] + cosines = [float(F.cosine_similarity( + feedback.flatten(), target.flatten(), dim=0)) + for feedback, target in pairs] + norm_ratios = [float(feedback.norm() / target.norm()) + for feedback, target in pairs] + assert sum(cosines) / len(cosines) > 0.985 + assert min(cosines) > 0.95 + assert min(norm_ratios) > 0.90 and max(norm_ratios) < 1.10 + + # The update consumes observations only. Changing every forward parameter + # after those observations were generated must not change the Q/R update. + left = CIFARHierarchicalFAResNet( + depth=8, base_width=2, seed=123, dtype=torch.float64) + right = CIFARHierarchicalFAResNet( + depth=8, base_width=2, seed=123, dtype=torch.float64) + shared_observations = hierarchical_mirror_observations( + left, batch_size=2, generator=torch.Generator().manual_seed(124)) + for value in right.W + [right.W_out]: + value.normal_(generator=torch.Generator().manual_seed(value.numel())) + normalized_response_mirror_update(left, shared_observations, eta=0.2) + normalized_response_mirror_update(right, shared_observations, eta=0.2) + independence_error = max(float((a - b).abs().max()) for a, b in zip( + left.Q[1:] + [left.R_out], right.Q[1:] + [right.R_out])) + assert independence_error == 0.0 + return { + "mirror_estimate_mean_forward_cosine": sum(cosines) / len(cosines), + "mirror_estimate_min_forward_cosine": min(cosines), + "mirror_estimate_min_norm_ratio": min(norm_ratios), + "mirror_estimate_max_norm_ratio": max(norm_ratios), + "mirror_update_forward_parameter_independence_error": independence_error, + "mirror_update_rms": metrics["mirror_update_rms"], + } + + def apical_learning_checks(): torch.manual_seed(11) net = CIFARSDILResNet(depth=8, base_width=2, seed=6) @@ -863,6 +911,7 @@ def main(): report.update(vectorizer_subspace_estimator_check()) report.update(hierarchical_feedback_checks()) report.update(hierarchical_parameter_calibration_checks()) + report.update(normalized_response_mirror_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 fb759b4..fcb6690 100644 --- a/experiments/conv_run.py +++ b/experiments/conv_run.py @@ -19,6 +19,7 @@ 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.data import DATA_DIR, get_cifar_image_splits @@ -105,7 +106,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"): + if args.mode in ("hfa", "lhfa", "wm"): net = CIFARHierarchicalFAResNet( **common, feedback_seed=args.apical_seed, feedback_scale=args.a_scale) @@ -156,6 +157,11 @@ def work_report(net, mode, counters): apical_regression = (regression_multiplier * counters["calibration_event_examples"] * apical_macs) + mirror_conv_macs = max(0, apical_macs - getattr(net, "R_out", torch.empty(0)).numel()) + 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_correlation = mirror_forward components = { "ordinary_forward_macs": normal_forward, "warmup_clean_forward_macs": warmup_forward, @@ -164,6 +170,8 @@ def work_report(net, mode, counters): "local_weight_correlation_macs": local_correlation, "apical_projection_macs": apical_inference, "apical_regression_macs": apical_regression, + "mirror_response_macs": mirror_forward, + "mirror_local_correlation_macs": mirror_correlation, } return { "forward_macs_per_example": forward_macs, @@ -180,6 +188,8 @@ def work_report(net, mode, counters): "logical_batch_loss_queries": counters["logical_batch_loss_queries"], "causal_scalar_observations": counters["causal_scalar_observations"], "per_example_cross_entropy_terms": counters["per_example_loss_terms"], + "mirror_probe_examples": counters["mirror_conv_examples"], + "mirror_readout_probe_examples": counters["mirror_readout_examples"], "definition": ( "multiply-accumulates in conv/linear maps; one local weight correlation " "equals one forward-weight MAC count; BP reverse is estimated as one " @@ -196,6 +206,12 @@ 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: + 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") + if not 0.0 < args.mirror_eta <= 1.0 or args.mirror_noise_std <= 0: + raise ValueError("invalid mirror learning hyperparameters") torch.manual_seed(args.seed) if str(args.device).startswith("cuda"): if not torch.cuda.is_available(): @@ -221,6 +237,8 @@ def run(args): args.perturb_seed) warmup_generator = torch.Generator(device=torch.device(args.device)).manual_seed( args.perturb_seed + 1) + mirror_generator = torch.Generator(device=torch.device(args.device)).manual_seed( + args.mirror_seed) counters = { "ordinary_examples": 0, @@ -232,6 +250,9 @@ def run(args): "causal_scalar_observations": 0, "per_example_loss_terms": 0, "perturbation_events": 0, + "mirror_conv_examples": 0, + "mirror_readout_examples": 0, + "mirror_events": 0, } log = { "schema_version": 1, @@ -239,10 +260,12 @@ def run(args): "calibration_metric_space": ( None if config is None or args.mode == "hfa" else "hierarchical_feedback_parameters" if args.mode == "lhfa" else { + "wm": "local_parent_child_response", "unit_targets": "full_hidden_field", "channel_subspace": "channel_basis_moments", "vectorizer_subspace": "vectorizer_parameter_gradients", - }[config.apical_calibration_mode]), + }[args.mode if args.mode == "wm" + else config.apical_calibration_mode]), "args": vars(args), "provenance": provenance(), "split": split, @@ -268,7 +291,7 @@ def run(args): if args.mode == "hfa" else 0), "adaptive_feedback_parameters": ( getattr(net, "n_fixed_feedback_parameters", 0) - if args.mode == "lhfa" else 0), + if args.mode in ("lhfa", "wm") else 0), }, "epochs": [], } @@ -277,7 +300,30 @@ def run(args): total_start = time.time() predictor_warmup_wall = 0.0 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: + sync(args.device) + mirror_start = time.time() + mirror_metrics = [] + for _ in range(args.mirror_warmup_steps): + metric = normalized_response_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(metric) + counters["mirror_conv_examples"] += args.mirror_batch_size + counters["mirror_readout_examples"] += metric["readout_batch_size"] + counters["mirror_events"] += 1 + log["mirror_warmup"] = { + "steps": args.mirror_warmup_steps, + "first": mirror_metrics[0], + "mean": {key: sum(value[key] for value in mirror_metrics) + / len(mirror_metrics) for key in mirror_metrics[0]}, + "last": mirror_metrics[-1], + } + sync(args.device) + mirror_warmup_wall = time.time() - mirror_start if config is not None and config.learn_P and args.predictor_warmup_steps: sync(args.device) warmup_start = time.time() @@ -365,6 +411,7 @@ def run(args): loss_sum = 0.0 examples = 0 calibration_metrics = [] + mirror_metrics = [] for x, y in train: batch = x.shape[0] if args.mode == "bp": @@ -383,6 +430,20 @@ def run(args): did_perturb = result["did_perturb"] if result["calibration"] is not None: calibration_metrics.append(result["calibration"]) + elif args.mode == "wm": + if step % args.mirror_every == 0: + mirror_metric = normalized_response_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) @@ -423,6 +484,11 @@ def run(args): / len(calibration_metrics) for key in calibration_metrics[0] } + if mirror_metrics: + record["mirror"] = { + key: sum(value[key] for value in mirror_metrics) + / len(mirror_metrics) for key in mirror_metrics[0] + } if args.eval_every and (epoch + 1) % args.eval_every == 0: sync(args.device) eval_start = time.time() @@ -461,7 +527,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") else + if args.mode in ("hfa", "lhfa", "wm") else conv_alignment_report( net, train.x[:probe], train.y[:probe], config)) sync(args.device) @@ -477,6 +543,7 @@ def run(args): "timing": { "predictor_warmup_wall_s": predictor_warmup_wall, "apical_warmup_wall_s": apical_warmup_wall, + "mirror_warmup_wall_s": mirror_warmup_wall, "train_wall_s": train_wall, "evaluation_wall_s": eval_wall, "total_timed_wall_s": total_wall, @@ -508,7 +575,8 @@ def run(args): def parse_args(): parser = argparse.ArgumentParser() parser.add_argument( - "--mode", choices=("bp", "dfa", "hfa", "lhfa", "sdil", "nodepert"), + "--mode", choices=( + "bp", "dfa", "hfa", "lhfa", "wm", "sdil", "nodepert"), required=True) parser.add_argument("--out", required=True) parser.add_argument("--device", default="cpu") @@ -561,6 +629,12 @@ def parse_args(): default="unit_targets") parser.add_argument("--predictor_warmup_steps", type=int, default=0) parser.add_argument("--a_warmup_steps", type=int, default=0) + parser.add_argument("--mirror_warmup_steps", type=int, default=0) + parser.add_argument("--mirror_every", type=int, default=16) + parser.add_argument("--mirror_batch_size", type=int, default=1) + parser.add_argument("--mirror_eta", type=float, default=0.1) + parser.add_argument("--mirror_noise_std", type=float, default=1.0) + parser.add_argument("--mirror_seed", type=int, default=3000) parser.add_argument("--alignment_probe", type=int, default=0) return parser.parse_args() -- cgit v1.2.3