diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 01:38:39 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 01:38:39 -0500 |
| commit | 98dddcfabb39c13f652b3aa44f0926ee04428c06 (patch) | |
| tree | 7f95c369e9c8ec88b9ae8a62034774bbe991f47a /experiments/synthetic_smoke.py | |
| parent | ee57b4ca69ace14ffadfcfde947112278001feef (diff) | |
experiments: isolate frozen validation protocols
Diffstat (limited to 'experiments/synthetic_smoke.py')
| -rw-r--r-- | experiments/synthetic_smoke.py | 3 |
1 files changed, 2 insertions, 1 deletions
diff --git a/experiments/synthetic_smoke.py b/experiments/synthetic_smoke.py index cd6e2af..d2325a8 100644 --- a/experiments/synthetic_smoke.py +++ b/experiments/synthetic_smoke.py @@ -20,7 +20,7 @@ def check_task(dataset, expected_in, expected_out, extra): args = get_args() finally: sys.argv = old - train, test, n_in, n_out = load_task(args, "cpu") + train, test, n_in, n_out, split = load_task(args, "cpu") assert (n_in, n_out) == (expected_in, expected_out) args.n_in, args.n_out = n_in, n_out x, y = next(iter(train)) @@ -31,6 +31,7 @@ def check_task(dataset, expected_in, expected_out, extra): net, _ = build(args, "cpu") assert net.logits(x).shape == (16, n_out) assert sum(batch_y.numel() for _, batch_y in test) == 32 + assert split["task_seed"] == args.task_seed print(f"{dataset}: input={n_in}, classes={n_out}, fixed task seed={args.task_seed}") |
