summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
Diffstat (limited to 'experiments')
-rw-r--r--experiments/conv_local_smoke.py49
-rw-r--r--experiments/conv_run.py84
2 files changed, 128 insertions, 5 deletions
diff --git a/experiments/conv_local_smoke.py b/experiments/conv_local_smoke.py
index 894f9d3..7d24154 100644
--- a/experiments/conv_local_smoke.py
+++ b/experiments/conv_local_smoke.py
@@ -12,7 +12,9 @@ from sdil.conv import (CIFARHierarchicalFAResNet, CIFARLocalResNet,
CIFARSDILResNet, ConvSDILConfig,
channel_subspace_apical_calibration,
conv_hierarchical_step, conv_local_step,
+ hierarchical_mirror_observations,
hierarchical_parameter_subspace_calibration,
+ normalized_response_mirror_update,
simultaneous_conv_node_perturbation,
vectorizer_subspace_apical_calibration)
@@ -756,6 +758,52 @@ def hierarchical_parameter_calibration_checks():
}
+def normalized_response_mirror_checks():
+ """Audit local response estimation and absence of W access in the update."""
+ net = CIFARHierarchicalFAResNet(
+ depth=8, base_width=2, seed=121, dtype=torch.float64,
+ normalization="batchnorm")
+ observations = hierarchical_mirror_observations(
+ net, batch_size=16, noise_std=1.0,
+ generator=torch.Generator().manual_seed(122))
+ metrics, _ = normalized_response_mirror_update(
+ net, observations, eta=1.0)
+ pairs = list(zip(net.Q[1:], net.W[1:])) + [
+ (net.R_out, -net.W_out.t())]
+ cosines = [float(F.cosine_similarity(
+ feedback.flatten(), target.flatten(), dim=0))
+ for feedback, target in pairs]
+ norm_ratios = [float(feedback.norm() / target.norm())
+ for feedback, target in pairs]
+ assert sum(cosines) / len(cosines) > 0.985
+ assert min(cosines) > 0.95
+ assert min(norm_ratios) > 0.90 and max(norm_ratios) < 1.10
+
+ # The update consumes observations only. Changing every forward parameter
+ # after those observations were generated must not change the Q/R update.
+ left = CIFARHierarchicalFAResNet(
+ depth=8, base_width=2, seed=123, dtype=torch.float64)
+ right = CIFARHierarchicalFAResNet(
+ depth=8, base_width=2, seed=123, dtype=torch.float64)
+ shared_observations = hierarchical_mirror_observations(
+ left, batch_size=2, generator=torch.Generator().manual_seed(124))
+ for value in right.W + [right.W_out]:
+ value.normal_(generator=torch.Generator().manual_seed(value.numel()))
+ normalized_response_mirror_update(left, shared_observations, eta=0.2)
+ normalized_response_mirror_update(right, shared_observations, eta=0.2)
+ 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
+ return {
+ "mirror_estimate_mean_forward_cosine": sum(cosines) / len(cosines),
+ "mirror_estimate_min_forward_cosine": min(cosines),
+ "mirror_estimate_min_norm_ratio": min(norm_ratios),
+ "mirror_estimate_max_norm_ratio": max(norm_ratios),
+ "mirror_update_forward_parameter_independence_error": independence_error,
+ "mirror_update_rms": metrics["mirror_update_rms"],
+ }
+
+
def apical_learning_checks():
torch.manual_seed(11)
net = CIFARSDILResNet(depth=8, base_width=2, seed=6)
@@ -863,6 +911,7 @@ def main():
report.update(vectorizer_subspace_estimator_check())
report.update(hierarchical_feedback_checks())
report.update(hierarchical_parameter_calibration_checks())
+ report.update(normalized_response_mirror_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 fb759b4..fcb6690 100644
--- a/experiments/conv_run.py
+++ b/experiments/conv_run.py
@@ -19,6 +19,7 @@ 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.data import DATA_DIR, get_cifar_image_splits
@@ -105,7 +106,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"):
+ if args.mode in ("hfa", "lhfa", "wm"):
net = CIFARHierarchicalFAResNet(
**common, feedback_seed=args.apical_seed,
feedback_scale=args.a_scale)
@@ -156,6 +157,11 @@ def work_report(net, mode, counters):
apical_regression = (regression_multiplier
* counters["calibration_event_examples"]
* apical_macs)
+ mirror_conv_macs = max(0, apical_macs - getattr(net, "R_out", torch.empty(0)).numel())
+ 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_correlation = mirror_forward
components = {
"ordinary_forward_macs": normal_forward,
"warmup_clean_forward_macs": warmup_forward,
@@ -164,6 +170,8 @@ def work_report(net, mode, counters):
"local_weight_correlation_macs": local_correlation,
"apical_projection_macs": apical_inference,
"apical_regression_macs": apical_regression,
+ "mirror_response_macs": mirror_forward,
+ "mirror_local_correlation_macs": mirror_correlation,
}
return {
"forward_macs_per_example": forward_macs,
@@ -180,6 +188,8 @@ def work_report(net, mode, counters):
"logical_batch_loss_queries": counters["logical_batch_loss_queries"],
"causal_scalar_observations": counters["causal_scalar_observations"],
"per_example_cross_entropy_terms": counters["per_example_loss_terms"],
+ "mirror_probe_examples": counters["mirror_conv_examples"],
+ "mirror_readout_probe_examples": counters["mirror_readout_examples"],
"definition": (
"multiply-accumulates in conv/linear maps; one local weight correlation "
"equals one forward-weight MAC count; BP reverse is estimated as one "
@@ -196,6 +206,12 @@ 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:
+ 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")
+ if not 0.0 < args.mirror_eta <= 1.0 or args.mirror_noise_std <= 0:
+ raise ValueError("invalid mirror learning hyperparameters")
torch.manual_seed(args.seed)
if str(args.device).startswith("cuda"):
if not torch.cuda.is_available():
@@ -221,6 +237,8 @@ def run(args):
args.perturb_seed)
warmup_generator = torch.Generator(device=torch.device(args.device)).manual_seed(
args.perturb_seed + 1)
+ mirror_generator = torch.Generator(device=torch.device(args.device)).manual_seed(
+ args.mirror_seed)
counters = {
"ordinary_examples": 0,
@@ -232,6 +250,9 @@ def run(args):
"causal_scalar_observations": 0,
"per_example_loss_terms": 0,
"perturbation_events": 0,
+ "mirror_conv_examples": 0,
+ "mirror_readout_examples": 0,
+ "mirror_events": 0,
}
log = {
"schema_version": 1,
@@ -239,10 +260,12 @@ def run(args):
"calibration_metric_space": (
None if config is None or args.mode == "hfa" else
"hierarchical_feedback_parameters" if args.mode == "lhfa" else {
+ "wm": "local_parent_child_response",
"unit_targets": "full_hidden_field",
"channel_subspace": "channel_basis_moments",
"vectorizer_subspace": "vectorizer_parameter_gradients",
- }[config.apical_calibration_mode]),
+ }[args.mode if args.mode == "wm"
+ else config.apical_calibration_mode]),
"args": vars(args),
"provenance": provenance(),
"split": split,
@@ -268,7 +291,7 @@ def run(args):
if args.mode == "hfa" else 0),
"adaptive_feedback_parameters": (
getattr(net, "n_fixed_feedback_parameters", 0)
- if args.mode == "lhfa" else 0),
+ if args.mode in ("lhfa", "wm") else 0),
},
"epochs": [],
}
@@ -277,7 +300,30 @@ def run(args):
total_start = time.time()
predictor_warmup_wall = 0.0
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:
+ sync(args.device)
+ mirror_start = time.time()
+ mirror_metrics = []
+ for _ in range(args.mirror_warmup_steps):
+ metric = normalized_response_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(metric)
+ counters["mirror_conv_examples"] += args.mirror_batch_size
+ counters["mirror_readout_examples"] += metric["readout_batch_size"]
+ counters["mirror_events"] += 1
+ log["mirror_warmup"] = {
+ "steps": args.mirror_warmup_steps,
+ "first": mirror_metrics[0],
+ "mean": {key: sum(value[key] for value in mirror_metrics)
+ / len(mirror_metrics) for key in mirror_metrics[0]},
+ "last": mirror_metrics[-1],
+ }
+ sync(args.device)
+ mirror_warmup_wall = time.time() - mirror_start
if config is not None and config.learn_P and args.predictor_warmup_steps:
sync(args.device)
warmup_start = time.time()
@@ -365,6 +411,7 @@ def run(args):
loss_sum = 0.0
examples = 0
calibration_metrics = []
+ mirror_metrics = []
for x, y in train:
batch = x.shape[0]
if args.mode == "bp":
@@ -383,6 +430,20 @@ def run(args):
did_perturb = result["did_perturb"]
if result["calibration"] is not None:
calibration_metrics.append(result["calibration"])
+ elif args.mode == "wm":
+ if step % args.mirror_every == 0:
+ mirror_metric = normalized_response_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)
@@ -423,6 +484,11 @@ def run(args):
/ len(calibration_metrics)
for key in calibration_metrics[0]
}
+ if mirror_metrics:
+ record["mirror"] = {
+ key: sum(value[key] for value in mirror_metrics)
+ / len(mirror_metrics) for key in mirror_metrics[0]
+ }
if args.eval_every and (epoch + 1) % args.eval_every == 0:
sync(args.device)
eval_start = time.time()
@@ -461,7 +527,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") else
+ if args.mode in ("hfa", "lhfa", "wm") else
conv_alignment_report(
net, train.x[:probe], train.y[:probe], config))
sync(args.device)
@@ -477,6 +543,7 @@ def run(args):
"timing": {
"predictor_warmup_wall_s": predictor_warmup_wall,
"apical_warmup_wall_s": apical_warmup_wall,
+ "mirror_warmup_wall_s": mirror_warmup_wall,
"train_wall_s": train_wall,
"evaluation_wall_s": eval_wall,
"total_timed_wall_s": total_wall,
@@ -508,7 +575,8 @@ def run(args):
def parse_args():
parser = argparse.ArgumentParser()
parser.add_argument(
- "--mode", choices=("bp", "dfa", "hfa", "lhfa", "sdil", "nodepert"),
+ "--mode", choices=(
+ "bp", "dfa", "hfa", "lhfa", "wm", "sdil", "nodepert"),
required=True)
parser.add_argument("--out", required=True)
parser.add_argument("--device", default="cpu")
@@ -561,6 +629,12 @@ def parse_args():
default="unit_targets")
parser.add_argument("--predictor_warmup_steps", type=int, default=0)
parser.add_argument("--a_warmup_steps", type=int, default=0)
+ parser.add_argument("--mirror_warmup_steps", type=int, default=0)
+ parser.add_argument("--mirror_every", type=int, default=16)
+ parser.add_argument("--mirror_batch_size", type=int, default=1)
+ parser.add_argument("--mirror_eta", type=float, default=0.1)
+ parser.add_argument("--mirror_noise_std", type=float, default=1.0)
+ parser.add_argument("--mirror_seed", type=int, default=3000)
parser.add_argument("--alignment_probe", type=int, default=0)
return parser.parse_args()