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