summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
Diffstat (limited to 'experiments')
-rw-r--r--experiments/conv_local_smoke.py84
-rw-r--r--experiments/conv_run.py63
2 files changed, 134 insertions, 13 deletions
diff --git a/experiments/conv_local_smoke.py b/experiments/conv_local_smoke.py
index 9207e04..894f9d3 100644
--- a/experiments/conv_local_smoke.py
+++ b/experiments/conv_local_smoke.py
@@ -1,5 +1,6 @@
#!/usr/bin/env python3
"""Prove the convolutional local eligibility matches exact BP when instructed."""
+import math
import os
import sys
@@ -11,6 +12,7 @@ from sdil.conv import (CIFARHierarchicalFAResNet, CIFARLocalResNet,
CIFARSDILResNet, ConvSDILConfig,
channel_subspace_apical_calibration,
conv_hierarchical_step, conv_local_step,
+ hierarchical_parameter_subspace_calibration,
simultaneous_conv_node_perturbation,
vectorizer_subspace_apical_calibration)
@@ -673,6 +675,87 @@ def hierarchical_feedback_checks():
}
+def hierarchical_parameter_calibration_checks():
+ """Audit the causal JVP and the exact local Q/R delta-rule moments."""
+ torch.manual_seed(109)
+ net = CIFARHierarchicalFAResNet(
+ depth=8, base_width=2, seed=110, dtype=torch.float64,
+ normalization="batchnorm", residual_scale=1.0)
+ x = torch.randn(3, 3, 32, 32, dtype=torch.float64)
+ y = torch.tensor([0, 4, 7])
+ parameters = net.W + net.gamma + net.beta + [net.W_out, net.b_out]
+ for parameter in parameters:
+ parameter.requires_grad_(True)
+ forward = net.forward(
+ x, return_cache=True, training=True, update_stats=False)
+ loss = F.cross_entropy(forward["logits"], y)
+ hidden_gradients = torch.autograd.grad(loss, forward["hiddens"])
+ output_signal = (torch.softmax(forward["logits"].detach(), dim=1)
+ - F.one_hot(y, 10).to(torch.float64))
+ _, diagnostic = hierarchical_parameter_subspace_calibration(
+ net, x, y, forward, output_signal, sigma=1e-5,
+ n_directions=1, eta=0.0,
+ generator=torch.Generator().manual_seed(111),
+ return_diagnostics=True)
+ directions = diagnostic["directions"][0]["hidden"]
+ exact_directional = x.shape[0] * sum(
+ (gradient * direction).sum()
+ for gradient, direction in zip(hidden_gradients, directions))
+ estimated_directional = diagnostic[
+ "directional_derivatives"][0]["scaled_directional"]
+ jvp_relative = float((estimated_directional - exact_directional).abs()
+ / exact_directional.abs().clamp_min(1e-30))
+ assert jvp_relative < 3e-7
+ for parameter in parameters:
+ parameter.requires_grad_(False)
+
+ # Under an audit-only symmetric copy, the hierarchical field is the exact
+ # negative gradient. Consequently every local Q/R predicted moment equals
+ # its exact causal regression target, including option-A shortcut terms.
+ exact = CIFARHierarchicalFAResNet(
+ depth=8, base_width=2, seed=112, dtype=torch.float64,
+ normalization="batchnorm", residual_scale=1.0)
+ exact.Q = [value.clone() for value in exact.W]
+ exact.R_out.copy_(-exact.W_out.t())
+ for parameter in exact.W + exact.gamma + exact.beta + [
+ exact.W_out, exact.b_out]:
+ parameter.requires_grad_(True)
+ clean = exact.forward(
+ x, return_cache=True, training=True, update_stats=False)
+ gradients = torch.autograd.grad(
+ F.cross_entropy(clean["logits"], y), clean["hiddens"])
+ negative = [-x.shape[0] * value.detach() for value in gradients]
+ signal = (torch.softmax(clean["logits"].detach(), dim=1)
+ - F.one_hot(y, 10).to(torch.float64))
+ teaching, contexts, recipients = exact.hierarchical_teaching(
+ signal, clean, return_edge_contexts=True)
+ numerator = 0.0
+ denominator = 0.0
+ for index in range(1, len(exact.Q)):
+ recipient = recipients[index]
+ spec = exact.layer_specs[index]
+ spatial = (negative[recipient].shape[2]
+ * negative[recipient].shape[3])
+ target = torch.nn.grad.conv2d_weight(
+ negative[recipient], exact.Q[index].shape, contexts[index],
+ stride=spec.stride, padding=spec.padding) / (x.shape[0] * spatial)
+ prediction = torch.nn.grad.conv2d_weight(
+ teaching[recipient], exact.Q[index].shape, contexts[index],
+ stride=spec.stride, padding=spec.padding) / (x.shape[0] * spatial)
+ numerator += float((target - prediction).square().sum())
+ denominator += float(target.square().sum())
+ target_r = negative[-1].mean(dim=(2, 3)).t() @ signal / x.shape[0]
+ prediction_r = teaching[-1].mean(dim=(2, 3)).t() @ signal / x.shape[0]
+ numerator += float((target_r - prediction_r).square().sum())
+ denominator += float(target_r.square().sum())
+ delta_rule_relative = math.sqrt(numerator / max(denominator, 1e-300))
+ assert delta_rule_relative < 2e-12
+ return {
+ "hierarchical_parameter_subspace_jvp_relative_error": jvp_relative,
+ "hierarchical_parameter_delta_rule_relative_error": delta_rule_relative,
+ }
+
+
def apical_learning_checks():
torch.manual_seed(11)
net = CIFARSDILResNet(depth=8, base_width=2, seed=6)
@@ -779,6 +862,7 @@ def main():
report.update(channel_subspace_estimator_check())
report.update(vectorizer_subspace_estimator_check())
report.update(hierarchical_feedback_checks())
+ report.update(hierarchical_parameter_calibration_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 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")