diff options
Diffstat (limited to 'experiments/diagnose_kp_traffic_nonfinite.py')
| -rw-r--r-- | experiments/diagnose_kp_traffic_nonfinite.py | 10 |
1 files changed, 8 insertions, 2 deletions
diff --git a/experiments/diagnose_kp_traffic_nonfinite.py b/experiments/diagnose_kp_traffic_nonfinite.py index 7b7389f..d438e01 100644 --- a/experiments/diagnose_kp_traffic_nonfinite.py +++ b/experiments/diagnose_kp_traffic_nonfinite.py @@ -91,6 +91,7 @@ def main(): default="nlms20") parser.add_argument("--predictor_every", type=int, default=16) parser.add_argument("--stability_margin", type=float, default=0.0) + parser.add_argument("--neutral_projection", action="store_true") parser.add_argument("--device", default="cuda") parser.add_argument("--max_steps", type=int, default=352) parser.add_argument("--out", required=True) @@ -176,7 +177,8 @@ def main(): break result = conv_kp_mixed_traffic_step( net, x, y, config, step, args.rule, - predictor_every=args.predictor_every) + predictor_every=args.predictor_every, + neutral_projection=args.neutral_projection) state = network_state(net) bad = nonfinite_groups(state) for name, indices in bad.items(): @@ -210,6 +212,7 @@ def main(): "traffic_rms": result["traffic_rms"], "predictor_updated": result["did_predictor_update"], "predictor_mse": result["predictor_mse"], + "neutral_projection": result["neutral_projection"], "parameter_state": state, }) if active_bad or active_metric_names: @@ -226,13 +229,16 @@ def main(): if row["predictor_mse"] is not None: numeric.append(float(row["predictor_mse"])) output = { - "protocol": "kp_mixed_traffic_nonfinite_diagnosis_v1", + "protocol": ("kp_dynamic_neutral_projection_diagnosis_v1" + if args.neutral_projection + else "kp_mixed_traffic_nonfinite_diagnosis_v1"), "scope": "training_only_no_validation_or_test_evaluation", "provenance": provenance(), "rule": args.rule, "predictor_mode": args.predictor_mode, "predictor_every": args.predictor_every, "stability_margin": args.stability_margin, + "neutral_projection": args.neutral_projection, "max_steps": args.max_steps, "split": split, "traffic_calibration": calibration, |
