Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
36 changes: 27 additions & 9 deletions src/art/trainer_rank/_checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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
Expand Down
12 changes: 8 additions & 4 deletions tests/unit/test_checkpoint_optimizer_copy.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
327 changes: 327 additions & 0 deletions tests/unit/test_checkpoint_selective_read.py
Original file line number Diff line number Diff line change
@@ -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)
Loading