diff options
Diffstat (limited to 'experiments/conv_run.py')
| -rw-r--r-- | experiments/conv_run.py | 39 |
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") |
