From 7ce17a2b8aa3e67f8f8b53409f357eca0fe65390 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 17:52:01 +0000 Subject: [PATCH] Validate canonical checkpoint optimizer counters --- src/art/trainer_rank/_checkpoint.py | 5 +++- tests/unit/test_trainer_rank_validation.py | 32 ++++++++++++++++++++++ 2 files changed, 36 insertions(+), 1 deletion(-) diff --git a/src/art/trainer_rank/_checkpoint.py b/src/art/trainer_rank/_checkpoint.py index 530a7f35e..a2b956679 100644 --- a/src/art/trainer_rank/_checkpoint.py +++ b/src/art/trainer_rank/_checkpoint.py @@ -372,7 +372,10 @@ def _validate_manifest( parameters[key] = normalized files.update(normalized) if any( - not isinstance(value, int | float) or isinstance(value, bool) + not isinstance(value, int | float) + or isinstance(value, bool) + or value < 0 + or (isinstance(value, float) and not value.is_integer()) for value in steps.values() ): raise RuntimeError("Checkpoint optimizer steps are invalid") diff --git a/tests/unit/test_trainer_rank_validation.py b/tests/unit/test_trainer_rank_validation.py index d6ed71048..6212a0613 100644 --- a/tests/unit/test_trainer_rank_validation.py +++ b/tests/unit/test_trainer_rank_validation.py @@ -1712,6 +1712,38 @@ def test_checkpoint_manifest_semantics_are_authenticated( prepare_checkpoint(str(root)) +@pytest.mark.parametrize("artifact_entries", (False, True)) +@pytest.mark.parametrize( + "step", (-1, -1.0, -0.5, 0.5, float("nan"), float("inf"), -float("inf"), True) +) +def test_checkpoint_rejects_invalid_optimizer_counter( + tmp_path: Path, step: float, artifact_entries: bool +) -> None: + root = tmp_path / "checkpoint" + manifest = _canonical_checkpoint(root) + manifest["steps"][next(iter(manifest["steps"]))] = step + # A valid digest authenticates the bytes, not the counter's semantics. + manifest["digest"] = _manifest_digest(manifest) + (root / "checkpoint.json").write_text(json.dumps(manifest)) + entries = [*manifest["files"], "checkpoint.json"] if artifact_entries else None + + with pytest.raises(RuntimeError, match="optimizer steps are invalid"): + prepare_checkpoint(str(root), artifact_entries=entries) + + +@pytest.mark.parametrize("step", (0, 0.0, -0.0, 50, 50.0, 2**53 + 1)) +def test_checkpoint_preserves_valid_optimizer_counter( + tmp_path: Path, step: float +) -> None: + root = tmp_path / "checkpoint" + manifest = _canonical_checkpoint(root) + manifest["steps"][next(iter(manifest["steps"]))] = step + manifest["digest"] = _manifest_digest(manifest) + (root / "checkpoint.json").write_text(json.dumps(manifest)) + + assert prepare_checkpoint(str(root)).manifest == manifest + + @pytest.mark.parametrize("extra", (True, False)) def test_checkpoint_optimizer_mapping_must_match_adapter( tmp_path: Path, extra: bool