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.py7
1 files changed, 6 insertions, 1 deletions
diff --git a/experiments/diagnose_kp_traffic_nonfinite.py b/experiments/diagnose_kp_traffic_nonfinite.py
index bdf79c3..99a4e99 100644
--- a/experiments/diagnose_kp_traffic_nonfinite.py
+++ b/experiments/diagnose_kp_traffic_nonfinite.py
@@ -89,12 +89,15 @@ def main():
parser.add_argument("--predictor_mode",
choices=("nlms20", "closed_form"),
default="nlms20")
+ parser.add_argument("--predictor_every", type=int, default=16)
parser.add_argument("--device", default="cuda")
parser.add_argument("--max_steps", type=int, default=352)
parser.add_argument("--out", required=True)
args = parser.parse_args()
if args.max_steps < 1:
raise ValueError("max_steps must be positive")
+ if args.predictor_every < 0:
+ raise ValueError("predictor_every must be nonnegative")
torch.manual_seed(0)
if str(args.device).startswith("cuda"):
@@ -168,7 +171,8 @@ def main():
if step >= args.max_steps:
break
result = conv_kp_mixed_traffic_step(
- net, x, y, config, step, args.rule, predictor_every=16)
+ net, x, y, config, step, args.rule,
+ predictor_every=args.predictor_every)
state = network_state(net)
bad = nonfinite_groups(state)
for name, indices in bad.items():
@@ -223,6 +227,7 @@ def main():
"provenance": provenance(),
"rule": args.rule,
"predictor_mode": args.predictor_mode,
+ "predictor_every": args.predictor_every,
"max_steps": args.max_steps,
"split": split,
"traffic_calibration": calibration,