summaryrefslogtreecommitdiff
path: root/experiments/conv_run.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 17:07:08 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 17:07:08 -0500
commitb56c6118d57e810cc1b2b16eb36c9c474df30174 (patch)
tree25d40fc06c04ce988d72aad5aef6f3d1a5e93648 /experiments/conv_run.py
parent6b85249c718cb366b13355a09e0765cb40c4399d (diff)
experiment: support audited dynamic projection endpoints
Diffstat (limited to 'experiments/conv_run.py')
-rw-r--r--experiments/conv_run.py91
1 files changed, 85 insertions, 6 deletions
diff --git a/experiments/conv_run.py b/experiments/conv_run.py
index a1ebb6a..b1372f5 100644
--- a/experiments/conv_run.py
+++ b/experiments/conv_run.py
@@ -195,7 +195,9 @@ def work_report(net, mode, counters):
+ counters["traffic_calibration_examples"]
* net.traffic_calibration_elementwise_ops_per_example
+ counters["traffic_audit_examples"]
- * net.predictor_audit_elementwise_ops_per_example)
+ * net.predictor_audit_elementwise_ops_per_example
+ + counters["neutral_projection_examples"]
+ * net.neutral_projection_elementwise_ops_per_example)
else:
elementwise_operations = 0
components = {
@@ -233,6 +235,8 @@ def work_report(net, mode, counters):
"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"],
+ "neutral_projection_observations": counters[
+ "neutral_projection_examples"],
"definition": (
"multiply-accumulates in conv/linear maps; one local weight correlation "
"equals one forward-weight MAC count; BP reverse is estimated as one "
@@ -254,12 +258,27 @@ def run(args):
if args.mode == "kp_traffic":
if not args.learn_P or args.predictor_warmup_steps < 1:
raise ValueError("mixed-traffic KP requires predictor warmup")
- if args.predictor_every < 1:
- raise ValueError("mixed-traffic KP requires a predictor cadence")
if args.traffic_calibration_examples < 1 or args.traffic_ratio <= 0:
raise ValueError("invalid mixed-traffic calibration")
+ if args.neutral_projection:
+ if (args.traffic_rule != "innovation"
+ or args.predictor_mode != "closed_form"
+ or args.predictor_warmup_steps != 1
+ or args.predictor_every != 0):
+ raise ValueError(
+ "neutral projection requires frozen one-step closed-form "
+ "innovation prediction")
+ elif args.predictor_every < 1:
+ raise ValueError("mixed-traffic KP requires a predictor cadence")
+ if (args.predictor_mode == "closed_form"
+ and args.predictor_warmup_steps != 1):
+ raise ValueError("closed-form predictor requires one warmup fit")
elif args.predictor_every:
raise ValueError("predictor cadence is restricted to mixed-traffic KP")
+ if args.mode != "kp_traffic" and args.neutral_projection:
+ raise ValueError("neutral projection is restricted to mixed-traffic KP")
+ if args.mode != "kp_traffic" and args.predictor_mode != "nlms":
+ raise ValueError("predictor mode is restricted to mixed-traffic KP")
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:
@@ -310,6 +329,7 @@ def run(args):
"traffic_calibration_examples": 0,
"traffic_audit_examples": 0,
"predictor_update_examples": 0,
+ "neutral_projection_examples": 0,
}
log = {
"schema_version": 1,
@@ -365,6 +385,7 @@ def run(args):
mirror_warmup_wall = 0.0
loader_state = train.g.get_state().clone()
traffic_calibration_batch = None
+ closed_form_predictor_ready = False
if args.mode == "kp_traffic":
count = min(args.traffic_calibration_examples, train.x.shape[0])
if count != args.traffic_calibration_examples:
@@ -391,6 +412,27 @@ def run(args):
"uses_validation_endpoint": False,
})
counters["traffic_calibration_examples"] += count
+ if args.predictor_mode == "closed_form":
+ sync(args.device)
+ closed_form_start = time.time()
+ fit = net.predictor_closed_form_fit(
+ calibration_forward["hiddens"], stability_margin=0.0)
+ sync(args.device)
+ predictor_warmup_wall += time.time() - closed_form_start
+ counters["predictor_update_examples"] += count
+ post_fit_ratio = net.predictor_traffic_residual_rms_ratio(
+ calibration_forward["hiddens"])
+ log["predictor_warmup"] = {
+ "mode": "closed_form",
+ "steps": 1,
+ "examples": count,
+ "closed_form_fit": fit,
+ "post_warmup_traffic_residual_rms_ratio": post_fit_ratio,
+ "instruction_present": False,
+ "task_loader_state_restored": True,
+ "reuses_traffic_calibration_forward": True,
+ }
+ closed_form_predictor_ready = True
del (calibration_forward, calibration_output_error,
calibration_instruction)
if args.mode in ("wm", "rrm") and args.mirror_warmup_steps:
@@ -418,7 +460,8 @@ def run(args):
}
sync(args.device)
mirror_warmup_wall = time.time() - mirror_start
- if config is not None and config.learn_P and args.predictor_warmup_steps:
+ if (config is not None and config.learn_P and args.predictor_warmup_steps
+ and not closed_form_predictor_ready):
sync(args.device)
warmup_start = time.time()
iterator = iter(train)
@@ -443,6 +486,7 @@ def run(args):
if not loader_state_restored:
raise AssertionError("predictor warmup did not restore loader state")
log["predictor_warmup"] = {
+ "mode": "nlms",
"steps": args.predictor_warmup_steps,
"first_mse": predictor_metrics[0],
"mean_mse": sum(predictor_metrics) / len(predictor_metrics),
@@ -534,6 +578,7 @@ def run(args):
calibration_metrics = []
mirror_metrics = []
signal_metrics = []
+ neutral_projection_metrics = []
for x, y in train:
batch = x.shape[0]
if args.mode == "bp":
@@ -587,7 +632,8 @@ def run(args):
elif args.mode == "kp_traffic":
result = conv_kp_mixed_traffic_step(
net, x, y, config, step, args.traffic_rule,
- args.predictor_every)
+ args.predictor_every,
+ neutral_projection=bool(args.neutral_projection))
loss = result["loss"]
did_perturb = False
signal_metrics.append({key: result[key] for key in (
@@ -595,6 +641,10 @@ def run(args):
"innovation_rms", "traffic_rms")})
if result["did_predictor_update"]:
counters["predictor_update_examples"] += batch
+ if result["neutral_projection"] is not None:
+ neutral_projection_metrics.append(
+ result["neutral_projection"])
+ counters["neutral_projection_examples"] += batch
else:
result = conv_local_step(
net, x, y, config, step, generator=perturb_generator)
@@ -645,6 +695,30 @@ def run(args):
key: sum(value[key] for value in signal_metrics)
/ len(signal_metrics) for key in signal_metrics[0]
}
+ if neutral_projection_metrics:
+ record["neutral_projection"] = {
+ "maximum_pre_projection_traffic_rms_ratio": max(
+ value["pre_projection_traffic_rms_ratio"]
+ for value in neutral_projection_metrics),
+ "maximum_post_projection_traffic_rms_ratio": max(
+ value["post_projection_traffic_rms_ratio"]
+ for value in neutral_projection_metrics),
+ "maximum_absolute_post_projection_soma_slope": max(
+ value["max_absolute_post_projection_soma_slope"]
+ for value in neutral_projection_metrics),
+ "maximum_absolute_correction_slope": max(
+ value["max_absolute_correction_slope"]
+ for value in neutral_projection_metrics),
+ "minimum_observations": min(
+ value["observations"]
+ for value in neutral_projection_metrics),
+ "maximum_observations": max(
+ value["observations"]
+ for value in neutral_projection_metrics),
+ "instruction_observations": sum(
+ value["instruction_observations"]
+ for value in neutral_projection_metrics),
+ }
if args.mode in ("kp", "kp_traffic"):
tracking = hierarchical_feedback_tracking_report(net)
record["feedback_tracking"] = {
@@ -695,7 +769,8 @@ def run(args):
diagnostic_start = time.time()
if args.mode == "kp_traffic":
diagnostics = conv_kp_mixed_traffic_alignment_report(
- net, train.x[:probe], train.y[:probe], args.traffic_rule)
+ net, train.x[:probe], train.y[:probe], args.traffic_rule,
+ neutral_projection=bool(args.neutral_projection))
elif args.mode in ("hfa", "lhfa", "wm", "rrm", "kp"):
diagnostics = conv_hierarchical_alignment_report(
net, train.x[:probe], train.y[:probe])
@@ -802,6 +877,10 @@ def parse_args():
default="unit_targets")
parser.add_argument("--predictor_warmup_steps", type=int, default=0)
parser.add_argument("--predictor_every", type=int, default=0)
+ parser.add_argument("--predictor_mode", choices=("nlms", "closed_form"),
+ default="nlms")
+ parser.add_argument("--neutral_projection", type=int, choices=(0, 1),
+ default=0)
parser.add_argument("--traffic_rule",
choices=("raw", "matched", "innovation"),
default="innovation")