summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
Diffstat (limited to 'experiments')
-rw-r--r--experiments/conv_local_smoke.py101
-rw-r--r--experiments/conv_run.py44
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")