summaryrefslogtreecommitdiff
path: root/experiments/diagnose_kp_traffic_nonfinite.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/diagnose_kp_traffic_nonfinite.py')
-rw-r--r--experiments/diagnose_kp_traffic_nonfinite.py10
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,