diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 17:07:08 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 17:07:08 -0500 |
| commit | b56c6118d57e810cc1b2b16eb36c9c474df30174 (patch) | |
| tree | 25d40fc06c04ce988d72aad5aef6f3d1a5e93648 /experiments | |
| parent | 6b85249c718cb366b13355a09e0765cb40c4399d (diff) | |
experiment: support audited dynamic projection endpoints
Diffstat (limited to 'experiments')
| -rw-r--r-- | experiments/conv_run.py | 91 |
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") |
