diff options
Diffstat (limited to 'experiments')
| -rw-r--r-- | experiments/conv_local_smoke.py | 101 | ||||
| -rw-r--r-- | experiments/conv_run.py | 44 |
2 files changed, 132 insertions, 13 deletions
diff --git a/experiments/conv_local_smoke.py b/experiments/conv_local_smoke.py index 4e388a0..e8cf294 100644 --- a/experiments/conv_local_smoke.py +++ b/experiments/conv_local_smoke.py @@ -8,10 +8,11 @@ import torch import torch.nn.functional as F sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) -from sdil.conv import (CIFARHierarchicalFAResNet, CIFARLocalResNet, - CIFARSDILResNet, ConvSDILConfig, +from sdil.conv import (CIFARHierarchicalFAResNet, CIFARKPResNet, + CIFARLocalResNet, CIFARSDILResNet, ConvSDILConfig, channel_subspace_apical_calibration, - conv_hierarchical_step, conv_local_step, + conv_hierarchical_step, conv_kolen_pollack_step, + conv_local_step, hierarchical_mirror_observations, hierarchical_parameter_subspace_calibration, normalized_residual_mirror_update, @@ -843,6 +844,99 @@ def normalized_response_mirror_checks(): } +def kolen_pollack_checks(): + """KP's reciprocal correlations are local and preserve exact symmetry.""" + torch.manual_seed(127) + common = dict( + depth=8, base_width=2, seed=53, dtype=torch.float64, + normalization="batchnorm", residual_scale=1.0) + net = CIFARKPResNet(**common) + for index in range(1, len(net.Q)): + net.Q[index].copy_(net.W[index]) + net.R_out.copy_(-net.W_out.t()) + x = torch.randn(3, 3, 32, 32, dtype=torch.float64) + y = torch.tensor([1, 4, 7]) + forward = net.forward( + x, return_cache=True, training=True, update_stats=False) + output_error = (torch.softmax(forward["logits"], dim=1) + - F.one_hot(y, 10).to(torch.float64)) + teaching = net.hierarchical_teaching(output_error, forward) + (forward_directions, gamma_directions, beta_directions, + output_weight, output_bias) = net.local_ascent_directions( + teaching, output_error, forward) + reciprocal_directions, reciprocal_readout = ( + net.reciprocal_feedback_directions( + teaching, output_error, forward)) + direction_error = max([ + float((left - right).abs().max()) + for left, right in zip(forward_directions[1:], reciprocal_directions[1:]) + ] + [float((reciprocal_readout + output_weight.t()).abs().max())]) + assert direction_error < 1e-14 + + # Once local activities have been observed, neither forward nor feedback + # parameter values may alter the independently formed reciprocal update. + before = [value.clone() for value in reciprocal_directions[1:]] + before_readout = reciprocal_readout.clone() + for value in net.W + net.Q: + value.add_(torch.randn_like(value)) + net.W_out.add_(torch.randn_like(net.W_out)) + net.R_out.add_(torch.randn_like(net.R_out)) + independent_directions, independent_readout = ( + net.reciprocal_feedback_directions( + teaching, output_error, forward)) + independence_error = max([ + float((left - right).abs().max()) + for left, right in zip(before, independent_directions[1:]) + ] + [float((before_readout - independent_readout).abs().max())]) + assert independence_error == 0.0 + + # A fresh symmetric state must remain symmetric under two momentum steps. + net = CIFARKPResNet(**common) + for index in range(1, len(net.Q)): + net.Q[index].copy_(net.W[index]) + net.R_out.copy_(-net.W_out.t()) + for _ in range(2): + forward = net.forward( + x, return_cache=True, training=True, update_stats=False) + output_error = (torch.softmax(forward["logits"], dim=1) + - F.one_hot(y, 10).to(torch.float64)) + teaching = net.hierarchical_teaching(output_error, forward) + (forward_directions, gamma_directions, beta_directions, + output_weight, output_bias) = net.local_ascent_directions( + teaching, output_error, forward) + reciprocal_directions, reciprocal_readout = ( + net.reciprocal_feedback_directions( + teaching, output_error, forward)) + net.apply_reciprocal_ascent( + reciprocal_directions, reciprocal_readout, + eta_hidden=0.013, eta_output=0.017, momentum=0.9, + weight_decay=1e-4) + net.apply_ascent( + forward_directions, output_weight, output_bias, + eta_hidden=0.013, eta_output=0.017, momentum=0.9, + weight_decay=1e-4, gamma_directions=gamma_directions, + beta_directions=beta_directions) + symmetry_error = max([ + float((net.Q[index] - net.W[index]).abs().max()) + for index in range(1, len(net.Q)) + ] + [float((net.R_out + net.W_out.t()).abs().max())]) + assert symmetry_error < 1e-14 + + # Exercise the public training step and ensure it remains graph-free. + result = conv_kolen_pollack_step( + net, x, y, ConvSDILConfig( + eta=1e-3, eta_output=1e-3, momentum=0.9, + weight_decay=1e-4, learn_A=False, learn_P=False)) + assert math.isfinite(result["loss"]) + assert all(not value.requires_grad for value in + net.W + net.Q + [net.W_out, net.R_out, net.b_out]) + return { + "kp_local_direction_absolute_error": direction_error, + "kp_forward_parameter_independence_error": independence_error, + "kp_symmetric_update_absolute_error": symmetry_error, + } + + def apical_learning_checks(): torch.manual_seed(11) net = CIFARSDILResNet(depth=8, base_width=2, seed=6) @@ -951,6 +1045,7 @@ def main(): report.update(hierarchical_feedback_checks()) report.update(hierarchical_parameter_calibration_checks()) report.update(normalized_response_mirror_checks()) + report.update(kolen_pollack_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 3b91da9..738c86d 100644 --- a/experiments/conv_run.py +++ b/experiments/conv_run.py @@ -12,13 +12,14 @@ import time import torch sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) -from sdil.conv import (CIFARHierarchicalFAResNet, CIFARLocalResNet, - CIFARSDILResNet, ConvSDILConfig, +from sdil.conv import (CIFARHierarchicalFAResNet, CIFARKPResNet, + CIFARLocalResNet, CIFARSDILResNet, ConvSDILConfig, conv_alignment_report, conv_apical_calibration_step, conv_hierarchical_alignment_report, - conv_hierarchical_step, + conv_hierarchical_step, conv_kolen_pollack_step, conv_learned_hierarchical_step, conv_local_step, - evaluate_conv, hierarchical_parameter_subspace_calibration) + evaluate_conv, hierarchical_feedback_tracking_report, + hierarchical_parameter_subspace_calibration) from sdil.conv import (normalized_residual_mirror_step, normalized_response_mirror_step) from sdil.data import DATA_DIR, get_cifar_image_splits @@ -107,8 +108,10 @@ 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", "rrm"): - net = CIFARHierarchicalFAResNet( + if args.mode in ("hfa", "lhfa", "wm", "rrm", "kp"): + network_class = ( + CIFARKPResNet if args.mode == "kp" else CIFARHierarchicalFAResNet) + net = network_class( **common, feedback_seed=args.apical_seed, feedback_scale=args.a_scale) config = ConvSDILConfig( @@ -164,6 +167,8 @@ def work_report(net, mode, counters): + counters["mirror_readout_examples"] * mirror_readout_macs) mirror_prediction = mirror_forward if mode == "rrm" else 0 mirror_correlation = mirror_forward + kp_feedback_correlation = ( + counters["ordinary_examples"] * apical_macs if mode == "kp" else 0) components = { "ordinary_forward_macs": normal_forward, "warmup_clean_forward_macs": warmup_forward, @@ -175,6 +180,7 @@ def work_report(net, mode, counters): "mirror_response_macs": mirror_forward, "mirror_feedback_prediction_macs": mirror_prediction, "mirror_local_correlation_macs": mirror_correlation, + "kp_reciprocal_correlation_macs": kp_feedback_correlation, } return { "forward_macs_per_example": forward_macs, @@ -262,7 +268,8 @@ def run(args): "protocol_family": "oral_a_cifar_local_resnet_development", "calibration_metric_space": ( None if config is None or args.mode == "hfa" else - "hierarchical_feedback_parameters" if args.mode == "lhfa" else { + "hierarchical_feedback_parameters" if args.mode == "lhfa" else + "reciprocal_local_activity_products" if args.mode == "kp" else { "wm": "local_parent_child_response", "rrm": "local_parent_child_response_residual", "unit_targets": "full_hidden_field", @@ -295,7 +302,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", "rrm") else 0), + if args.mode in ("lhfa", "wm", "rrm", "kp") else 0), }, "epochs": [], } @@ -465,6 +472,10 @@ def run(args): result = conv_hierarchical_step(net, x, y, config) loss = result["loss"] did_perturb = False + elif args.mode == "kp": + result = conv_kolen_pollack_step(net, x, y, config) + loss = result["loss"] + did_perturb = False else: result = conv_local_step( net, x, y, config, step, generator=perturb_generator) @@ -510,6 +521,18 @@ def run(args): key: sum(value[key] for value in mirror_metrics) / len(mirror_metrics) for key in mirror_metrics[0] } + if args.mode == "kp": + tracking = hierarchical_feedback_tracking_report(net) + record["feedback_tracking"] = { + "mean_feedback_forward_cosine": tracking[ + "mean_feedback_forward_cosine"], + "mean_feedback_forward_relative_error": tracking[ + "mean_feedback_forward_relative_error"], + "min_feedback_forward_cosine": min( + tracking["feedback_forward_cosine"]), + "max_feedback_forward_relative_error": max( + tracking["feedback_forward_relative_error"]), + } if args.eval_every and (epoch + 1) % args.eval_every == 0: sync(args.device) eval_start = time.time() @@ -548,7 +571,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", "rrm") else + if args.mode in ("hfa", "lhfa", "wm", "rrm", "kp") else conv_alignment_report( net, train.x[:probe], train.y[:probe], config)) sync(args.device) @@ -597,7 +620,8 @@ def parse_args(): parser = argparse.ArgumentParser() parser.add_argument( "--mode", choices=( - "bp", "dfa", "hfa", "lhfa", "wm", "rrm", "sdil", "nodepert"), + "bp", "dfa", "hfa", "lhfa", "wm", "rrm", "kp", "sdil", + "nodepert"), required=True) parser.add_argument("--out", required=True) parser.add_argument("--device", default="cpu") |
