diff --git a/src/art/trainer_rank/_checkpoint.py b/src/art/trainer_rank/_checkpoint.py index 0193d534f..085c592ed 100644 --- a/src/art/trainer_rank/_checkpoint.py +++ b/src/art/trainer_rank/_checkpoint.py @@ -2,7 +2,8 @@ from __future__ import annotations -from collections.abc import Callable, Iterable, Mapping, Sequence +from collections.abc import Callable, Iterable, Iterator, Mapping, Sequence +from contextlib import ExitStack, contextmanager from copy import deepcopy from dataclasses import dataclass, field import hashlib @@ -21,6 +22,8 @@ import torch.distributed as dist if TYPE_CHECKING: + from safetensors import safe_open + from art.megatron.lora import LoRA, LoraShardMeta, LoRASlotRef from art.trainer_rank._impl import ( TrainerRank, @@ -1021,24 +1024,63 @@ def prepare_checkpoint_save( trainer._checkpoint_save_condition.notify_all() +@dataclass +class _SnapshotBlock: + opened: ExitStack + payloads: dict[str, tuple[safe_open, list[str], set[str]]] = field( + default_factory=dict + ) + shards: dict[str, list[_LocalShard]] | None = None + + +@contextmanager +def _snapshot_block(group: dist.ProcessGroup | None) -> Iterator[_SnapshotBlock]: + opened = ExitStack() + primary: BaseException | None = None + try: + yield _SnapshotBlock(opened) + except BaseException as exc: + primary = exc + raise + finally: + error: BaseException | None = None + try: + opened.close() + except BaseException as exc: + error = exc + try: + raise_distributed(error, "close checkpoint snapshot block", group) + except BaseException as exc: + if primary is None: + raise + primary.add_note(f"Checkpoint snapshot close also failed: {exc!r}") + + def _read_snapshot( prepared: _PreparedSave, relative: str, prefix: str, keys: Iterable[str] | None = None, + *, + snapshot: _SnapshotBlock | None = None, ) -> dict[str, torch.Tensor]: safe_open = importlib.import_module("safetensors").safe_open - with safe_open( - prepared.snapshot / relative, framework="pt", device="cpu" - ) as payload: - names = payload.offset_keys() + with ExitStack() as opened: + if snapshot is None: + snapshot = _SnapshotBlock(opened) + if relative not in snapshot.payloads: + payload = snapshot.opened.enter_context( + safe_open(prepared.snapshot / relative, framework="pt", device="cpu") + ) + names = payload.offset_keys() + snapshot.payloads[relative] = payload, names, set(names) + payload, names, available = snapshot.payloads[relative] 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}" @@ -1049,8 +1091,13 @@ def _read_snapshot( def _matching_shards( - prepared: _PreparedSave, metadata: Sequence[LoraShardMeta] + prepared: _PreparedSave, + metadata: Sequence[LoraShardMeta], + *, + snapshot: _SnapshotBlock | None = None, ) -> dict[str, list[_LocalShard]]: + if snapshot is not None and snapshot.shards is not None: + return snapshot.shards by_key: dict[str, list[LoraShardMeta]] = {} for item in metadata: by_key.setdefault(item.key, []).append(item) @@ -1058,6 +1105,8 @@ def _matching_shards( for record in prepared.shards: if record.metadata in by_key.get(record.metadata.key, ()): matched.setdefault(record.metadata.key, []).append(record) + if snapshot is not None: + snapshot.shards = matched return matched @@ -1066,6 +1115,8 @@ def _merge_component( metadata: Sequence[LoraShardMeta], component: str, group: dist.ProcessGroup | None, + *, + snapshot: _SnapshotBlock | None = None, ) -> dict[str, torch.Tensor]: from art.megatron.weights.lora_publish import merge_sharded_adapter_entries @@ -1074,7 +1125,7 @@ def _merge_component( error: BaseException | None = None try: if owned: - shards = _matching_shards(prepared, owned) + shards = _matching_shards(prepared, owned, snapshot=snapshot) files = {record.file for records in shards.values() for record in records} for relative in files: keys = [ @@ -1087,7 +1138,11 @@ def _merge_component( ) == relative ] - local.update(_read_snapshot(prepared, relative, component, keys)) + local.update( + _read_snapshot( + prepared, relative, component, keys, snapshot=snapshot + ) + ) except BaseException as exc: error = exc raise_distributed(error, f"read checkpoint {component} block", group) @@ -1204,67 +1259,76 @@ def _finish(trainer: TrainerRank, prepared: _PreparedSave) -> None: lora_shards: list[Path] = [] try: for index, block in enumerate(blocks): - block_metadata = [item for item in selected if item.block == block] - lora = _merge_component(prepared, block_metadata, "lora", group) - relative = f".adapter-{index:06d}.safetensors" - _rank_zero_phase( - lambda: importlib.import_module("safetensors.torch").save_file( - lora, temporary / relative - ), - "write checkpoint adapter block", - group, - ) - if _rank() == 0: - lora_shards.append(temporary / relative) - if prepared.optimizer is None: - continue - files: list[str] = [] - for component in ("master", "exp_avg", "exp_avg_sq"): - tensors = _merge_component(prepared, block_metadata, component, group) - relative = f"optimizer/{component}-{index:06d}.safetensors" - - def write_optimizer_block() -> None: - (temporary / "optimizer").mkdir(exist_ok=True) - importlib.import_module("safetensors.torch").save_file( - tensors, temporary / relative - ) - + with _snapshot_block(group) as snapshot: + block_metadata = [item for item in selected if item.block == block] + lora = _merge_component( + prepared, block_metadata, "lora", group, snapshot=snapshot + ) + relative = f".adapter-{index:06d}.safetensors" _rank_zero_phase( - write_optimizer_block, "write checkpoint optimizer block", group + lambda: importlib.import_module("safetensors.torch").save_file( + lora, temporary / relative + ), + "write checkpoint adapter block", + group, ) - files.append(relative) - if _rank() == 0: - for key in (item.key for item in block_metadata): - parameters[key] = list(files) - owned = [item for item in block_metadata if item.owner_rank == _rank()] - local_steps: dict[str, float] = {} - error: BaseException | None = None - try: - for relative in { - record.file - for records in _matching_shards(prepared, owned).values() - for record in records + if _rank() == 0: + lora_shards.append(temporary / relative) + if prepared.optimizer is None: + continue + files: list[str] = [] + for component in ("master", "exp_avg", "exp_avg_sq"): + tensors = _merge_component( + prepared, block_metadata, component, group, snapshot=snapshot + ) + relative = f"optimizer/{component}-{index:06d}.safetensors" + + def write_optimizer_block() -> None: + (temporary / "optimizer").mkdir(exist_ok=True) + importlib.import_module("safetensors.torch").save_file( + tensors, temporary / relative + ) + + _rank_zero_phase( + write_optimizer_block, "write checkpoint optimizer block", group + ) + files.append(relative) + if _rank() == 0: + for key in (item.key for item in block_metadata): + parameters[key] = list(files) + owned = [item for item in block_metadata if item.owner_rank == _rank()] + local_steps: dict[str, float] = {} + error: BaseException | None = None + try: + for relative in { + record.file + for records in _matching_shards( + prepared, owned, snapshot=snapshot + ).values() + for record in records + }: + local_steps.update( + (key, float(value.item())) + for key, value in _read_snapshot( + prepared, relative, "step", snapshot=snapshot + ).items() + ) + except BaseException as exc: + error = exc + raise_distributed(error, "read checkpoint optimizer steps", group) + step_values: dict[str, set[float]] = {} + for values in _gather(local_steps, group): + for key, value in values.items(): + step_values.setdefault(key, set()).add(value) + if mismatched := { + key: values + for key, values in step_values.items() + if len(values) != 1 }: - local_steps.update( - (key, float(value.item())) - for key, value in _read_snapshot( - prepared, relative, "step" - ).items() + raise trainer._slot_state_error( + f"Optimizer shard steps differ: {mismatched}" ) - except BaseException as exc: - error = exc - raise_distributed(error, "read checkpoint optimizer steps", group) - step_values: dict[str, set[float]] = {} - for values in _gather(local_steps, group): - for key, value in values.items(): - step_values.setdefault(key, set()).add(value) - if mismatched := { - key: values for key, values in step_values.items() if len(values) != 1 - }: - raise trainer._slot_state_error( - f"Optimizer shard steps differ: {mismatched}" - ) - steps.update((key, values.pop()) for key, values in step_values.items()) + steps.update((key, values.pop()) for key, values in step_values.items()) def commit() -> None: safe_open = importlib.import_module("safetensors").safe_open diff --git a/tests/unit/test_checkpoint_block_gloo.py b/tests/unit/test_checkpoint_block_gloo.py new file mode 100644 index 000000000..073bc890b --- /dev/null +++ b/tests/unit/test_checkpoint_block_gloo.py @@ -0,0 +1,296 @@ +"""Real CPU finalizer collectives; only public metadata and shard merge are fixtures.""" + +from __future__ import annotations + +import asyncio +from dataclasses import dataclass, replace +from datetime import timedelta +import os +from pathlib import Path +import subprocess +import sys +import time +from typing import Any, cast + +import pytest +import safetensors +from safetensors.torch import load_file, save_file +from test_checkpoint_selective_read import ( + _dependencies, + _eager, + _files, + _Meta, + _prepared, + _trainer, +) +import torch +import torch.distributed as dist + +from art.trainer_rank import _checkpoint + + +@dataclass(frozen=True) +class _ShardMeta(_Meta): + shard_rank: int = 0 + world: int = 2 + + @property + def manifest(self): + return { + "sharded": True, + "shard_world_size": self.world, + "shard_rank": self.shard_rank, + "export_shard_strategy": "uniform", + "export_shard_dim": 1, + } + + +def test_block_close_real_gloo_fanout_and_cleanup(tmp_path): + if not dist.is_gloo_available(): + pytest.skip("PyTorch was built without Gloo") + children, logs = [], [] + deadline = time.monotonic() + 90 + try: + for rank in range(2): + log = (tmp_path / f"rank-{rank}.log").open("w") + logs.append(log) + children.append( + subprocess.Popen( + [sys.executable, __file__, str(rank), str(tmp_path)], + stdout=log, + stderr=subprocess.STDOUT, + env=os.environ + | { + "PYTHONPATH": os.pathsep.join( + ( + str(Path(__file__).resolve().parents[2] / "src"), + os.environ.get("PYTHONPATH", ""), + ) + ), + "CUDA_VISIBLE_DEVICES": "", + "OMP_NUM_THREADS": "1", + "OPENBLAS_NUM_THREADS": "1", + "MKL_NUM_THREADS": "1", + "PYTHONHASHSEED": "0", + }, + ) + ) + for index, child in enumerate(children): + assert child.wait(timeout=max(0.1, deadline - time.monotonic())) == 0, ( + tmp_path / f"rank-{index}.log" + ).read_text() + finally: + for child in children: + if child.poll() is None: + child.terminate() + for child in children: + try: + child.wait(timeout=3) + except subprocess.TimeoutExpired: + child.kill() + child.wait(timeout=3) + for log in logs: + log.close() + assert all(child.poll() == 0 for child in children) + + +def _worker(rank, root): + torch.set_num_threads(1) + dist.init_process_group( + "gloo", + rank=rank, + world_size=2, + init_method=f"file://{root / 'rendezvous'}", + timeout=timedelta(seconds=10), + ) + try: + with pytest.MonkeyPatch.context() as patch: + # Reuse public merge/config fixtures, leaving group helpers untouched. + _dependencies(patch) + publish = sys.modules["art.megatron.weights.lora_publish"] + + def merge(entries): + out = {} + for key, parts in entries.items(): + if parts[0][0]["sharded"]: + assert [m["shard_rank"] for m, _ in parts] == [0, 1] + out[key] = torch.cat([value for _, value in parts], dim=1) + else: + assert len(parts) == 1 + out[key] = parts[0][1] + return out + + patch.setattr(publish, "merge_sharded_adapter_entries", merge) + outputs = [] + for case in ( + "eager", + "reuse", + "eager-cp", + "reuse-cp", + "read", + "close", + "read-close", + "cancel-close", + ): + prepared = _prepared(root / case / f"rank{rank}") + prepared.reservation.rmdir() + if rank == 0: + (root / case / "reservation").mkdir() + prepared = replace( + prepared, + destination=root / case / "result", + reservation=root / case / "reservation", + shards=tuple( + _checkpoint._LocalShard( + cast( + Any, + _ShardMeta( + shard.metadata.key, + shard.metadata.block, + shard.metadata.shape, + shard.metadata.dtype_name, + owner_rank=rank, + shard_rank=rank, + ), + ), + shard.file, + ) + for shard in prepared.shards + ), + ) + if case.endswith("cp"): + # CP replicas share a shard identity. Only owner 0 may be read. + prepared = replace( + prepared, + shards=tuple( + _checkpoint._LocalShard( + cast( + Any, + _Meta( + s.metadata.key, + s.metadata.block, + s.metadata.shape, + s.metadata.dtype_name, + owner_rank=rank, + ), + ), + "never-open" if rank else s.file, + ) + for s in prepared.shards + ), + ) + if rank: + for path in prepared.snapshot.iterdir(): + save_file( + { + k: v if k.startswith("step/") else v + 1 + for k, v in load_file(path).items() + }, + path, + ) + dist.barrier() + trainer = _trainer(prepared) + primary = ( + asyncio.CancelledError("rank1 cancelled read") + if case == "cancel-close" + else KeyError("rank1 read") + ) + close = OSError("rank1 close") + real = safetensors.safe_open + handles = [] + + class Reader: + def __init__(self, *args, **kwargs): + self.inner = real(*args, **kwargs) + self.name = Path(args[0]).name + + def __enter__(self): + self.inner.__enter__() + handles.append(self) + return self + + def __exit__(self, *args): + try: + self.inner.__exit__(*args) + finally: + handles.remove(self) + if ( + rank == 1 + and "close" in case + and self.name == "block0.safetensors" + ): + raise close + + def offset_keys(self): + return self.inner.offset_keys() + + def keys(self): + return self.inner.keys() + + def get_tensor(self, key): + if ( + rank == 1 + and ("read" in case or "cancel" in case) + and key.startswith("master/") + ): + raise primary + return self.inner.get_tensor(key) + + error = None + with pytest.MonkeyPatch.context() as reading: + reading.setattr(safetensors, "safe_open", Reader) + if case.startswith("eager"): + reading.setattr(_checkpoint, "_read_snapshot", _eager) + try: + _checkpoint.finish_checkpoint_save( + trainer, str(prepared.destination) + ) + except BaseException as exc: + error = exc + success = case.startswith(("eager", "reuse")) + if success: + assert error is None + if rank == 0: + outputs.append(_files(prepared.destination)) + else: + assert error is not None + if rank == 1: + assert error is (close if case == "close" else primary) + else: + assert isinstance( + error, RuntimeError + ) and "Another rank failed" in str(error) + if case in ("read-close", "cancel-close"): + assert any( + "snapshot close also failed" in note.lower() + for note in error.__notes__ + ) + assert not prepared.destination.exists() + assert not handles + assert ( + not prepared.snapshot.exists() and not prepared.reservation.exists() + ) + assert ( + not trainer._prepared_checkpoint_saves + and not trainer._checkpoint_finalizing_saves + ) + assert trainer._finalized_checkpoint_saves[ + str(prepared.destination) + ].outcome == ("finish" if success else "abort") + assert not list((root / case).glob(".result.tmp-*")) + assert ( + dist.get_backend(trainer._checkpoint_finalize_process_group) + == "gloo" + ) + dist.destroy_process_group(trainer._checkpoint_finalize_process_group) + dist.destroy_process_group(trainer._checkpoint_process_group) + dist.barrier() + if rank == 0: + assert len(outputs) == 4 + assert outputs[0] == outputs[1] and outputs[2] == outputs[3] + finally: + dist.destroy_process_group() + + +if __name__ == "__main__": + _worker(int(sys.argv[1]), Path(sys.argv[2])) diff --git a/tests/unit/test_checkpoint_selective_read.py b/tests/unit/test_checkpoint_selective_read.py index 1974bd6f0..a5b6d2a67 100644 --- a/tests/unit/test_checkpoint_selective_read.py +++ b/tests/unit/test_checkpoint_selective_read.py @@ -18,7 +18,7 @@ from art.trainer_rank import _checkpoint -def _eager(prepared, relative, prefix, keys=None): +def _eager(prepared, relative, prefix, keys=None, *, snapshot=None): payload = load_file(prepared.snapshot / relative) if keys is None: return { @@ -44,6 +44,11 @@ def manifest(self): @pytest.fixture def dependencies(monkeypatch): + _dependencies(monkeypatch) + monkeypatch.setattr(_checkpoint, "_ensure_finalize_group", lambda trainer: None) + + +def _dependencies(monkeypatch): # Only the unchanged single-rank replicated merge and config writer are stand-ins. publish = ModuleType("art.megatron.weights.lora_publish") @@ -65,7 +70,6 @@ def merge(entries): ) 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): @@ -325,3 +329,71 @@ def get_tensors(self): results[name] = (len(reads), sum(size for _, size in reads)) assert results["eager"] == (100, 5160) assert results["selective"] == (20, 1032) + + +def test_finalizer_opens_and_indexes_each_snapshot_once( + dependencies, monkeypatch, tmp_path +): + real = safetensors.safe_open + events = [] + + class Tracking: + def __init__(self, filename, **kwargs): + self.reader = real(filename, **kwargs) + self.path = Path(filename) + + def __enter__(self): + self.reader.__enter__() + if self.path.parent.name == "snapshot": + events.append(("open", self.path.name)) + return self + + def __exit__(self, *args): + try: + return self.reader.__exit__(*args) + finally: + if self.path.parent.name == "snapshot": + events.append(("close", self.path.name)) + + def offset_keys(self): + events.append(("index", self.path.name)) + return self.reader.offset_keys() + + def keys(self): + return self.reader.keys() + + def get_tensor(self, key): + return self.reader.get_tensor(key) + + monkeypatch.setattr(safetensors, "safe_open", Tracking) + prepared = _prepared(tmp_path) + _checkpoint.finish_checkpoint_save(_trainer(prepared), str(prepared.destination)) + assert events == [ + (action, f"block{block}.safetensors") + for block in range(2) + for action in ("open", "index", "close") + ] + + +def test_block_reader_multiple_files_selected_keys_and_tensor_lifetime(tmp_path): + payload = {"lora/a": torch.tensor([-0.0, 3.0]), "step/a": torch.tensor(350.0)} + for name in ("one", "two"): + save_file(payload, tmp_path / name) + prepared = cast(_checkpoint._PreparedSave, SimpleNamespace(snapshot=tmp_path)) + with _checkpoint._snapshot_block(None) as snapshot: + first = _checkpoint._read_snapshot( + prepared, "one", "lora", iter(["a", "a"]), snapshot=snapshot + ) + second = _checkpoint._read_snapshot(prepared, "two", "step", snapshot=snapshot) + with pytest.raises(KeyError): + _checkpoint._read_snapshot( + prepared, "one", "lora", ["missing"], snapshot=snapshot + ) + for name in ("one", "two"): + (tmp_path / name).unlink() + assert list(first) == ["a"] + assert ( + first["a"].view(torch.int32).tolist() + == payload["lora/a"].view(torch.int32).tolist() + ) + assert second["a"].item() == 350.0 diff --git a/tests/unit/test_trainer_rank_validation.py b/tests/unit/test_trainer_rank_validation.py index d6ed71048..4c280a357 100644 --- a/tests/unit/test_trainer_rank_validation.py +++ b/tests/unit/test_trainer_rank_validation.py @@ -1902,7 +1902,7 @@ def test_checkpoint_merge_rejects_same_key_with_different_metadata( ) tensor = torch.ones(2, 3) - def read(_prepared, relative, component, keys): + def read(_prepared, relative, component, keys, *, snapshot=None): assert (relative, component, list(keys)) == ("selected", "master", ["weight"]) return {"weight": tensor}