From d4a82c55d66b59da1abc4ef5cb8abbcc8ddab26c Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 08:00:15 +0000 Subject: [PATCH 1/2] Selectively read checkpoint snapshot components during finalization --- src/art/trainer_rank/_checkpoint.py | 36 +- tests/unit/test_checkpoint_selective_read.py | 327 +++++++++++++++++++ 2 files changed, 354 insertions(+), 9 deletions(-) create mode 100644 tests/unit/test_checkpoint_selective_read.py diff --git a/src/art/trainer_rank/_checkpoint.py b/src/art/trainer_rank/_checkpoint.py index 7b8a67e05..0193d534f 100644 --- a/src/art/trainer_rank/_checkpoint.py +++ b/src/art/trainer_rank/_checkpoint.py @@ -1022,11 +1022,30 @@ def prepare_checkpoint_save( def _read_snapshot( - prepared: _PreparedSave, relative: str, prefix: str, keys: Iterable[str] + prepared: _PreparedSave, + relative: str, + prefix: str, + keys: Iterable[str] | None = None, ) -> dict[str, torch.Tensor]: - load = importlib.import_module("safetensors.torch").load_file - payload = load(prepared.snapshot / relative) - return {key: payload[f"{prefix}/{key}"] for key in keys} + safe_open = importlib.import_module("safetensors").safe_open + with safe_open( + prepared.snapshot / relative, framework="pt", device="cpu" + ) as payload: + names = payload.offset_keys() + if keys is None: + keys = [ + key.removeprefix(f"{prefix}/") + for key in names + if key.startswith(f"{prefix}/") + ] + available = set(names) + tensors = {} + for key in keys: + name = f"{prefix}/{key}" + if name not in available: + raise KeyError(name) + tensors[key] = payload.get_tensor(name) + return tensors def _matching_shards( @@ -1226,12 +1245,11 @@ def write_optimizer_block() -> None: for records in _matching_shards(prepared, owned).values() for record in records }: - load = importlib.import_module("safetensors.torch").load_file - payload = load(prepared.snapshot / relative) local_steps.update( - (key.removeprefix("step/"), float(value.item())) - for key, value in payload.items() - if key.startswith("step/") + (key, float(value.item())) + for key, value in _read_snapshot( + prepared, relative, "step" + ).items() ) except BaseException as exc: error = exc diff --git a/tests/unit/test_checkpoint_selective_read.py b/tests/unit/test_checkpoint_selective_read.py new file mode 100644 index 000000000..1974bd6f0 --- /dev/null +++ b/tests/unit/test_checkpoint_selective_read.py @@ -0,0 +1,327 @@ +from __future__ import annotations + +import asyncio +from dataclasses import dataclass +import json +from pathlib import Path +import sys +import threading +from types import ModuleType, SimpleNamespace +from typing import Any, cast + +import pytest +import safetensors +import safetensors.torch +from safetensors.torch import load_file, save_file +import torch + +from art.trainer_rank import _checkpoint + + +def _eager(prepared, relative, prefix, keys=None): + payload = load_file(prepared.snapshot / relative) + if keys is None: + return { + k.removeprefix(prefix + "/"): v + for k, v in payload.items() + if k.startswith(prefix + "/") + } + return {key: payload[f"{prefix}/{key}"] for key in keys} + + +@dataclass(frozen=True) +class _Meta: + key: str + block: str + shape: tuple[int, ...] + dtype_name: str + owner_rank: int = 0 + + @property + def manifest(self): + return {"sharded": False, "shard_world_size": 1, "shard_rank": 0} + + +@pytest.fixture +def dependencies(monkeypatch): + # Only the unchanged single-rank replicated merge and config writer are stand-ins. + publish = ModuleType("art.megatron.weights.lora_publish") + + def merge(entries): + assert all( + len(parts) == 1 and not parts[0][0]["sharded"] for parts in entries.values() + ) + return {key: parts[0][1] for key, parts in entries.items()} + + monkeypatch.setattr(publish, "merge_sharded_adapter_entries", merge, raising=False) + disk = ModuleType("art.megatron.model_support.lora_disk") + monkeypatch.setattr( + disk, + "save_adapter_config", + lambda path, config: (path / "adapter_config.json").write_text( + json.dumps(config, sort_keys=True, indent=2) + "\n" + ), + raising=False, + ) + monkeypatch.setitem(sys.modules, publish.__name__, publish) + monkeypatch.setitem(sys.modules, disk.__name__, disk) + monkeypatch.setattr(_checkpoint, "_ensure_finalize_group", lambda trainer: None) + + +def _prepared(root, *, dtype=torch.bfloat16, rank=1, optimizer=True, damage=None): + snapshot = root / "snapshot" + reservation = root / "reservation" + snapshot.mkdir(parents=True) + reservation.mkdir() + shards = [] + for block in range(2): + payload = {} + for index in range(2): + key = f"layer{block}.expert{index}.lora_A.weight" + tensor = ( + torch.arange(6 * rank).reshape(6, rank).T + block * 32 + index + ).to(dtype) + payload[f"lora/{key}"] = tensor.contiguous() + if optimizer: + for offset, name in enumerate(("master", "exp_avg", "exp_avg_sq")): + payload[f"{name}/{key}"] = tensor.float().contiguous() + offset + # Mixed scalar dtypes ensure the old offset ordering is preserved. + payload[f"step/{key}"] = torch.tensor( + 350.0, dtype=torch.float64 if index else torch.float32 + ) + shards.append( + _checkpoint._LocalShard( + cast( + Any, + _Meta( + key, + f"layer{block}", + tuple(tensor.shape), + str(dtype).removeprefix("torch."), + ), + ), + f"block{block}.safetensors", + ) + ) + if block == 0 and damage: + key = "layer0.expert0.lora_A.weight" + if damage == "missing_master": + del payload[f"master/{key}"] + elif damage == "missing_step": + del payload[f"step/{key}"] + elif damage == "nonscalar_step": + payload[f"step/{key}"] = torch.ones(2) + path = snapshot / f"block{block}.safetensors" + save_file(payload, path) + if block == 0 and damage == "truncated": + path.write_bytes(path.read_bytes()[:-1]) + opt: _checkpoint.OptimizerConfig | None = ( + dict(learning_rate=0.001, beta1=0.9, beta2=0.99, eps=1e-8, weight_decay=0.01) + if optimizer + else None + ) + return _checkpoint._PreparedSave( + 0, + snapshot, + reservation, + root / "result", + { + "base_model_name_or_path": "public/model", + "r": rank, + "lora_alpha": 32, + "target_modules": ["q_proj"], + }, + tuple(shards), + opt, + ) + + +def _trainer(prepared): + return SimpleNamespace( + _slot_state_error=ValueError, + _checkpoint_finalize_lock=threading.Lock(), + _checkpoint_save_condition=threading.Condition(), + _prepared_checkpoint_saves={str(prepared.destination): prepared}, + _finalized_checkpoint_saves={}, + _checkpoint_save_outcomes={}, + _checkpoint_finalizing_saves={}, + _checkpoint_save_next=0, + _checkpoint_save_skipped=set(), + ) + + +def _files(root): + return { + str(p.relative_to(root)): p.read_bytes() for p in root.rglob("*") if p.is_file() + } + + +@pytest.mark.parametrize( + "dtype", [torch.float16, torch.bfloat16, torch.float32, torch.float64] +) +@pytest.mark.parametrize("rank", [1, 3]) +@pytest.mark.parametrize("optimizer", [False, True]) +def test_final_checkpoint_bytes_match_eager( + dependencies, monkeypatch, tmp_path, dtype, rank, optimizer +): + actual = _prepared(tmp_path / "actual", dtype=dtype, rank=rank, optimizer=optimizer) + expected = _prepared( + tmp_path / "expected", dtype=dtype, rank=rank, optimizer=optimizer + ) + _checkpoint.finish_checkpoint_save(_trainer(actual), str(actual.destination)) + with monkeypatch.context() as patch: + patch.setattr(_checkpoint, "_read_snapshot", _eager) + _checkpoint.finish_checkpoint_save( + _trainer(expected), str(expected.destination) + ) + assert _files(actual.destination) == _files(expected.destination) + assert not actual.snapshot.exists() and not actual.reservation.exists() + manifest = json.loads((actual.destination / "checkpoint.json").read_text()) + assert set(manifest["steps"].values()) == ({350.0} if optimizer else set()) + for path in actual.destination.rglob("*.safetensors"): + for key, value in load_file(path).items(): + assert value.dtype == ( + torch.float32 if "optimizer" in path.parts else dtype + ) + + +@pytest.mark.parametrize( + "damage", ["missing_master", "missing_step", "nonscalar_step", "truncated"] +) +def test_failure_matches_eager_and_closes_finalization( + dependencies, monkeypatch, tmp_path, damage +): + errors = [] + for name, reader in (("eager", _eager), ("selective", _checkpoint._read_snapshot)): + prepared = _prepared(tmp_path / name, damage=damage) + trainer = _trainer(prepared) + with monkeypatch.context() as patch: + patch.setattr(_checkpoint, "_read_snapshot", reader) + with pytest.raises(Exception) as caught: + _checkpoint.finish_checkpoint_save(trainer, str(prepared.destination)) + errors.append((type(caught.value), str(caught.value))) + assert not list((tmp_path / name).iterdir()) + assert ( + not trainer._prepared_checkpoint_saves + and not trainer._checkpoint_finalizing_saves + ) + assert ( + trainer._finalized_checkpoint_saves[str(prepared.destination)].outcome + == "abort" + ) + assert errors[0] == errors[1] + + +def test_reader_exact_keys_order_and_closed_handle(tmp_path): + payload = { + "lora/a/b": torch.tensor([-0.0, 3.0]), + "lora/a": torch.tensor([5.0]), + "step/a": torch.tensor(350.0), + "steps/not_step": torch.ones(2), + } + save_file(payload, tmp_path / "block") + prepared = cast(_checkpoint._PreparedSave, SimpleNamespace(snapshot=tmp_path)) + out = _checkpoint._read_snapshot( + prepared, "block", "lora", iter(["a/b", "a", "a/b"]) + ) + assert list(out) == ["a/b", "a"] + steps = _checkpoint._read_snapshot(prepared, "block", "step") + assert list(steps) == ["a"] and steps["a"].item() == 350.0 + (tmp_path / "block").unlink() + assert ( + out["a/b"].view(torch.int32).tolist() + == payload["lora/a/b"].view(torch.int32).tolist() + ) + with pytest.raises(FileNotFoundError): + _checkpoint._read_snapshot(prepared, "block", "lora", []) + + +@pytest.mark.parametrize("error_type", [OSError, asyncio.CancelledError]) +def test_get_tensor_error_identity_and_cleanup( + dependencies, monkeypatch, tmp_path, error_type +): + prepared = _prepared(tmp_path) + real = safetensors.safe_open + error = error_type("public injected tensor-read failure") + closed = [] + + class Failing: + def __init__(self, *args, **kwargs): + self.reader = real(*args, **kwargs) + + def __enter__(self): + self.reader.__enter__() + return self + + def __exit__(self, *args): + closed.append(True) + return self.reader.__exit__(*args) + + def offset_keys(self): + return self.reader.offset_keys() + + def get_tensor(self, key): + raise error + + monkeypatch.setattr(safetensors, "safe_open", Failing) + with pytest.raises(error_type) as caught: + _checkpoint.finish_checkpoint_save( + _trainer(prepared), str(prepared.destination) + ) + assert caught.value is error and closed == [True] + assert not list(tmp_path.iterdir()) + + +def test_only_requested_component_bytes_materialized( + dependencies, monkeypatch, tmp_path +): + real = safetensors.safe_open + reads = [] + + class Tracking: + def __init__(self, filename, **kwargs): + self.reader = real(filename, **kwargs) + self.snapshot = Path(filename).parent.name == "snapshot" + + def __enter__(self): + self.reader.__enter__() + return self + + def __exit__(self, *args): + return self.reader.__exit__(*args) + + def keys(self): + return self.reader.keys() + + def offset_keys(self): + return self.reader.offset_keys() + + def get_tensor(self, key): + value = self.reader.get_tensor(key) + if self.snapshot: + reads.append((key, value.numel() * value.element_size())) + return value + + def get_tensors(self): + values = cast(Any, self.reader).get_tensors() + if self.snapshot: + reads.extend( + (key, value.numel() * value.element_size()) + for key, value in values.items() + ) + return values + + monkeypatch.setattr(safetensors, "safe_open", Tracking) + monkeypatch.setattr(safetensors.torch, "safe_open", Tracking) + results = {} + for name, reader in (("eager", _eager), ("selective", _checkpoint._read_snapshot)): + prepared = _prepared(tmp_path / name, rank=3) + reads.clear() + with monkeypatch.context() as patch: + patch.setattr(_checkpoint, "_read_snapshot", reader) + _checkpoint._finish( + cast(Any, SimpleNamespace(_slot_state_error=ValueError)), prepared + ) + results[name] = (len(reads), sum(size for _, size in reads)) + assert results["eager"] == (100, 5160) + assert results["selective"] == (20, 1032) From 5fc36d9579c8466ace362c9b7307c9f8db490c21 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 08:09:53 +0000 Subject: [PATCH 2/2] Make inherited checkpoint test module fixtures type-check --- tests/unit/test_checkpoint_optimizer_copy.py | 12 ++++++++---- 1 file changed, 8 insertions(+), 4 deletions(-) diff --git a/tests/unit/test_checkpoint_optimizer_copy.py b/tests/unit/test_checkpoint_optimizer_copy.py index a81b7fbb2..d2d720a49 100644 --- a/tests/unit/test_checkpoint_optimizer_copy.py +++ b/tests/unit/test_checkpoint_optimizer_copy.py @@ -65,11 +65,15 @@ def _trainer(monkeypatch, *, dtype=torch.float32, rank=1, mode="full", missing=( if mode == "metadata_subset": metadata = metadata[::2] lora = ModuleType("art.megatron.lora") - lora.LoRA = _Exports # type: ignore[attr-defined] + setattr(lora, "LoRA", _Exports) publish = ModuleType("art.megatron.weights.lora_publish") - publish.collect_local_lora_entries = lambda *args, **kwargs: ( # type: ignore[attr-defined] - {item.key: torch.zeros(3, rank) for item in metadata}, - metadata, + setattr( + publish, + "collect_local_lora_entries", + lambda *args, **kwargs: ( + {item.key: torch.zeros(3, rank) for item in metadata}, + metadata, + ), ) monkeypatch.setitem(sys.modules, lora.__name__, lora) monkeypatch.setitem(sys.modules, publish.__name__, publish)