diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 15:21:02 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 15:21:02 -0500 |
| commit | 9e755189e9fe5fe72781d356fced1223ce235e53 (patch) | |
| tree | f88ede06613a68bf3e03cd124f6ad0355fc52e12 /experiments/conv_run.py | |
| parent | 7455c126ab8a5cbfd3bb1ebe8b029a04f858e824 (diff) | |
audit: enforce paired task order after neutral warmup
Diffstat (limited to 'experiments/conv_run.py')
| -rw-r--r-- | experiments/conv_run.py | 4 |
1 files changed, 4 insertions, 0 deletions
diff --git a/experiments/conv_run.py b/experiments/conv_run.py index c25c48c..a1ebb6a 100644 --- a/experiments/conv_run.py +++ b/experiments/conv_run.py @@ -439,12 +439,16 @@ def run(args): counters["predictor_warmup_examples"] += x.shape[0] del forward train.g.set_state(loader_state) + loader_state_restored = torch.equal(train.g.get_state(), loader_state) + if not loader_state_restored: + raise AssertionError("predictor warmup did not restore loader state") log["predictor_warmup"] = { "steps": args.predictor_warmup_steps, "first_mse": predictor_metrics[0], "mean_mse": sum(predictor_metrics) / len(predictor_metrics), "last_mse": predictor_metrics[-1], "instruction_present": False, + "task_loader_state_restored": loader_state_restored, } if args.mode == "kp_traffic": audit_x, _ = traffic_calibration_batch |
