summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
Diffstat (limited to 'experiments')
-rw-r--r--experiments/conv_local_smoke.py39
-rw-r--r--experiments/conv_run.py39
2 files changed, 69 insertions, 9 deletions
diff --git a/experiments/conv_local_smoke.py b/experiments/conv_local_smoke.py
index 7d24154..4e388a0 100644
--- a/experiments/conv_local_smoke.py
+++ b/experiments/conv_local_smoke.py
@@ -14,6 +14,7 @@ from sdil.conv import (CIFARHierarchicalFAResNet, CIFARLocalResNet,
conv_hierarchical_step, conv_local_step,
hierarchical_mirror_observations,
hierarchical_parameter_subspace_calibration,
+ normalized_residual_mirror_update,
normalized_response_mirror_update,
simultaneous_conv_node_perturbation,
vectorizer_subspace_apical_calibration)
@@ -794,6 +795,38 @@ def normalized_response_mirror_checks():
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
+
+ # Residual-response LMS has a per-observation exact fixed point: its update
+ # is zero, not merely zero in expectation, when Q/R match the forward maps.
+ fixed = CIFARHierarchicalFAResNet(
+ depth=8, base_width=2, seed=125, dtype=torch.float64)
+ fixed.Q = [value.clone() for value in fixed.W]
+ fixed.R_out.copy_(-fixed.W_out.t())
+ fixed_observations = hierarchical_mirror_observations(
+ fixed, batch_size=2, generator=torch.Generator().manual_seed(126))
+ fixed_metrics = normalized_residual_mirror_update(
+ fixed, fixed_observations, eta=1.0)
+ assert fixed_metrics["mirror_update_rms"] < 1e-14
+ assert fixed_metrics["mirror_response_residual_fraction"] < 1e-14
+
+ residual_left = CIFARHierarchicalFAResNet(
+ depth=8, base_width=2, seed=127, dtype=torch.float64)
+ residual_right = CIFARHierarchicalFAResNet(
+ depth=8, base_width=2, seed=127, dtype=torch.float64)
+ residual_observations = hierarchical_mirror_observations(
+ residual_left, batch_size=2,
+ generator=torch.Generator().manual_seed(128))
+ for value in residual_right.W + [residual_right.W_out]:
+ value.normal_(generator=torch.Generator().manual_seed(value.numel() + 1))
+ normalized_residual_mirror_update(
+ residual_left, residual_observations, eta=0.2)
+ normalized_residual_mirror_update(
+ residual_right, residual_observations, eta=0.2)
+ residual_independence_error = max(float((a - b).abs().max())
+ for a, b in zip(
+ residual_left.Q[1:] + [residual_left.R_out],
+ residual_right.Q[1:] + [residual_right.R_out]))
+ assert residual_independence_error == 0.0
return {
"mirror_estimate_mean_forward_cosine": sum(cosines) / len(cosines),
"mirror_estimate_min_forward_cosine": min(cosines),
@@ -801,6 +834,12 @@ def normalized_response_mirror_checks():
"mirror_estimate_max_norm_ratio": max(norm_ratios),
"mirror_update_forward_parameter_independence_error": independence_error,
"mirror_update_rms": metrics["mirror_update_rms"],
+ "residual_mirror_exact_fixed_point_update_rms": fixed_metrics[
+ "mirror_update_rms"],
+ "residual_mirror_exact_fixed_point_fraction": fixed_metrics[
+ "mirror_response_residual_fraction"],
+ "residual_mirror_forward_parameter_independence_error": (
+ residual_independence_error),
}
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")