From 639231941ae05de42228833d857c37b0e736302f Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 11 Aug 2026 05:18:35 +0000 Subject: [PATCH 1/8] Implement topology-portable trainer checkpoints --- dev/trainer_rank.py | 20 +- dev/trainer_rank_checkpoint_benchmark.py | 341 +++ dev/trainer_rank_support.py | 17 +- src/art/megatron/_collective.py | 74 + src/art/megatron/lora.py | 444 ++-- src/art/megatron/model_support/lora_disk.py | 72 +- src/art/megatron/provider.py | 29 +- src/art/megatron/train.py | 26 +- src/art/megatron/weights/lora_publish.py | 762 ++++--- src/art/trainer_rank/__init__.py | 87 +- src/art/trainer_rank/_checkpoint.py | 1685 +++++++++++++++ src/art/trainer_rank/_impl.py | 924 +++++---- .../megatron/lora/test_dynamic_lora_slots.py | 217 +- .../megatron/lora/test_lora_disk_codecs.py | 1824 +++++++++++++++-- .../megatron/model_support/oracle_worker.py | 5 +- .../model_support/test_provider_support.py | 54 +- .../train_inf_mismatch/output_parity.py | 12 +- .../unit/test_megatron_reference_logprobs.py | 3 +- tests/unit/test_trainer_rank_validation.py | 510 +++-- tests/unit/test_trainer_rank_weird_shapes.py | 19 +- 20 files changed, 5682 insertions(+), 1443 deletions(-) create mode 100644 dev/trainer_rank_checkpoint_benchmark.py create mode 100644 src/art/megatron/_collective.py create mode 100644 src/art/trainer_rank/_checkpoint.py diff --git a/dev/trainer_rank.py b/dev/trainer_rank.py index 0dd8b24bc..f1d6518ea 100644 --- a/dev/trainer_rank.py +++ b/dev/trainer_rank.py @@ -1,5 +1,7 @@ +from collections.abc import Iterable, Mapping from itertools import islice import os +from typing import cast import torch import torch.distributed as dist @@ -35,7 +37,12 @@ def main( tokenizer = AutoTokenizer.from_pretrained(model, trust_remote_code=True) inputs: list[ForwardInput[torch.Tensor, None, None, None]] = [] - rows = load_dataset("roneneldan/TinyStories", split="train", streaming=True) + rows = cast( + Iterable[Mapping[str, object]], + load_dataset("roneneldan/TinyStories", split="train", streaming=True), + ) + if not callable(tokenizer): + raise TypeError("Tokenizer backend is not callable") for row in islice(rows, samples): token_ids = tokenizer( str(row["text"]), # type: ignore[index] @@ -45,9 +52,12 @@ def main( return_tensors="pt", )["input_ids"].reshape(-1) inputs.append( - ForwardInput( - input_tokens=token_ids[:-1], - target_tokens=token_ids[1:], + cast( + ForwardInput[torch.Tensor, None, None, None], + ForwardInput( + input_tokens=token_ids[:-1], + target_tokens=token_ids[1:], + ), ) ) @@ -58,7 +68,7 @@ def main( ) rank = TrainerRank(runtime) (slot,) = load_random_checkpoint_slots(runtime, rank, 1, lora_rank=lora_rank) - rank.set_checkpoint(slot) + rank._set_default_slot(rank._slot_ref(slot)) for step in range(steps): loss_sum = torch.tensor(0.0, device=rank.device) diff --git a/dev/trainer_rank_checkpoint_benchmark.py b/dev/trainer_rank_checkpoint_benchmark.py new file mode 100644 index 000000000..f17621a1a --- /dev/null +++ b/dev/trainer_rank_checkpoint_benchmark.py @@ -0,0 +1,341 @@ +"""Profile portable TrainerRank checkpoint save/load under ``torchrun``. + +The JSON schema is intentionally stable so results from different commits and +topologies can be compared. Caladan augments ``post_queue_upload_seconds``. +""" + +from __future__ import annotations + +import argparse +import asyncio +from collections.abc import Sequence +import json +import math +import os +from pathlib import Path +import threading +import time +from typing import Any, Literal + +import torch +import torch.distributed as dist + +from art.megatron.model_support.lora_disk import load_adapter_config +from art.megatron.weights import lora_publish +from art.trainer_rank import AdamParams, TrainerRank + +Operation = Literal["load", "save", "roundtrip"] + + +def _tensor_bytes(shape: Sequence[int], dtype_name: str) -> int: + dtype = getattr(torch, dtype_name) + return math.prod(shape) * torch.empty((), dtype=dtype).element_size() + + +class _CheckpointProbe: + def __init__(self, rank: int) -> None: + self.rank = rank + self.gather_seconds = 0.0 + self.sent_bytes = 0 + self._exchange = lora_publish._exchange_tensors + + def install(self) -> None: + def exchange(metadata, **kwargs): + started = time.perf_counter() + try: + return self._exchange(metadata, **kwargs) + finally: + self.gather_seconds += time.perf_counter() - started + self.sent_bytes += sum( + _tensor_bytes(meta.shape, meta.dtype_name) + for meta in metadata + if meta.owner_rank == self.rank and self.rank != 0 + ) + + setattr(lora_publish, "_exchange_tensors", exchange) + + def restore(self) -> None: + setattr(lora_publish, "_exchange_tensors", self._exchange) + + +class _RssSampler: + def __init__(self) -> None: + self.initial = self._rss() + self.peak = self.initial + self._stop = threading.Event() + self._thread = threading.Thread(target=self._sample, daemon=True) + + @staticmethod + def _rss() -> int: + with Path("/proc/self/statm").open() as handle: + resident_pages = int(handle.read().split()[1]) + return resident_pages * os.sysconf("SC_PAGE_SIZE") + + def _sample(self) -> None: + while not self._stop.wait(0.01): + self.peak = max(self.peak, self._rss()) + + def __enter__(self) -> _RssSampler: + self._thread.start() + return self + + def __exit__(self, *_args: object) -> None: + self._stop.set() + self._thread.join() + self.peak = max(self.peak, self._rss()) + + +def _artifact_size(path: str | None) -> int | None: + if path is None or not Path(path).is_dir(): + return None + return sum(item.stat().st_size for item in Path(path).rglob("*") if item.is_file()) + + +def _topology(args: argparse.Namespace, world_size: int) -> dict[str, int]: + model_parallel = args.tp * args.pp * args.cp + if world_size % model_parallel: + raise ValueError( + f"world_size={world_size} is not divisible by tp*pp*cp={model_parallel}" + ) + return { + "world_size": world_size, + "tp": args.tp, + "pp": args.pp, + "cp": args.cp, + "ep": args.ep, + "etp": args.etp, + "dp": world_size // model_parallel, + } + + +def _initialize_optimizer_state( + trainer: TrainerRank, checkpoint: str, learning_rate: float +) -> int: + dynamic = trainer._dynamic_optimizers.get(checkpoint) + if dynamic is None: + dynamic = trainer._new_dynamic_optimizer( + checkpoint, AdamParams(learning_rate=learning_rate) + ) + trainer._dynamic_optimizers[checkpoint] = dynamic + for index, master in enumerate(dynamic.master_params, start=1): + master.grad = torch.full_like(master, 1e-3 * index) + dynamic.optimizer.step() + dynamic.optimizer.zero_grad(set_to_none=True) + with torch.no_grad(): + for model, master in zip( + trainer._checkpoint_slot_params_by_name[checkpoint], + dynamic.master_params, + strict=True, + ): + model.copy_(master) + model.grad = None + + resident_bytes = 0 + for master in dynamic.master_params: + state = dynamic.optimizer.state[master] + moments = (state.get("exp_avg"), state.get("exp_avg_sq")) + if not all( + isinstance(moment, torch.Tensor) and bool(torch.count_nonzero(moment)) + for moment in moments + ): + raise RuntimeError("benchmark optimizer moments were not initialized") + step = state.get("step") + if not isinstance(step, torch.Tensor) or float(step.item()) < 1: + raise RuntimeError("benchmark optimizer step was not initialized") + resident_bytes += master.numel() * master.element_size() + resident_bytes += sum( + moment.numel() * moment.element_size() + for moment in moments + if isinstance(moment, torch.Tensor) + ) + return resident_bytes + + +async def _exercise( + trainer: TrainerRank, + *, + operation: Operation, + source: str, + output: str | None, + learning_rate: float, +) -> dict[str, float | int | None]: + started = time.perf_counter() + await trainer.load_checkpoint(source) + input_load = time.perf_counter() - started + save_total: float | None = None + queue_pause: float | None = None + post_queue_serialization: float | None = None + restored_load: float | None = None + resident_optimizer_bytes: int | None = None + if operation in {"save", "roundtrip"}: + if output is None: + raise ValueError("--output is required for save and roundtrip") + resident_optimizer_bytes = _initialize_optimizer_state( + trainer, source, learning_rate + ) + started = time.perf_counter() + trainer._prepare_checkpoint_save(output, source) + queue_pause = time.perf_counter() - started + started = time.perf_counter() + trainer._finish_checkpoint_save(output) + post_queue_serialization = time.perf_counter() - started + save_total = queue_pause + post_queue_serialization + if operation == "roundtrip": + assert output is not None + started = time.perf_counter() + await trainer.load_checkpoint(output) + restored_load = time.perf_counter() - started + return { + "input_load_seconds": input_load, + "save_total_seconds": save_total, + "queue_pause_seconds": queue_pause, + "post_queue_serialization_seconds": post_queue_serialization, + "load_total_seconds": restored_load + if restored_load is not None + else input_load, + "ready_seconds": restored_load if restored_load is not None else input_load, + "resident_optimizer_state_bytes": resident_optimizer_bytes, + } + + +def _parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--source", required=True) + parser.add_argument("--output") + parser.add_argument("--output-json") + parser.add_argument( + "--operation", choices=("load", "save", "roundtrip"), default="roundtrip" + ) + parser.add_argument("--model") + parser.add_argument("--layers", type=int, default=0) + parser.add_argument("--learning-rate", type=float, default=3e-4) + parser.add_argument("--tp", type=int, default=1) + parser.add_argument("--pp", type=int, default=1) + parser.add_argument("--cp", type=int, default=1) + parser.add_argument("--ep", type=int, default=1) + parser.add_argument("--etp", type=int, default=1) + return parser.parse_args() + + +def main() -> None: + args = _parse_args() + if not torch.cuda.is_available(): + raise RuntimeError("checkpoint benchmark requires CUDA") + device = int(os.environ["LOCAL_RANK"]) + torch.cuda.set_device(device) + dist.init_process_group("nccl") + rank = dist.get_rank() + world_size = dist.get_world_size() + topology = _topology(args, world_size) + for key, value in ( + ("TENSOR_MODEL", args.tp), + ("PIPELINE_MODEL", args.pp), + ("CONTEXT", args.cp), + ("EXPERT_MODEL", args.ep), + ("EXPERT_TENSOR", args.etp), + ): + os.environ[f"ART_MEGATRON_{key}_PARALLEL_SIZE"] = str(value) + + try: + from art.megatron import train as megatron_train + + config = load_adapter_config(args.source) + model = args.model or config.get("base_model_name_or_path") + if not isinstance(model, str) or not model: + raise ValueError( + "--model or checkpoint base_model_name_or_path is required" + ) + runtime = megatron_train.build_training_runtime( + model_identifier=model, + provider_configure=( + (lambda provider: setattr(provider, "num_layers", args.layers)) + if args.layers > 0 + else None + ), + print_env=rank == 0, + ) + trainer = TrainerRank(runtime) + probe = _CheckpointProbe(rank) + torch.cuda.reset_peak_memory_stats(device) + initial_gpu = torch.cuda.memory_allocated(device) + probe.install() + try: + with _RssSampler() as rss: + timings = asyncio.run( + _exercise( + trainer, + operation=args.operation, + source=args.source, + output=args.output, + learning_rate=args.learning_rate, + ) + ) + finally: + probe.restore() + rank_metrics = { + **timings, + "rank": rank, + "snapshot_gather_seconds": probe.gather_seconds, + "communication_sent_bytes": probe.sent_bytes, + "peak_cpu_rss_bytes": rss.peak, + "peak_cpu_rss_delta_bytes": max(rss.peak - rss.initial, 0), + "peak_gpu_allocated_bytes": torch.cuda.max_memory_allocated(device), + "peak_gpu_allocated_delta_bytes": max( + torch.cuda.max_memory_allocated(device) - initial_gpu, 0 + ), + "peak_gpu_reserved_bytes": torch.cuda.max_memory_reserved(device), + } + all_metrics: list[dict[str, Any] | None] = [None] * world_size + dist.all_gather_object(all_metrics, rank_metrics) + if rank == 0: + ranks = [value for value in all_metrics if value is not None] + + def maximum(key: str) -> float | int | None: + values = [value[key] for value in ranks if value[key] is not None] + return max(values) if values else None + + aggregate = { + "queue_pause_seconds": maximum("queue_pause_seconds"), + "snapshot_gather_seconds": maximum("snapshot_gather_seconds"), + "save_total_seconds": maximum("save_total_seconds"), + "load_total_seconds": maximum("load_total_seconds"), + "post_queue_serialization_seconds": maximum( + "post_queue_serialization_seconds" + ), + "post_queue_upload_seconds": None, + "distributed_communication_bytes": sum( + value["communication_sent_bytes"] for value in ranks + ), + "artifact_size_bytes": _artifact_size(args.output or args.source), + "resident_optimizer_state_bytes": maximum( + "resident_optimizer_state_bytes" + ), + "peak_cpu_rss_bytes": maximum("peak_cpu_rss_bytes"), + "peak_cpu_rss_delta_bytes": maximum("peak_cpu_rss_delta_bytes"), + "peak_gpu_allocated_bytes": maximum("peak_gpu_allocated_bytes"), + "peak_gpu_allocated_delta_bytes": maximum( + "peak_gpu_allocated_delta_bytes" + ), + "peak_gpu_reserved_bytes": maximum("peak_gpu_reserved_bytes"), + "ready_seconds": maximum("ready_seconds"), + } + payload = { + "schema_version": 1, + "operation": args.operation, + "model": model, + "source": args.source, + "output": args.output, + "topology": topology, + "aggregate": aggregate, + "ranks": ranks, + } + encoded = json.dumps(payload, indent=2, sort_keys=True) + if args.output_json: + Path(args.output_json).write_text(encoded + "\n") + print(encoded, flush=True) + finally: + dist.destroy_process_group() + + +if __name__ == "__main__": + main() diff --git a/dev/trainer_rank_support.py b/dev/trainer_rank_support.py index 2d9e07423..2dbb89e77 100644 --- a/dev/trainer_rank_support.py +++ b/dev/trainer_rank_support.py @@ -17,12 +17,13 @@ def load_random_checkpoint_slots( assert count >= 0, "slots must be >= 0" if count == 0: return () - from art.megatron.lora import LoRAPublishPlanner + from art.megatron.weights.lora_publish import collect_local_lora_entries - gathered: list[list[Any] | None] = [None] * dist.get_world_size() - dist.all_gather_object( - gathered, LoRAPublishPlanner(runtime.model).global_metadata({}) + _tensors, local_metadata = collect_local_lora_entries( + runtime.model, {}, owner_rank=dist.get_rank() ) + gathered: list[list[Any] | None] = [None] * dist.get_world_size() + dist.all_gather_object(gathered, local_metadata) metadata = {meta.key: meta for values in gathered if values for meta in values} selected = sorted(metadata.values(), key=lambda item: item.key) if site_limit is not None: @@ -58,7 +59,11 @@ def load_random_checkpoint_slots( shape, device=rank.device, dtype=dtype, generator=generator ) adapter[meta.key] = tensor if is_a else tensor.mul_(1e-3) - assert rank.load_checkpoint_slot(name, adapter) > 0, ( - "TrainerRank check requires installed LoRA adapter sites" + loaded = rank._load_checkpoint_slot(name, adapter, alpha=lora_rank) + assert loaded > 0, "TrainerRank check requires installed LoRA adapter sites" + ref = rank._slot_ref(name) + rank._checkpoint_slot_params_by_name[name] = tuple( + rank._iter_slot_parameters(ref) ) + rank._checkpoint_revisions[name] = 0 return names diff --git a/src/art/megatron/_collective.py b/src/art/megatron/_collective.py new file mode 100644 index 000000000..afd605acc --- /dev/null +++ b/src/art/megatron/_collective.py @@ -0,0 +1,74 @@ +"""Small collective helpers shared by Megatron checkpoint exporters.""" + +from collections.abc import Iterator +from contextlib import contextmanager +from typing import cast + +import torch +import torch.distributed as dist + + +def distributed() -> bool: + return dist.is_available() and dist.is_initialized() + + +def gather_objects[T]( + value: T, + *, + group: dist.ProcessGroup | None = None, +) -> tuple[T, ...]: + if not distributed(): + return (value,) + values: list[T | None] = [None] * dist.get_world_size(group) + dist.all_gather_object(values, value, group=group) + return tuple(cast(T, item) for item in values) + + +def raise_distributed( + error: BaseException | None, + phase: str, + *, + group: dist.ProcessGroup | None = None, +) -> None: + errors = gather_objects(None if error is None else repr(error), group=group) + if not any(errors): + return + if error is not None: + raise error + raise RuntimeError( + f"Another rank failed to {phase}: {next(item for item in errors if item)}" + ) + + +@contextmanager +def collective_errors( + phase: str, *, group: dist.ProcessGroup | None = None +) -> Iterator[None]: + error: BaseException | None = None + try: + yield + except BaseException as exc: + error = exc + raise_distributed(error, phase, group=group) + + +def rank() -> int: + return dist.get_rank() if distributed() else 0 + + +def device() -> torch.device: + return ( + torch.device("cuda", torch.cuda.current_device()) + if torch.cuda.is_available() + else torch.device("cpu") + ) + + +def dtype_from_name(name: str) -> torch.dtype: + if isinstance(dtype := getattr(torch, name, None), torch.dtype): + return dtype + raise RuntimeError(f"Unsupported tensor dtype: {name!r}") + + +def dtype_name(dtype: torch.dtype) -> str: + return str(dtype).removeprefix("torch.") diff --git a/src/art/megatron/lora.py b/src/art/megatron/lora.py index 5bc10ed93..9ae5e1e15 100644 --- a/src/art/megatron/lora.py +++ b/src/art/megatron/lora.py @@ -8,7 +8,16 @@ import math import os import re -from typing import Any, Callable, Literal, NamedTuple, TypeVar, cast +from typing import ( + Any, + Callable, + Literal, + NamedTuple, + NotRequired, + TypedDict, + TypeVar, + cast, +) from megatron.bridge.models.gpt_provider import GPTModelProvider from megatron.core import parallel_state as ps @@ -30,6 +39,8 @@ from megatron.core.transformer.transformer_layer import TransformerLayer import torch +from art.megatron._collective import dtype_name as _dtype_name + from .kernels.cute_grouped_lora_quack import ( quack_grouped_lora, quack_grouped_lora_dual, @@ -45,7 +56,6 @@ ShardDomain = Literal["tp", "expert_tp"] GradSyncDomain = Literal["tp_default", "expert_tp"] GradSyncOp = Literal["none", "sum", "avg"] -LoraSlotKind = Literal["checkpoint", "lora"] _F = TypeVar("_F", bound=Callable[..., Any]) TP_DEFAULT_GRAD_SYNC_DOMAIN: GradSyncDomain = "tp_default" @@ -57,7 +67,6 @@ @dataclass(frozen=True) class LoRASlotRef: - kind: LoraSlotKind name: str | None @@ -170,12 +179,23 @@ class LoRAParallelSpec: grad_sync_op: GradSyncOp = GRAD_SYNC_OP_NONE +class LoraShardManifest(TypedDict): + domain: NotRequired[ShardDomain] + sharded: bool + shard_dim: NotRequired[int | None] + shard_world_size: int + shard_rank: int + export_shard_dim: NotRequired[int] + export_shard_strategy: NotRequired[str] + component_sizes: NotRequired[Sequence[int]] + + class LoraShardMeta(NamedTuple): key: str owner_rank: int shape: tuple[int, ...] dtype_name: str - manifest: dict[str, Any] + manifest: LoraShardManifest block: str @property @@ -183,20 +203,6 @@ def numel(self) -> int: return math.prod(self.shape) -class _LoraPublishTemplate(NamedTuple): - adapter_model_prefix: str - suffix: str - shape: tuple[int, ...] - dtype_name: str - num_local_experts: int - shard_domain: ShardDomain - sharded: bool - shard_world_size: int - export_shard_dim: int - export_shard_strategy: str | None - component_sizes: tuple[int, ...] - - def _distributed_initialized() -> bool: is_initialized = getattr(torch.distributed, "is_initialized", None) return ( @@ -228,10 +234,6 @@ def _get_shard_rank(domain: ShardDomain) -> int: return group.rank() -def _dtype_name(dtype: torch.dtype) -> str: - return str(dtype).removeprefix("torch.") - - def _block_for_key(key: str) -> str: match = _LAYER_BLOCK_RE.match(key) if match is not None: @@ -239,19 +241,6 @@ def _block_for_key(key: str) -> str: return "__global__" -def _process_group_ranks(group: Any | None) -> tuple[int, ...]: - if group is None or not _distributed_initialized(): - return (0,) - get_process_group_ranks = getattr( - torch.distributed, - "get_process_group_ranks", - None, - ) - if not callable(get_process_group_ranks): - raise RuntimeError("torch.distributed.get_process_group_ranks is unavailable") - return tuple(int(rank) for rank in get_process_group_ranks(group)) - - def _normalize_axis(axis: int, ndim: int) -> int: if axis < 0: axis += ndim @@ -429,13 +418,12 @@ def __init__( alpha: float, a_template: torch.nn.Parameter, b_template: torch.nn.Parameter, - requires_grad: bool, ) -> None: super().__init__() self.ref = ref self.alpha = float(alpha) - self.A_T = torch.nn.Parameter(a_t.detach().clone(), requires_grad=requires_grad) - self.B_T = torch.nn.Parameter(b_t.detach().clone(), requires_grad=requires_grad) + self.A_T = torch.nn.Parameter(a_t.detach().clone()) + self.B_T = torch.nn.Parameter(b_t.detach().clone()) _copy_lora_param_metadata(a_template, self.A_T) _copy_lora_param_metadata(b_template, self.B_T) @@ -473,6 +461,7 @@ def __init__( self.out_features = int(out_features) self.scale = alpha / rank self._slot_modules = torch.nn.ModuleDict() + self._next_slot_key = 0 self._slot_keys: dict[LoRASlotRef, str] = {} self.A_T = torch.nn.Parameter( torch.zeros( @@ -538,7 +527,7 @@ def reset_lora_parameters(self) -> None: self._broadcast_if_replicated(self.B_T) def _expected_weight_keys(self, suffix: str) -> list[str]: - if self.num_local_experts > 1: + if "{expert}" in self.adapter_model_prefix: return [ f"{self.adapter_model_prefix.format(expert=expert + self._expert_offset)}.{suffix}.weight" for expert in range(self.num_local_experts) @@ -551,33 +540,55 @@ def load_lora_slot( adapter_model: dict[str, torch.Tensor], *, alpha: float = LORA_ALPHA, - requires_grad: bool, + localized: bool = False, ) -> bool: if ref.name is None: raise ValueError("base-model slot refs do not own LoRA tensors") weights = self._adapter_weights(adapter_model, require=False) if weights is None: return False - a_t = self._localized_weight(weights[0], into=self.A_T) - b_t = self._localized_weight(weights[1], into=self.B_T) + a_t = ( + weights[0].contiguous() + if localized + else self._localized_weight(weights[0], into=self.A_T) + ) + b_t = ( + weights[1].contiguous() + if localized + else self._localized_weight(weights[1], into=self.B_T) + ) + if ( + a_t.ndim != self.A_T.ndim + or tuple(a_t.shape[:-1]) != tuple(self.A_T.shape[:-1]) + or b_t.ndim != self.B_T.ndim + or tuple(b_t.shape[:-2]) != tuple(self.B_T.shape[:-2]) + or int(a_t.shape[-1]) != int(b_t.shape[-2]) + or int(b_t.shape[-1]) != int(self.B_T.shape[-1]) + ): + raise ValueError( + f"{self.adapter_model_prefix}: local LoRA checkpoint shape mismatch" + ) slot_key = self._slot_keys.get(ref) - if slot_key is None: - slot_key = f"slot_{len(self._slot_keys)}" - self._slot_keys[ref] = slot_key - elif self._has_live_slot_grads(ref): + if slot_key is not None and self._has_live_slot_grads(ref): raise RuntimeError( - f"Cannot overwrite live LoRA slot {ref.kind}:{ref.name} for " + f"Cannot overwrite live LoRA slot {ref.name!r} for " f"{self.adapter_model_prefix}; clear grads/backward graph first." ) - self._slot_modules[slot_key] = LoRASlot( + slot = LoRASlot( ref=ref, a_t=a_t, b_t=b_t, alpha=alpha, a_template=self.A_T, b_template=self.B_T, - requires_grad=requires_grad, ) + if slot_key is None: + while f"slot_{self._next_slot_key}" in self._slot_modules: + self._next_slot_key += 1 + slot_key = f"slot_{self._next_slot_key}" + self._next_slot_key += 1 + self._slot_modules[slot_key] = slot + self._slot_keys[ref] = slot_key return True def lora_slot_params(self, ref: LoRASlotRef) -> list[torch.nn.Parameter]: @@ -700,7 +711,7 @@ def _should_export_parameter(self, param: torch.nn.Parameter) -> bool: Determine if the given LoRA param should be exported in the sharded LoRA state dict (drop replicated ranks/params). """ - if self.num_local_experts > 1: # self is a MoE layer + if "{expert}" in self.adapter_model_prefix: if ps.get_expert_data_parallel_rank() != 0: return False else: # self is a non-MoE layer @@ -714,8 +725,8 @@ def _should_export_parameter(self, param: torch.nn.Parameter) -> bool: # param is replicated, tp rank 0 or etp rank 0 participates return _get_shard_rank(param.lora_shard_domain) == 0 # ty: ignore[unresolved-attribute] - def _manifest_for_param(self, param: torch.nn.Parameter) -> dict[str, Any]: - manifest = { + def _manifest_for_param(self, param: torch.nn.Parameter) -> LoraShardManifest: + manifest: LoraShardManifest = { "domain": param.lora_shard_domain, # ty: ignore[unresolved-attribute] "sharded": param.lora_tp_sharded, # ty: ignore[unresolved-attribute] "shard_dim": param.lora_tp_shard_dim, # ty: ignore[unresolved-attribute] @@ -761,17 +772,19 @@ def _export_items( for key, param in self._lora_params(ref): if not self._should_export_parameter(param): continue - if self.num_local_experts > 1: + if "{expert}" in self.adapter_model_prefix: for expert in range(self.num_local_experts): full_key = f"{self.adapter_model_prefix.format(expert=expert + self._expert_offset)}.{key}" - export_items.append((full_key, param, expert)) + export_items.append( + (full_key, param, expert if param.ndim == 3 else None) + ) else: export_items.append((f"{self.adapter_model_prefix}.{key}", param, None)) return export_items def sharded_lora_manifest( self, ref: LoRASlotRef | None = None - ) -> dict[str, dict[str, Any]]: + ) -> dict[str, LoraShardManifest]: return { key: self._manifest_for_param(param) for key, param, _expert in self._export_items(ref) @@ -835,209 +848,6 @@ def forward( return out if scale == 1.0 else out * scale -class LoRAPublishPlanner: - def __init__( - self, - model_chunks: Sequence[torch.nn.Module], - slot_ref: LoRASlotRef | None = None, - ) -> None: - self.templates = tuple(self._collect_templates(model_chunks, slot_ref)) - - def global_metadata( - self, - adapter_dtypes: dict[str, torch.dtype], - ) -> list[LoraShardMeta]: - if _distributed_initialized(): - pp_world_size = ps.get_pipeline_model_parallel_world_size() - if pp_world_size != 1: - raise RuntimeError( - "LoRA publish planner requires pipeline_model_parallel_size=1; " - f"got {pp_world_size}. Rank-local modules cannot describe remote " - "pipeline stages without exchanging templates." - ) - return [ - meta - for template in self.templates - for meta in self._metadata_for_template(template, adapter_dtypes) - ] - - @staticmethod - def _collect_templates( - model_chunks: Sequence[torch.nn.Module], - slot_ref: LoRASlotRef | None = None, - ) -> list[_LoraPublishTemplate]: - templates: list[_LoraPublishTemplate] = [] - for chunk in model_chunks: - for module in chunk.modules(): - if not isinstance(module, LoRA): - continue - for suffix, param in module._lora_params(slot_ref): - if not module._should_export_parameter(param): - continue - sharded = bool(getattr(param, "lora_tp_sharded")) - shard_domain = getattr(param, "lora_shard_domain") - if shard_domain not in ("tp", "expert_tp"): - raise RuntimeError( - f"invalid LoRA shard domain: {shard_domain!r}" - ) - templates.append( - _LoraPublishTemplate( - adapter_model_prefix=module.adapter_model_prefix, - suffix=suffix, - shape=_exported_param_shape(module, param), - dtype_name=_dtype_name(param.dtype), - num_local_experts=module.num_local_experts, - shard_domain=shard_domain, - sharded=sharded, - shard_world_size=( - _get_shard_world_size(shard_domain) if sharded else 1 - ), - export_shard_dim=( - _exported_shard_dim(param) if sharded else -1 - ), - export_shard_strategy=( - getattr(param, "lora_tp_shard_strategy", "uniform") - if sharded - else None - ), - component_sizes=tuple( - int(size) - for size in getattr( - param, - "lora_tp_component_sizes", - (), - ) - ), - ) - ) - return templates - - def _metadata_for_template( - self, - template: _LoraPublishTemplate, - adapter_dtypes: dict[str, torch.dtype], - ) -> list[LoraShardMeta]: - shard_ranks = range(template.shard_world_size) if template.sharded else (0,) - if template.num_local_experts <= 1: - tp_ranks = ( - _process_group_ranks(ps.get_tensor_model_parallel_group()) - if _distributed_initialized() - else (0,) - ) - owners = [ - ( - f"{template.adapter_model_prefix}.{template.suffix}", - tp_ranks[shard_rank], - shard_rank, - ) - for shard_rank in shard_ranks - ] - else: - ep_world_size = 1 - if _distributed_initialized(): - ep_world_size = ps.get_expert_model_parallel_world_size() - owners = [ - ( - f"{template.adapter_model_prefix.format(expert=expert)}.{template.suffix}", - self._expert_owner_rank(ep_rank, shard_rank), - shard_rank, - ) - for ep_rank in range(ep_world_size) - for local_expert in range(template.num_local_experts) - for expert in [ep_rank * template.num_local_experts + local_expert] - for shard_rank in shard_ranks - ] - return [ - self._make_metadata( - template, - key=key, - owner_rank=owner_rank, - shard_rank=shard_rank, - adapter_dtypes=adapter_dtypes, - ) - for key, owner_rank, shard_rank in owners - ] - - @staticmethod - def _make_metadata( - template: _LoraPublishTemplate, - *, - key: str, - owner_rank: int, - shard_rank: int, - adapter_dtypes: dict[str, torch.dtype], - ) -> LoraShardMeta: - manifest: dict[str, Any] = { - "sharded": template.sharded, - "shard_world_size": template.shard_world_size if template.sharded else 1, - "shard_rank": shard_rank if template.sharded else 0, - } - if template.sharded: - manifest["export_shard_dim"] = template.export_shard_dim - manifest["export_shard_strategy"] = ( - template.export_shard_strategy or "uniform" - ) - if template.component_sizes: - manifest["component_sizes"] = list(template.component_sizes) - return LoraShardMeta( - key=key, - owner_rank=owner_rank, - shape=template.shape, - dtype_name=( - _dtype_name(adapter_dtypes[key]) - if key in adapter_dtypes - else template.dtype_name - ), - manifest=manifest, - block=_block_for_key(key), - ) - - @staticmethod - def _expert_owner_rank(ep_rank: int, shard_rank: int) -> int: - if not _distributed_initialized(): - return 0 - joint_ranks = _process_group_ranks( - ps.get_expert_tensor_and_model_parallel_group(check_initialized=False) - ) - ep_world_size = ps.get_expert_model_parallel_world_size() - etp_world_size = _get_shard_world_size("expert_tp") - expected_size = ep_world_size * etp_world_size - if len(joint_ranks) != expected_size: - raise RuntimeError( - "Unexpected expert TP x EP group size: " - f"got {len(joint_ranks)}, expected {expected_size}" - ) - if shard_rank >= etp_world_size: - raise RuntimeError( - f"Invalid expert tensor shard rank {shard_rank} for world size {etp_world_size}" - ) - if ep_rank >= ep_world_size: - raise RuntimeError( - f"Invalid expert parallel rank {ep_rank} for world size {ep_world_size}" - ) - - ep_group_ranks = _process_group_ranks(ps.get_expert_model_parallel_group()) - etp_group = ps.get_expert_tensor_parallel_group(check_initialized=False) - etp_group_ranks = _process_group_ranks(etp_group) - ep_positions = [joint_ranks.index(rank) for rank in ep_group_ranks] - etp_positions = [joint_ranks.index(rank) for rank in etp_group_ranks] - - if etp_positions == list(range(etp_world_size)): - return joint_ranks[ep_rank * etp_world_size + shard_rank] - if ep_positions == list(range(ep_world_size)): - return joint_ranks[shard_rank * ep_world_size + ep_rank] - raise RuntimeError( - "Unsupported expert TP x EP group rank order: " - f"joint={joint_ranks}, ep_positions={ep_positions}, etp_positions={etp_positions}" - ) - - -def _exported_param_shape(module: LoRA, param: torch.nn.Parameter) -> tuple[int, ...]: - if module.num_local_experts > 1: - return tuple(int(dim) for dim in param[0].T.shape) - return tuple(int(dim) for dim in param.T.shape) - - @torch.compiler.disable def _expert_grouped_lora_forward( lora: LoRA, @@ -1935,7 +1745,7 @@ def load_lora_slot_into_model( adapter_model: dict[str, torch.Tensor], *, alpha: float = LORA_ALPHA, - requires_grad: bool, + localized: bool = False, ) -> int: loaded = 0 for chunk in model: @@ -1944,14 +1754,126 @@ def load_lora_slot_into_model( ref, adapter_model, alpha=alpha, - requires_grad=requires_grad, + localized=localized, ): loaded += 1 if loaded == 0 and ref.name is not None: - raise RuntimeError(f"LoRA slot {ref.kind}:{ref.name} loaded no adapter sites") + raise RuntimeError(f"LoRA slot {ref.name!r} loaded no adapter sites") return loaded +def delete_lora_slot_from_model( + model: Sequence[torch.nn.Module], ref: LoRASlotRef +) -> None: + for chunk in model: + for module in chunk.modules(): + if not isinstance(module, LoRA): + continue + key = module._slot_keys.pop(ref, None) + if key is not None and key in module._slot_modules: + del module._slot_modules[key] + + +def _lora_slot_replacements( + model: Sequence[torch.nn.Module], + source: LoRASlotRef, + destination: LoRASlotRef, +) -> list[tuple[LoRA, str | None, str | None]]: + replacements: list[tuple[LoRA, str | None, str | None]] = [] + for chunk in model: + for module in chunk.modules(): + if not isinstance(module, LoRA): + continue + source_key = module._slot_keys.get(source) + destination_key = module._slot_keys.get(destination) + for ref, key in ((source, source_key), (destination, destination_key)): + if key is None: + continue + if key not in module._slot_modules: + raise RuntimeError( + f"LoRA slot {ref.name!r} maps to missing module {key!r}" + ) + slot = cast(LoRASlot, module._slot_modules[key]) + if slot.ref != ref: + raise RuntimeError( + f"LoRA slot {ref.name!r} has inconsistent module metadata" + ) + replacements.append((module, source_key, destination_key)) + return replacements + + +type _LoraSlotSnapshot = tuple[ + tuple[ + LoRA, + dict[LoRASlotRef, str], + dict[str, LoRASlot], + dict[str, LoRASlotRef], + ], + ..., +] + + +def _snapshot_lora_slots( + model: Sequence[torch.nn.Module], +) -> _LoraSlotSnapshot: + snapshots = [] + seen: set[int] = set() + for chunk in model: + for module in chunk.modules(): + if not isinstance(module, LoRA) or id(module) in seen: + continue + seen.add(id(module)) + slots = { + key: cast(LoRASlot, slot) for key, slot in module._slot_modules.items() + } + snapshots.append( + ( + module, + dict(module._slot_keys), + slots, + {key: slot.ref for key, slot in slots.items()}, + ) + ) + return tuple(snapshots) + + +def _restore_lora_slots(snapshot: _LoraSlotSnapshot) -> None: + for module, keys, slots, refs in snapshot: + module._slot_keys.clear() + module._slot_keys.update(keys) + module._slot_modules = torch.nn.ModuleDict(slots) + for key, ref in refs.items(): + slots[key].ref = ref + + +def validate_lora_slot_replacement( + model: Sequence[torch.nn.Module], + source: LoRASlotRef, + destination: LoRASlotRef, +) -> None: + _lora_slot_replacements(model, source, destination) + + +def replace_lora_slot_in_model( + model: Sequence[torch.nn.Module], + source: LoRASlotRef, + destination: LoRASlotRef, +) -> None: + replacements = _lora_slot_replacements(model, source, destination) + for module, source_key, destination_key in replacements: + if source_key is None: + if destination_key is not None: + module._slot_keys.pop(destination) + del module._slot_modules[destination_key] + continue + slot = cast(LoRASlot, module._slot_modules[source_key]) + slot.ref = destination + module._slot_keys.pop(source) + module._slot_keys[destination] = source_key + if destination_key is not None and destination_key != source_key: + del module._slot_modules[destination_key] + + def iter_lora_slot_parameters( model: Sequence[torch.nn.Module], ref: LoRASlotRef, diff --git a/src/art/megatron/model_support/lora_disk.py b/src/art/megatron/model_support/lora_disk.py index 46741d50a..28de9e6aa 100644 --- a/src/art/megatron/model_support/lora_disk.py +++ b/src/art/megatron/model_support/lora_disk.py @@ -1,6 +1,7 @@ import importlib import json from pathlib import Path +import struct from typing import Any import torch @@ -8,6 +9,7 @@ from art.megatron.model_support.spec import ModelSupportHandler ART_LORA_FORMAT_CONFIG_KEY = "art_lora_format" +ART_LORA_FORMAT_MEGATRON = "megatron" ART_LORA_FORMAT_VLLM = "vllm" safetensors = importlib.import_module("safetensors") @@ -88,6 +90,68 @@ def save_vllm_lora_tensors( ) +def _consolidate_safetensors( + shards: list[Path], output: Path, *, chunk_size: int = 8 * 1024 * 1024 +) -> None: + """Join safetensors shards without materializing their tensor payloads.""" + sources: dict[str, tuple[Path, int, int, int, dict[str, object]]] = {} + for shard in shards: + with shard.open("rb") as handle: + encoded_length = handle.read(8) + if len(encoded_length) != 8: + raise RuntimeError(f"Invalid safetensors header: {shard}") + header_length = struct.unpack(" dict[str, torch.Tensor]: + adapter_config = load_adapter_config(lora_path) + tensors = load_vllm_lora_tensors(lora_path) + if adapter_config.get(ART_LORA_FORMAT_CONFIG_KEY) == ART_LORA_FORMAT_MEGATRON: + return tensors resolved_handler = resolve_lora_handler( lora_path, handler, allow_unvalidated_arch=allow_unvalidated_arch, ) return resolved_handler.from_vllm_lora_tensors( - load_vllm_lora_tensors(lora_path), - adapter_config=load_adapter_config(lora_path), + tensors, + adapter_config=adapter_config, ) diff --git a/src/art/megatron/provider.py b/src/art/megatron/provider.py index 52bc7ff7c..657f66742 100644 --- a/src/art/megatron/provider.py +++ b/src/art/megatron/provider.py @@ -18,7 +18,10 @@ get_model_support_handler_for_spec, get_model_support_spec, ) -from art.megatron.model_support.spec import ModelSupportSpec +from art.megatron.model_support.spec import ( + ModelSupportHandler, + ModelSupportSpec, +) from art.megatron.runtime.bridge_runtime import install_art_bridge_runtime_patches install_art_bridge_runtime_patches() @@ -77,9 +80,10 @@ class ProviderBundle(BaseModel): model_config = ConfigDict(arbitrary_types_allowed=True) - provider: Any - bridge: Any - handler: Any + model_identifier: str + provider: GPTModelProvider + bridge: AutoBridge + handler: ModelSupportHandler spec: ModelSupportSpec @@ -345,13 +349,13 @@ def _resolve_default_hybridep_num_sms() -> int: return 24 -def _handler_cp_supported(handler: Any) -> bool: +def _handler_cp_supported(handler: ModelSupportHandler) -> bool: return bool(getattr(handler, "cp_supported", True)) def _apply_default_parallel_topology( provider: GPTModelProvider, - handler: Any, + handler: ModelSupportHandler, ) -> None: visible_gpu_count = max(torch.cuda.device_count(), 1) cp_supported = _handler_cp_supported(handler) @@ -368,7 +372,7 @@ def _apply_default_parallel_topology( def _apply_art_training_runtime_prepare_defaults( provider: GPTModelProvider, - handler: Any, + handler: ModelSupportHandler, ) -> None: provider.recompute_granularity = "full" provider.recompute_method = "uniform" @@ -378,7 +382,7 @@ def _apply_art_training_runtime_prepare_defaults( def _validate_context_parallel_support( - handler: Any, + handler: ModelSupportHandler, runtime_env: _ProviderRuntimeEnv, ) -> None: if _handler_cp_supported(handler): @@ -625,6 +629,7 @@ def _build_provider_bundle( provider = bridge.to_megatron_provider() handler.patch_bridge(bridge) return ProviderBundle( + model_identifier=model, provider=provider, bridge=bridge, handler=handler, @@ -661,8 +666,10 @@ def prepare_provider_bundle( bundle.handler.configure_provider_for_runtime(provider) _validate_context_parallel_support(bundle.handler, runtime_env) _apply_runtime_env_overrides(provider, runtime_env) - provider.art_flex_compile_crash_config = ( - bundle.handler.flex_attention_compile_crash_config(provider) + setattr( + provider, + "art_flex_compile_crash_config", + bundle.handler.flex_attention_compile_crash_config(provider), ) provider.sequence_parallel = provider.tensor_model_parallel_size > 1 _install_art_training_flex_attention(provider) @@ -673,7 +680,7 @@ def prepare_provider_bundle( def finalize_provider_bundle(provider_bundle: ProviderBundle) -> ProviderBundle: runtime_env = _ProviderRuntimeEnv.from_environ() - provider = cast(GPTModelProvider, provider_bundle.provider) + provider = provider_bundle.provider _apply_art_training_runtime_finalize_defaults(provider) _enforce_art_moe_grouped_gemm_fast_path(provider) _finalize_provider_with_art_overrides(provider) diff --git a/src/art/megatron/train.py b/src/art/megatron/train.py index 284620115..5519a60a1 100644 --- a/src/art/megatron/train.py +++ b/src/art/megatron/train.py @@ -49,6 +49,10 @@ load_adapter_config, load_lora_tensors_for_megatron, ) +from art.megatron.model_support.spec import ( + ModelSupportHandler, + ModelSupportSpec, +) from art.megatron.optimizer_state import ( commit_optimizer_generation, optimizer_generation_files, @@ -180,16 +184,20 @@ def _validate_model(cls, value: ModelChunks) -> ModelChunks: validate_model_chunks(value) return value + @property + def model_identifier(self) -> str: + return self.provider_bundle.model_identifier + @property def bridge(self) -> Any: return self.provider_bundle.bridge @property - def model_support_handler(self) -> Any: + def model_support_handler(self) -> ModelSupportHandler: return self.provider_bundle.handler @property - def model_support_spec(self) -> Any: + def model_support_spec(self) -> ModelSupportSpec: return self.provider_bundle.spec @@ -1000,7 +1008,7 @@ def _load_adapter_into_model( lora_path: str, rank: int, *, - handler: Any | None = None, + handler: ModelSupportHandler | None = None, optimizer: Any | None = None, ) -> dict[str, torch.Tensor]: print0(rank, "Loading adapter model from", lora_path) @@ -1278,7 +1286,7 @@ def load_adapter_into_model( adapter_model: dict[str, torch.Tensor], optimizer: Any | None = None, *, - model_support_handler: Any | None = None, + model_support_handler: ModelSupportHandler | None = None, ) -> None: with torch.no_grad(): for chunk in model_chunks: @@ -1306,7 +1314,7 @@ def _optimizer_step( optimizer: Any, learning_rate: float, *, - model_support_handler: Any | None = None, + model_support_handler: ModelSupportHandler | None = None, model_chunks: ModelChunks | None = None, ) -> tuple[bool, float, int | None]: for param_group in optimizer.param_groups: @@ -1492,7 +1500,7 @@ def _select_next_ref_logprobs( def _forward_prepared_rl_micro( *, model_chunks: ModelChunks, - model_support_handler: Any, + model_support_handler: ModelSupportHandler, prepared_micro: PreparedRLMicroInputs, device: torch.device, ) -> torch.Tensor: @@ -1561,7 +1569,7 @@ def _calculate_megatron_logprobs( *, model_chunks: ModelChunks, provider: Any, - model_support_handler: Any, + model_support_handler: ModelSupportHandler, inputs: PackedTensors, moe_routing_replay_controller: MoeRoutingReplayController | None = None, step_index: int | None = None, @@ -1766,7 +1774,7 @@ def run_megatron_sft_step( *, model_chunks: ModelChunks, provider: Any, - model_support_handler: Any, + model_support_handler: ModelSupportHandler, optimizer: Any, learning_rate: float, inputs: dict[str, torch.Tensor] | list[dict[str, torch.Tensor]], @@ -1910,7 +1918,7 @@ def run_training_step( *, model_chunks: ModelChunks, provider: Any, - model_support_handler: Any, + model_support_handler: ModelSupportHandler, optimizer: Any, learning_rate: float, inputs: PackedTensors | list[PackedTensors], diff --git a/src/art/megatron/weights/lora_publish.py b/src/art/megatron/weights/lora_publish.py index e9d8b4b08..8cb2e76ac 100644 --- a/src/art/megatron/weights/lora_publish.py +++ b/src/art/megatron/weights/lora_publish.py @@ -1,21 +1,52 @@ from collections.abc import Iterable, Sequence +import hashlib +from pathlib import Path from typing import Any, NamedTuple +from safetensors.torch import save_file import torch +from art.megatron._collective import ( + collective_errors as _collective_errors, +) +from art.megatron._collective import ( + device as _device, +) +from art.megatron._collective import ( + distributed as _distributed_ready, +) +from art.megatron._collective import ( + dtype_from_name as _dtype_from_name, +) +from art.megatron._collective import ( + gather_objects as _gather_objects, +) +from art.megatron._collective import ( + raise_distributed as _raise_rank_errors, +) +from art.megatron._collective import ( + rank as _rank, +) from art.megatron.lora import ( LoRA, - LoRAPublishPlanner, + LoraShardManifest, LoraShardMeta, LoRASlotRef, _block_for_key, _dtype_name, ) -from art.megatron.lora import ( - _distributed_initialized as _distributed_ready, +from art.megatron.model_support.lora_disk import ( + ART_LORA_FORMAT_CONFIG_KEY, + ART_LORA_FORMAT_MEGATRON, + ART_LORA_FORMAT_VLLM, + _consolidate_safetensors, + save_adapter_config, +) +from art.megatron.model_support.spec import ( + ExpertPackedLoraGroup, + ExpertPackedLoraSlot, + ModelSupportHandler, ) -from art.megatron.model_support.lora_disk import save_vllm_lora_tensors -from art.megatron.model_support.spec import ExpertPackedLoraGroup, ExpertPackedLoraSlot from art.megatron.training.model_chunks import ModelChunks @@ -24,7 +55,7 @@ class PackedExpertShardMeta(NamedTuple): owner_rank: int shape: tuple[int, ...] dtype_name: str - manifest: dict[str, Any] + manifest: LoraShardManifest expert_start: int expert_count: int pack_layout: str @@ -37,34 +68,6 @@ def numel(self) -> int: return total -class _PinnedCpuStager: - def __init__(self) -> None: - self._events: list[torch.cuda.Event] = [] - self._stream = torch.cuda.Stream() if torch.cuda.is_available() else None - - def stage(self, tensor: torch.Tensor) -> torch.Tensor: - source = tensor.detach() - if self._stream is None or not source.is_cuda: - return source.cpu() - - source = source.contiguous() - target = torch.empty_like(source, device="cpu", pin_memory=True) - source_stream = torch.cuda.current_stream(source.device) - self._stream.wait_stream(source_stream) - with torch.cuda.stream(self._stream): - target.copy_(source, non_blocking=True) - source.record_stream(self._stream) - event = torch.cuda.Event() - event.record(self._stream) - self._events.append(event) - return target - - def finish(self) -> None: - for event in self._events: - event.synchronize() - self._events.clear() - - def iter_lora_modules(model_chunks: ModelChunks) -> Iterable[LoRA]: for chunk in model_chunks: for module in chunk.modules(): @@ -72,12 +75,6 @@ def iter_lora_modules(model_chunks: ModelChunks) -> Iterable[LoRA]: yield module -def _dtype_from_name(name: str) -> torch.dtype: - if isinstance(dtype := getattr(torch, name, None), torch.dtype): - return dtype - raise RuntimeError(f"Unsupported LoRA tensor dtype={name!r}") - - def _packed_expert_slot( adapter_model_prefix: str, suffix: str, @@ -119,13 +116,13 @@ def collect_local_lora_entries( slot_ref: LoRASlotRef | None = None, ) -> tuple[dict[str, torch.Tensor], list[LoraShardMeta]]: local_tensors: dict[str, torch.Tensor] = {} - local_manifest: dict[str, dict[str, Any]] = {} + local_manifest: dict[str, LoraShardManifest] = {} for module in iter_lora_modules(model_chunks): if _uses_packed_expert_publish(module, packed_expert_groups, slot_ref): continue for key, value in module.sharded_lora_state_dict(slot_ref).items(): target_dtype = adapter_dtypes[key] if key in adapter_dtypes else value.dtype - local_tensors[key] = value.to(target_dtype).contiguous() + local_tensors[key] = value.to(target_dtype) local_manifest.update(module.sharded_lora_manifest(slot_ref)) if set(local_tensors) != set(local_manifest): @@ -173,14 +170,14 @@ def collect_local_packed_expert_entries( continue group_prefix, slot = slot_match key = f"{group_prefix}.{slot.output_suffix}" - tensor = param.data.transpose(1, 2).contiguous() + tensor = param.data.transpose(1, 2) source_keys = module._expected_weight_keys(suffix.removesuffix(".weight")) target_dtype = ( adapter_dtypes[source_keys[0]] if source_keys and source_keys[0] in adapter_dtypes else tensor.dtype ) - tensor = tensor.to(target_dtype).contiguous() + tensor = tensor.to(target_dtype) if key in local_tensors: raise RuntimeError(f"Duplicate packed expert LoRA tensor: {key}") local_tensors[key] = tensor @@ -199,96 +196,11 @@ def collect_local_packed_expert_entries( return local_tensors, metadata -def _global_packed_expert_metadata( - planner: LoRAPublishPlanner, - adapter_dtypes: dict[str, torch.dtype], - packed_expert_groups: Sequence[ExpertPackedLoraGroup], -) -> list[PackedExpertShardMeta]: - metadata: list[PackedExpertShardMeta] = [] - for template in planner.templates: - if int(template.num_local_experts) <= 1: - continue - slot_match = _packed_expert_slot( - template.adapter_model_prefix, - template.suffix, - packed_expert_groups, - ) - if slot_match is None: - continue - group_prefix, slot = slot_match - shard_ranks = range(template.shard_world_size) if template.sharded else (0,) - ep_world_size = 1 - if _distributed_ready(): - from megatron.core import parallel_state as ps - - ep_world_size = ps.get_expert_model_parallel_world_size() - for ep_rank in range(ep_world_size): - expert_start = ep_rank * template.num_local_experts - expert_key = ( - f"{template.adapter_model_prefix.format(expert=expert_start)}." - f"{template.suffix}" - ) - for shard_rank in shard_ranks: - owner_rank = planner._expert_owner_rank(ep_rank, shard_rank) - per_expert_meta = planner._make_metadata( - template, - key=expert_key, - owner_rank=owner_rank, - shard_rank=shard_rank, - adapter_dtypes=adapter_dtypes, - ) - metadata.append( - PackedExpertShardMeta( - key=f"{group_prefix}.{slot.output_suffix}", - owner_rank=owner_rank, - shape=(template.num_local_experts, *per_expert_meta.shape), - dtype_name=per_expert_meta.dtype_name, - manifest=per_expert_meta.manifest, - expert_start=expert_start, - expert_count=template.num_local_experts, - pack_layout=slot.pack_layout, - ) - ) - return metadata - - -def _global_regular_metadata( - planner: LoRAPublishPlanner, - adapter_dtypes: dict[str, torch.dtype], - packed_expert_groups: Sequence[ExpertPackedLoraGroup], -) -> list[LoraShardMeta]: - if not packed_expert_groups: - return planner.global_metadata(adapter_dtypes) - if _distributed_ready(): - from megatron.core import parallel_state as ps - - pp_world_size = ps.get_pipeline_model_parallel_world_size() - if pp_world_size != 1: - raise RuntimeError( - "LoRA publish planner requires pipeline_model_parallel_size=1; " - f"got {pp_world_size}. Rank-local modules cannot describe remote " - "pipeline stages without exchanging templates." - ) - metadata: list[LoraShardMeta] = [] - for template in planner.templates: - if ( - _packed_expert_slot( - template.adapter_model_prefix, - template.suffix, - packed_expert_groups, - ) - is not None - ): - continue - metadata.extend(planner._metadata_for_template(template, adapter_dtypes)) - return metadata - - def _merge_sharded_tensor( key: str, *, ordered_shards: Sequence[torch.Tensor], - manifest: dict[str, Any], + manifest: LoraShardManifest, ) -> torch.Tensor: strategy = manifest.get("export_shard_strategy") assert strategy is not None @@ -322,9 +234,9 @@ def _merge_sharded_tensor( def _merge_manifest_entries( key: str, - key_entries: Sequence[tuple[dict[str, Any], torch.Tensor]], + key_entries: Sequence[tuple[LoraShardManifest, torch.Tensor]], *, - manifest: dict[str, Any] | None = None, + manifest: LoraShardManifest | None = None, ) -> torch.Tensor: first_manifest = key_entries[0][0] sharded = bool(first_manifest["sharded"]) @@ -365,7 +277,7 @@ def _merge_manifest_entries( def merge_sharded_adapter_entries( - entries_by_key: dict[str, list[tuple[dict[str, Any], torch.Tensor]]], + entries_by_key: dict[str, list[tuple[LoraShardManifest, torch.Tensor]]], ) -> dict[str, torch.Tensor]: return { key: _merge_manifest_entries(key, key_entries) @@ -373,102 +285,190 @@ def merge_sharded_adapter_entries( } -def _rank_and_device() -> tuple[int, torch.device]: - return ( - torch.distributed.get_rank() if _distributed_ready() else 0, # type: ignore[possibly-missing-attribute] - torch.device("cuda", torch.cuda.current_device()) - if torch.cuda.is_available() - else torch.device("cpu"), - ) +def _gather_metadata[T]( + local: list[T], *, group: torch.distributed.ProcessGroup | None = None +) -> list[T]: + return [item for values in _gather_objects(local, group=group) for item in values] -def _metadata_by_owner_dtype( - metadata: Sequence[Any], -) -> dict[tuple[int, str], list[Any]]: - grouped: dict[tuple[int, str], list[Any]] = {} - for meta in metadata: - grouped.setdefault((meta.owner_rank, meta.dtype_name), []).append(meta) - return { - key: sorted(group, key=lambda meta: meta.key) - for key, group in sorted(grouped.items()) - } +def _tensor_digest(tensor: torch.Tensor) -> str: + """Hash a tensor with bounded host staging and a bounded collective payload.""" + digest = hashlib.blake2b(digest_size=16) + digest.update(_dtype_name(tensor.dtype).encode()) + digest.update(repr(tuple(tensor.shape)).encode()) + data = tensor.detach().contiguous().reshape(-1).view(torch.uint8) + chunk_bytes = 1024 * 1024 + for offset in range(0, data.numel(), chunk_bytes): + chunk = data.narrow(0, offset, min(chunk_bytes, data.numel() - offset)) + digest.update(chunk.cpu().numpy().tobytes()) + return digest.hexdigest() -def _pack_metadata_tensors( - metadata: Sequence[Any], - tensors: dict[str, torch.Tensor], -) -> torch.Tensor: - return torch.cat( - [tensors[meta.key].detach().contiguous().view(-1) for meta in metadata] - ) +def _contributor_identity( + metadata: LoraShardMeta | PackedExpertShardMeta, +) -> tuple[object, ...]: + shard_rank = int(metadata.manifest.get("shard_rank", 0)) + if isinstance(metadata, PackedExpertShardMeta): + return ("packed", metadata.key, metadata.expert_start, shard_rank) + return ("lora", metadata.key, shard_rank) + + +def _validate_replica_digests( + records: Iterable[tuple[LoraShardMeta | PackedExpertShardMeta, str]], +) -> None: + digests: dict[tuple[object, ...], str] = {} + for metadata, digest in records: + identity = _contributor_identity(metadata) + previous = digests.setdefault(identity, digest) + if previous != digest: + raise RuntimeError( + f"Inconsistent replicated tensor contents for {metadata.key!r}" + ) + + +def _elect_lora_contributors( + metadata: list[LoraShardMeta], +) -> list[LoraShardMeta]: + elected: dict[tuple[str, int], LoraShardMeta] = {} + for candidate in metadata: + identity = (candidate.key, int(candidate.manifest.get("shard_rank", 0))) + current = elected.get(identity) + if current is not None and ( + current.shape != candidate.shape + or current.dtype_name != candidate.dtype_name + or current.manifest != candidate.manifest + or current.block != candidate.block + ): + raise RuntimeError( + f"Inconsistent replicated LoRA metadata for {candidate.key!r}" + ) + if current is None or candidate.owner_rank < current.owner_rank: + elected[identity] = candidate + return list(elected.values()) + + +def _elect_packed_expert_contributors( + metadata: list[PackedExpertShardMeta], +) -> list[PackedExpertShardMeta]: + elected: dict[tuple[str, int, int], PackedExpertShardMeta] = {} + for candidate in metadata: + identity = ( + candidate.key, + candidate.expert_start, + int(candidate.manifest.get("shard_rank", 0)), + ) + current = elected.get(identity) + if ( + current is not None + and current._replace(owner_rank=candidate.owner_rank) != candidate + ): + raise RuntimeError( + f"Inconsistent replicated packed LoRA metadata for {candidate.key!r}" + ) + if current is None or candidate.owner_rank < current.owner_rank: + elected[identity] = candidate + return list(elected.values()) + + +def _rank_and_device() -> tuple[int, torch.device]: + return _rank(), _device() -def _views_from_flat( +def _prepare_exchange_buffers( + metadata: Sequence[LoraShardMeta | PackedExpertShardMeta], *, - owner_rank: int, - metadata: Sequence[Any], - flat: torch.Tensor, -) -> dict[tuple[int, str], torch.Tensor]: - views: dict[tuple[int, str], torch.Tensor] = {} - offset = 0 + local_tensors: dict[str, torch.Tensor], + rank: int, + device: torch.device, +) -> tuple[ + dict[tuple[int, str], torch.Tensor], + dict[tuple[int, str], torch.Tensor], + dict[tuple[int, str], torch.Tensor], +]: + sends: dict[tuple[int, str], torch.Tensor] = {} + receives: dict[tuple[int, str], torch.Tensor] = {} + received: dict[tuple[int, str], torch.Tensor] = {} for meta in metadata: - views[(owner_rank, meta.key)] = flat.narrow(0, offset, meta.numel).view( - meta.shape - ) - offset += meta.numel - return views + identity = (meta.owner_rank, meta.key) + if rank == meta.owner_rank: + tensor = local_tensors[meta.key].detach().contiguous() + if rank == 0: + received[identity] = tensor.cpu().contiguous() + else: + sends[identity] = tensor + elif rank == 0: + receives[identity] = torch.empty( + meta.shape, + dtype=_dtype_from_name(meta.dtype_name), + device=device, + ) + return sends, receives, received -def _exchange_batched_tensors( - metadata: Sequence[Any], +def _exchange_tensors( + metadata: Sequence[LoraShardMeta | PackedExpertShardMeta], *, local_tensors: dict[str, torch.Tensor], rank: int, device: torch.device, + group: torch.distributed.ProcessGroup | None = None, ) -> dict[tuple[int, str], torch.Tensor]: + ordered = sorted( + metadata, + key=lambda meta: ( + meta.owner_rank, + meta.key, + int(meta.manifest.get("shard_rank", 0)), + ), + ) if not _distributed_ready(): return { - (rank, meta.key): local_tensors[meta.key].contiguous() for meta in metadata + (rank, meta.key): local_tensors[meta.key].detach().cpu().contiguous() + for meta in ordered } + sends: dict[tuple[int, str], torch.Tensor] = {} + receives: dict[tuple[int, str], torch.Tensor] = {} received: dict[tuple[int, str], torch.Tensor] = {} - for (owner_rank, dtype_name), group_metadata in _metadata_by_owner_dtype( - metadata - ).items(): - if rank == owner_rank: - flat = _pack_metadata_tensors(group_metadata, local_tensors) - if rank == 0: - received.update( - _views_from_flat( - owner_rank=owner_rank, - metadata=group_metadata, - flat=flat, - ) - ) - else: - torch.distributed.send(flat, dst=0) # type: ignore[possibly-missing-attribute] - elif rank == 0: - flat = torch.empty( - sum(meta.numel for meta in group_metadata), - dtype=_dtype_from_name(dtype_name), - device=device, + error: BaseException | None = None + try: + sends, receives, received = _prepare_exchange_buffers( + ordered, + local_tensors=local_tensors, + rank=rank, + device=device, + ) + except BaseException as exc: + error = exc + _raise_rank_errors(error, "prepare tensor exchange", group=group) + + for meta in ordered: + identity = (meta.owner_rank, meta.key) + if rank == meta.owner_rank and rank != 0: + torch.distributed.send(sends[identity], dst=0, group=group) # type: ignore[possibly-missing-attribute] + elif rank == 0 and meta.owner_rank != 0: + torch.distributed.recv( # type: ignore[possibly-missing-attribute] + receives[identity], src=meta.owner_rank, group=group ) - torch.distributed.recv(flat, src=owner_rank) # type: ignore[possibly-missing-attribute] + + error = None + try: + if rank == 0: received.update( - _views_from_flat( - owner_rank=owner_rank, - metadata=group_metadata, - flat=flat, - ) + (identity, tensor.cpu().contiguous()) + for identity, tensor in receives.items() ) + except BaseException as exc: + error = exc + _raise_rank_errors(error, "finalize tensor exchange", group=group) return received def _entries_by_key( metadata: list[LoraShardMeta], tensors_by_owner_key: dict[tuple[int, str], torch.Tensor], -) -> dict[str, list[tuple[dict[str, Any], torch.Tensor]]]: - entries: dict[str, list[tuple[dict[str, Any], torch.Tensor]]] = {} +) -> dict[str, list[tuple[LoraShardManifest, torch.Tensor]]]: + entries: dict[str, list[tuple[LoraShardManifest, torch.Tensor]]] = {} for meta in metadata: entries.setdefault(meta.key, []).append( (meta.manifest, tensors_by_owner_key[(meta.owner_rank, meta.key)]) @@ -476,13 +476,41 @@ def _entries_by_key( return entries +def _gather_merged_adapter_tensors( + metadata: Sequence[LoraShardMeta], + *, + local_tensors: dict[str, torch.Tensor], + rank: int, + device: torch.device, + group: torch.distributed.ProcessGroup | None = None, +) -> dict[str, torch.Tensor]: + exchanged = _exchange_tensors( + metadata, + local_tensors=local_tensors, + rank=rank, + device=device, + group=group, + ) + merged: dict[str, torch.Tensor] = {} + error: BaseException | None = None + try: + if rank == 0: + merged = merge_sharded_adapter_entries( + _entries_by_key(list(metadata), exchanged) + ) + except BaseException as exc: + error = exc + _raise_rank_errors(error, "merge adapter tensors", group=group) + return merged + + def _merge_packed_expert_block( key: str, - key_entries: list[tuple[dict[str, Any], torch.Tensor]], + key_entries: list[tuple[LoraShardManifest, torch.Tensor]], ) -> torch.Tensor: - manifest = dict(key_entries[0][0]) - if bool(manifest["sharded"]): - manifest["export_shard_dim"] = int(manifest["export_shard_dim"]) + 1 + manifest: LoraShardManifest = {**key_entries[0][0]} + if manifest["sharded"]: + manifest["export_shard_dim"] = manifest.get("export_shard_dim", 0) + 1 return _merge_manifest_entries(key, key_entries, manifest=manifest) @@ -556,7 +584,7 @@ def merge_packed_expert_adapter_entries( ) -> dict[str, torch.Tensor]: entries_by_key_start: dict[ tuple[str, int], - list[tuple[PackedExpertShardMeta, dict[str, Any], torch.Tensor]], + list[tuple[PackedExpertShardMeta, LoraShardManifest, torch.Tensor]], ] = {} for meta in metadata: entries_by_key_start.setdefault((meta.key, meta.expert_start), []).append( @@ -582,60 +610,6 @@ def merge_packed_expert_adapter_entries( } -def _stage_published_tensors( - tensors: dict[str, torch.Tensor], - stager: _PinnedCpuStager, -) -> dict[str, torch.Tensor]: - grouped: dict[tuple[str, int | None, str], list[tuple[str, torch.Tensor]]] = {} - for key, tensor in tensors.items(): - dtype_name = _dtype_name(tensor.dtype) - group_key = (tensor.device.type, tensor.device.index, dtype_name) - grouped.setdefault(group_key, []).append((key, tensor)) - - staged: dict[str, torch.Tensor] = {} - for _group_key, group in sorted(grouped.items()): - flat = torch.cat( - [tensor.detach().contiguous().view(-1) for _key, tensor in sorted(group)] - ) - staged_flat = stager.stage(flat) - offset = 0 - for key, tensor in sorted(group): - numel = tensor.numel() - if key in staged: - raise RuntimeError( - f"Duplicate vLLM LoRA tensor after conversion: {key}" - ) - staged[key] = staged_flat.narrow(0, offset, numel).view(tensor.shape) - offset += numel - return staged - - -def _save_rank0_vllm_lora( - *, - metadata: list[LoraShardMeta], - tensors_by_owner_key: dict[tuple[int, str], torch.Tensor], - packed_expert_metadata: list[PackedExpertShardMeta] | None = None, - packed_expert_tensors_by_owner_key: ( - dict[tuple[int, str], torch.Tensor] | None - ) = None, - handler: Any, - adapter_config: dict[str, Any], - output_dir: str, -) -> None: - vllm_tensors, published_config = _rank0_vllm_lora_tensors( - metadata=metadata, - tensors_by_owner_key=tensors_by_owner_key, - packed_expert_metadata=packed_expert_metadata, - packed_expert_tensors_by_owner_key=packed_expert_tensors_by_owner_key, - handler=handler, - adapter_config=adapter_config, - ) - stager = _PinnedCpuStager() - published_tensors = _stage_published_tensors(vllm_tensors, stager) - stager.finish() - save_vllm_lora_tensors(output_dir, published_tensors, published_config) - - def _rank0_vllm_lora_tensors( *, metadata: list[LoraShardMeta], @@ -644,7 +618,7 @@ def _rank0_vllm_lora_tensors( packed_expert_tensors_by_owner_key: ( dict[tuple[int, str], torch.Tensor] | None ) = None, - handler: Any, + handler: ModelSupportHandler, adapter_config: dict[str, Any], ) -> tuple[dict[str, torch.Tensor], dict[str, Any]]: merged_tensors = merge_sharded_adapter_entries( @@ -667,80 +641,144 @@ def _rank0_vllm_lora_tensors( ) -def build_vllm_lora_tensors_from_model( +class _LocalLoraExport(NamedTuple): + rank: int + device: torch.device + tensors: dict[str, torch.Tensor] + metadata: list[LoraShardMeta] + packed_tensors: dict[str, torch.Tensor] + packed_metadata: list[PackedExpertShardMeta] + blocks: tuple[str, ...] + + +def _prepare_local_lora_export( *, model: ModelChunks, adapter_dtypes: dict[str, torch.dtype], - handler: Any, - adapter_config: dict[str, Any], + handler: ModelSupportHandler, rank: int, world_size: int, slot_ref: LoRASlotRef | None = None, -) -> tuple[dict[str, torch.Tensor], dict[str, Any]] | None: + native_format: bool = False, +) -> _LocalLoraExport: actual_rank, device = _rank_and_device() - if _distributed_ready(): - actual_world_size = torch.distributed.get_world_size() # type: ignore[possibly-missing-attribute] - if actual_rank != rank or actual_world_size != world_size: - raise RuntimeError( - "LoRA publisher rank/world-size mismatch: " - f"runtime=({rank}, {world_size}) distributed=({actual_rank}, {actual_world_size})" - ) - else: - if rank != 0 or world_size != 1: + local_tensors: dict[str, torch.Tensor] = {} + local_metadata: list[LoraShardMeta] = [] + local_digests: dict[str, str] = {} + local_packed_tensors: dict[str, torch.Tensor] = {} + local_packed_metadata: list[PackedExpertShardMeta] = [] + local_packed_digests: dict[str, str] = {} + with _collective_errors("prepare LoRA export"): + if _distributed_ready(): + actual_world_size = torch.distributed.get_world_size() # type: ignore[possibly-missing-attribute] + if actual_rank != rank or actual_world_size != world_size: + raise RuntimeError( + "LoRA publisher rank/world-size mismatch: " + f"runtime=({rank}, {world_size}) " + f"distributed=({actual_rank}, {actual_world_size})" + ) + elif rank != 0 or world_size != 1: raise RuntimeError( "Non-distributed LoRA publish requires rank=0 and world_size=1, " f"got rank={rank} world_size={world_size}" ) - rank = 0 - packed_expert_groups = tuple(handler.expert_packed_lora_groups()) - planner = LoRAPublishPlanner(model, slot_ref) - local_tensors, local_metadata = collect_local_lora_entries( - model, - adapter_dtypes, - owner_rank=rank, - packed_expert_groups=packed_expert_groups, - slot_ref=slot_ref, + packed_expert_groups = ( + () if native_format else tuple(handler.expert_packed_lora_groups()) + ) + local_tensors, local_metadata = collect_local_lora_entries( + model, + adapter_dtypes, + owner_rank=rank, + packed_expert_groups=packed_expert_groups, + slot_ref=slot_ref, + ) + (local_packed_tensors, local_packed_metadata) = ( + collect_local_packed_expert_entries( + model, + adapter_dtypes, + owner_rank=rank, + packed_expert_groups=packed_expert_groups, + slot_ref=slot_ref, + ) + ) + local_digests = { + key: _tensor_digest(tensor) for key, tensor in local_tensors.items() + } + local_packed_digests = { + key: _tensor_digest(tensor) for key, tensor in local_packed_tensors.items() + } + gathered_metadata = _gather_metadata( + [(metadata, local_digests[metadata.key]) for metadata in local_metadata] ) - local_packed_tensors, local_packed_metadata = collect_local_packed_expert_entries( - model, - adapter_dtypes, - owner_rank=rank, - packed_expert_groups=packed_expert_groups, - slot_ref=slot_ref, + gathered_packed_metadata = _gather_metadata( + [ + (metadata, local_packed_digests[metadata.key]) + for metadata in local_packed_metadata + ] ) - all_packed_metadata = ( - _global_packed_expert_metadata(planner, adapter_dtypes, packed_expert_groups) - if rank == 0 - else local_packed_metadata + _validate_replica_digests(gathered_metadata) + _validate_replica_digests(gathered_packed_metadata) + metadata = _elect_lora_contributors( + [metadata for metadata, _digest in gathered_metadata] ) - if rank == 0: - all_metadata = _global_regular_metadata( - planner, - adapter_dtypes, - packed_expert_groups if all_packed_metadata else (), + packed_metadata = _elect_packed_expert_contributors( + [metadata for metadata, _digest in gathered_packed_metadata] + ) + blocks = tuple( + sorted( + {item.block for item in metadata} + | {_block_for_key(item.key) for item in packed_metadata} ) - else: - all_metadata = local_metadata - exchanged_tensors = _exchange_batched_tensors( - all_metadata, - local_tensors=local_tensors, + ) + return _LocalLoraExport( + rank, + device, + local_tensors, + metadata, + local_packed_tensors, + packed_metadata, + blocks, + ) + + +def build_vllm_lora_tensors_from_model( + *, + model: ModelChunks, + adapter_dtypes: dict[str, torch.dtype], + handler: ModelSupportHandler, + adapter_config: dict[str, Any], + rank: int, + world_size: int, + slot_ref: LoRASlotRef | None = None, +) -> tuple[dict[str, torch.Tensor], dict[str, Any]] | None: + local = _prepare_local_lora_export( + model=model, + adapter_dtypes=adapter_dtypes, + handler=handler, rank=rank, - device=device, + world_size=world_size, + slot_ref=slot_ref, ) - exchanged_packed_tensors = _exchange_batched_tensors( - all_packed_metadata, - local_tensors=local_packed_tensors, + exchanged_tensors = _exchange_tensors( + local.metadata, + local_tensors=local.tensors, rank=rank, - device=device, + device=local.device, + ) + exchanged_packed_tensors = _exchange_tensors( + local.packed_metadata, + local_tensors=local.packed_tensors, + rank=rank, + device=local.device, ) if rank != 0: return None return _rank0_vllm_lora_tensors( - metadata=all_metadata, + metadata=local.metadata, tensors_by_owner_key=exchanged_tensors, - packed_expert_metadata=all_packed_metadata, + packed_expert_metadata=local.packed_metadata, packed_expert_tensors_by_owner_key=exchanged_packed_tensors, handler=handler, adapter_config=adapter_config, @@ -751,26 +789,110 @@ def save_vllm_lora_from_model( *, model: ModelChunks, adapter_dtypes: dict[str, torch.dtype], - handler: Any, + handler: ModelSupportHandler, adapter_config: dict[str, Any], output_dir: str, rank: int, world_size: int, slot_ref: LoRASlotRef | None = None, -) -> None: - result = build_vllm_lora_tensors_from_model( + native_format: bool = False, +) -> _LocalLoraExport: + local = _prepare_local_lora_export( model=model, adapter_dtypes=adapter_dtypes, handler=handler, - adapter_config=adapter_config, rank=rank, world_size=world_size, slot_ref=slot_ref, + native_format=native_format, ) - if result is None: - return - vllm_tensors, published_config = result - stager = _PinnedCpuStager() - published_tensors = _stage_published_tensors(vllm_tensors, stager) - stager.finish() - save_vllm_lora_tensors(output_dir, published_tensors, published_config) + blocks = local.blocks + root = Path(output_dir) + shards: list[Path] = [] + published_config = dict(adapter_config) + written: set[str] = set() + with _collective_errors("prepare LoRA output"): + if rank == 0: + root.mkdir(parents=True, exist_ok=True) + + try: + for index, block in enumerate(blocks): + metadata = [meta for meta in local.metadata if meta.block == block] + packed_metadata = [ + meta + for meta in local.packed_metadata + if _block_for_key(meta.key) == block + ] + exchanged = _exchange_tensors( + metadata, + local_tensors=local.tensors, + rank=rank, + device=local.device, + ) + exchanged_packed = _exchange_tensors( + packed_metadata, + local_tensors=local.packed_tensors, + rank=rank, + device=local.device, + ) + with _collective_errors(f"serialize LoRA block {block}"): + if rank == 0: + if native_format: + tensors = merge_sharded_adapter_entries( + _entries_by_key(metadata, exchanged) + ) + block_config = adapter_config + else: + tensors, block_config = _rank0_vllm_lora_tensors( + metadata=metadata, + tensors_by_owner_key=exchanged, + packed_expert_metadata=packed_metadata, + packed_expert_tensors_by_owner_key=exchanged_packed, + handler=handler, + adapter_config=adapter_config, + ) + if duplicates := written & tensors.keys(): + raise RuntimeError( + "Duplicate LoRA tensors across model blocks: " + f"{sorted(duplicates)}" + ) + written.update(tensors) + if block_config != adapter_config: + if ( + published_config != adapter_config + and published_config != block_config + ): + raise RuntimeError( + "Model blocks produced inconsistent LoRA configs" + ) + published_config = block_config + shard = root / f".adapter_model-{index:06d}.safetensors" + shards.append(shard) + save_file( + { + key: tensor.detach().cpu().contiguous() + for key, tensor in tensors.items() + }, + shard, + ) + with _collective_errors("finalize LoRA export"): + if rank == 0: + if not shards: + raise RuntimeError("No LoRA tensors were available to export") + _consolidate_safetensors(shards, root / "adapter_model.safetensors") + save_adapter_config( + root, + { + **published_config, + ART_LORA_FORMAT_CONFIG_KEY: ( + ART_LORA_FORMAT_MEGATRON + if native_format + else ART_LORA_FORMAT_VLLM + ), + }, + ) + finally: + if rank == 0: + for shard in shards: + shard.unlink(missing_ok=True) + return local diff --git a/src/art/trainer_rank/__init__.py b/src/art/trainer_rank/__init__.py index 15c413f56..b506532e5 100644 --- a/src/art/trainer_rank/__init__.py +++ b/src/art/trainer_rank/__init__.py @@ -1,27 +1,12 @@ from __future__ import annotations -from collections.abc import Callable, Iterable, Iterator, Mapping, Sequence -from typing import TYPE_CHECKING, Literal, TypedDict, cast, overload +import asyncio +from collections.abc import Callable, Iterable, Iterator, Sequence +from typing import TYPE_CHECKING, Literal, cast, overload import torch import torch.distributed as dist - -class TrainerRankOptimizerLayout(TypedDict): - parallel: tuple[int, int, int, int, int, int, int, int] - parameters: tuple[ - tuple[tuple[int, ...], str, str, bool, int | None, str, tuple[int, ...]], - ..., - ] - - -class TrainerRankOptimizerState(TypedDict): - format_version: Literal[1] - layout: TrainerRankOptimizerLayout - master_params: tuple[torch.Tensor, ...] - optimizer: dict[str, object] - - from . import _impl # noqa: E402 AdapterSelection = _impl.AdapterSelection @@ -42,7 +27,7 @@ class TrainerRankOptimizerState(TypedDict): TrainerRankMemoryError = _impl.TrainerRankMemoryError TrainerRankSlotStateError = _impl.TrainerRankSlotStateError Unset = _impl.Unset -_PushedSlot = _impl._PushedSlot +PushedCheckpoint = _impl.PushedCheckpoint if TYPE_CHECKING: from art.megatron.train import TrainingRuntime @@ -56,8 +41,7 @@ class TrainerRankOptimizerState(TypedDict): MicroBatchStats, TopK, TrainerRankMemoryError, - TrainerRankOptimizerLayout, - TrainerRankOptimizerState, + PushedCheckpoint, TrainerRankSlotStateError, ): _public_type.__module__ = __name__ @@ -85,55 +69,31 @@ def __init__( def zero_grad(self) -> None: super().zero_grad() - def set_checkpoint(self, name: str | None) -> None: - super().set_checkpoint(name) - - def set_lora(self, name: str | None) -> None: - super().set_lora(name) + def prefetch_checkpoints(self, *paths: str) -> asyncio.Task[None]: + return super().prefetch_checkpoints(*paths) - def push_checkpoint(self, name: str | None) -> _PushedSlot: - return super().push_checkpoint(name) + def load_checkpoint(self, path: str | None) -> asyncio.Task[None]: + return super().load_checkpoint(path) - def push_lora(self, name: str | None) -> _PushedSlot: - return super().push_lora(name) + def push_checkpoint(self, path: str | None) -> PushedCheckpoint: + return super().push_checkpoint(path) - def pop_pushed_lora_or_checkpoint(self) -> None: - super().pop_pushed_lora_or_checkpoint() + def pop_checkpoint(self) -> None: + super().pop_checkpoint() - def load_checkpoint_slot( + def save_checkpoint( self, - name: str, - adapter_model: Mapping[str, torch.Tensor], - *, - optimizer_state: TrainerRankOptimizerState | None = None, - alpha: float | None = None, - adapter_config: Mapping[str, object] | None = None, - ) -> int: - return super().load_checkpoint_slot( - name, - adapter_model, - optimizer_state=optimizer_state, - alpha=alpha, - adapter_config=adapter_config, - ) - - def checkpoint_slot_optimizer_state( - self, name: str - ) -> TrainerRankOptimizerState | None: - return super().checkpoint_slot_optimizer_state(name) - - def save_checkpoint_slot_lora(self, name: str, output_dir: str) -> None: - """Collectively publish a trained checkpoint slot as a vLLM LoRA.""" - super().save_checkpoint_slot_lora(name, output_dir) + output_dir: str, + checkpoint_path: str | Literal["active"] = "active", + ) -> None: + super().save_checkpoint(output_dir, checkpoint_path) - def load_lora_slot( + def export_lora( self, - name: str, - adapter_model: Mapping[str, torch.Tensor], - *, - alpha: float | None = None, + output_dir: str, + checkpoint_path: str | Literal["active"] = "active", ) -> int: - return super().load_lora_slot(name, adapter_model, alpha=alpha) + return super().export_lora(output_dir, checkpoint_path) @overload def forward_micro_batches( @@ -290,8 +250,7 @@ def optim_step( "TopK", "TrainerRank", "TrainerRankMemoryError", - "TrainerRankOptimizerLayout", - "TrainerRankOptimizerState", + "PushedCheckpoint", "TrainerRankSlotStateError", "Unset", ] diff --git a/src/art/trainer_rank/_checkpoint.py b/src/art/trainer_rank/_checkpoint.py new file mode 100644 index 000000000..9d46e26b8 --- /dev/null +++ b/src/art/trainer_rank/_checkpoint.py @@ -0,0 +1,1685 @@ +"""Topology-independent dynamic LoRA checkpoint persistence.""" + +from __future__ import annotations + +from collections.abc import Iterable, Mapping, Sequence +from copy import deepcopy +from dataclasses import dataclass +import hashlib +import importlib +import json +import os +from pathlib import Path, PurePosixPath, PureWindowsPath +import re +import shutil +import threading +from typing import ( + TYPE_CHECKING, + Annotated, + Literal, + Protocol, + cast, +) +import uuid + +from pydantic import BaseModel, ConfigDict, Field +import torch +import torch.distributed as dist + +from art.megatron._collective import ( + collective_errors as _collective_errors, +) +from art.megatron._collective import ( + distributed as _distributed, +) +from art.megatron._collective import ( + dtype_name as _dtype_name, +) +from art.megatron._collective import ( + gather_objects as _gather_objects, +) +from art.megatron._collective import ( + raise_distributed as _raise_distributed, +) +from art.megatron._collective import ( + rank as _rank, +) + +if TYPE_CHECKING: + from art.megatron.lora import LoraShardManifest, LoraShardMeta + from art.megatron.weights.lora_publish import _LocalLoraExport + from art.trainer_rank._impl import TrainerRank, _DynamicOptimizer + + +MANIFEST_FILE = "checkpoint.json" +_LAYER_RE = re.compile(r"\.layers\.(?P\d+)\.") + + +class TensorRecord(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid") + + file: str + tensor: str + shape: tuple[int, ...] + dtype: str + + +class ParameterRecord(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid") + + master: TensorRecord + exp_avg: TensorRecord + exp_avg_sq: TensorRecord + + +class AdamWRecord(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid") + + learning_rate: float + beta1: float + beta2: float + eps: float + weight_decay: float + amsgrad: Literal[False] = False + + +class CheckpointManifest(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid") + + format_version: Literal[1] = 1 + optimizer_format: Literal["adamw"] = "adamw" + base_model_name_or_path: str + optimizer: AdamWRecord | None + parameters: dict[str, ParameterRecord] + steps: dict[str, Annotated[float, Field(ge=0)]] + digest: str + + +@dataclass(frozen=True) +class PreparedCheckpoint: + path: Path + adapter_config: dict[str, object] + manifest: CheckpointManifest | None + artifact_keys: tuple[str, ...] + digest: str + + +@dataclass(frozen=True) +class LocalOptimizerState: + masters: tuple[torch.Tensor, ...] + exp_avgs: tuple[torch.Tensor, ...] + exp_avg_sqs: tuple[torch.Tensor, ...] + steps: tuple[float, ...] + config: AdamWRecord + + +@dataclass(frozen=True) +class _LocalShard: + metadata: LoraShardMeta + master: torch.Tensor + exp_avg: torch.Tensor | None + exp_avg_sq: torch.Tensor | None + expert: int | None + step: float + + +@dataclass(frozen=True) +class _SnapshotShard: + metadata: LoraShardMeta + lora_file: str + optimizer_files: Mapping[str, str] | None + step: float | None + + +@dataclass(frozen=True) +class _PreparedSave: + sequence: int + snapshot: Path + destination: Path + adapter_config: dict[str, object] + shards: tuple[_SnapshotShard, ...] + optimizer: AdamWRecord | None + + +class _SafeSlice(Protocol): + def get_shape(self) -> list[int]: ... + + def __getitem__(self, slices: tuple[slice, ...]) -> torch.Tensor: ... + + +class _SafeTensorFile(Protocol): + def keys(self) -> list[str]: ... + + def get_slice(self, key: str) -> _SafeSlice: ... + + def get_tensor(self, key: str) -> torch.Tensor: ... + + +def _checkpoint_metadata( + path: str | Path, + *, + require_optimizer: bool = False, + verify_payload: bool = True, + artifact_entries: Iterable[str] | None = None, + expected_digest: str | None = None, +) -> tuple[Path, dict[str, object], tuple[str, ...], CheckpointManifest | None, str]: + if (artifact_entries is None) != (expected_digest is None): + raise ValueError( + "artifact_entries and expected_digest must be provided together" + ) + root = Path(path).resolve(strict=True) + if not root.is_dir(): + raise FileNotFoundError(f"Checkpoint directory does not exist: {path}") + from art.megatron.model_support.lora_disk import load_adapter_config, safe_open + + adapter_config = cast(dict[str, object], load_adapter_config(root)) + with safe_open(root / "adapter_model.safetensors", framework="pt") as handle: + artifact_keys = tuple(sorted(handle.keys())) + manifest_path = root / MANIFEST_FILE + manifest = ( + CheckpointManifest.model_validate_json(manifest_path.read_text()) + if manifest_path.is_file() + else None + ) + if require_optimizer and (manifest is None or manifest.optimizer is None): + raise RuntimeError("Checkpoint does not contain canonical optimizer state") + if manifest is not None: + from art.megatron.model_support.lora_disk import ( + ART_LORA_FORMAT_CONFIG_KEY, + ART_LORA_FORMAT_MEGATRON, + ) + + if adapter_config.get(ART_LORA_FORMAT_CONFIG_KEY) != ART_LORA_FORMAT_MEGATRON: + raise RuntimeError("Exact checkpoint is not in ART-native Megatron format") + _validate_manifest(manifest, artifact_keys) + if artifact_entries is not None: + assert expected_digest is not None + _validate_artifact_view(manifest, artifact_entries, expected_digest) + elif verify_payload: + _validate_files(root, manifest) + actual = _digest(root, manifest.model_copy(update={"digest": ""})) + if actual != manifest.digest: + raise RuntimeError( + f"Checkpoint digest mismatch for {path}: {actual} != {manifest.digest}" + ) + digest = ( + manifest.digest + if manifest is not None + else _hash_files(root, ("adapter_config.json", "adapter_model.safetensors")) + ) + return root, adapter_config, artifact_keys, manifest, digest + + +def validate_checkpoint( + path: str | Path, *, require_optimizer: bool = False +) -> CheckpointManifest | None: + """Validate an ART checkpoint without materializing trainer-local state.""" + return _checkpoint_metadata(path, require_optimizer=require_optimizer)[3] + + +def materialize_lora( + path: str | Path, + output_dir: str | Path, + *, + require_optimizer: bool = False, + artifact_entries: Iterable[str] | None = None, + expected_digest: str | None = None, +) -> None: + """Copy only the inference adapter view from a validated checkpoint.""" + root, _config, _keys, _manifest, _digest_value = _checkpoint_metadata( + path, + require_optimizer=require_optimizer, + verify_payload=artifact_entries is None, + artifact_entries=artifact_entries, + expected_digest=expected_digest, + ) + inference_files = {"adapter_config.json", "adapter_model.safetensors"} + destination = Path(output_dir) + if destination.exists() and any(destination.iterdir()): + raise FileExistsError(f"LoRA output directory is not empty: {destination}") + shutil.copytree( + root, + destination, + dirs_exist_ok=True, + ignore=lambda _root, names: [ + name for name in names if name not in inference_files + ], + ) + from art.megatron.model_support.lora_disk import ( + normalize_lora_checkpoint_to_vllm, + ) + + normalize_lora_checkpoint_to_vllm(destination) + + +def prepare_checkpoint(path: str) -> PreparedCheckpoint: + root, adapter_config, artifact_keys, manifest, digest = _checkpoint_metadata( + path, verify_payload=_is_node_validator() + ) + + return PreparedCheckpoint( + root, + adapter_config, + manifest, + artifact_keys, + digest, + ) + + +def save_checkpoint( + trainer: TrainerRank, + output_dir: str, + checkpoint_name: str, +) -> None: + prepare_checkpoint_save(trainer, output_dir, checkpoint_name) + finish_checkpoint_save(trainer, output_dir) + + +def prepare_checkpoint_save( + trainer: TrainerRank, + output_dir: str, + checkpoint_name: str, +) -> None: + """Capture immutable rank-local state without retaining live tensors.""" + _ensure_checkpoint_save_state(trainer) + destination = Path(output_dir) + if output_dir in trainer._prepared_checkpoint_saves: + raise RuntimeError(f"Checkpoint save is already pending: {output_dir}") + with trainer._checkpoint_save_condition: + trainer._completed_checkpoint_saves.discard(output_dir) + _ensure_checkpoint_group(trainer) + _validate_save_state(trainer, checkpoint_name) + adapter_config = deepcopy(dict(_checkpoint_config(trainer, checkpoint_name))) + dynamic = trainer._dynamic_optimizers.get(checkpoint_name) + initialized = dynamic is not None + if any(value != initialized for value in _gather_objects(initialized)): + raise trainer._slot_state_error( + "Checkpoint optimizer initialization differs across ranks" + ) + optimizer: AdamWRecord | None = None + with _collective_errors("prepare optimizer snapshot metadata"): + if dynamic is not None: + optimizer = _optimizer_config(dynamic) + if any(value != optimizer for value in _gather_objects(optimizer)): + raise trainer._slot_state_error( + "Checkpoint optimizer config differs across ranks" + ) + + sequence = trainer._checkpoint_save_sequence + if any(value != sequence for value in _gather_objects(sequence)): + raise RuntimeError("Checkpoint save sequence differs across ranks") + snapshot = destination.with_name( + f".{destination.name}.snapshot-r{_rank()}-{uuid.uuid4().hex}" + ) + prepared: _PreparedSave | None = None + error: BaseException | None = None + try: + snapshot.mkdir(parents=True) + from art.megatron.weights.lora_publish import collect_local_lora_entries + + local_tensors, metadata = collect_local_lora_entries( + trainer.runtime.model, + {}, + owner_rank=_rank(), + slot_ref=trainer._slot_ref(checkpoint_name), + ) + local_optimizer = ( + () + if dynamic is None + else _local_optimizer_shards(trainer, checkpoint_name, dynamic) + ) + optimizer_by_key = {item.metadata.key: item for item in local_optimizer} + if dynamic is not None and set(optimizer_by_key) != set(local_tensors): + raise trainer._slot_state_error( + "Local optimizer tensors differ from the LoRA snapshot: " + f"optimizer={sorted(optimizer_by_key)} " + f"lora={sorted(local_tensors)}" + ) + + records: list[_SnapshotShard] = [] + blocks = sorted({item.block for item in metadata}) + for index, block in enumerate(blocks): + block_metadata = [item for item in metadata if item.block == block] + lora_file = f"lora-{index:06d}.safetensors" + _save_file( + { + item.key: local_tensors[item.key].detach().cpu().contiguous() + for item in block_metadata + }, + snapshot / lora_file, + ) + optimizer_files: dict[str, str] | None = None + if dynamic is not None: + optimizer_files = {} + for component in ("master", "exp_avg", "exp_avg_sq"): + relative = f"{component}-{index:06d}.safetensors" + _save_file( + { + item.key: _optimizer_component( + optimizer_by_key[item.key], component + ) + .detach() + .cpu() + .contiguous() + for item in block_metadata + }, + snapshot / relative, + ) + optimizer_files[component] = relative + records.extend( + _SnapshotShard( + item, + lora_file, + None if optimizer_files is None else dict(optimizer_files), + (None if dynamic is None else optimizer_by_key[item.key].step), + ) + for item in block_metadata + ) + + manifest = { + "sequence": sequence, + "adapter_config": adapter_config, + "optimizer": None if optimizer is None else optimizer.model_dump(), + "shards": [ + { + "metadata": dict(record.metadata._asdict()), + "lora_file": record.lora_file, + "optimizer_files": record.optimizer_files, + "step": record.step, + } + for record in records + ], + } + (snapshot / "snapshot.json").write_text( + json.dumps(manifest, indent=2, sort_keys=True) + "\n" + ) + prepared = _PreparedSave( + sequence, + snapshot, + destination, + adapter_config, + tuple(records), + optimizer, + ) + except BaseException as exc: + error = exc + try: + _raise_distributed(error, "prepare immutable checkpoint snapshot") + except BaseException: + shutil.rmtree(snapshot, ignore_errors=True) + raise + assert prepared is not None + trainer._prepared_checkpoint_saves[output_dir] = prepared + trainer._checkpoint_save_sequence += 1 + + +def finish_checkpoint_save(trainer: TrainerRank, output_dir: str) -> None: + """Finalize a prepared checkpoint without reading mutable trainer state.""" + _ensure_checkpoint_save_state(trainer) + condition = trainer._checkpoint_save_condition + with condition: + condition.wait_for( + lambda: output_dir not in trainer._checkpoint_finishing_saves + ) + if output_dir in trainer._completed_checkpoint_saves: + return + prepared = trainer._prepared_checkpoint_saves.get(output_dir) + if prepared is None: + raise RuntimeError(f"Checkpoint save was not prepared: {output_dir}") + condition.wait_for( + lambda: prepared.sequence == trainer._checkpoint_finish_sequence + ) + trainer._checkpoint_finishing_saves.add(output_dir) + + error: BaseException | None = None + cleanup_error: BaseException | None = None + try: + _finish_prepared_save(trainer, prepared) + except BaseException as exc: + error = exc + try: + shutil.rmtree(prepared.snapshot) + except FileNotFoundError: + pass + except BaseException as exc: + cleanup_error = exc + finally: + trainer._prepared_checkpoint_saves.pop(output_dir, None) + with condition: + trainer._checkpoint_finishing_saves.discard(output_dir) + if error is None and cleanup_error is None: + trainer._completed_checkpoint_saves.add(output_dir) + trainer._checkpoint_finish_sequence += 1 + condition.notify_all() + if error is not None and cleanup_error is not None: + raise BaseExceptionGroup( + "checkpoint finalization and snapshot cleanup both failed", + [error, cleanup_error], + ) + if error is not None: + raise error + if cleanup_error is not None: + raise cleanup_error + + +def export_lora( + trainer: TrainerRank, + output_dir: str, + checkpoint_name: str, + *, + native_format: bool = False, +) -> int: + config = _checkpoint_config(trainer, checkpoint_name) + _export_lora_files( + trainer, output_dir, checkpoint_name, config, native_format=native_format + ) + return trainer._checkpoint_revisions.get(checkpoint_name, 0) + + +def _export_lora_files( + trainer: TrainerRank, + output_dir: str, + checkpoint_name: str, + config: Mapping[str, object], + *, + native_format: bool, +) -> "_LocalLoraExport": + from art.megatron.weights.lora_publish import save_vllm_lora_from_model + + result: _LocalLoraExport | None = None + error: BaseException | None = None + try: + result = save_vllm_lora_from_model( + model=trainer.runtime.model, + adapter_dtypes={}, + handler=trainer.runtime.model_support_handler, + adapter_config=dict(config), + output_dir=output_dir, + rank=trainer.runtime.rank, + world_size=trainer.runtime.world_size, + slot_ref=trainer._slot_ref(checkpoint_name), + native_format=native_format, + ) + except BaseException as exc: + error = exc + _raise_distributed(error, "export LoRA") + assert result is not None + return result + + +def _load_stage_lora( + trainer: TrainerRank, source: PreparedCheckpoint +) -> dict[str, torch.Tensor]: + if _native_checkpoint(source): + return _load_native_lora(trainer, source) + local_layers = { + match.group("layer") + for key in trainer._local_adapter_keys() + if (match := _LAYER_RE.search(key)) is not None + } + selected = { + key + for key in source.artifact_keys + if (match := _LAYER_RE.search(key)) is None + or match.group("layer") in local_layers + } + from art.megatron.model_support.lora_disk import safe_open + + with safe_open(source.path / "adapter_model.safetensors", framework="pt") as file: + tensors = {key: file.get_tensor(key) for key in selected} + return trainer.runtime.model_support_handler.from_vllm_lora_tensors( + tensors, adapter_config=dict(source.adapter_config) + ) + + +def _native_checkpoint(source: PreparedCheckpoint) -> bool: + from art.megatron.model_support.lora_disk import ( + ART_LORA_FORMAT_CONFIG_KEY, + ART_LORA_FORMAT_MEGATRON, + ) + + return ( + source.adapter_config.get(ART_LORA_FORMAT_CONFIG_KEY) + == ART_LORA_FORMAT_MEGATRON + ) + + +def _local_tensor_plan( + trainer: TrainerRank, +) -> dict[str, tuple[LoraShardManifest, tuple[int, ...]]]: + from art.megatron.lora import LoRA + + plan: dict[str, tuple[LoraShardManifest, tuple[int, ...]]] = {} + for chunk in trainer.runtime.model: + for module in chunk.modules(): + if not isinstance(module, LoRA): + continue + for suffix, parameter in (("lora_A", module.A_T), ("lora_B", module.B_T)): + manifest = module._manifest_for_param(parameter) + keys = module._expected_weight_keys(suffix) + for expert, key in enumerate(keys): + local = parameter[expert] if parameter.ndim == 3 else parameter + plan[key] = (manifest, tuple(reversed(local.shape))) + return plan + + +def _read_local_slice( + handle: _SafeTensorFile, + key: str, + manifest: LoraShardManifest, + expected_shape: tuple[int, ...], +) -> torch.Tensor: + view = handle.get_slice(key) + shape = tuple(view.get_shape()) + slices = [slice(None)] * len(shape) + if not bool(manifest["sharded"]): + tensor = view[tuple(slices)] + else: + axis = int(manifest["export_shard_dim"]) + world_size = int(manifest["shard_world_size"]) + rank = int(manifest["shard_rank"]) + strategy = str(manifest.get("export_shard_strategy", "uniform")) + if strategy == "uniform": + if shape[axis] % world_size: + raise RuntimeError( + f"Checkpoint tensor {key!r} cannot be sharded across {world_size} ranks" + ) + size = shape[axis] // world_size + slices[axis] = slice(rank * size, (rank + 1) * size) + tensor = view[tuple(slices)] + elif strategy == "componentwise": + components = tuple(int(size) for size in manifest["component_sizes"]) + if sum(components) != shape[axis] or any( + size % world_size for size in components + ): + raise RuntimeError( + f"Checkpoint tensor {key!r} has incompatible component shards" + ) + parts: list[torch.Tensor] = [] + offset = 0 + for component in components: + size = component // world_size + slices[axis] = slice(offset + rank * size, offset + (rank + 1) * size) + parts.append(view[tuple(slices)]) + offset += component + tensor = torch.cat(parts, dim=axis) + else: + raise RuntimeError( + f"Checkpoint tensor {key!r} has unsupported shard strategy {strategy!r}" + ) + if tuple(tensor.shape) != expected_shape: + raise RuntimeError( + f"Checkpoint tensor {key!r} has local shape {tuple(tensor.shape)}; " + f"expected {expected_shape}" + ) + return tensor + + +def _load_native_lora( + trainer: TrainerRank, source: PreparedCheckpoint +) -> dict[str, torch.Tensor]: + from art.megatron.model_support.lora_disk import safe_open + + plan = _local_tensor_plan(trainer) + tensors: dict[str, torch.Tensor] = {} + with safe_open(source.path / "adapter_model.safetensors", framework="pt") as file: + keys = set(file.keys()) + for key, (manifest, shape) in plan.items(): + if key in keys: + tensors[key] = _read_local_slice(file, key, manifest, shape) + return tensors + + +def load_checkpoint( + trainer: TrainerRank, source: PreparedCheckpoint, name: str +) -> None: + config = trainer._validate_checkpoint_adapter_config( + name, source.adapter_config, alpha=None + ) + assert config is not None + native = _native_checkpoint(source) + if any(digest != source.digest for digest in _gather_objects(source.digest)): + raise trainer._slot_state_error( + f"Checkpoint {name!r} content differs across ranks" + ) + adapter_model: dict[str, torch.Tensor] = {} + read_error: BaseException | None = None + try: + adapter_model = _load_stage_lora(trainer, source) + except BaseException as exc: + read_error = exc + _raise_distributed(read_error, "read checkpoint") + prepared: dict[str, torch.Tensor] = {} + validation_error: BaseException | None = None + try: + prepared = trainer._preflight_adapter( + name, + adapter_model, + config, + localized=native, + canonicalized=native, + ) + _validate_base_model(trainer, source, config) + trainer._guard_slot_can_load(trainer._slot_ref(name)) + except BaseException as exc: + validation_error = exc + _raise_distributed(validation_error, "validate checkpoint") + expected_keys = { + key for rank_keys in _gather_objects(tuple(prepared)) for key in rank_keys + } + coverage_error: BaseException | None = None + try: + artifact_keys = set(source.artifact_keys) + if native and expected_keys != artifact_keys: + raise trainer._slot_state_error( + "Canonical checkpoint coverage differs from the target runtime: " + f"missing={sorted(artifact_keys - expected_keys)[:8]} " + f"unexpected={sorted(expected_keys - artifact_keys)[:8]}" + ) + except BaseException as exc: + coverage_error = exc + _raise_distributed(coverage_error, "validate checkpoint coverage") + + target = trainer._slot_ref(name) + temporary_name = f"__art_loading_{uuid.uuid4().hex}" + temporary_ref = trainer._slot_ref(temporary_name) + dynamic: _DynamicOptimizer | None = None + loaded = 0 + stage_error: BaseException | None = None + try: + loaded = trainer._load_checkpoint_slot( + temporary_name, + prepared, + alpha=float(config["lora_alpha"]), + _prepared=True, + _localized=native, + ) + trainer._validate_loaded_checkpoint_config(temporary_name, config) + staged_params = tuple(trainer._iter_slot_parameters(temporary_ref)) + trainer._checkpoint_slot_params_by_name[temporary_name] = staged_params + if source.manifest is not None and source.manifest.optimizer is not None: + local_optimizer = _load_local_optimizer(trainer, source, temporary_name) + dynamic = trainer._restore_canonical_optimizer( + temporary_name, local_optimizer + ) + except BaseException as exc: + stage_error = exc + stage_errors = _gather_objects(None if stage_error is None else repr(stage_error)) + if any(stage_errors): + _discard_staged_checkpoint(trainer, temporary_name) + if stage_error is not None: + raise stage_error + raise RuntimeError( + "Another rank failed to stage checkpoint: " + f"{next(item for item in stage_errors if item)}" + ) + + consistency_error: BaseException | None = None + params: tuple[torch.nn.Parameter, ...] = () + try: + params = trainer._validate_checkpoint_consistency( + temporary_name, loaded, expected_keys + ) + except BaseException as exc: + consistency_error = exc + consistency_errors = _gather_objects( + None if consistency_error is None else repr(consistency_error) + ) + if any(consistency_errors): + _discard_staged_checkpoint(trainer, temporary_name) + if consistency_error is not None: + raise consistency_error + raise RuntimeError( + "Another rank rejected staged checkpoint: " + f"{next(item for item in consistency_errors if item)}" + ) + + from art.megatron.lora import ( + _restore_lora_slots, + _snapshot_lora_slots, + replace_lora_slot_in_model, + validate_lora_slot_replacement, + ) + + replacement_error: BaseException | None = None + try: + validate_lora_slot_replacement(trainer.runtime.model, temporary_ref, target) + except BaseException as exc: + replacement_error = exc + replacement_errors = _gather_objects( + None if replacement_error is None else repr(replacement_error) + ) + if any(replacement_errors): + _discard_staged_checkpoint(trainer, temporary_name) + if replacement_error is not None: + raise replacement_error + raise RuntimeError( + "Another rank rejected checkpoint commit: " + f"{next(item for item in replacement_errors if item)}" + ) + + slot_snapshot = _snapshot_lora_slots(trainer.runtime.model) + had_params = name in trainer._checkpoint_slot_params_by_name + previous_params = trainer._checkpoint_slot_params_by_name.get(name) + had_dynamic = name in trainer._dynamic_optimizers + previous_dynamic = trainer._dynamic_optimizers.get(name) + had_config = name in trainer._checkpoint_slot_adapter_configs + previous_config = trainer._checkpoint_slot_adapter_configs.get(name) + had_revision = name in trainer._checkpoint_revisions + previous_revision = trainer._checkpoint_revisions.get(name) + commit_error: BaseException | None = None + try: + replace_lora_slot_in_model(trainer.runtime.model, temporary_ref, target) + trainer._checkpoint_slot_params_by_name.pop(temporary_name, None) + trainer._checkpoint_slot_params_by_name[name] = params + if dynamic is None: + trainer._dynamic_optimizers.pop(name, None) + else: + trainer._dynamic_optimizers[name] = dynamic + trainer._checkpoint_slot_adapter_configs[name] = config + trainer._checkpoint_revisions[name] = ( + trainer._checkpoint_revisions.get(name, -1) + 1 + ) + except BaseException as exc: + commit_error = exc + commit_errors = _gather_objects( + None if commit_error is None else repr(commit_error) + ) + if not any(commit_errors): + return + + rollback_error: BaseException | None = None + try: + _restore_lora_slots(slot_snapshot) + _discard_staged_checkpoint(trainer, temporary_name) + if had_params: + assert previous_params is not None + trainer._checkpoint_slot_params_by_name[name] = previous_params + else: + trainer._checkpoint_slot_params_by_name.pop(name, None) + if had_dynamic: + assert previous_dynamic is not None + trainer._dynamic_optimizers[name] = previous_dynamic + else: + trainer._dynamic_optimizers.pop(name, None) + if had_config: + assert previous_config is not None + trainer._checkpoint_slot_adapter_configs[name] = previous_config + else: + trainer._checkpoint_slot_adapter_configs.pop(name, None) + if had_revision: + assert previous_revision is not None + trainer._checkpoint_revisions[name] = previous_revision + else: + trainer._checkpoint_revisions.pop(name, None) + except BaseException as exc: + rollback_error = exc + rollback_errors = _gather_objects( + None if rollback_error is None else repr(rollback_error) + ) + if any(rollback_errors): + if rollback_error is not None: + raise rollback_error + raise RuntimeError( + "Another rank failed to roll back checkpoint commit: " + f"{next(item for item in rollback_errors if item)}" + ) + if commit_error is not None: + raise commit_error + raise RuntimeError( + "Another rank failed to commit checkpoint: " + f"{next(item for item in commit_errors if item)}" + ) + + +def _discard_staged_checkpoint(trainer: TrainerRank, name: str) -> None: + from art.megatron.lora import delete_lora_slot_from_model + + delete_lora_slot_from_model(trainer.runtime.model, trainer._slot_ref(name)) + trainer._checkpoint_slot_params_by_name.pop(name, None) + + +def _ensure_checkpoint_save_state(trainer: TrainerRank) -> None: + if hasattr(trainer, "_prepared_checkpoint_saves"): + return + trainer._checkpoint_process_group = None + trainer._checkpoint_save_condition = threading.Condition() + trainer._checkpoint_save_sequence = 0 + trainer._checkpoint_finish_sequence = 0 + trainer._prepared_checkpoint_saves = {} + trainer._checkpoint_finishing_saves = set() + trainer._completed_checkpoint_saves = set() + + +def _ensure_checkpoint_group(trainer: TrainerRank) -> dist.ProcessGroup | None: + if _distributed() and trainer._checkpoint_process_group is None: + trainer._checkpoint_process_group = dist.new_group(backend="gloo") + return trainer._checkpoint_process_group + + +def _snapshot_identity(metadata: LoraShardMeta) -> tuple[str, int]: + return metadata.key, int(metadata.manifest.get("shard_rank", 0)) + + +def _load_snapshot_file( + prepared: _PreparedSave, relative: str +) -> dict[str, torch.Tensor]: + load_file = importlib.import_module("safetensors.torch").load_file + return load_file(prepared.snapshot / relative) + + +def _local_snapshot_digests( + prepared: _PreparedSave, +) -> tuple[ + list[tuple[LoraShardMeta, str]], + list[tuple[tuple[str, int], float, str]], +]: + from art.megatron.weights.lora_publish import _tensor_digest + + lora_records: list[tuple[LoraShardMeta, str]] = [] + optimizer_records: list[tuple[tuple[str, int], float, str]] = [] + by_file: dict[str, list[_SnapshotShard]] = {} + for record in prepared.shards: + by_file.setdefault(record.lora_file, []).append(record) + for relative, records in sorted(by_file.items()): + tensors = _load_snapshot_file(prepared, relative) + expected = {record.metadata.key for record in records} + if set(tensors) != expected: + raise RuntimeError( + f"Checkpoint snapshot tensor coverage differs: {relative}" + ) + lora_records.extend( + (record.metadata, _tensor_digest(tensors[record.metadata.key])) + for record in records + ) + if prepared.optimizer is None: + continue + hashers = { + record.metadata.key: hashlib.blake2b(digest_size=16) for record in records + } + for component in ("master", "exp_avg", "exp_avg_sq"): + files = { + cast(Mapping[str, str], record.optimizer_files)[component] + for record in records + } + if len(files) != 1: + raise RuntimeError( + f"Checkpoint snapshot {component} files differ within a block" + ) + component_tensors = _load_snapshot_file(prepared, files.pop()) + if set(component_tensors) != expected: + raise RuntimeError(f"Checkpoint snapshot {component} coverage differs") + for key, hasher in hashers.items(): + hasher.update(_tensor_digest(component_tensors[key]).encode()) + for record in records: + if record.step is None: + raise RuntimeError( + f"Checkpoint snapshot optimizer step is missing for " + f"{record.metadata.key!r}" + ) + optimizer_records.append( + ( + _snapshot_identity(record.metadata), + record.step, + hashers[record.metadata.key].hexdigest(), + ) + ) + return lora_records, optimizer_records + + +def _snapshot_plan( + trainer: TrainerRank, + prepared: _PreparedSave, + group: dist.ProcessGroup | None, +) -> tuple[list[LoraShardMeta], tuple[str, ...], dict[str, float]]: + from art.megatron.weights.lora_publish import ( + _elect_lora_contributors, + _validate_replica_digests, + ) + + local_lora: list[tuple[LoraShardMeta, str]] = [] + local_optimizer: list[tuple[tuple[str, int], float, str]] = [] + error: BaseException | None = None + try: + local_lora, local_optimizer = _local_snapshot_digests(prepared) + except BaseException as exc: + error = exc + _raise_distributed(error, "validate local checkpoint snapshot", group=group) + + gathered_lora = [ + item + for rank_items in _gather_objects(tuple(local_lora), group=group) + for item in rank_items + ] + metadata: list[LoraShardMeta] = [] + with _collective_errors("elect checkpoint contributors", group=group): + _validate_replica_digests(gathered_lora) + metadata = _elect_lora_contributors( + [item for item, _digest_value in gathered_lora] + ) + + optimizers = _gather_objects(prepared.optimizer, group=group) + if any(value != prepared.optimizer for value in optimizers): + raise trainer._slot_state_error( + "Checkpoint snapshot optimizer config differs across ranks" + ) + steps: dict[str, float] = {} + if prepared.optimizer is not None: + gathered_optimizer = [ + item + for rank_items in _gather_objects(tuple(local_optimizer), group=group) + for item in rank_items + ] + expected = {_snapshot_identity(item) for item in metadata} + records: dict[tuple[str, int], list[tuple[float, str]]] = {} + for identity, step, digest in gathered_optimizer: + records.setdefault(identity, []).append((step, digest)) + if set(records) != expected: + raise trainer._slot_state_error( + "Optimizer shard plan differs from snapshotted LoRA: " + f"missing={sorted(expected - set(records))[:8]} " + f"unexpected={sorted(set(records) - expected)[:8]}" + ) + steps_by_key: dict[str, set[float]] = {} + for (key, _shard_rank), replicas in records.items(): + if len({digest for _step, digest in replicas}) != 1: + raise RuntimeError( + f"Inconsistent replicated tensor contents for {key!r}" + ) + shard_steps = {step for step, _digest_value in replicas} + if len(shard_steps) != 1: + raise trainer._slot_state_error( + f"Replicated optimizer step differs for {key!r}" + ) + steps_by_key.setdefault(key, set()).update(shard_steps) + if mismatched := { + key: values for key, values in steps_by_key.items() if len(values) != 1 + }: + raise trainer._slot_state_error( + f"Optimizer shard steps differ: {mismatched}" + ) + steps = {key: values.pop() for key, values in steps_by_key.items()} + + blocks = tuple(sorted({item.block for item in metadata})) + return metadata, blocks, steps + + +def _snapshot_block_tensors( + prepared: _PreparedSave, + metadata: Sequence[LoraShardMeta], + component: Literal["lora", "master", "exp_avg", "exp_avg_sq"], +) -> dict[str, torch.Tensor]: + selected = [item for item in metadata if item.owner_rank == _rank()] + if not selected: + return {} + records = { + _snapshot_identity(record.metadata): record for record in prepared.shards + } + selected_records = [records[_snapshot_identity(item)] for item in selected] + if component == "lora": + files = {record.lora_file for record in selected_records} + else: + files = { + cast(Mapping[str, str], record.optimizer_files)[component] + for record in selected_records + } + if len(files) != 1: + raise RuntimeError( + f"Checkpoint snapshot {component} tensors span unexpected files" + ) + tensors = _load_snapshot_file(prepared, files.pop()) + return {item.key: tensors[item.key] for item in selected} + + +def _serialize_snapshot_lora( + prepared: _PreparedSave, + output: Path, + metadata: list[LoraShardMeta], + blocks: Sequence[str], + group: dist.ProcessGroup | None, +) -> None: + from art.megatron.model_support.lora_disk import ( + ART_LORA_FORMAT_CONFIG_KEY, + ART_LORA_FORMAT_MEGATRON, + _consolidate_safetensors, + save_adapter_config, + ) + from art.megatron.weights.lora_publish import ( + _gather_merged_adapter_tensors, + ) + + shards: list[Path] = [] + try: + for index, block in enumerate(blocks): + block_metadata = [item for item in metadata if item.block == block] + local_tensors: dict[str, torch.Tensor] = {} + error: BaseException | None = None + try: + local_tensors = _snapshot_block_tensors( + prepared, block_metadata, "lora" + ) + except BaseException as exc: + error = exc + _raise_distributed( + error, + f"read checkpoint LoRA block {block}", + group=group, + ) + canonical = _gather_merged_adapter_tensors( + block_metadata, + local_tensors=local_tensors, + rank=_rank(), + device=torch.device("cpu"), + group=group, + ) + with _collective_errors( + f"serialize checkpoint LoRA block {block}", + group=group, + ): + if _rank() == 0: + shard = output / f".adapter_model-{index:06d}.safetensors" + _save_file(canonical, shard) + shards.append(shard) + with _collective_errors("finalize checkpoint LoRA", group=group): + if _rank() == 0: + if not shards: + raise RuntimeError("No LoRA tensors were available to checkpoint") + _consolidate_safetensors(shards, output / "adapter_model.safetensors") + save_adapter_config( + output, + { + **prepared.adapter_config, + ART_LORA_FORMAT_CONFIG_KEY: ART_LORA_FORMAT_MEGATRON, + }, + ) + finally: + if _rank() == 0: + for shard in shards: + shard.unlink(missing_ok=True) + + +def _serialize_snapshot_optimizer( + trainer: TrainerRank, + prepared: _PreparedSave, + output: Path, + metadata: list[LoraShardMeta], + blocks: Sequence[str], + steps: dict[str, float], + group: dist.ProcessGroup | None, +) -> CheckpointManifest: + base_model = str(prepared.adapter_config["base_model_name_or_path"]) + if prepared.optimizer is None: + return CheckpointManifest( + base_model_name_or_path=base_model, + optimizer=None, + parameters={}, + steps={}, + digest="", + ) + + from art.megatron.weights.lora_publish import ( + _gather_merged_adapter_tensors, + ) + + records_by_component: dict[str, dict[str, TensorRecord]] = {} + for component in ("master", "exp_avg", "exp_avg_sq"): + component_records: dict[str, TensorRecord] = {} + for index, block in enumerate(blocks): + block_metadata = [item for item in metadata if item.block == block] + local_tensors: dict[str, torch.Tensor] = {} + error: BaseException | None = None + try: + local_tensors = _snapshot_block_tensors( + prepared, + block_metadata, + component, + ) + except BaseException as exc: + error = exc + _raise_distributed( + error, + f"read checkpoint optimizer {component} block {block}", + group=group, + ) + canonical = _gather_merged_adapter_tensors( + block_metadata, + local_tensors=local_tensors, + rank=_rank(), + device=torch.device("cpu"), + group=group, + ) + with _collective_errors( + f"serialize checkpoint optimizer {component} block {block}", + group=group, + ): + if _rank() == 0: + relative = f"optimizer/{component}-{index:06d}.safetensors" + (output / "optimizer").mkdir(parents=True, exist_ok=True) + _save_file(canonical, output / relative) + component_records.update( + ( + key, + TensorRecord( + file=relative, + tensor=key, + shape=tuple(int(dim) for dim in tensor.shape), + dtype=_dtype_name(tensor.dtype), + ), + ) + for key, tensor in canonical.items() + ) + records_by_component[component] = component_records + + parameters: dict[str, ParameterRecord] = {} + with _collective_errors("validate checkpoint optimizer coverage", group=group): + if _rank() == 0: + coverage = { + component: set(records) + for component, records in records_by_component.items() + } + expected = next(iter(coverage.values()), set()) + if any(keys != expected for keys in coverage.values()): + raise RuntimeError( + f"Canonical optimizer component coverage differs: {coverage}" + ) + from art.megatron.model_support.lora_disk import safe_open + + with safe_open( + output / "adapter_model.safetensors", framework="pt" + ) as handle: + artifact_keys = set(handle.keys()) + if expected != artifact_keys: + raise RuntimeError( + "Canonical optimizer coverage differs from exported LoRA: " + f"optimizer={sorted(expected)} lora={sorted(artifact_keys)}" + ) + parameters = { + key: ParameterRecord( + master=records_by_component["master"][key], + exp_avg=records_by_component["exp_avg"][key], + exp_avg_sq=records_by_component["exp_avg_sq"][key], + ) + for key in sorted(expected) + } + return CheckpointManifest( + base_model_name_or_path=base_model, + optimizer=prepared.optimizer, + parameters=parameters, + steps=steps, + digest="", + ) + + +def _finish_prepared_save( + trainer: TrainerRank, + prepared: _PreparedSave, +) -> None: + group = trainer._checkpoint_process_group + temporary = _temporary_output(str(prepared.destination), group=group) + error: BaseException | None = None + try: + metadata, blocks, steps = _snapshot_plan(trainer, prepared, group) + _serialize_snapshot_lora( + prepared, + temporary, + metadata, + blocks, + group, + ) + manifest = _serialize_snapshot_optimizer( + trainer, + prepared, + temporary, + metadata, + blocks, + steps, + group, + ) + if _rank() == 0: + digest = _digest(temporary, manifest) + _write_manifest( + temporary, + manifest.model_copy(update={"digest": digest}), + ) + _commit_output(temporary, prepared.destination, digest) + except BaseException as exc: + error = exc + finally: + if _rank() == 0 and temporary.exists(): + try: + shutil.rmtree(temporary) + except BaseException as exc: + if error is None: + error = exc + _raise_distributed(error, "finish checkpoint save", group=group) + + +def _optimizer_component(local: _LocalShard, component: str) -> torch.Tensor: + value = getattr(local, component) + if value is None: + value = torch.zeros_like(local.master) + if local.expert is not None: + value = value[local.expert] + return value.T.contiguous() + + +def _local_optimizer_shards( + trainer: TrainerRank, + name: str, + dynamic: _DynamicOptimizer, +) -> tuple[_LocalShard, ...]: + from art.megatron.lora import LoRA, LoraShardMeta, _block_for_key + + params = trainer._checkpoint_slot_params_by_name[name] + masters = { + id(param): master + for param, master in zip(params, dynamic.master_params, strict=True) + } + ref = trainer._slot_ref(name) + items: list[_LocalShard] = [] + for chunk in trainer.runtime.model: + for module in chunk.modules(): + if not isinstance(module, LoRA): + continue + for key, param, expert in module._export_items(ref): + master = masters.get(id(param)) + if master is None: + raise trainer._slot_state_error( + f"Cannot map optimizer parameter for {key!r}" + ) + state = dynamic.optimizer.state.get(master, {}) + exp_avg = state.get("exp_avg") + exp_avg_sq = state.get("exp_avg_sq") + step = state.get("step") + if exp_avg is exp_avg_sq is step is None: + step = 0.0 + elif not all( + isinstance(value, torch.Tensor) for value in (exp_avg, exp_avg_sq) + ): + raise trainer._slot_state_error( + f"AdamW state for {key!r} is incomplete" + ) + local_master = master if expert is None else master[expert] + exported_shape = tuple(reversed(local_master.shape)) + items.append( + _LocalShard( + LoraShardMeta( + key=key, + owner_rank=_rank(), + shape=exported_shape, + dtype_name=_dtype_name(local_master.dtype), + manifest=module._manifest_for_param(param), + block=_block_for_key(key), + ), + master, + cast(torch.Tensor | None, exp_avg), + cast(torch.Tensor | None, exp_avg_sq), + expert, + _scalar(step), + ) + ) + return tuple(items) + + +def _load_local_optimizer( + trainer: TrainerRank, + source: PreparedCheckpoint, + name: str, +) -> LocalOptimizerState: + manifest = source.manifest + assert manifest is not None and manifest.optimizer is not None + plan = _local_tensor_plan(trainer) + localized: dict[str, tuple[torch.Tensor, ...]] = {} + for component in ("master", "exp_avg", "exp_avg_sq"): + records = { + key: cast(TensorRecord, getattr(record, component)) + for key, record in manifest.parameters.items() + if key in plan + } + artifact_tensors = _load_tensors( + source.path, + records, + plans={key: plan[key] for key in records}, + ) + localized[component] = trainer._localize_adapter_tensors( + artifact_tensors, name, localized=True + ) + + lengths = {len(values) for values in localized.values()} + if len(lengths) != 1: + raise trainer._slot_state_error( + f"Canonical optimizer component lengths differ: {lengths}" + ) + steps = tuple( + _parameter_group_step(group, manifest.steps) + for group in trainer._local_parameter_key_groups(name) + ) + return LocalOptimizerState( + masters=localized["master"], + exp_avgs=localized["exp_avg"], + exp_avg_sqs=localized["exp_avg_sq"], + steps=steps, + config=manifest.optimizer, + ) + + +def _parameter_group_step(keys: Sequence[str], steps: Mapping[str, float]) -> float: + missing = [key for key in keys if key not in steps] + if missing: + raise RuntimeError(f"Canonical optimizer is missing steps: {missing}") + values = {steps[key] for key in keys} + if len(values) != 1: + raise RuntimeError( + f"Target optimizer parameter combines different steps for {tuple(keys)!r}" + ) + return values.pop() + + +def _checkpoint_config(trainer: TrainerRank, name: str) -> Mapping[str, object]: + loaded = name in trainer._checkpoint_slot_params_by_name + if any(value != loaded for value in _gather_objects(loaded)): + raise trainer._slot_state_error( + f"Checkpoint {name!r} is not loaded consistently across ranks" + ) + if not loaded: + raise ValueError(f"Unknown checkpoint: {name!r}") + config = trainer._checkpoint_slot_adapter_configs.get(name) + configs = _gather_objects(config) + if any(value != config for value in configs): + raise trainer._slot_state_error( + f"Checkpoint {name!r} adapter config differs across ranks" + ) + if config is None: + raise trainer._slot_state_error( + f"Checkpoint {name!r} was loaded without adapter_config" + ) + return config + + +def _validate_save_state(trainer: TrainerRank, name: str) -> None: + _checkpoint_config(trainer, name) + if any(trainer._checkpoint_grad_flags((name,))): + raise trainer._slot_state_error( + f"Cannot save checkpoint {name!r} with accumulated gradients" + ) + local_graph = trainer._has_live_slot_graph(trainer._slot_ref(name)) + if any(bool(value) for value in _gather_objects(local_graph)): + raise trainer._slot_state_error( + f"Cannot save checkpoint {name!r} with a live backward graph" + ) + + +def _validate_base_model( + trainer: TrainerRank, + source: PreparedCheckpoint, + config: Mapping[str, object], +) -> None: + configured = str(config["base_model_name_or_path"]) + if ( + source.manifest is not None + and source.manifest.base_model_name_or_path != configured + ): + raise trainer._slot_state_error( + "Checkpoint manifest and adapter config name different base models" + ) + runtime_model = getattr(trainer.runtime, "model_identifier", None) + if runtime_model is not None and runtime_model != configured: + kind = "Exact checkpoint" if source.manifest is not None else "Checkpoint" + raise trainer._slot_state_error( + f"{kind} base model {configured!r} differs from runtime model " + f"{runtime_model!r}" + ) + supported = tuple( + getattr(getattr(trainer.runtime, "model_support_spec", None), "model_names", ()) + ) + if supported and configured not in supported: + raise trainer._slot_state_error( + f"Checkpoint base model {configured!r} is incompatible with this runtime" + ) + + +def _optimizer_config(dynamic: _DynamicOptimizer) -> AdamWRecord: + group = dynamic.optimizer.param_groups[0] + if bool(group["amsgrad"]): + raise RuntimeError("Canonical checkpoints do not support AdamW amsgrad state") + beta1, beta2 = cast(tuple[float, float], group["betas"]) + return AdamWRecord( + learning_rate=float(group["lr"]), + beta1=float(beta1), + beta2=float(beta2), + eps=float(group["eps"]), + weight_decay=float(group["weight_decay"]), + amsgrad=False, + ) + + +def _temporary_output( + output_dir: str, *, group: dist.ProcessGroup | None = None +) -> Path: + destination = Path(output_dir) + value = str( + destination.with_name(f".{destination.name}.tmp-{uuid.uuid4().hex}") + if _rank() == 0 + else "" + ) + values = [value] + if _distributed(): + torch.distributed.broadcast_object_list(values, src=0, group=group) + temporary = Path(values[0]) + error: BaseException | None = None + try: + if _rank() == 0: + temporary.mkdir(parents=True) + except BaseException as exc: + error = exc + _raise_distributed(error, "create checkpoint staging directory", group=group) + return temporary + + +def _commit_output(temporary: Path, destination: Path, digest: str) -> None: + if destination.exists(): + existing_manifest = destination / MANIFEST_FILE + if existing_manifest.is_file(): + existing = CheckpointManifest.model_validate_json( + existing_manifest.read_text() + ) + if existing.digest == digest: + return + if any(destination.iterdir()): + raise FileExistsError(f"Checkpoint output is not empty: {destination}") + destination.rmdir() + destination.parent.mkdir(parents=True, exist_ok=True) + os.replace(temporary, destination) + + +def _write_manifest(path: Path, manifest: CheckpointManifest) -> None: + target = path / MANIFEST_FILE + temporary = target.with_suffix(".tmp") + temporary.write_text(manifest.model_dump_json(indent=2) + "\n") + os.replace(temporary, target) + + +def _digest(root: Path, manifest: CheckpointManifest) -> str: + seed = json.dumps( + manifest.model_dump(mode="json"), sort_keys=True, separators=(",", ":") + ).encode() + files = { + "adapter_config.json", + "adapter_model.safetensors", + *_manifest_files(manifest), + } + return _hash_files(root, files, seed=seed) + + +def _hash_files(root: Path, files: Iterable[str], *, seed: bytes = b"") -> str: + digest = hashlib.sha256(seed) + for relative in sorted(set(files)): + digest.update(relative.encode()) + with (root / _safe_relative_path(relative)).open("rb") as handle: + while chunk := handle.read(1024 * 1024): + digest.update(chunk) + return digest.hexdigest() + + +def _manifest_files(manifest: CheckpointManifest) -> set[str]: + return { + record.file + for value in manifest.parameters.values() + for record in (value.master, value.exp_avg, value.exp_avg_sq) + } + + +def _safe_relative_path(relative: str) -> PurePosixPath: + path = PurePosixPath(relative) + windows_path = PureWindowsPath(relative) + if ( + not relative + or chr(0) in relative + or chr(92) in relative + or ":" in relative + or path.is_absolute() + or windows_path.is_absolute() + or bool(windows_path.drive) + or any(part in {"", ".", ".."} for part in relative.split("/")) + or path.as_posix() != relative + ): + raise RuntimeError(f"Unsafe checkpoint tensor path: {relative!r}") + return path + + +def _validate_artifact_view( + manifest: CheckpointManifest, + artifact_entries: Iterable[str], + expected_digest: str, +) -> None: + if not expected_digest or manifest.digest != expected_digest: + raise RuntimeError( + "Checkpoint manifest digest differs from the durable artifact digest" + ) + required = { + MANIFEST_FILE, + "adapter_config.json", + "adapter_model.safetensors", + *_manifest_files(manifest), + } + missing = sorted(required - set(artifact_entries)) + if missing: + raise RuntimeError( + f"Checkpoint artifact is missing referenced files: {missing[:8]}" + ) + + +def _validate_manifest( + manifest: CheckpointManifest, artifact_keys: Sequence[str] +) -> None: + for relative in _manifest_files(manifest): + _safe_relative_path(relative) + parameter_keys = set(manifest.parameters) + step_keys = set(manifest.steps) + if manifest.optimizer is None: + if parameter_keys or step_keys: + raise RuntimeError( + "LoRA-only checkpoint unexpectedly contains optimizer parameters" + ) + return + for key, parameter in manifest.parameters.items(): + for component, record in ( + ("master", parameter.master), + ("exp_avg", parameter.exp_avg), + ("exp_avg_sq", parameter.exp_avg_sq), + ): + if record.dtype != "float32": + raise RuntimeError( + f"Checkpoint optimizer tensor {key!r} {component} must use " + f"float32, got {record.dtype!r}" + ) + if parameter_keys != set(artifact_keys): + missing = sorted(set(artifact_keys) - parameter_keys) + extra = sorted(parameter_keys - set(artifact_keys)) + raise RuntimeError( + "Checkpoint optimizer coverage differs from its LoRA tensors: " + f"missing={missing[:8]} extra={extra[:8]}" + ) + if step_keys != parameter_keys: + raise RuntimeError( + "Checkpoint optimizer step coverage differs from its tensors: " + f"missing={sorted(parameter_keys - step_keys)[:8]} " + f"extra={sorted(step_keys - parameter_keys)[:8]}" + ) + + +def _validate_files(root: Path, manifest: CheckpointManifest) -> None: + safe_open = importlib.import_module("safetensors").safe_open + files: dict[str, dict[str, TensorRecord]] = {} + for key, value in manifest.parameters.items(): + for record in (value.master, value.exp_avg, value.exp_avg_sq): + if record.tensor != key or record.tensor in files.setdefault( + record.file, {} + ): + raise RuntimeError(f"Invalid checkpoint tensor index: {record.file}") + files[record.file][record.tensor] = record + for relative, records in files.items(): + candidate = (root / _safe_relative_path(relative)).resolve() + if root.resolve() not in candidate.parents or not candidate.is_file(): + raise RuntimeError(f"Invalid checkpoint tensor path: {relative}") + with safe_open(candidate, framework="pt") as handle: + if set(handle.keys()) != set(records): + raise RuntimeError(f"Checkpoint tensor index mismatch: {relative}") + for key, record in records.items(): + view = handle.get_slice(key) + shape = tuple(view.get_shape()) + empty = view[tuple(slice(0, 0) for _ in shape)] + if shape != record.shape or _dtype_name(empty.dtype) != record.dtype: + raise RuntimeError( + f"Checkpoint tensor metadata mismatch: {record.file}" + ) + + +def _load_tensors( + root: Path, + records: Mapping[str, TensorRecord], + *, + plans: Mapping[str, tuple[LoraShardManifest, tuple[int, ...]]] | None = None, +) -> dict[str, torch.Tensor]: + safe_open = importlib.import_module("safetensors").safe_open + files: dict[str, list[tuple[str, TensorRecord]]] = {} + for key, record in records.items(): + files.setdefault(record.file, []).append((key, record)) + tensors: dict[str, torch.Tensor] = {} + for relative, file_records in files.items(): + with safe_open(root / relative, framework="pt") as handle: + for key, record in file_records: + if plans is None: + tensor = handle.get_tensor(record.tensor) + shape_matches = tuple(tensor.shape) == record.shape + else: + manifest, expected_shape = plans[key] + tensor = _read_local_slice( + handle, record.tensor, manifest, expected_shape + ) + shape_matches = True + if not shape_matches or _dtype_name(tensor.dtype) != record.dtype: + raise RuntimeError( + f"Checkpoint tensor metadata mismatch: {record.file}" + ) + tensors[record.tensor] = tensor + return tensors + + +def _save_file(tensors: dict[str, torch.Tensor], path: Path) -> None: + importlib.import_module("safetensors.torch").save_file(tensors, path) + + +def _is_node_validator() -> bool: + if not _distributed(): + return True + local_rank = os.environ.get("LOCAL_RANK") + return True if local_rank is None else int(local_rank) == 0 + + +def _scalar(value: object) -> float: + if isinstance(value, torch.Tensor) and value.numel() == 1: + return float(value.item()) + if isinstance(value, int | float) and not isinstance(value, bool): + return float(value) + raise RuntimeError("AdamW optimizer step is not scalar") diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index afae6f79f..0d7663661 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -2,7 +2,10 @@ from __future__ import annotations +import asyncio from collections.abc import ( + Awaitable, + Callable, Iterable, Iterator, Mapping, @@ -11,7 +14,8 @@ from copy import deepcopy from dataclasses import dataclass import os -from types import TracebackType +from pathlib import Path +import threading from typing import ( TYPE_CHECKING, Generic, @@ -47,7 +51,11 @@ from art.megatron.lora import LoRASlotRef from art.megatron.prefix_tree_state import PrefixTreeAttentionState from art.megatron.train import TrainingRuntime - from art.trainer_rank import TrainerRankOptimizerLayout, TrainerRankOptimizerState + from art.trainer_rank._checkpoint import ( + LocalOptimizerState, + PreparedCheckpoint, + _PreparedSave, + ) @dataclass(frozen=True) @@ -91,7 +99,6 @@ class _Unset: @dataclass(frozen=True) class _LocalLoRASlotRef: - kind: Literal["checkpoint", "lora"] name: str | None @@ -111,7 +118,6 @@ class ForwardInput(Generic[LogprobsT, TopKT, LogitsT, HiddenStatesT]): logits: bool = False hidden_states: bool = False checkpoint: AdapterSelection = Unset - lora: AdapterSelection = Unset @overload def __new__( @@ -123,7 +129,6 @@ def __new__( logits: Literal[False] = False, hidden_states: Literal[False] = False, checkpoint: AdapterSelection = Unset, - lora: AdapterSelection = Unset, ) -> "ForwardInput[None, None, None, None]": ... @overload @@ -136,7 +141,6 @@ def __new__( logits: Literal[False] = False, hidden_states: Literal[False] = False, checkpoint: AdapterSelection = Unset, - lora: AdapterSelection = Unset, ) -> "ForwardInput[torch.Tensor, None, None, None]": ... @overload @@ -149,7 +153,6 @@ def __new__( logits: Literal[False] = False, hidden_states: Literal[False] = False, checkpoint: AdapterSelection = Unset, - lora: AdapterSelection = Unset, ) -> "ForwardInput[None, TopK, None, None]": ... @overload @@ -162,7 +165,6 @@ def __new__( logits: Literal[True], hidden_states: Literal[False] = False, checkpoint: AdapterSelection = Unset, - lora: AdapterSelection = Unset, ) -> "ForwardInput[None, None, torch.Tensor, None]": ... @overload @@ -175,7 +177,6 @@ def __new__( logits: Literal[False] = False, hidden_states: Literal[True], checkpoint: AdapterSelection = Unset, - lora: AdapterSelection = Unset, ) -> "ForwardInput[None, None, None, torch.Tensor]": ... @overload @@ -188,7 +189,6 @@ def __new__( logits: Literal[False] = False, hidden_states: Literal[False] = False, checkpoint: AdapterSelection = Unset, - lora: AdapterSelection = Unset, ) -> "ForwardInput[torch.Tensor, TopK, None, None]": ... @overload @@ -201,7 +201,6 @@ def __new__( logits: Literal[True], hidden_states: Literal[False] = False, checkpoint: AdapterSelection = Unset, - lora: AdapterSelection = Unset, ) -> "ForwardInput[torch.Tensor, None, torch.Tensor, None]": ... @overload @@ -214,7 +213,6 @@ def __new__( logits: Literal[False] = False, hidden_states: Literal[True], checkpoint: AdapterSelection = Unset, - lora: AdapterSelection = Unset, ) -> "ForwardInput[torch.Tensor, None, None, torch.Tensor]": ... @overload @@ -227,7 +225,6 @@ def __new__( logits: Literal[True], hidden_states: Literal[False] = False, checkpoint: AdapterSelection = Unset, - lora: AdapterSelection = Unset, ) -> "ForwardInput[None, TopK, torch.Tensor, None]": ... @overload @@ -240,7 +237,6 @@ def __new__( logits: Literal[False] = False, hidden_states: Literal[True], checkpoint: AdapterSelection = Unset, - lora: AdapterSelection = Unset, ) -> "ForwardInput[None, TopK, None, torch.Tensor]": ... @overload @@ -253,7 +249,6 @@ def __new__( logits: Literal[True], hidden_states: Literal[True], checkpoint: AdapterSelection = Unset, - lora: AdapterSelection = Unset, ) -> "ForwardInput[None, None, torch.Tensor, torch.Tensor]": ... @overload @@ -266,7 +261,6 @@ def __new__( logits: Literal[True], hidden_states: Literal[False] = False, checkpoint: AdapterSelection = Unset, - lora: AdapterSelection = Unset, ) -> "ForwardInput[torch.Tensor, TopK, torch.Tensor, None]": ... @overload @@ -279,7 +273,6 @@ def __new__( logits: Literal[False] = False, hidden_states: Literal[True], checkpoint: AdapterSelection = Unset, - lora: AdapterSelection = Unset, ) -> "ForwardInput[torch.Tensor, TopK, None, torch.Tensor]": ... @overload @@ -292,7 +285,6 @@ def __new__( logits: Literal[True], hidden_states: Literal[True], checkpoint: AdapterSelection = Unset, - lora: AdapterSelection = Unset, ) -> "ForwardInput[torch.Tensor, None, torch.Tensor, torch.Tensor]": ... @overload @@ -305,7 +297,6 @@ def __new__( logits: Literal[True], hidden_states: Literal[True], checkpoint: AdapterSelection = Unset, - lora: AdapterSelection = Unset, ) -> "ForwardInput[None, TopK, torch.Tensor, torch.Tensor]": ... @overload @@ -318,7 +309,6 @@ def __new__( logits: Literal[True], hidden_states: Literal[True], checkpoint: AdapterSelection = Unset, - lora: AdapterSelection = Unset, ) -> "ForwardInput[torch.Tensor, TopK, torch.Tensor, torch.Tensor]": ... @overload @@ -331,7 +321,6 @@ def __new__( logits: bool = False, hidden_states: bool = False, checkpoint: AdapterSelection = Unset, - lora: AdapterSelection = Unset, ) -> "ForwardInput[torch.Tensor | None, TopK | None, torch.Tensor | None, torch.Tensor | None]": ... def __new__( @@ -343,15 +332,12 @@ def __new__( logits: bool = False, hidden_states: bool = False, checkpoint: AdapterSelection = Unset, - lora: AdapterSelection = Unset, ) -> Self: return object.__new__(cls) def __post_init__(self) -> None: if self.top_k is not None and self.top_k < 1: raise ValueError("top_k must be >= 1") - if self.checkpoint is not Unset and self.lora is not Unset: - raise ValueError("ForwardInput cannot set both checkpoint and lora") type AnyForwardInput = ForwardInput[ @@ -455,26 +441,37 @@ class _DynamicOptimizer: @dataclass(frozen=True) -class _PushedSlot: +class PushedCheckpoint: trainer: "TrainerRank" - ref: "LoRASlotRef" + path: str | None + task: asyncio.Task[None] + + def __await__(self): + return self.task.__await__() - def __enter__(self) -> "_PushedSlot": + async def __aenter__(self) -> "PushedCheckpoint": + try: + await self.task + except asyncio.CancelledError: + if ( + self.task.done() + and not self.task.cancelled() + and self.task.exception() is None + ): + self._pop() + raise return self - def __exit__( - self, - exc_type: type[BaseException] | None, - exc_value: BaseException | None, - traceback: TracebackType | None, - ) -> bool: - if not self.trainer._slot_stack or self.trainer._slot_stack[-1] != self.ref: - raise RuntimeError( - "Pushed LoRA/checkpoint stack changed before context exit" - ) - self.trainer.pop_pushed_lora_or_checkpoint() + async def __aexit__(self, *exc_info: object) -> bool: + self._pop() return False + def _pop(self) -> None: + ref = self.trainer._slot_ref(self.path) + if not self.trainer._slot_stack or self.trainer._slot_stack[-1] != ref: + raise RuntimeError("Pushed checkpoint stack changed before context exit") + self.trainer.pop_checkpoint() + @dataclass(frozen=True) class _ForwardItem: @@ -567,6 +564,15 @@ def __init__( str, tuple[torch.nn.Parameter, ...] ] = {} self._checkpoint_slot_adapter_configs: dict[str, _AdapterConfig] = {} + self._checkpoint_revisions: dict[str, int] = {} + self._checkpoint_process_group: dist.ProcessGroup | None = None + self._checkpoint_mutation_tail: asyncio.Task[None] | None = None + self._checkpoint_save_condition = threading.Condition() + self._checkpoint_save_sequence = 0 + self._checkpoint_finish_sequence = 0 + self._prepared_checkpoint_saves: dict[str, _PreparedSave] = {} + self._checkpoint_finishing_saves: set[str] = set() + self._completed_checkpoint_saves: set[str] = set() self._pending_slot_graphs: dict[ LoRASlotRef, list[weakref.ReferenceType[torch.Tensor]] ] = {} @@ -591,121 +597,220 @@ def zero_grad(self) -> None: param.grad = None self._prune_slot_graphs() - def set_checkpoint(self, name: str | None) -> None: - self._set_default_slot(self._slot_ref("checkpoint", name)) + def prefetch_checkpoints(self, *paths: str) -> asyncio.Task[None]: + return self._prefetch_checkpoint_paths(paths) + + def _prefetch_checkpoints_from_sources( + self, *sources: tuple[str, str] + ) -> asyncio.Task[None]: + return self._prefetch_checkpoint_paths( + source_path for _logical_path, source_path in sources + ) + + def _prefetch_checkpoint_paths(self, paths: Iterable[str]) -> asyncio.Task[None]: + async def prefetch() -> None: + await asyncio.gather(*(self._prefetch_checkpoint(path) for path in paths)) - def set_lora(self, name: str | None) -> None: - self._set_default_slot(self._slot_ref("lora", name)) + return asyncio.create_task(prefetch()) - def push_checkpoint(self, name: str | None) -> _PushedSlot: - ref = self._slot_ref("checkpoint", name) - self._slot_stack.append(ref) - return _PushedSlot(self, ref) + def load_checkpoint(self, path: str | None) -> asyncio.Task[None]: + return self._load_checkpoint(path, path) - def push_lora(self, name: str | None) -> _PushedSlot: - ref = self._slot_ref("lora", name) - self._slot_stack.append(ref) - return _PushedSlot(self, ref) + def _load_checkpoint_from_source( + self, logical_path: str, source_path: str + ) -> asyncio.Task[None]: + return self._load_checkpoint(logical_path, source_path) - def pop_pushed_lora_or_checkpoint(self) -> None: + def _load_checkpoint( + self, logical_path: str | None, source_path: str | None + ) -> asyncio.Task[None]: + prefetch = ( + None + if source_path is None + else asyncio.create_task(self._prefetch_checkpoint(source_path)) + ) + + async def load() -> None: + if self._slot_stack: + raise RuntimeError("Cannot load a checkpoint while one is pushed") + if logical_path is None: + self._set_default_slot(self._slot_ref(None)) + return + assert source_path is not None and prefetch is not None + await self._load_checkpoint_path( + logical_path, source_path=source_path, prefetch=prefetch + ) + self._set_default_slot(self._slot_ref(logical_path)) + + return self._checkpoint_mutation_task(load) + + def push_checkpoint(self, path: str | None) -> PushedCheckpoint: + return self._push_checkpoint(path, path) + + def _push_checkpoint_from_source( + self, logical_path: str, source_path: str + ) -> PushedCheckpoint: + return self._push_checkpoint(logical_path, source_path) + + def _push_checkpoint( + self, logical_path: str | None, source_path: str | None + ) -> PushedCheckpoint: + prefetch = ( + asyncio.create_task(self._prefetch_checkpoint(source_path)) + if source_path is not None + and logical_path not in self._checkpoint_slot_params_by_name + else None + ) + + async def push() -> None: + if prefetch is not None: + assert logical_path is not None and source_path is not None + await self._load_checkpoint_path( + logical_path, + source_path=source_path, + prefetch=prefetch, + ) + self._slot_stack.append(self._slot_ref(logical_path)) + + return PushedCheckpoint( + self, logical_path, self._checkpoint_mutation_task(push) + ) + + def _checkpoint_mutation_task( + self, operation: Callable[[], Awaitable[None]] + ) -> asyncio.Task[None]: + predecessor = getattr(self, "_checkpoint_mutation_tail", None) + + async def ordered() -> None: + if predecessor is not None: + try: + await asyncio.shield(predecessor) + except asyncio.CancelledError: + current = asyncio.current_task() + if current is not None and current.cancelling(): + raise + except Exception: + pass + await operation() + + task = asyncio.create_task(ordered()) + self._checkpoint_mutation_tail = task + return task + + def pop_checkpoint(self) -> None: if not self._slot_stack: - raise RuntimeError("No pushed LoRA or checkpoint to pop") + raise RuntimeError("No pushed checkpoint to pop") self._slot_stack.pop() - def load_checkpoint_slot( + def save_checkpoint( self, - name: str, - adapter_model: Mapping[str, torch.Tensor], - *, - optimizer_state: TrainerRankOptimizerState | None = None, - alpha: float | None = None, - adapter_config: Mapping[str, object] | None = None, - ) -> int: - config = self._validate_checkpoint_slot_adapter_config( - name, adapter_config, alpha=alpha + output_dir: str, + checkpoint_path: str | Literal["active"] = "active", + ) -> None: + from . import _checkpoint + + _checkpoint.save_checkpoint( + self, output_dir, self._resolve_checkpoint_name(checkpoint_path) ) - loaded = self._load_slot( - "checkpoint", - name, - adapter_model, - trainable=True, - alpha=alpha if config is None else float(config["lora_alpha"]), + + def _prepare_checkpoint_save( + self, + output_dir: str, + checkpoint_path: str | Literal["active"] = "active", + ) -> None: + from . import _checkpoint + + _checkpoint.prepare_checkpoint_save( + self, output_dir, self._resolve_checkpoint_name(checkpoint_path) ) - slot_params = self._validate_dynamic_slot_consistency( - "checkpoint", name, loaded + + def _finish_checkpoint_save(self, output_dir: str) -> None: + from . import _checkpoint + + _checkpoint.finish_checkpoint_save(self, output_dir) + + def export_lora( + self, + output_dir: str, + checkpoint_path: str | Literal["active"] = "active", + ) -> int: + from . import _checkpoint + + return _checkpoint.export_lora( + self, output_dir, self._resolve_checkpoint_name(checkpoint_path) ) - if config is not None: - self._validate_loaded_checkpoint_slot_config(name, config) - self._checkpoint_slot_params_by_name[name] = slot_params - if optimizer_state is None: - self._dynamic_optimizers.pop(name, None) - else: - self._dynamic_optimizers[name] = self._restore_dynamic_optimizer( - name, optimizer_state + + @staticmethod + def _checkpoint_source_key(path: str) -> str: + return str(Path(path).resolve()) + + async def _prefetch_checkpoint(self, source_path: str) -> PreparedCheckpoint: + key = self._checkpoint_source_key(source_path) + sources = getattr(self, "_checkpoint_sources", None) + if sources is None: + sources = self._checkpoint_sources = {} + if key in sources: + return sources[key] + tasks = getattr(self, "_checkpoint_prefetch_tasks", None) + if tasks is None: + tasks = self._checkpoint_prefetch_tasks = {} + task = tasks.get(key) + if task is None: + from ._checkpoint import prepare_checkpoint + + task = tasks[key] = asyncio.create_task( + asyncio.to_thread(prepare_checkpoint, key) ) - configs = getattr(self, "_checkpoint_slot_adapter_configs", None) - if configs is None: - configs = self._checkpoint_slot_adapter_configs = {} - if config is None: - configs.pop(name, None) - else: - configs[name] = config - return loaded - - def checkpoint_slot_optimizer_state( - self, name: str - ) -> TrainerRankOptimizerState | None: - if name not in self._checkpoint_slot_params_by_name: - raise ValueError(f"Unknown checkpoint slot: {name!r}") - dynamic = self._dynamic_optimizers.get(name) - if dynamic is None: - return None - state: TrainerRankOptimizerState = { - "format_version": 1, - "layout": self._dynamic_optimizer_layout(name), - "master_params": tuple( - param.detach().cpu().clone() for param in dynamic.master_params - ), - "optimizer": cast( - dict[str, object], - _state_to_cpu(dynamic.optimizer.state_dict()), - ), - } - return state + try: + source = await asyncio.shield(task) + except BaseException: + if task.done(): + tasks.pop(key, None) + raise + sources[key] = source + tasks.pop(key, None) + return source - def save_checkpoint_slot_lora(self, name: str, output_dir: str) -> None: - """Collectively publish a trained checkpoint slot as a vLLM LoRA.""" - known = name in self._checkpoint_slot_params_by_name - if dist.is_available() and dist.is_initialized(): - gathered: list[tuple[str, bool] | None] = [None] * dist.get_world_size() - dist.all_gather_object(gathered, (name, known)) - if any(state != (name, True) for state in gathered): - raise ValueError( - "Checkpoint slot publish requires the same loaded name on all " - f"ranks; got {gathered}" - ) - if not known: - raise ValueError(f"Unknown checkpoint slot: {name!r}") - config = getattr(self, "_checkpoint_slot_adapter_configs", {}).get(name) - if config is None: - raise TrainerRankSlotStateError( - f"Checkpoint slot {name!r} was loaded without adapter_config; " - "reload it with adapter_config=... before publishing." + async def _load_checkpoint_path( + self, + logical_path: str, + *, + source_path: str | None = None, + prefetch: asyncio.Task[PreparedCheckpoint] | None = None, + ) -> None: + from . import _checkpoint + + key = self._checkpoint_source_key(source_path or logical_path) + source: PreparedCheckpoint | None = None + error: BaseException | None = None + try: + source = await asyncio.shield( + prefetch + if prefetch is not None + else asyncio.create_task(self._prefetch_checkpoint(key)) ) - from art.megatron.weights.lora_publish import save_vllm_lora_from_model + except BaseException as exc: + error = exc + _checkpoint._raise_distributed(error, "prepare checkpoint") + assert source is not None + _checkpoint.load_checkpoint(self, source, logical_path) + sources = getattr(self, "_checkpoint_sources", None) + if sources is not None and sources.get(key) is source: + sources.pop(key) + + def _resolve_checkpoint_name(self, checkpoint_path: str | Literal["active"]) -> str: + if checkpoint_path != "active": + return checkpoint_path + ref = self._slot_stack[-1] if self._slot_stack else self._default_slot_ref + if ref is None or ref.name is None: + raise TrainerRankSlotStateError("No active trainable checkpoint") + return ref.name - save_vllm_lora_from_model( - model=self.runtime.model, - adapter_dtypes={}, - handler=self.runtime.model_support_handler, - adapter_config=config, - output_dir=output_dir, - rank=self.runtime.rank, - world_size=self.runtime.world_size, - slot_ref=self._slot_ref("checkpoint", name), - ) + @staticmethod + def _slot_state_error(message: str) -> TrainerRankSlotStateError: + return TrainerRankSlotStateError(message) - def _validate_checkpoint_slot_adapter_config( + def _validate_checkpoint_adapter_config( self, name: str, adapter_config: Mapping[str, object] | None, @@ -757,12 +862,12 @@ def _validate_checkpoint_slot_adapter_config( ) return cast(_AdapterConfig, config) - def _validate_loaded_checkpoint_slot_config( + def _validate_loaded_checkpoint_config( self, name: str, config: _AdapterConfig ) -> None: from art.megatron.lora import LoRA - ref = self._slot_ref("checkpoint", name) + ref = self._slot_ref(name) slots = [ slot for chunk in self.runtime.model @@ -778,19 +883,6 @@ def _validate_loaded_checkpoint_slot_config( f"rank/alpha={expected}, loaded weights use {sorted(actual)}" ) - def load_lora_slot( - self, - name: str, - adapter_model: Mapping[str, torch.Tensor], - *, - alpha: float | None = None, - ) -> int: - loaded = self._load_slot( - "lora", name, adapter_model, trainable=False, alpha=alpha - ) - self._validate_dynamic_slot_consistency("lora", name, loaded) - return loaded - @overload def forward_micro_batches( self, @@ -859,6 +951,7 @@ def forward_micro_batches( ) -> Iterator[MicroBatch[ForwardInputs, ForwardOutputs]]: items = [_materialize(item) for item in inputs] requests = list(_flatten(items)) + self._validate_distributed_checkpoint_selections(requests) for _, indices in self._group_active_request_indices(requests): for index in indices: self._forward_item(requests[index]) @@ -942,7 +1035,9 @@ def dp_rank_forward( def dp_rank_forward(self, inputs: ForwardInputs) -> ForwardOutputs: materialized = _materialize(inputs) - plan = self._plan_flat_forward(list(_flatten(materialized))) + requests = list(_flatten(materialized)) + self._validate_distributed_checkpoint_selections(requests) + plan = self._plan_flat_forward(requests) check = self._memory_check(plan) if not check.fits: self._raise_memory_error( @@ -987,35 +1082,38 @@ def optim_step( scale_grads=scale_grads, ) - def _load_slot( + def _load_checkpoint_slot( self, - kind: Literal["checkpoint", "lora"], name: str, adapter_model: Mapping[str, torch.Tensor], *, - trainable: bool, - alpha: float | None, + alpha: float, + _prepared: bool = False, + _localized: bool = False, ) -> int: - if self._slot_stack: - raise RuntimeError("Cannot load a LoRA/checkpoint while a slot is pushed") - adapter_model = self._prepare_adapter_model(kind, name, adapter_model) - from art.megatron.lora import LORA_ALPHA, load_lora_slot_into_model + adapter_model = ( + dict(adapter_model) + if _prepared + else self._prepare_adapter_model(name, adapter_model) + ) + from art.megatron.lora import load_lora_slot_into_model - ref = self._slot_ref(kind, name) + ref = self._slot_ref(name) self._guard_slot_can_load(ref) return load_lora_slot_into_model( self.runtime.model, ref, adapter_model, - alpha=LORA_ALPHA if alpha is None else alpha, - requires_grad=trainable, + alpha=alpha, + localized=_localized, ) def _prepare_adapter_model( self, - kind: Literal["checkpoint", "lora"], name: str, adapter_model: Mapping[str, torch.Tensor], + *, + canonicalized: bool = False, ) -> dict[str, torch.Tensor]: templates = self._local_lora_adapter_templates() keys = set(adapter_model) @@ -1028,22 +1126,24 @@ def _prepare_adapter_model( preview = ", ".join(repr(key) for key in unknown[:8]) more = "" if len(unknown) <= 8 else f", ... +{len(unknown) - 8} more" raise ValueError( - f"Adapter for {kind} slot {name!r} contains keys that do not match " - f"installed LoRA wrapper sites: {preview}{more}. Configure the " - "Megatron runtime with matching LoRA target modules before loading." + f"Checkpoint {name!r} contains keys that do not match installed " + f"LoRA wrapper sites: {preview}{more}. Configure the runtime with " + "matching LoRA target modules before loading." ) local_state = { key: tensor for key, tensor in adapter_model.items() if key in templates } adapter_model = ( - self.runtime.model_support_handler.canonicalize_loaded_lora_state( + local_state + if canonicalized + else self.runtime.model_support_handler.canonicalize_loaded_lora_state( local_state, self.runtime.model ) ) if set(adapter_model) != set(local_state): raise TrainerRankSlotStateError( "Model-specific LoRA canonicalization changed the adapter key set " - f"for {kind} slot {name!r}." + f"for checkpoint {name!r}." ) return { key: tensor.to( @@ -1054,6 +1154,65 @@ def _prepare_adapter_model( for key, tensor in adapter_model.items() } + def _preflight_adapter( + self, + name: str, + adapter_model: Mapping[str, torch.Tensor], + config: _AdapterConfig, + *, + localized: bool = False, + canonicalized: bool = False, + ) -> dict[str, torch.Tensor]: + prepared = self._prepare_adapter_model( + name, adapter_model, canonicalized=canonicalized + ) + from art.megatron.lora import LoRA + + rank = int(config["r"]) + loaded = 0 + for chunk in self.runtime.model: + for module in chunk.modules(): + if not isinstance(module, LoRA): + continue + weights = module._adapter_weights(prepared, require=False) + if weights is None: + continue + a_t = ( + weights[0] + if localized + else module._localized_weight(weights[0], into=module.A_T) + ) + b_t = ( + weights[1] + if localized + else module._localized_weight(weights[1], into=module.B_T) + ) + if ( + a_t.ndim != module.A_T.ndim + or tuple(a_t.shape[:-1]) != tuple(module.A_T.shape[:-1]) + or int(a_t.shape[-1]) != rank + ): + raise TrainerRankSlotStateError( + f"Checkpoint {name!r} has incompatible LoRA-A shape " + f"{tuple(a_t.shape)} for {module.adapter_model_prefix}" + ) + if ( + b_t.ndim != module.B_T.ndim + or tuple(b_t.shape[:-2]) != tuple(module.B_T.shape[:-2]) + or int(b_t.shape[-2]) != rank + or int(b_t.shape[-1]) != int(module.B_T.shape[-1]) + ): + raise TrainerRankSlotStateError( + f"Checkpoint {name!r} has incompatible LoRA-B shape " + f"{tuple(b_t.shape)} for {module.adapter_model_prefix}" + ) + loaded += 1 + if loaded == 0: + raise TrainerRankSlotStateError( + f"Checkpoint {name!r} loaded no adapter sites" + ) + return prepared + def _local_lora_adapter_templates(self) -> dict[str, torch.Tensor]: templates: dict[str, torch.Tensor] = {} for chunk in self.runtime.model: @@ -1066,79 +1225,185 @@ def _local_lora_adapter_templates(self) -> dict[str, torch.Tensor]: ("lora_B", "B_T"), ): parameter = getattr(module, parameter_name, None) - if not isinstance(parameter, torch.Tensor): + if isinstance(parameter, torch.Tensor): + templates.update( + (str(key), parameter) + for key in expected_weight_keys(suffix) + ) + return templates + + def _local_adapter_keys(self) -> tuple[str, ...]: + return tuple(self._local_lora_adapter_templates()) + + def _local_parameter_key_groups(self, name: str) -> tuple[tuple[str, ...], ...]: + from art.megatron.lora import LoRA + + ref = self._slot_ref(name) + return tuple( + tuple(str(key) for key in module._expected_weight_keys(suffix)) + for chunk in self.runtime.model + for module in chunk.modules() + if isinstance(module, LoRA) and module._slot(ref) is not None + for suffix in ("lora_A", "lora_B") + ) + + def _localize_adapter_tensors( + self, + tensors: Mapping[str, torch.Tensor], + name: str, + *, + localized: bool = False, + ) -> tuple[torch.Tensor, ...]: + from art.megatron.lora import LoRA + + key_groups = self._local_parameter_key_groups(name) + if missing := sorted( + key for group in key_groups for key in group if key not in tensors + ): + raise TrainerRankSlotStateError( + f"Canonical optimizer is missing local tensors: {missing[:8]}" + ) + values: list[torch.Tensor] = [] + ref = self._slot_ref(name) + for chunk in self.runtime.model: + for module in chunk.modules(): + if not isinstance(module, LoRA): + continue + slot = module._slot(ref) + if slot is None: + continue + for suffix, template in ( + ("lora_A", slot.A_T), + ("lora_B", slot.B_T), + ): + keys = module._expected_weight_keys(suffix) + if not all(key in tensors for key in keys): continue - templates.update( - (str(key), parameter) for key in expected_weight_keys(suffix) + weight = ( + torch.stack([tensors[key].T for key in keys]) + if module.num_local_experts > 1 + else tensors[keys[0]].T ) - return templates + value = ( + weight.contiguous() + if localized + else module._localized_weight(weight, into=template) + ) + if tuple(value.shape) != tuple(template.shape): + raise TrainerRankSlotStateError( + f"Canonical optimizer tensor for {keys[0]!r} has an " + f"incompatible local shape {tuple(value.shape)}; expected " + f"{tuple(template.shape)}" + ) + values.append(value) + if len(values) != len(key_groups): + raise TrainerRankSlotStateError( + "Canonical optimizer did not map every local checkpoint parameter" + ) + return tuple(values) def _set_default_slot(self, ref: "LoRASlotRef") -> None: if self._slot_stack: - raise RuntimeError("Cannot set a LoRA/checkpoint while a slot is pushed") + raise RuntimeError("Cannot select a checkpoint while one is pushed") self._default_slot_ref = ref @staticmethod - def _slot_ref( - kind: Literal["checkpoint", "lora"], name: str | None - ) -> "LoRASlotRef": + def _slot_ref(name: str | None) -> "LoRASlotRef": try: from art.megatron.lora import LoRASlotRef except ModuleNotFoundError as exc: if exc.name is None or not exc.name.startswith("megatron"): raise + return cast("LoRASlotRef", _LocalLoRASlotRef(name=name)) + return LoRASlotRef(name=name) - return cast("LoRASlotRef", _LocalLoRASlotRef(kind=kind, name=name)) + def _iter_slot_parameters(self, ref: "LoRASlotRef") -> Iterator[torch.nn.Parameter]: + from art.megatron.lora import iter_lora_slot_parameters - return LoRASlotRef(kind=kind, name=name) + return iter_lora_slot_parameters(self.runtime.model, ref) - def _validate_dynamic_slot_consistency( + def _validate_checkpoint_consistency( self, - kind: Literal["checkpoint", "lora"], name: str, loaded_sites: int, + expected_keys: set[str], ) -> tuple[torch.nn.Parameter, ...]: - from art.megatron.lora import iter_lora_slot_parameters + ref = self._slot_ref(name) + params = tuple(self._iter_slot_parameters(ref)) + local_keys = { + key for group in self._local_parameter_key_groups(name) for key in group + } + from art.megatron.lora import LoRA - ref = self._slot_ref(kind, name) - params = tuple(iter_lora_slot_parameters(self.runtime.model, ref)) + actual_sites = sum( + module._slot(ref) is not None + for chunk in self.runtime.model + for module in chunk.modules() + if isinstance(module, LoRA) + ) + local = (loaded_sites, actual_sites, local_keys, len(params)) if not (dist.is_available() and dist.is_initialized()): - return params + ranks = (local,) + else: + gathered: list[tuple[int, int, set[str], int] | None] = [ + None + ] * dist.get_world_size() + dist.all_gather_object(gathered, local) + ranks = tuple(state for state in gathered if state is not None) + if any(loaded != actual for loaded, actual, _keys, _params in ranks): + raise RuntimeError( + f"Checkpoint {name!r} loaded an inconsistent number of sites: " + f"{[(loaded, actual) for loaded, actual, _keys, _params in ranks]}" + ) + covered = set().union(*(keys for _loaded, _actual, keys, _params in ranks)) + if covered != expected_keys: + raise RuntimeError( + f"Checkpoint {name!r} logical-key coverage differs: " + f"missing={sorted(expected_keys - covered)[:8]} " + f"extra={sorted(covered - expected_keys)[:8]}" + ) + if any(keys and count == 0 for _loaded, _actual, keys, count in ranks): + raise RuntimeError( + f"Checkpoint {name!r} did not create parameters on every owning rank" + ) + return params - signature = tuple( - ( - tuple(param.shape), - str(param.dtype), - bool(getattr(param, "allreduce", True)), - str(getattr(param, "grad_sync_domain", "tp_default")), - str(getattr(param, "grad_sync_op", "none")), + def _validate_distributed_checkpoint_selections( + self, requests: Sequence[AnyForwardInput] + ) -> None: + missing = tuple( + sorted( + { + name + for request in requests + if request.checkpoint is not Unset + if (name := cast(str | None, request.checkpoint)) is not None + if name not in self._checkpoint_slot_params_by_name + } ) - for param in params ) - local = (int(loaded_sites), signature) - gathered: list[tuple[int, object] | None] = [None] * dist.get_world_size() - dist.all_gather_object(gathered, local) - ranks = [state for state in gathered if state is not None] - if all(state == ranks[0] for state in ranks[1:]): - return params - raise RuntimeError( - f"Dynamic LoRA slot {kind}:{name} is not loaded consistently across " - "distributed ranks. This usually means a sharded/exported LoRA state " - "dict was passed directly to TrainerRank; gather or materialize the " - "full adapter state before loading a dynamic slot. " - f"Loaded-site counts by rank: {[state[0] for state in ranks]}." + if self._all_ranks_true(not missing): + return + location = f": {list(missing)}" if missing else " on another distributed rank" + raise TrainerRankSlotStateError( + f"Forward inputs select unloaded checkpoint slots{location}. " + "Call load_checkpoint(...) before selecting them." ) def _resolve_slot_ref(self, request: AnyForwardInput) -> "LoRASlotRef | None": if request.checkpoint is not Unset: - return self._slot_ref("checkpoint", cast(str | None, request.checkpoint)) - if request.lora is not Unset: - return self._slot_ref("lora", cast(str | None, request.lora)) + name = cast(str | None, request.checkpoint) + if name is not None and name not in self._checkpoint_slot_params_by_name: + raise TrainerRankSlotStateError( + f"Forward input selects unloaded checkpoint slot {name!r}. " + "Call load_checkpoint(...) before selecting it." + ) + return self._slot_ref(name) if self._slot_stack: return self._slot_stack[-1] if self._default_slot_ref is not None: return self._default_slot_ref - return self._slot_ref("checkpoint", None) + return self._slot_ref(None) def _selected_dynamic_checkpoints( self, @@ -1148,7 +1413,7 @@ def _selected_dynamic_checkpoints( if not loaded: raise TrainerRankSlotStateError( "TrainerRank.optim_step requires a loaded checkpoint slot. Call " - "load_checkpoint_slot(...) and run backward on outputs produced by " + "load_checkpoint(...) and run backward on outputs produced by " "that slot before stepping." ) requested = ( @@ -1251,7 +1516,10 @@ def _dynamic_optim_step( ): model.copy_(master) model.grad = None - self._prune_slot_graphs(self._slot_ref("checkpoint", name)) + self._prune_slot_graphs(self._slot_ref(name)) + self._checkpoint_revisions[name] = ( + self._checkpoint_revisions.get(name, 0) + 1 + ) return { "learning_rate": float(params.learning_rate), "grad_norm": float(grad_norm), @@ -1291,6 +1559,13 @@ def _new_dynamic_optimizer( f"Optimizer state for checkpoint slot {name!r} has " f"{len(sources)} master parameters; expected {len(model_params)}." ) + if any( + tuple(source.shape) != tuple(model.shape) + for source, model in zip(sources, model_params, strict=True) + ): + raise TrainerRankSlotStateError( + f"Optimizer master parameter shape does not match checkpoint {name!r}" + ) masters = tuple( torch.nn.Parameter( source.detach().to(device=model.device, dtype=torch.float32).clone() @@ -1309,134 +1584,54 @@ def _new_dynamic_optimizer( ) return _DynamicOptimizer(optimizer, masters) - def _restore_dynamic_optimizer( + def _restore_canonical_optimizer( self, name: str, - state: TrainerRankOptimizerState, + state: LocalOptimizerState, ) -> _DynamicOptimizer: - if state.get("format_version") != 1: - raise TrainerRankSlotStateError( - f"Unsupported optimizer state format for checkpoint slot {name!r}." - ) - if state.get("layout") != self._dynamic_optimizer_layout(name): - raise TrainerRankSlotStateError( - f"Optimizer state for checkpoint slot {name!r} was saved for a " - "different topology or parameter layout. Save and restore one " - "optimizer shard per TrainerRank with matching TP/EP/ETP ranks." - ) - master_params = state.get("master_params") - optimizer_state = state.get("optimizer") - if not isinstance(master_params, Sequence) or not isinstance( - optimizer_state, Mapping - ): - raise TrainerRankSlotStateError( - f"Optimizer state for checkpoint slot {name!r} is incomplete." - ) dynamic = self._new_dynamic_optimizer( name, - AdamParams(learning_rate=0.0), - master_params=cast(Sequence[torch.Tensor], master_params), + AdamParams( + learning_rate=state.config.learning_rate, + beta1=state.config.beta1, + beta2=state.config.beta2, + weight_decay=state.config.weight_decay, + ), + master_params=state.masters, ) - try: - dynamic.optimizer.load_state_dict( - {str(key): value for key, value in optimizer_state.items()} - ) - except ValueError as exc: - raise TrainerRankSlotStateError( - f"Optimizer state for checkpoint slot {name!r} does not match the " - "loaded slot parameter groups." - ) from exc - for param in dynamic.master_params: - for state_name, value in dynamic.optimizer.state.get(param, {}).items(): - if ( - isinstance(value, torch.Tensor) - and int(value.ndim) > 0 - and tuple(value.shape) != tuple(param.shape) - ): - raise TrainerRankSlotStateError( - f"Optimizer state {state_name!r} for checkpoint slot " - f"{name!r} has shape {tuple(value.shape)}, but the loaded " - f"slot parameter has shape {tuple(param.shape)}." - ) - self._zero_dynamic_optimizer_padding(name, dynamic) - return dynamic - - def _zero_dynamic_optimizer_padding( - self, - name: str, - dynamic: _DynamicOptimizer, - ) -> None: - masks = self._dynamic_optimizer_padding_masks(name) - with torch.no_grad(): - for param, mask in zip(dynamic.master_params, masks, strict=True): - param.masked_fill_(mask, 0) - for value in dynamic.optimizer.state.get(param, {}).values(): - if isinstance(value, torch.Tensor) and value.shape == param.shape: - value.masked_fill_(mask, 0) - - def _dynamic_optimizer_padding_masks(self, name: str) -> tuple[torch.Tensor, ...]: - params = self._checkpoint_slot_params_by_name[name] - masks = tuple(torch.zeros_like(param, dtype=torch.bool) for param in params) - param_indices = {id(param): index for index, param in enumerate(params)} - exported: dict[str, torch.Tensor] = {} - owners: dict[str, tuple[int, int | None]] = {} - mapped_indices: set[int] = set() - ref = self._slot_ref("checkpoint", name) - - for chunk in self.runtime.model: - for module in chunk.modules(): - lora_params = getattr(module, "_lora_params", None) - expected_keys = getattr(module, "_expected_weight_keys", None) - if not callable(lora_params) or not callable(expected_keys): - continue - for suffix, param in lora_params(ref): - index = param_indices.get(id(param)) - if index is None: - continue - mapped_indices.add(index) - keys = expected_keys(str(suffix).removesuffix(".weight")) - if int(param.ndim) == 3: - if len(keys) != int(param.shape[0]): - raise TrainerRankSlotStateError( - f"Cannot map optimizer padding for checkpoint slot " - f"{name!r}: {len(keys)} adapter keys describe " - f"{int(param.shape[0])} local experts." - ) - for expert, key in enumerate(keys): - exported[str(key)] = torch.ones_like(param[expert].T) - owners[str(key)] = (index, expert) - else: - if len(keys) != 1: - raise TrainerRankSlotStateError( - f"Cannot map optimizer padding for checkpoint slot " - f"{name!r}: expected one adapter key, got {len(keys)}." - ) - key = str(keys[0]) - exported[key] = torch.ones_like(param.T) - owners[key] = (index, None) - - if mapped_indices and ( - missing := sorted(set(range(len(params))) - mapped_indices) + if not ( + len(dynamic.master_params) + == len(state.exp_avgs) + == len(state.exp_avg_sqs) + == len(state.steps) ): raise TrainerRankSlotStateError( - f"Cannot map optimizer padding for checkpoint slot {name!r}: " - f"parameter indices {missing} do not belong to installed LoRA sites." + f"Canonical optimizer state for {name!r} has inconsistent lengths" ) - - canonical = self.runtime.model_support_handler.canonicalize_loaded_lora_state( - exported, self.runtime.model - ) - for key, value in canonical.items(): - owner = owners.get(key) - if owner is None or not isinstance(value, torch.Tensor): - continue - index, expert = owner - mask = value.T == 0 - if expert is None: - masks[index].copy_(mask) - else: - masks[index][expert].copy_(mask) - return masks + dynamic.optimizer.param_groups[0]["eps"] = state.config.eps + for master, exp_avg, exp_avg_sq, step in zip( + dynamic.master_params, + state.exp_avgs, + state.exp_avg_sqs, + state.steps, + strict=True, + ): + if tuple(exp_avg.shape) != tuple(master.shape) or tuple( + exp_avg_sq.shape + ) != tuple(master.shape): + raise TrainerRankSlotStateError( + f"Canonical optimizer moment shape does not match {name!r}" + ) + dynamic.optimizer.state[master] = { + "step": torch.tensor(step, dtype=torch.float32), + "exp_avg": exp_avg.to( + device=master.device, dtype=torch.float32 + ).clone(), + "exp_avg_sq": exp_avg_sq.to( + device=master.device, dtype=torch.float32 + ).clone(), + } + return dynamic def _reduce_dynamic_grads( self, @@ -1489,39 +1684,6 @@ def add( coalesced_all_reduce(bucket_grads, group=group, op=op) return grads - def _dynamic_optimizer_layout(self, name: str) -> TrainerRankOptimizerLayout: - parameters = cast( - tuple[ - tuple[ - tuple[int, ...], - str, - str, - bool, - int | None, - str, - tuple[int, ...], - ], - ..., - ], - tuple( - ( - tuple(param.shape), - str(param.dtype), - str(getattr(param, "lora_shard_domain", "tp")), - bool(getattr(param, "lora_tp_sharded", False)), - getattr(param, "lora_tp_shard_dim", None), - str(getattr(param, "lora_tp_shard_strategy", "uniform")), - tuple(getattr(param, "lora_tp_component_sizes", ())), - ) - for param in self._checkpoint_slot_params_by_name[name] - ), - ) - layout: TrainerRankOptimizerLayout = { - "parallel": _parallel_optimizer_coordinates(), - "parameters": parameters, - } - return layout - def _select_next_micro_batch( self, items: Sequence[ForwardInputsT], @@ -1928,7 +2090,7 @@ def _guard_slot_can_load(self, ref: "LoRASlotRef") -> None: if not self._has_live_slot_graph(ref): return raise TrainerRankSlotStateError( - f"Cannot load {ref.kind} slot {ref.name!r} while outputs from an " + f"Cannot load checkpoint {ref.name!r} while outputs from an " "earlier forward using that slot still have a live backward graph. " "Activation checkpoint recompute resolves slots by name, so replacing " "the slot before backward can compute gradients with different LoRA " @@ -1938,7 +2100,7 @@ def _guard_slot_can_load(self, ref: "LoRASlotRef") -> None: ) def _guard_checkpoint_can_step(self, name: str) -> None: - ref = self._slot_ref("checkpoint", name) + ref = self._slot_ref(name) if not self._has_live_slot_graph(ref): return raise TrainerRankSlotStateError( @@ -2907,36 +3069,6 @@ def _include_in_distributed_grad_norm(param: torch.nn.Parameter) -> bool: return shard_group is None or shard_group.size() <= 1 or shard_group.rank() == 0 -def _parallel_optimizer_coordinates() -> tuple[int, int, int, int, int, int, int, int]: - if not (dist.is_available() and dist.is_initialized()): - return (1, 0, 1, 0, 1, 0, 1, 0) - from megatron.core import parallel_state as ps - - expert_tp_group = ps.get_expert_tensor_parallel_group(check_initialized=False) - return ( - int(ps.get_tensor_model_parallel_world_size()), - int(ps.get_tensor_model_parallel_rank()), - int(ps.get_expert_model_parallel_world_size()), - int(ps.get_expert_model_parallel_rank()), - 1 if expert_tp_group is None else int(expert_tp_group.size()), - 0 if expert_tp_group is None else int(expert_tp_group.rank()), - int(ps.get_pipeline_model_parallel_world_size()), - int(ps.get_pipeline_model_parallel_rank()), - ) - - -def _state_to_cpu(value: object) -> object: - if isinstance(value, torch.Tensor): - return value.detach().cpu().clone() - if isinstance(value, Mapping): - return {key: _state_to_cpu(item) for key, item in value.items()} - if isinstance(value, tuple): - return tuple(_state_to_cpu(item) for item in value) - if isinstance(value, list): - return [_state_to_cpu(item) for item in value] - return value - - def _vocab_parallel_target_logprobs( local_logits: torch.Tensor, labels: torch.Tensor, diff --git a/tests/integration/megatron/lora/test_dynamic_lora_slots.py b/tests/integration/megatron/lora/test_dynamic_lora_slots.py index 2a3ee6da0..4d128b235 100644 --- a/tests/integration/megatron/lora/test_dynamic_lora_slots.py +++ b/tests/integration/megatron/lora/test_dynamic_lora_slots.py @@ -1,5 +1,6 @@ from __future__ import annotations +import asyncio from contextlib import contextmanager import os from pathlib import Path @@ -19,8 +20,11 @@ from art.megatron.lora import LoRA, LoRASlotRef, use_lora_slot # noqa: E402 from art.trainer_rank import ( # noqa: E402 AdamParams, + ForwardInput, TrainerRank, + TrainerRankSlotStateError, ) +from art.trainer_rank._checkpoint import AdamWRecord, LocalOptimizerState # noqa: E402 from art.trainer_rank._impl import ( # noqa: E402 _distributed_grad_norm, _vocab_parallel_log_z, @@ -42,19 +46,15 @@ def test_dynamic_lora_slots_capture_recompute_context_and_step_independently() - dtype=torch.float32, device=device, ) - ref_a = LoRASlotRef("checkpoint", "A") - ref_b = LoRASlotRef("checkpoint", "B") - lora.load_lora_slot( - ref_a, _adapter("dense", rank=1, seed=1), requires_grad=True - ) - lora.load_lora_slot( - ref_b, _adapter("dense", rank=4, seed=2), requires_grad=True - ) + ref_a = LoRASlotRef("A") + ref_b = LoRASlotRef("B") + lora.load_lora_slot(ref_a, _adapter("dense", rank=1, seed=1)) + lora.load_lora_slot(ref_b, _adapter("dense", rank=4, seed=2)) x = torch.randn(7, 4, device=device) - with use_lora_slot(LoRASlotRef("checkpoint", None)): + with use_lora_slot(LoRASlotRef(None)): assert torch.equal(lora(x), torch.zeros(7, 5, device=device)) - with use_lora_slot(LoRASlotRef("lora", "missing")): + with use_lora_slot(LoRASlotRef("missing")): assert torch.equal(lora(x), torch.zeros(7, 5, device=device)) slot_a = lora._slot(ref_a) @@ -74,20 +74,23 @@ def test_dynamic_lora_slots_capture_recompute_context_and_step_independently() - key: value.cpu().double() for key, value in _adapter("dense", rank=3, seed=7).items() } - trainer.load_checkpoint_slot("CPU", cpu_adapter) - cpu_slot = lora._slot(LoRASlotRef("checkpoint", "CPU")) + _install_checkpoint(trainer, "CPU", cpu_adapter) + cpu_slot = lora._slot(LoRASlotRef("CPU")) assert cpu_slot is not None assert cpu_slot.A_T.device == lora.A_T.device assert cpu_slot.A_T.dtype == lora.A_T.dtype - with use_lora_slot(LoRASlotRef("checkpoint", "CPU")): + with use_lora_slot(LoRASlotRef("CPU")): assert lora(x).is_cuda - with trainer.push_checkpoint("A"): - assert trainer._slot_stack[-1] == ref_a - with trainer.push_lora(None): - assert trainer._slot_stack[-1].name is None - assert trainer._slot_stack[-1] == ref_a - assert trainer._slot_stack == [] + async def assert_checkpoint_stack() -> None: + async with trainer.push_checkpoint("A"): + assert trainer._slot_stack[-1] == ref_a + async with trainer.push_checkpoint(None): + assert trainer._slot_stack[-1].name is None + assert trainer._slot_stack[-1] == ref_a + assert trainer._slot_stack == [] + + asyncio.run(assert_checkpoint_stack()) from megatron.core.tensor_parallel.random import ( checkpoint as megatron_checkpoint, @@ -102,6 +105,47 @@ def test_dynamic_lora_slots_capture_recompute_context_and_step_independently() - _assert_reload_replaces_slot_optimizer(ref_a, lora, trainer) +def _asymmetric_checkpoint_selection_worker( + rank: int, world_size: int, init_method: str +) -> None: + init_process_group( + "gloo", rank=rank, world_size=world_size, init_method=init_method + ) + try: + trainer = TrainerRank.__new__(TrainerRank) + trainer.device = torch.device("cpu") + trainer._checkpoint_slot_params_by_name = {"student": ()} + request = ForwardInput( + input_tokens=torch.tensor([1]), + target_tokens=torch.tensor([1]), + checkpoint="typo" if rank == 0 else None, + ) + + for operation in ("dp_rank_forward", "forward_micro_batches"): + with pytest.raises(TrainerRankSlotStateError, match="unloaded checkpoint"): + if operation == "dp_rank_forward": + trainer.dp_rank_forward([request]) + else: + list(trainer.forward_micro_batches([request])) + + completed = torch.tensor(1) + torch.distributed.all_reduce(completed) + assert int(completed.item()) == world_size + finally: + destroy_process_group() + + +def test_forwards_reject_asymmetric_unloaded_checkpoint( + tmp_path: Path, +) -> None: + mp.spawn( + _asymmetric_checkpoint_selection_worker, + args=(2, f"file://{tmp_path / 'checkpoint-selection-init'}"), + nprocs=2, + join=True, + ) + + @pytest.mark.parametrize("tp_size", (2, 4)) def test_trainer_rank_tp_head_backward_matches_unsharded_oracle( tp_size: int, @@ -188,7 +232,7 @@ def _tp_head_backward_worker(rank: int, world: int, init_method: str) -> None: from megatron.core import tensor_parallel - local_hidden = torch.randn(2, 1, 3, device=device, requires_grad=True) + local_hidden = torch.randn(2, 1, 3, device=device) gathered_hidden = tensor_parallel.gather_from_sequence_parallel_region( local_hidden, tensor_parallel_output_grad=False, @@ -252,10 +296,10 @@ def _assert_replica_grad_reduction( def _assert_distributed_optimizer_restore(device: torch.device) -> None: - ref = LoRASlotRef("checkpoint", "A") + ref = LoRASlotRef("A") adapter = _adapter("dense", rank=2, seed=11) lora = LoRA("dense", 4, 5, 2, 32, torch.float32, device) - lora.load_lora_slot(ref, adapter, requires_grad=True) + lora.load_lora_slot(ref, adapter) trainer = _trainer_for(lora, device) params = AdamParams(learning_rate=1e-3, weight_decay=0.0, grad_clip_norm=0.0) x = torch.randn(3, 4, device=device) @@ -263,8 +307,7 @@ def _assert_distributed_optimizer_restore(device: torch.device) -> None: with use_lora_slot(ref): lora(x).sum().backward() trainer.optim_step(params=params, checkpoints=["A"]) - state = trainer.checkpoint_slot_optimizer_state("A") - assert state is not None + state = _optimizer_state(trainer, "A") slot = lora._slot(ref) assert slot is not None adapter = { @@ -277,7 +320,10 @@ def _assert_distributed_optimizer_restore(device: torch.device) -> None: restored_lora = LoRA("dense", 4, 5, 2, 32, torch.float32, device) restored = _trainer_for(restored_lora, device) - restored.load_checkpoint_slot("A", adapter, optimizer_state=state) + _install_checkpoint(restored, "A", adapter) + restored._dynamic_optimizers["A"] = restored._restore_canonical_optimizer( + "A", state + ) with use_lora_slot(ref): restored_lora(x).sum().backward() restored.optim_step(params=params, checkpoints=["A"]) @@ -287,76 +333,6 @@ def _assert_distributed_optimizer_restore(device: torch.device) -> None: torch.testing.assert_close(actual, expected, atol=0, rtol=0) -@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required.") -def test_restored_dynamic_optimizer_canonicalizes_internal_padding() -> None: - with _single_rank_model_parallel(): - for num_local_experts in (1, 2): - _assert_restored_dynamic_optimizer_canonicalizes_internal_padding( - num_local_experts - ) - - -def _assert_restored_dynamic_optimizer_canonicalizes_internal_padding( - num_local_experts: int, -) -> None: - device = torch.device("cuda") - ref = LoRASlotRef("checkpoint", "A") - prefix = "dense" if num_local_experts == 1 else "experts.{expert}" - adapter = { - key: value - for expert in range(num_local_experts) - for key, value in _adapter( - prefix.format(expert=expert), rank=2, seed=17 + expert - ).items() - } - lora = LoRA( - prefix, - 4, - 5, - 2, - 32, - torch.float32, - device, - num_local_experts=num_local_experts, - ) - lora.load_lora_slot(ref, adapter, requires_grad=True) - trainer = _trainer_for(lora, device) - - def canonicalize( - state: dict[str, torch.Tensor], _model: object - ) -> dict[str, torch.Tensor]: - result = {key: value.clone() for key, value in state.items()} - for value in result.values(): - value[..., -1] = 0 - return result - - trainer.runtime.model_support_handler.canonicalize_loaded_lora_state = canonicalize - for param in trainer._checkpoint_slot_params_by_name["A"]: - param.grad = torch.ones_like(param) - trainer.optim_step( - params=AdamParams(learning_rate=1e-3, weight_decay=0.1, grad_clip_norm=0.0), - checkpoints=["A"], - ) - state = trainer.checkpoint_slot_optimizer_state("A") - assert state is not None - masks = trainer._dynamic_optimizer_padding_masks("A") - masters = cast(tuple[torch.Tensor, ...], state["master_params"]) - optimizer = state["optimizer"] - optimizer_states = cast(dict[int, dict[str, object]], optimizer["state"]) - for index, (master, mask) in enumerate(zip(masters, masks, strict=True)): - master.masked_fill_(mask.cpu(), 5) - for value in optimizer_states[index].values(): - if isinstance(value, torch.Tensor) and value.shape == master.shape: - value.masked_fill_(mask.cpu(), 5) - - restored = trainer._restore_dynamic_optimizer("A", state) - for master, mask in zip(restored.master_params, masks, strict=True): - assert torch.count_nonzero(master[mask]) == 0 - for value in restored.optimizer.state[master].values(): - if isinstance(value, torch.Tensor) and value.shape == master.shape: - assert torch.count_nonzero(value[mask]) == 0 - - def _local_shard(full: torch.Tensor, rank: int, size: int) -> torch.Tensor: return full[:, rank * size : (rank + 1) * size].clone().requires_grad_() @@ -448,7 +424,7 @@ def _assert_reload_replaces_slot_optimizer( old_params = trainer._checkpoint_slot_params_by_name[ref.name] assert ref.name in trainer._dynamic_optimizers - trainer.load_checkpoint_slot(ref.name, _adapter("dense", rank=3, seed=9)) + _install_checkpoint(trainer, ref.name, _adapter("dense", rank=3, seed=9)) new_params = trainer._checkpoint_slot_params_by_name[ref.name] assert ref.name not in trainer._dynamic_optimizers @@ -459,6 +435,47 @@ def _assert_reload_replaces_slot_optimizer( assert slot.rank == 3 +def _install_checkpoint( + trainer: TrainerRank, name: str, adapter: dict[str, torch.Tensor] +) -> int: + loaded = trainer._load_checkpoint_slot(name, adapter, alpha=32.0) + trainer._checkpoint_slot_params_by_name[name] = tuple( + trainer._iter_slot_parameters(trainer._slot_ref(name)) + ) + trainer._dynamic_optimizers.pop(name, None) + trainer._checkpoint_revisions[name] = ( + trainer._checkpoint_revisions.get(name, -1) + 1 + ) + return loaded + + +def _optimizer_state(trainer: TrainerRank, name: str) -> LocalOptimizerState: + dynamic = trainer._dynamic_optimizers[name] + states = [ + cast(dict[str, torch.Tensor], dynamic.optimizer.state[master]) + for master in dynamic.master_params + ] + group = dynamic.optimizer.param_groups[0] + beta1, beta2 = cast(tuple[float, float], group["betas"]) + return LocalOptimizerState( + masters=tuple( + master.detach().cpu().clone() for master in dynamic.master_params + ), + exp_avgs=tuple(state["exp_avg"].detach().cpu().clone() for state in states), + exp_avg_sqs=tuple( + state["exp_avg_sq"].detach().cpu().clone() for state in states + ), + steps=tuple(float(state["step"].item()) for state in states), + config=AdamWRecord( + learning_rate=float(group["lr"]), + beta1=beta1, + beta2=beta2, + eps=float(group["eps"]), + weight_decay=float(group["weight_decay"]), + ), + ) + + def _trainer_for(lora: LoRA, device: torch.device) -> TrainerRank: trainer = TrainerRank.__new__(TrainerRank) trainer.runtime = SimpleNamespace( @@ -475,9 +492,13 @@ def _trainer_for(lora: LoRA, device: torch.device) -> TrainerRank: trainer._default_slot_ref = None trainer._dynamic_optimizers = {} trainer._checkpoint_slot_params_by_name = { - "A": tuple(lora.lora_slot_params(LoRASlotRef("checkpoint", "A"))), - "B": tuple(lora.lora_slot_params(LoRASlotRef("checkpoint", "B"))), + "A": tuple(lora.lora_slot_params(LoRASlotRef("A"))), + "B": tuple(lora.lora_slot_params(LoRASlotRef("B"))), } + trainer._checkpoint_slot_adapter_configs = {} + trainer._checkpoint_revisions = {"A": 0, "B": 0} + trainer._checkpoint_sources = {} + trainer._pending_slot_graphs = {} return trainer diff --git a/tests/integration/megatron/lora/test_lora_disk_codecs.py b/tests/integration/megatron/lora/test_lora_disk_codecs.py index 074c297a6..fdab49760 100644 --- a/tests/integration/megatron/lora/test_lora_disk_codecs.py +++ b/tests/integration/megatron/lora/test_lora_disk_codecs.py @@ -1,3 +1,5 @@ +import asyncio +from collections.abc import Sequence import importlib.util import json import os @@ -5,17 +7,26 @@ import shutil import subprocess import sys +import threading +import time from types import SimpleNamespace -from typing import Any, cast +from typing import Any, Literal, cast import pytest from safetensors.torch import load_file, save_file import torch +import torch.multiprocessing as mp +from torch.multiprocessing.spawn import ProcessRaisedException pytest.importorskip("megatron.bridge.models.gpt_provider") from art.megatron import lora as lora_module -from art.megatron.lora import LoRA, LoRAParallelSpec, LoRAPublishPlanner, LoRASlotRef +from art.megatron.lora import ( + LoRA, + LoRAParallelSpec, + LoraShardManifest, + LoRASlotRef, +) from art.megatron.model_support.handlers import ( DEFAULT_DENSE_HANDLER, GPT_OSS_MOE_HANDLER, @@ -26,18 +37,31 @@ from art.megatron.model_support.handlers.gemma4 import GEMMA4_MOE_HANDLER from art.megatron.model_support.lora_disk import ( ART_LORA_FORMAT_CONFIG_KEY, + ART_LORA_FORMAT_MEGATRON, ART_LORA_FORMAT_VLLM, load_lora_tensors_for_megatron, normalize_lora_checkpoint_to_vllm, + save_adapter_config, save_vllm_lora_tensors, ) +from art.megatron.model_support.spec import ModelSupportHandler from art.megatron.weights import lora_publish from art.megatron.weights.lora_publish import ( LoraShardMeta, merge_sharded_adapter_entries, save_vllm_lora_from_model, ) -from art.trainer_rank import TrainerRank +from art.trainer_rank import AdamParams, TrainerRank +from art.trainer_rank._checkpoint import ( + _PreparedSave, + materialize_lora, + prepare_checkpoint, + validate_checkpoint, +) +from art.trainer_rank._checkpoint import ( + load_checkpoint as load_trainer_checkpoint, +) +from art.trainer_rank._impl import _AdapterConfig from art.utils.convert_moe_lora import convert_checkpoint_if_needed REPO_ROOT = Path(__file__).parents[4] @@ -126,6 +150,10 @@ def _config(base_model: str, rank: int = 2, alpha: int = 4) -> dict: } +def _manifest(**values: object) -> LoraShardManifest: + return cast(LoraShardManifest, values) + + def _qwen35_config(base_model: str, rank: int = 2, alpha: int = 4) -> dict: config = _config(base_model, rank=rank, alpha=alpha) config.update( @@ -156,10 +184,10 @@ def _save_adapter(path: Path, tensors: dict[str, torch.Tensor], config: dict) -> def _old_merge_shard_files_to_vllm( lora_path: Path, *, - handler, + handler: ModelSupportHandler, adapter_config: dict, ) -> None: - entries_by_key: dict[str, list[tuple[dict, torch.Tensor]]] = {} + entries_by_key: dict[str, list[tuple[LoraShardManifest, torch.Tensor]]] = {} shard_paths = sorted(lora_path.glob("adapter_model-*-of-*.safetensors")) manifest_paths = sorted(lora_path.glob("adapter_manifest-*-of-*.json")) for shard_path in shard_paths: @@ -172,7 +200,9 @@ def _old_merge_shard_files_to_vllm( shard_tensors = load_file(shard_path) assert set(shard_tensors) == set(manifest) for key, tensor in shard_tensors.items(): - entries_by_key.setdefault(key, []).append((manifest[key], tensor)) + entries_by_key.setdefault(key, []).append( + (cast(LoraShardManifest, manifest[key]), tensor) + ) merged = merge_sharded_adapter_entries(entries_by_key) vllm_tensors, adapter_config = handler.to_vllm_lora_tensors( @@ -1185,17 +1215,17 @@ def test_qwen35_megatron_shards_merge_to_vllm_checkpoint_and_roundtrip( + 30, } - def unsharded() -> dict: - return {"sharded": False, "shard_world_size": 1, "shard_rank": 0} + def unsharded() -> LoraShardManifest: + return _manifest(sharded=False, shard_world_size=1, shard_rank=0) - def sharded(rank_id: int, dim: int) -> dict: - return { - "sharded": True, - "shard_world_size": 2, - "shard_rank": rank_id, - "export_shard_dim": dim, - "export_shard_strategy": "uniform", - } + def sharded(rank_id: int, dim: int) -> LoraShardManifest: + return _manifest( + sharded=True, + shard_world_size=2, + shard_rank=rank_id, + export_shard_dim=dim, + export_shard_strategy="uniform", + ) shard0 = { f"{prefix}.gate_up_proj.lora_A.weight": full[ @@ -1270,7 +1300,7 @@ def test_lora_publish_keeps_same_key_shards_separate(): owner_rank=0, shape=tuple(shard0.shape), dtype_name="float32", - manifest={**manifest, "shard_rank": 0}, + manifest=_manifest(**manifest, shard_rank=0), block="base_model.model.model.layers.0", ), LoraShardMeta( @@ -1278,7 +1308,7 @@ def test_lora_publish_keeps_same_key_shards_separate(): owner_rank=1, shape=tuple(shard1.shape), dtype_name="float32", - manifest={**manifest, "shard_rank": 1}, + manifest=_manifest(**manifest, shard_rank=1), block="base_model.model.model.layers.0", ), ] @@ -1295,102 +1325,6 @@ def test_lora_publish_keeps_same_key_shards_separate(): assert torch.equal(merged[key], torch.tensor([[1.0], [2.0], [3.0], [4.0]])) -def test_lora_publish_planner_derives_metadata_from_lora_modules(): - prefix = "base_model.model.model.layers.0.self_attn.q_proj" - b_parallel_spec = LoRAParallelSpec(sharded=True, shard_dim=-1) - lora = LoRA( - adapter_model_prefix=prefix, - in_features=4, - out_features=6, - rank=2, - alpha=4, - dtype=torch.bfloat16, - device=torch.device("cpu"), - b_parallel_spec=b_parallel_spec, - ) - adapter_dtypes = { - f"{prefix}.lora_A.weight": torch.float32, - f"{prefix}.lora_B.weight": torch.float32, - } - - metadata = LoRAPublishPlanner([torch.nn.Sequential(lora)]).global_metadata( - adapter_dtypes - ) - by_key = {meta.key: meta for meta in metadata} - - a_meta = by_key[f"{prefix}.lora_A.weight"] - assert a_meta.shape == (2, 4) - assert a_meta.dtype_name == "float32" - assert a_meta.owner_rank == 0 - assert a_meta.manifest == { - "sharded": False, - "shard_world_size": 1, - "shard_rank": 0, - } - assert a_meta.block == "base_model.model.model.layers.0" - - b_meta = by_key[f"{prefix}.lora_B.weight"] - assert b_meta.shape == (6, 2) - assert b_meta.dtype_name == "float32" - assert b_meta.owner_rank == 0 - assert b_meta.manifest == { - "sharded": True, - "shard_world_size": 1, - "shard_rank": 0, - "export_shard_dim": 0, - "export_shard_strategy": "uniform", - } - - -def test_lora_publish_planner_maps_expert_owner_ranks(monkeypatch): - monkeypatch.setattr(lora_module, "_distributed_initialized", lambda: True) - monkeypatch.setattr( - lora_module, - "_get_shard_world_size", - lambda domain: 2 if domain == "expert_tp" else 1, - ) - monkeypatch.setattr( - lora_module.ps, - "get_expert_model_parallel_world_size", - lambda: 4, - ) - monkeypatch.setattr( - lora_module.ps, - "get_expert_tensor_and_model_parallel_group", - lambda check_initialized=False: "joint", - ) - monkeypatch.setattr( - lora_module.ps, - "get_expert_model_parallel_group", - lambda: "ep", - ) - monkeypatch.setattr( - lora_module.ps, - "get_expert_tensor_parallel_group", - lambda check_initialized=False: "etp", - ) - - row_major = {"joint": (0, 1, 2, 3, 4, 5, 6, 7), "ep": (0, 2, 4, 6), "etp": (0, 1)} - monkeypatch.setattr( - lora_module, - "_process_group_ranks", - lambda group: row_major[group], - ) - assert LoRAPublishPlanner._expert_owner_rank(ep_rank=3, shard_rank=1) == 7 - - column_major = { - "joint": (0, 1, 2, 3, 4, 5, 6, 7), - "ep": (0, 1, 2, 3), - "etp": (0, 4), - } - monkeypatch.setattr( - lora_module, - "_process_group_ranks", - lambda group: column_major[group], - ) - assert LoRAPublishPlanner._expert_owner_rank(ep_rank=3, shard_rank=1) == 7 - - def test_batched_lora_publish_matches_old_shard_merge_exactly(tmp_path: Path): uniform_key = "base_model.model.model.layers.0.self_attn.q_proj.lora_B.weight" componentwise_key = ( @@ -1410,7 +1344,7 @@ def test_batched_lora_publish_matches_old_shard_merge_exactly(tmp_path: Path): uniform_key: full_uniform[2:], componentwise_key: torch.tensor([[10.0], [11.0], [12.0], [13.0]]), } - unsharded_manifest = {"sharded": False, "shard_world_size": 1, "shard_rank": 0} + unsharded_manifest = _manifest(sharded=False, shard_world_size=1, shard_rank=0) uniform_manifest = { "sharded": True, "shard_world_size": 2, @@ -1426,12 +1360,12 @@ def test_batched_lora_publish_matches_old_shard_merge_exactly(tmp_path: Path): } manifest0 = { unsharded_key: unsharded_manifest, - uniform_key: {**uniform_manifest, "shard_rank": 0}, - componentwise_key: {**componentwise_manifest, "shard_rank": 0}, + uniform_key: _manifest(**uniform_manifest, shard_rank=0), + componentwise_key: _manifest(**componentwise_manifest, shard_rank=0), } manifest1 = { - uniform_key: {**uniform_manifest, "shard_rank": 1}, - componentwise_key: {**componentwise_manifest, "shard_rank": 1}, + uniform_key: _manifest(**uniform_manifest, shard_rank=1), + componentwise_key: _manifest(**componentwise_manifest, shard_rank=1), } class IdentityHandler: @@ -1453,7 +1387,7 @@ def to_vllm_lora_tensors(self, tensors, *, adapter_config): handler = IdentityHandler() _old_merge_shard_files_to_vllm( old_dir, - handler=handler, + handler=cast(ModelSupportHandler, handler), adapter_config=adapter_config, ) @@ -1478,16 +1412,16 @@ def to_vllm_lora_tensors(self, tensors, *, adapter_config): ) for key, tensor in shard1.items() ] - lora_publish._save_rank0_vllm_lora( + current_tensors, current_config = lora_publish._rank0_vllm_lora_tensors( metadata=metadata, tensors_by_owner_key={ **{(0, key): tensor for key, tensor in shard0.items()}, **{(1, key): tensor for key, tensor in shard1.items()}, }, - handler=handler, + handler=cast(ModelSupportHandler, handler), adapter_config=adapter_config, - output_dir=str(current_dir), ) + save_vllm_lora_tensors(current_dir, current_tensors, current_config) old_tensors = load_file(old_dir / "adapter_model.safetensors") current_tensors = load_file(current_dir / "adapter_model.safetensors") @@ -1584,11 +1518,17 @@ def test_trainer_rank_publishes_named_checkpoint_slot_without_mutating_base( trainer._dynamic_optimizers = {} trainer._checkpoint_slot_params_by_name = {} trainer._checkpoint_slot_adapter_configs = {} + trainer._checkpoint_revisions = {} config = _config("Qwen/Qwen3-8B", rank=2, alpha=2) - assert trainer.load_checkpoint_slot("student", adapter, adapter_config=config) == 1 + assert trainer._load_checkpoint_slot("student", adapter, alpha=2) == 1 + trainer._checkpoint_slot_params_by_name["student"] = tuple( + trainer._iter_slot_parameters(trainer._slot_ref("student")) + ) + trainer._checkpoint_slot_adapter_configs["student"] = config + trainer._checkpoint_revisions["student"] = 0 output_dir = tmp_path / "checkpoint" - trainer.save_checkpoint_slot_lora("student", str(output_dir)) + assert trainer.export_lora(str(output_dir), "student") == 0 _assert_tensors_equal(load_file(output_dir / "adapter_model.safetensors"), adapter) assert json.loads((output_dir / "adapter_config.json").read_text()) == { @@ -1597,6 +1537,1630 @@ def test_trainer_rank_publishes_named_checkpoint_slot_without_mutating_base( } assert torch.equal(lora.A_T, baseline[0]) assert torch.equal(lora.B_T, baseline[1]) + with pytest.raises(RuntimeError, match="canonical optimizer state"): + validate_checkpoint(output_dir, require_optimizer=True) + + +def test_model_lora_disk_export_streams_one_layer_block_at_a_time( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + prefixes = [ + f"base_model.model.model.layers.{index}.self_attn.q_proj" for index in range(2) + ] + loras = [ + LoRA(prefix, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + for prefix in prefixes + ] + expected: dict[str, torch.Tensor] = {} + for index, lora in enumerate(loras): + lora.A_T.data.fill_(index + 1) + lora.B_T.data.fill_(index + 2) + expected.update(lora.sharded_lora_state_dict()) + + exchanged_blocks: list[set[str]] = [] + original_exchange = lora_publish._exchange_tensors + + def record_exchange(metadata, **kwargs): + if metadata: + exchanged_blocks.append({meta.block for meta in metadata}) + return original_exchange(metadata, **kwargs) + + monkeypatch.setattr(lora_publish, "_exchange_tensors", record_exchange) + save_vllm_lora_from_model( + model=cast(Any, [torch.nn.Sequential(*loras)]), + adapter_dtypes={}, + handler=DEFAULT_DENSE_HANDLER, + adapter_config=_config("Qwen/Qwen3-8B", rank=2, alpha=2), + output_dir=str(tmp_path), + rank=0, + world_size=1, + ) + + assert exchanged_blocks == [ + {prefixes[0].rsplit(".", 2)[0]}, + {prefixes[1].rsplit(".", 2)[0]}, + ] + _assert_tensors_equal(load_file(tmp_path / "adapter_model.safetensors"), expected) + assert not list(tmp_path.glob(".adapter_model-*.safetensors")) + + +def test_checkpoint_load_fetches_only_local_pipeline_layer( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + local_prefix = "base_model.model.model.layers.0.self_attn.q_proj" + remote_prefix = "base_model.model.model.layers.1.self_attn.q_proj" + tensors: dict[str, torch.Tensor] = { + f"{prefix}.lora_{side}.weight": torch.ones(shape) + for prefix in (local_prefix, remote_prefix) + for side, shape in (("A", (2, 3)), ("B", (4, 2))) + } + source = tmp_path / "two-layers" + save_vllm_lora_tensors(source, tensors, _config("Qwen/Qwen3-8B", rank=2, alpha=2)) + prepared = prepare_checkpoint(str(source)) + trainer = _portable_trainer( + LoRA(local_prefix, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + + lora_disk = importlib.import_module("art.megatron.model_support.lora_disk") + original_safe_open = lora_disk.safe_open + fetched: list[str] = [] + + class TrackedSafeOpen: + def __init__(self, *args, **kwargs) -> None: + self._context = original_safe_open(*args, **kwargs) + self._file = None + + def __enter__(self): + self._file = self._context.__enter__() + return self + + def __exit__(self, *args): + return self._context.__exit__(*args) + + def get_tensor(self, key: str) -> torch.Tensor: + fetched.append(key) + assert self._file is not None + return self._file.get_tensor(key) + + monkeypatch.setattr(lora_disk, "safe_open", TrackedSafeOpen) + load_trainer_checkpoint(trainer, prepared, "student") + + assert set(fetched) == { + f"{local_prefix}.lora_A.weight", + f"{local_prefix}.lora_B.weight", + } + assert not any(remote_prefix in key for key in fetched) + + +def test_native_megatron_lora_load_bypasses_vllm_decoder(tmp_path: Path) -> None: + source = tmp_path / "native" + source.mkdir() + save_file(_PORTABLE_ADAPTER, source / "adapter_model.safetensors") + save_adapter_config( + source, + { + **_PORTABLE_CONFIG, + ART_LORA_FORMAT_CONFIG_KEY: ART_LORA_FORMAT_MEGATRON, + }, + ) + handler = SimpleNamespace( + from_vllm_lora_tensors=lambda *_args, **_kwargs: pytest.fail( + "native checkpoint must not be decoded as vLLM" + ) + ) + + _assert_tensors_equal( + load_lora_tensors_for_megatron(source, handler=cast(Any, handler)), + _PORTABLE_ADAPTER, + ) + + +def test_native_checkpoint_reader_materializes_only_local_shards() -> None: + checkpoint_module = importlib.import_module("art.trainer_rank._checkpoint") + tensor = torch.arange(24, dtype=torch.float32).reshape(12, 2) + requests: list[tuple[slice, ...]] = [] + + class TrackedSlice: + def get_shape(self) -> list[int]: + return list(tensor.shape) + + def __getitem__(self, slices: tuple[slice, ...]) -> torch.Tensor: + requests.append(slices) + return tensor[slices] + + class SliceOnlyFile: + def keys(self) -> list[str]: + return ["weight"] + + def get_slice(self, key: str) -> TrackedSlice: + assert key == "weight" + return TrackedSlice() + + def get_tensor(self, key: str) -> torch.Tensor: + del key + raise AssertionError("full canonical tensors must not be materialized") + + uniform = checkpoint_module._read_local_slice( + SliceOnlyFile(), + "weight", + cast( + LoraShardManifest, + { + "sharded": True, + "shard_world_size": 2, + "shard_rank": 1, + "export_shard_dim": 0, + "export_shard_strategy": "uniform", + }, + ), + (6, 2), + ) + torch.testing.assert_close(uniform, tensor[6:12]) + assert requests == [(slice(6, 12), slice(None))] + + requests.clear() + componentwise = checkpoint_module._read_local_slice( + SliceOnlyFile(), + "weight", + cast( + LoraShardManifest, + { + "sharded": True, + "shard_world_size": 2, + "shard_rank": 1, + "export_shard_dim": 0, + "export_shard_strategy": "componentwise", + "component_sizes": (4, 8), + }, + ), + (6, 2), + ) + torch.testing.assert_close(componentwise, torch.cat((tensor[2:4], tensor[8:12]))) + assert requests == [ + (slice(2, 4), slice(None)), + (slice(8, 12), slice(None)), + ] + + +_PORTABLE_PREFIX = "base_model.model.model.layers.0.self_attn.q_proj" +_PORTABLE_CONFIG = _config("Qwen/Qwen3-8B", rank=2, alpha=2) +_PORTABLE_ADAPTER = { + f"{_PORTABLE_PREFIX}.lora_A.weight": torch.arange(6, dtype=torch.float32).reshape( + 2, 3 + ), + f"{_PORTABLE_PREFIX}.lora_B.weight": torch.arange(8, dtype=torch.float32).reshape( + 4, 2 + ), +} + + +_PIPELINE_PREFIXES = tuple( + f"base_model.model.model.layers.{index}.self_attn.q_proj" for index in range(2) +) +_PIPELINE_ADAPTER = { + f"{prefix}.lora_{side}.weight": tensor + 100 * layer + for layer, prefix in enumerate(_PIPELINE_PREFIXES) + for side, tensor in ( + ("A", torch.arange(6, dtype=torch.float32).reshape(2, 3)), + ("B", torch.arange(8, dtype=torch.float32).reshape(4, 2)), + ) +} +_MOE_PREFIX = "base_model.model.model.layers.0.mlp.experts.{expert}.gate_up_proj" +_MOE_CONFIG = _config("Qwen/Qwen3.5-35B-A3B", rank=2, alpha=2) +_MOE_ADAPTER = { + f"{_MOE_PREFIX.format(expert=expert)}.lora_{side}.weight": (tensor + 100 * expert) + for expert in range(2) + for side, tensor in ( + ("A", torch.arange(6, dtype=torch.float32).reshape(2, 3)), + ("B", torch.arange(16, dtype=torch.float32).reshape(8, 2)), + ) +} + + +def _pipeline_loras(*layers: int) -> list[LoRA]: + return [ + LoRA( + _PIPELINE_PREFIXES[layer], + 3, + 4, + 2, + 2, + torch.float32, + torch.device("cpu"), + ) + for layer in layers + ] + + +def _moe_lora(*, num_local_experts: int, expert_tp: bool = False) -> LoRA: + a_parallel_spec = LoRAParallelSpec(shard_domain="expert_tp" if expert_tp else "tp") + b_parallel_spec = LoRAParallelSpec( + shard_domain="expert_tp" if expert_tp else "tp", + sharded=expert_tp, + shard_dim=-1 if expert_tp else None, + ) + return LoRA( + _MOE_PREFIX, + 3, + 4 if expert_tp else 8, + 2, + 2, + torch.float32, + torch.device("cpu"), + num_local_experts=num_local_experts, + a_parallel_spec=a_parallel_spec, + b_parallel_spec=b_parallel_spec, + ) + + +def _portable_trainer( + lora: torch.nn.Module, + *, + rank: int = 0, + world_size: int = 1, + model_identifier: str = str(_PORTABLE_CONFIG["base_model_name_or_path"]), + model_names: tuple[str, ...] = (), +) -> TrainerRank: + trainer = TrainerRank.__new__(TrainerRank) + trainer.runtime = SimpleNamespace( + model=[lora], + model_support_handler=DEFAULT_DENSE_HANDLER, + rank=rank, + world_size=world_size, + model_identifier=model_identifier, + model_support_spec=SimpleNamespace(model_names=model_names), + ) + trainer.device = torch.device("cpu") + trainer._slot_stack = [] + trainer._default_slot_ref = None + trainer._pending_slot_graphs = {} + trainer._dynamic_optimizers = {} + trainer._checkpoint_slot_params_by_name = {} + trainer._checkpoint_slot_adapter_configs = {} + trainer._checkpoint_revisions = {} + return trainer + + +def _install_checkpoint( + trainer: TrainerRank, + adapter: dict[str, torch.Tensor], + config: dict[str, object], +) -> None: + alpha = config["lora_alpha"] + assert isinstance(alpha, int | float) + assert trainer._load_checkpoint_slot("student", adapter, alpha=float(alpha)) > 0 + trainer._checkpoint_slot_params_by_name["student"] = tuple( + trainer._iter_slot_parameters(trainer._slot_ref("student")) + ) + trainer._checkpoint_slot_adapter_configs["student"] = cast(_AdapterConfig, config) + trainer._checkpoint_revisions["student"] = 0 + + +def _install_portable_checkpoint(trainer: TrainerRank) -> None: + _install_checkpoint(trainer, _PORTABLE_ADAPTER, _PORTABLE_CONFIG) + + +def _step_portable_checkpoint(trainer: TrainerRank, gradient: float) -> None: + dynamic = trainer._dynamic_optimizers["student"] + for master in dynamic.master_params: + master.grad = torch.full_like(master, gradient) + dynamic.optimizer.step() + dynamic.optimizer.zero_grad(set_to_none=True) + with torch.no_grad(): + for model, master in zip( + trainer._checkpoint_slot_params_by_name["student"], + dynamic.master_params, + strict=True, + ): + model.copy_(master) + model.grad = None + + +def _assert_checkpoint_state_equal(expected: TrainerRank, actual: TrainerRank) -> None: + expected_params = expected._checkpoint_slot_params_by_name["student"] + actual_params = actual._checkpoint_slot_params_by_name["student"] + for expected_param, actual_param in zip( + expected_params, actual_params, strict=True + ): + torch.testing.assert_close(actual_param, expected_param, atol=0, rtol=0) + expected_optimizer = expected._dynamic_optimizers["student"] + actual_optimizer = actual._dynamic_optimizers["student"] + for expected_master, actual_master in zip( + expected_optimizer.master_params, + actual_optimizer.master_params, + strict=True, + ): + torch.testing.assert_close(actual_master, expected_master, atol=0, rtol=0) + for key in ("step", "exp_avg", "exp_avg_sq"): + torch.testing.assert_close( + actual_optimizer.optimizer.state[actual_master][key], + expected_optimizer.optimizer.state[expected_master][key], + atol=0, + rtol=0, + ) + + +def _save_before_next_step(trainer: TrainerRank, path: Path) -> None: + trainer._dynamic_optimizers["student"] = trainer._new_dynamic_optimizer( + "student", + AdamParams(learning_rate=3e-4, beta1=0.8, beta2=0.95, weight_decay=0.1), + ) + _step_portable_checkpoint(trainer, 0.25) + trainer.save_checkpoint(str(path), "student") + _step_portable_checkpoint(trainer, -0.125) + + +def test_single_local_expert_uses_global_expert_keys( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(lora_module.ps, "get_expert_model_parallel_rank", lambda: 1) + monkeypatch.setattr(lora_module.ps, "get_expert_data_parallel_rank", lambda: 0) + module = _moe_lora(num_local_experts=1) + ref = LoRASlotRef("student") + expected = { + key: tensor for key, tensor in _MOE_ADAPTER.items() if ".experts.1." in key + } + + assert module._expected_weight_keys("lora_A") == [ + f"{_MOE_PREFIX.format(expert=1)}.lora_A.weight" + ] + assert module.load_lora_slot(ref, expected, alpha=2) + _assert_tensors_equal(module.sharded_lora_state_dict(ref), expected) + assert set(module.sharded_lora_manifest(ref)) == set(expected) + + +def test_checkpoint_load_coordinates_rank_local_read_failure( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + source = tmp_path / "source" + save_vllm_lora_tensors(source, _PORTABLE_ADAPTER, _PORTABLE_CONFIG) + prepared = prepare_checkpoint(str(source)) + trainer = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + checkpoint_module = importlib.import_module("art.trainer_rank._checkpoint") + monkeypatch.setattr( + checkpoint_module, + "_load_stage_lora", + lambda *_args: (_ for _ in ()).throw(OSError("rank-local read failed")), + ) + + def coordinated(error: BaseException | None, phase: str) -> None: + assert phase == "read checkpoint" + assert isinstance(error, OSError) + raise RuntimeError("coordinated read failure") from error + + monkeypatch.setattr(checkpoint_module, "_raise_distributed", coordinated) + with pytest.raises(RuntimeError, match="coordinated read failure"): + load_trainer_checkpoint(trainer, prepared, "student") + + +def test_checkpoint_load_rejects_same_keys_with_different_rank_content( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + paths = [tmp_path / "first", tmp_path / "second"] + save_vllm_lora_tensors(paths[0], _PORTABLE_ADAPTER, _PORTABLE_CONFIG) + save_vllm_lora_tensors( + paths[1], + {key: value + 1 for key, value in _PORTABLE_ADAPTER.items()}, + _PORTABLE_CONFIG, + ) + first, second = (prepare_checkpoint(str(path)) for path in paths) + assert first.artifact_keys == second.artifact_keys + assert first.digest != second.digest + trainer = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + checkpoint_module = importlib.import_module("art.trainer_rank._checkpoint") + monkeypatch.setattr( + checkpoint_module, + "_gather_objects", + lambda value: (value, second.digest), + ) + monkeypatch.setattr( + checkpoint_module, + "_load_stage_lora", + lambda *_args: pytest.fail("mismatched content must fail before reading"), + ) + + with pytest.raises(RuntimeError, match="content differs across ranks"): + load_trainer_checkpoint(trainer, first, "student") + + +def test_checkpoint_payload_validation_elects_one_rank_per_node( + monkeypatch: pytest.MonkeyPatch, +) -> None: + checkpoint_module = importlib.import_module("art.trainer_rank._checkpoint") + monkeypatch.setattr(checkpoint_module, "_distributed", lambda: True) + monkeypatch.setattr(checkpoint_module, "_rank", lambda: 7) + monkeypatch.delenv("LOCAL_RANK", raising=False) + assert checkpoint_module._is_node_validator() + monkeypatch.setenv("LOCAL_RANK", "1") + assert not checkpoint_module._is_node_validator() + monkeypatch.setenv("LOCAL_RANK", "0") + assert checkpoint_module._is_node_validator() + + +def test_prepared_checkpoint_pins_a_symlink_target(tmp_path: Path) -> None: + first = tmp_path / "first" + second = tmp_path / "second" + save_vllm_lora_tensors(first, _PORTABLE_ADAPTER, _PORTABLE_CONFIG) + save_vllm_lora_tensors( + second, + {key: tensor + 100 for key, tensor in _PORTABLE_ADAPTER.items()}, + _PORTABLE_CONFIG, + ) + latest = tmp_path / "latest" + latest.symlink_to(first, target_is_directory=True) + prepared = prepare_checkpoint(str(latest)) + latest.unlink() + latest.symlink_to(second, target_is_directory=True) + + trainer = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + checkpoint_module = importlib.import_module("art.trainer_rank._checkpoint") + assert prepared.path == first.resolve() + _assert_tensors_equal( + checkpoint_module._load_stage_lora(trainer, prepared), _PORTABLE_ADAPTER + ) + + +def test_native_checkpoint_rejects_unmapped_optimizer_tensor( + tmp_path: Path, +) -> None: + original = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + _install_portable_checkpoint(original) + original._dynamic_optimizers["student"] = original._new_dynamic_optimizer( + "student", AdamParams(learning_rate=3e-4) + ) + _step_portable_checkpoint(original, 0.25) + output = tmp_path / "exact-with-extra" + original.save_checkpoint(str(output), "student") + + checkpoint_module = importlib.import_module("art.trainer_rank._checkpoint") + manifest = validate_checkpoint(output, require_optimizer=True) + assert manifest is not None and manifest.optimizer is not None + source_key = next(iter(manifest.parameters)) + extra_key = f"{source_key}.unmapped" + adapter = load_file(output / "adapter_model.safetensors") + adapter[extra_key] = adapter[source_key].clone() + save_file(adapter, output / "adapter_model.safetensors") + + source_record = manifest.parameters[source_key] + component_updates = {} + for component in ("master", "exp_avg", "exp_avg_sq"): + record = getattr(source_record, component) + tensors = load_file(output / record.file) + tensors[extra_key] = tensors[source_key].clone() + save_file(tensors, output / record.file) + component_updates[component] = record.model_copy(update={"tensor": extra_key}) + parameters = { + **manifest.parameters, + extra_key: source_record.model_copy(update=component_updates), + } + unsigned = manifest.model_copy( + update={ + "parameters": parameters, + "steps": {**manifest.steps, extra_key: manifest.steps[source_key]}, + "digest": "", + } + ) + checkpoint_module._write_manifest( + output, + unsigned.model_copy( + update={"digest": checkpoint_module._digest(output, unsigned)} + ), + ) + + prepared = prepare_checkpoint(str(output)) + restored = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + with pytest.raises(RuntimeError, match="coverage differs from the target runtime"): + load_trainer_checkpoint(restored, prepared, "student") + assert not restored._checkpoint_slot_params_by_name + + +@pytest.mark.parametrize("with_optimizer", [False, True]) +def test_canonical_checkpoint_rejects_different_model_in_same_support_spec( + tmp_path: Path, + with_optimizer: bool, +) -> None: + original = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + _install_portable_checkpoint(original) + if with_optimizer: + original._dynamic_optimizers["student"] = original._new_dynamic_optimizer( + "student", AdamParams(learning_rate=3e-4) + ) + output = tmp_path / f"canonical-model-{with_optimizer}" + original.save_checkpoint(str(output), "student") + + restored = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")), + model_identifier="Qwen/Qwen3-8B-Base", + model_names=("Qwen/Qwen3-8B", "Qwen/Qwen3-8B-Base"), + ) + with pytest.raises(RuntimeError, match="Exact checkpoint base model"): + load_trainer_checkpoint(restored, prepare_checkpoint(str(output)), "student") + assert not restored._checkpoint_slot_params_by_name + + +def test_lora_only_checkpoint_rejects_different_model_in_same_support_spec( + tmp_path: Path, +) -> None: + source = tmp_path / "lora-only-different-model" + save_vllm_lora_tensors(source, _PORTABLE_ADAPTER, _PORTABLE_CONFIG) + prepared = prepare_checkpoint(str(source)) + assert prepared.manifest is None + + restored = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")), + model_identifier="Qwen/Qwen3-8B-Base", + model_names=("Qwen/Qwen3-8B", "Qwen/Qwen3-8B-Base"), + ) + with pytest.raises(RuntimeError, match="Checkpoint base model"): + load_trainer_checkpoint(restored, prepared, "student") + assert not restored._checkpoint_slot_params_by_name + + +def test_checkpoint_prepare_captures_immutable_state_and_finish_is_idempotent( + tmp_path: Path, +) -> None: + trainer = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + _install_portable_checkpoint(trainer) + trainer._dynamic_optimizers["student"] = trainer._new_dynamic_optimizer( + "student", + AdamParams(learning_rate=3e-4, beta1=0.8, beta2=0.95), + ) + _step_portable_checkpoint(trainer, 0.25) + expected_params = tuple( + value.detach().clone() + for value in trainer._checkpoint_slot_params_by_name["student"] + ) + dynamic = trainer._dynamic_optimizers["student"] + expected_masters = tuple(value.detach().clone() for value in dynamic.master_params) + expected_states = tuple( + { + key: value.detach().clone() + for key, value in dynamic.optimizer.state[master].items() + if isinstance(value, torch.Tensor) + } + for master in dynamic.master_params + ) + + output = tmp_path / "immutable" + trainer._prepare_checkpoint_save(str(output), "student") + with torch.no_grad(): + for value in trainer._checkpoint_slot_params_by_name["student"]: + value.add_(100) + for master in dynamic.master_params: + master.add_(200) + for value in dynamic.optimizer.state[master].values(): + if isinstance(value, torch.Tensor): + value.add_(300) + trainer._finish_checkpoint_save(str(output)) + trainer._finish_checkpoint_save(str(output)) + + restored = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + load_trainer_checkpoint( + restored, + prepare_checkpoint(str(output)), + "student", + ) + for expected, actual in zip( + expected_params, + restored._checkpoint_slot_params_by_name["student"], + strict=True, + ): + torch.testing.assert_close(actual, expected, atol=0, rtol=0) + restored_dynamic = restored._dynamic_optimizers["student"] + for expected, expected_state, actual in zip( + expected_masters, + expected_states, + restored_dynamic.master_params, + strict=True, + ): + torch.testing.assert_close(actual, expected, atol=0, rtol=0) + for key, value in expected_state.items(): + torch.testing.assert_close( + restored_dynamic.optimizer.state[actual][key], + value, + atol=0, + rtol=0, + ) + assert not list(tmp_path.glob(".immutable.snapshot-*")) + + +def test_checkpoint_finalizers_run_fifo_and_clean_snapshots( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + trainer = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + _install_portable_checkpoint(trainer) + trainer._dynamic_optimizers["student"] = trainer._new_dynamic_optimizer( + "student", AdamParams(learning_rate=3e-4) + ) + _step_portable_checkpoint(trainer, 0.25) + first = tmp_path / "fifo-first" + second = tmp_path / "fifo-second" + trainer._prepare_checkpoint_save(str(first), "student") + trainer._prepare_checkpoint_save(str(second), "student") + + checkpoint_module = importlib.import_module("art.trainer_rank._checkpoint") + original_finish = checkpoint_module._finish_prepared_save + order: list[str] = [] + + def record_finish(trainer_rank: TrainerRank, prepared: _PreparedSave) -> None: + order.append(prepared.destination.name) + original_finish(trainer_rank, prepared) + + monkeypatch.setattr(checkpoint_module, "_finish_prepared_save", record_finish) + second_started = threading.Event() + second_done = threading.Event() + + def finish_second() -> None: + second_started.set() + trainer._finish_checkpoint_save(str(second)) + second_done.set() + + second_thread = threading.Thread(target=finish_second) + second_thread.start() + assert second_started.wait(1) + assert not second_done.wait(0.1) + first_thread = threading.Thread( + target=trainer._finish_checkpoint_save, + args=(str(first),), + ) + first_thread.start() + first_thread.join(10) + second_thread.join(10) + assert not first_thread.is_alive() + assert not second_thread.is_alive() + assert order == ["fifo-first", "fifo-second"] + assert first.is_dir() and second.is_dir() + assert not list(tmp_path.glob(".*.snapshot-*")) + + +def test_checkpoint_prepare_failure_does_not_advance_or_wedge_fifo( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + trainer = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + _install_portable_checkpoint(trainer) + checkpoint_module = importlib.import_module("art.trainer_rank._checkpoint") + original_save_file = checkpoint_module._save_file + + def fail_snapshot(*_args: object) -> None: + raise OSError("injected checkpoint snapshot failure") + + monkeypatch.setattr(checkpoint_module, "_save_file", fail_snapshot) + failed = tmp_path / "prepare-failure" + with pytest.raises(OSError, match="injected checkpoint snapshot failure"): + trainer._prepare_checkpoint_save(str(failed), "student") + assert trainer._checkpoint_save_sequence == 0 + assert trainer._checkpoint_finish_sequence == 0 + assert not trainer._prepared_checkpoint_saves + assert not list(tmp_path.glob(".prepare-failure.snapshot-*")) + + monkeypatch.setattr(checkpoint_module, "_save_file", original_save_file) + recovered = tmp_path / "prepare-recovered" + trainer.save_checkpoint(str(recovered), "student") + assert recovered.is_dir() + assert trainer._checkpoint_save_sequence == 1 + assert trainer._checkpoint_finish_sequence == 1 + + +def test_checkpoint_finalizer_failure_cleans_and_can_retry( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + trainer = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + _install_portable_checkpoint(trainer) + output = tmp_path / "retry" + trainer._prepare_checkpoint_save(str(output), "student") + checkpoint_module = importlib.import_module("art.trainer_rank._checkpoint") + original_finish = checkpoint_module._finish_prepared_save + + def fail_finish(*_args: object) -> None: + raise OSError("injected checkpoint finalization failure") + + monkeypatch.setattr(checkpoint_module, "_finish_prepared_save", fail_finish) + with pytest.raises(OSError, match="injected checkpoint finalization failure"): + trainer._finish_checkpoint_save(str(output)) + assert not output.exists() + assert not list(tmp_path.glob(".retry.snapshot-*")) + + monkeypatch.setattr(checkpoint_module, "_finish_prepared_save", original_finish) + trainer._prepare_checkpoint_save(str(output), "student") + trainer._finish_checkpoint_save(str(output)) + assert output.is_dir() + assert not list(tmp_path.glob(".retry.snapshot-*")) + + +def test_trainer_rank_checkpoint_roundtrip_preserves_next_adam_step( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + + original = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + _install_portable_checkpoint(original) + original._dynamic_optimizers["student"] = original._new_dynamic_optimizer( + "student", + AdamParams(learning_rate=3e-4, beta1=0.8, beta2=0.95, weight_decay=0.1), + ) + unstepped_output = tmp_path / "unstepped" + original.save_checkpoint(str(unstepped_output), "student") + unstepped = prepare_checkpoint(str(unstepped_output)) + assert unstepped.manifest is not None + assert set(unstepped.manifest.steps.values()) == {0.0} + _step_portable_checkpoint(original, 0.25) + original_optimizer = original._dynamic_optimizers["student"] + for step, master in enumerate(original_optimizer.master_params, start=2): + original_optimizer.optimizer.state[master]["step"].fill_(step) + output = tmp_path / "exact" + + original.save_checkpoint(str(output), "student") + original.save_checkpoint(str(output), "student") + manifest = validate_checkpoint(output, require_optimizer=True) + assert manifest is not None and manifest.optimizer is not None + assert set(manifest.steps.values()) == {2.0, 3.0} + checkpoint_module = importlib.import_module("art.trainer_rank._checkpoint") + parameter_key, parameter = next(iter(manifest.parameters.items())) + for invalid_path in ( + "", + ".", + "../escape.safetensors", + "/tmp/escape.safetensors", + "optimizer/../escape.safetensors", + "optimizer\\escape.safetensors", + "C:/escape.safetensors", + "optimizer//escape.safetensors", + "./optimizer/escape.safetensors", + ): + invalid_parameter = parameter.model_copy( + update={ + "master": parameter.master.model_copy(update={"file": invalid_path}) + } + ) + invalid_manifest = manifest.model_copy( + update={ + "parameters": { + **manifest.parameters, + parameter_key: invalid_parameter, + } + } + ) + with pytest.raises(RuntimeError, match="Unsafe checkpoint tensor path"): + checkpoint_module._validate_manifest( + invalid_manifest, tuple(manifest.parameters) + ) + source_adapter = (output / "adapter_model.safetensors").read_bytes() + source_config = (output / "adapter_config.json").read_bytes() + assert json.loads(source_config)[ART_LORA_FORMAT_CONFIG_KEY] == ( + ART_LORA_FORMAT_MEGATRON + ) + _assert_tensors_equal( + load_lora_tensors_for_megatron(output), + load_file(output / "adapter_model.safetensors"), + ) + assert len(list((output / "optimizer").glob("*.safetensors"))) == 3 + malformed = tmp_path / "malformed" + shutil.copytree(output, malformed) + malformed_manifest = json.loads((malformed / "checkpoint.json").read_text()) + first = next(iter(malformed_manifest["parameters"].values())) + first["master"]["shape"][0] += 1 + (malformed / "checkpoint.json").write_text(json.dumps(malformed_manifest)) + with pytest.raises(RuntimeError, match="tensor metadata mismatch"): + validate_checkpoint(malformed, require_optimizer=True) + inference = tmp_path / "inference" + materialize_lora(output, inference, require_optimizer=True) + assert (output / "adapter_model.safetensors").read_bytes() == source_adapter + assert (output / "adapter_config.json").read_bytes() == source_config + assert ( + json.loads((inference / "adapter_config.json").read_text())[ + ART_LORA_FORMAT_CONFIG_KEY + ] + == ART_LORA_FORMAT_VLLM + ) + assert {item.name for item in inference.iterdir()} == { + "adapter_config.json", + "adapter_model.safetensors", + } + + entries = { + item.relative_to(output).as_posix() + for item in output.rglob("*") + if item.is_file() + } + staged = tmp_path / "selective-stage" + staged.mkdir() + for relative in ( + "adapter_config.json", + "adapter_model.safetensors", + "checkpoint.json", + ): + shutil.copy2(output / relative, staged / relative) + selective = tmp_path / "selective-inference" + materialize_lora( + staged, + selective, + require_optimizer=True, + artifact_entries=entries, + expected_digest=manifest.digest, + ) + assert {item.name for item in selective.iterdir()} == { + "adapter_config.json", + "adapter_model.safetensors", + } + with pytest.raises(RuntimeError, match="missing referenced files"): + materialize_lora( + staged, + tmp_path / "missing-entry", + require_optimizer=True, + artifact_entries=entries + - {next(iter(manifest.parameters.values())).master.file}, + expected_digest=manifest.digest, + ) + with pytest.raises(RuntimeError, match="durable artifact digest"): + materialize_lora( + staged, + tmp_path / "wrong-digest", + require_optimizer=True, + artifact_entries=entries, + expected_digest="wrong", + ) + with pytest.raises(FileExistsError, match="not empty"): + materialize_lora(output, inference, require_optimizer=True) + + restored = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + source = prepare_checkpoint(str(output)) + load_trainer_checkpoint(restored, source, "student") + assert source.manifest is not None and source.manifest.optimizer is not None + assert restored._checkpoint_revisions["student"] == 0 + + _step_portable_checkpoint(original, -0.125) + _step_portable_checkpoint(restored, -0.125) + for expected, actual in zip( + original._checkpoint_slot_params_by_name["student"], + restored._checkpoint_slot_params_by_name["student"], + strict=True, + ): + torch.testing.assert_close(actual, expected, atol=0, rtol=0) + + with pytest.raises(FileExistsError, match="not empty"): + original.save_checkpoint(str(output), "student") + + failed = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + _install_portable_checkpoint(failed) + before = tuple( + parameter.detach().clone() + for parameter in failed._checkpoint_slot_params_by_name["student"] + ) + + tampered_dtype = tmp_path / "tampered-dtype" + shutil.copytree(output, tampered_dtype) + tampered_manifest = json.loads((tampered_dtype / "checkpoint.json").read_text()) + tampered_parameter = next(iter(tampered_manifest["parameters"].values())) + tampered_parameter["exp_avg"]["dtype"] = "float16" + (tampered_dtype / "checkpoint.json").write_text(json.dumps(tampered_manifest)) + + async def load_tampered_dtype() -> None: + await failed.load_checkpoint(str(tampered_dtype)) + + with pytest.raises(RuntimeError, match="must use float32"): + asyncio.run(load_tampered_dtype()) + assert set(failed._checkpoint_slot_params_by_name) == {"student"} + assert failed._default_slot_ref is None + for expected, actual in zip( + before, + failed._checkpoint_slot_params_by_name["student"], + strict=True, + ): + torch.testing.assert_close(actual, expected, atol=0, rtol=0) + + def reject_optimizer(*_args: object) -> None: + raise RuntimeError("injected optimizer failure") + + monkeypatch.setattr(failed, "_restore_canonical_optimizer", reject_optimizer) + with pytest.raises(RuntimeError, match="injected optimizer failure"): + load_trainer_checkpoint(failed, source, "student") + assert set(failed._checkpoint_slot_params_by_name) == {"student"} + for expected, actual in zip( + before, + failed._checkpoint_slot_params_by_name["student"], + strict=True, + ): + torch.testing.assert_close(actual, expected, atol=0, rtol=0) + + +def _portable_topology_worker( + rank: int, + world_size: int, + init_method: str, + source_path: str, + output_path: str, +) -> None: + torch.distributed.init_process_group( + "gloo", rank=rank, world_size=world_size, init_method=init_method + ) + try: + lora_module.ps.initialize_model_parallel( + tensor_model_parallel_size=world_size, + pipeline_model_parallel_size=1, + context_parallel_size=1, + expert_model_parallel_size=1, + ) + import art.trainer_rank._checkpoint as checkpoint_module + + setattr(checkpoint_module, "_device", lambda: torch.device("cpu")) + setattr( + lora_publish, + "_rank_and_device", + lambda: (rank, torch.device("cpu")), + ) + lora = LoRA( + _PORTABLE_PREFIX, + 3, + 4 // world_size, + 2, + 2, + torch.float32, + torch.device("cpu"), + b_parallel_spec=LoRAParallelSpec(sharded=True, shard_dim=-1), + ) + trainer = _portable_trainer(lora, rank=rank, world_size=world_size) + load_trainer_checkpoint( + trainer, + prepare_checkpoint(source_path), + "student", + ) + _step_portable_checkpoint(trainer, -0.125) + trainer.save_checkpoint(output_path, "student") + finally: + if lora_module.ps.model_parallel_is_initialized(): + lora_module.ps.destroy_model_parallel() + torch.distributed.destroy_process_group() + + +def _advanced_topology_worker( + rank: int, + world_size: int, + init_method: str, + source_path: str, + output_path: str, + topology: Literal["pp", "ep", "etp", "cp"], + counts_path: str | None, +) -> None: + torch.distributed.init_process_group( + "gloo", rank=rank, world_size=world_size, init_method=init_method + ) + try: + lora_module.ps.initialize_model_parallel( + tensor_model_parallel_size=1, + pipeline_model_parallel_size=2 if topology == "pp" else 1, + context_parallel_size=2 if topology == "cp" else 1, + expert_model_parallel_size=2 if topology == "ep" else 1, + expert_tensor_parallel_size=2 if topology == "etp" else None, + ) + if topology == "pp": + modules = _pipeline_loras(rank) + model_identifier = str(_PORTABLE_CONFIG["base_model_name_or_path"]) + elif topology == "ep": + modules = [_moe_lora(num_local_experts=1)] + model_identifier = str(_MOE_CONFIG["base_model_name_or_path"]) + elif topology == "etp": + modules = [_moe_lora(num_local_experts=2, expert_tp=True)] + model_identifier = str(_MOE_CONFIG["base_model_name_or_path"]) + else: + modules = [ + LoRA( + _PORTABLE_PREFIX, + 3, + 4, + 2, + 2, + torch.float32, + torch.device("cpu"), + ) + ] + model_identifier = str(_PORTABLE_CONFIG["base_model_name_or_path"]) + trainer = _portable_trainer( + torch.nn.Sequential(*modules), + rank=rank, + world_size=world_size, + model_identifier=model_identifier, + ) + setattr(lora_publish, "_rank_and_device", lambda: (rank, torch.device("cpu"))) + load_trainer_checkpoint(trainer, prepare_checkpoint(source_path), "student") + _step_portable_checkpoint(trainer, -0.125) + + sends = 0 + original_send = torch.distributed.send + + def counted_send(*args, **kwargs): + nonlocal sends + sends += 1 + return original_send(*args, **kwargs) + + torch.distributed.send = counted_send + trainer.save_checkpoint(output_path, "student") + if counts_path is not None: + counts: list[int | None] = [None] * world_size + torch.distributed.all_gather_object(counts, sends) + if rank == 0: + Path(counts_path).write_text(json.dumps(counts)) + finally: + if lora_module.ps.model_parallel_is_initialized(): + lora_module.ps.destroy_model_parallel() + torch.distributed.destroy_process_group() + + +def _reversed_prefetch_order_worker( + rank: int, + world_size: int, + init_method: str, + first_path: str, + second_path: str, +) -> None: + torch.distributed.init_process_group( + "gloo", rank=rank, world_size=world_size, init_method=init_method + ) + try: + lora_module.ps.initialize_model_parallel( + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + context_parallel_size=1, + expert_model_parallel_size=1, + ) + lora = LoRA( + _PORTABLE_PREFIX, + 3, + 4, + 2, + 2, + torch.float32, + torch.device("cpu"), + ) + trainer = _portable_trainer(lora, rank=rank, world_size=world_size) + import art.trainer_rank._checkpoint as checkpoint_module + + original_prepare = checkpoint_module.prepare_checkpoint + preparation_barrier = threading.Barrier(2) + completed: list[str] = [] + + def reversed_prepare(path: str): + preparation_barrier.wait(timeout=10) + slow_path = first_path if rank == 0 else second_path + if path == slow_path: + time.sleep(0.25) + source = original_prepare(path) + completed.append(path) + return source + + setattr(checkpoint_module, "prepare_checkpoint", reversed_prepare) + + async def load_both() -> None: + first = trainer.load_checkpoint(first_path) + second = trainer.load_checkpoint(second_path) + await asyncio.wait_for(asyncio.gather(first, second), timeout=30) + + asyncio.run(load_both()) + assert trainer._default_slot_ref == trainer._slot_ref(second_path) + assert set(trainer._checkpoint_slot_params_by_name) == { + first_path, + second_path, + } + orders: list[list[str] | None] = [None] * world_size + torch.distributed.all_gather_object(orders, completed) + assert orders[0] is not None and orders[0][0] == second_path + assert orders[1] is not None and orders[1][0] == first_path + finally: + if lora_module.ps.model_parallel_is_initialized(): + lora_module.ps.destroy_model_parallel() + torch.distributed.destroy_process_group() + + +def _transactional_load_failure_worker( + rank: int, + world_size: int, + init_method: str, + source_path: str, +) -> None: + torch.distributed.init_process_group( + "gloo", rank=rank, world_size=world_size, init_method=init_method + ) + try: + lora_module.ps.initialize_model_parallel( + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + context_parallel_size=1, + expert_model_parallel_size=1, + ) + lora = LoRA( + _PORTABLE_PREFIX, + 3, + 4, + 2, + 2, + torch.float32, + torch.device("cpu"), + ) + trainer = _portable_trainer(lora, rank=rank, world_size=world_size) + _install_portable_checkpoint(trainer) + trainer._dynamic_optimizers["student"] = trainer._new_dynamic_optimizer( + "student", AdamParams(learning_rate=1e-3) + ) + _step_portable_checkpoint(trainer, -0.5) + previous_params = trainer._checkpoint_slot_params_by_name["student"] + previous_values = tuple( + parameter.detach().clone() for parameter in previous_params + ) + previous_dynamic = trainer._dynamic_optimizers["student"] + original_replace = lora_module.replace_lora_slot_in_model + if rank == 1: + + def fail_after_replacement( + model: Sequence[torch.nn.Module], + source: LoRASlotRef, + destination: LoRASlotRef, + ) -> None: + original_replace(model, source, destination) + raise RuntimeError("injected checkpoint commit failure") + + setattr(lora_module, "replace_lora_slot_in_model", fail_after_replacement) + + with pytest.raises( + RuntimeError, match="checkpoint commit failure|failed to commit" + ): + load_trainer_checkpoint(trainer, prepare_checkpoint(source_path), "student") + + restored_params = trainer._checkpoint_slot_params_by_name["student"] + assert tuple(map(id, restored_params)) == tuple(map(id, previous_params)) + for expected, actual in zip(previous_values, restored_params, strict=True): + torch.testing.assert_close(actual, expected, atol=0, rtol=0) + assert trainer._dynamic_optimizers["student"] is previous_dynamic + assert trainer._checkpoint_slot_adapter_configs["student"] == _PORTABLE_CONFIG + assert trainer._checkpoint_revisions["student"] == 0 + assert all( + not (ref.name or "").startswith("__art_loading_") for ref in lora._slot_keys + ) + finally: + if lora_module.ps.model_parallel_is_initialized(): + lora_module.ps.destroy_model_parallel() + torch.distributed.destroy_process_group() + + +def _portable_replica_worker( + rank: int, + world_size: int, + init_method: str, + output_path: str, + counts_path: str, + diverge: Literal["lora", "optimizer"] | None = None, +) -> None: + torch.distributed.init_process_group( + "gloo", rank=rank, world_size=world_size, init_method=init_method + ) + try: + lora_module.ps.initialize_model_parallel( + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + context_parallel_size=1, + expert_model_parallel_size=1, + ) + import art.trainer_rank._checkpoint as checkpoint_module + + setattr(checkpoint_module, "_device", lambda: torch.device("cpu")) + setattr( + lora_publish, + "_rank_and_device", + lambda: (rank, torch.device("cpu")), + ) + lora = LoRA( + _PORTABLE_PREFIX, + 3, + 4, + 2, + 2, + torch.float32, + torch.device("cpu"), + ) + setattr(lora, "_should_export_parameter", lambda _param: True) + trainer = _portable_trainer(lora, rank=rank, world_size=world_size) + _install_portable_checkpoint(trainer) + trainer._dynamic_optimizers["student"] = trainer._new_dynamic_optimizer( + "student", + AdamParams(learning_rate=3e-4, beta1=0.8, beta2=0.95, weight_decay=0.1), + ) + _step_portable_checkpoint(trainer, 0.25) + if diverge == "lora" and rank == 1: + with torch.no_grad(): + trainer._checkpoint_slot_params_by_name["student"][0].add_(1) + elif diverge == "optimizer" and rank == 1: + dynamic = trainer._dynamic_optimizers["student"] + master = dynamic.master_params[0] + dynamic.optimizer.state[master]["exp_avg"].add_(1) + + sends = 0 + original_send = torch.distributed.send + + def counted_send(*args, **kwargs): + nonlocal sends + sends += 1 + return original_send(*args, **kwargs) + + torch.distributed.send = counted_send + trainer.save_checkpoint(output_path, "student") + counts: list[int | None] = [None] * world_size + torch.distributed.all_gather_object(counts, sends) + if rank == 0: + Path(counts_path).write_text(json.dumps(counts)) + finally: + if lora_module.ps.model_parallel_is_initialized(): + lora_module.ps.destroy_model_parallel() + torch.distributed.destroy_process_group() + + +def _exchange_preparation_failure_worker( + rank: int, + world_size: int, + init_method: str, +) -> None: + torch.distributed.init_process_group( + "gloo", rank=rank, world_size=world_size, init_method=init_method + ) + try: + if rank == 1: + + def fail_preparation(*_args: object, **_kwargs: object) -> None: + raise OSError("injected tensor exchange preparation failure") + + setattr( + lora_publish, + "_prepare_exchange_buffers", + fail_preparation, + ) + key = "base_model.model.model.layers.0.self_attn.q_proj.lora_A.weight" + metadata = ( + LoraShardMeta( + key=key, + owner_rank=1, + shape=(2,), + dtype_name="float32", + manifest={"sharded": False, "shard_world_size": 1, "shard_rank": 0}, + block="base_model.model.model.layers.0", + ), + ) + with pytest.raises( + (OSError, RuntimeError), + match="injected tensor exchange preparation failure", + ): + lora_publish._exchange_tensors( + metadata, + local_tensors={key: torch.ones(2)} if rank == 1 else {}, + rank=rank, + device=torch.device("cpu"), + ) + finally: + torch.distributed.destroy_process_group() + + +def test_lora_exchange_coordinates_asymmetric_preparation_failure( + tmp_path: Path, +) -> None: + mp.spawn( + _exchange_preparation_failure_worker, + args=(2, f"file://{tmp_path / 'exchange-failure-init'}"), + nprocs=2, + join=True, + ) + + +def test_checkpoint_loads_follow_call_order_across_ranks(tmp_path: Path) -> None: + first = tmp_path / "prefetch-first" + second = tmp_path / "prefetch-second" + save_vllm_lora_tensors(first, _PORTABLE_ADAPTER, _PORTABLE_CONFIG) + save_vllm_lora_tensors( + second, + {key: value + 1 for key, value in _PORTABLE_ADAPTER.items()}, + _PORTABLE_CONFIG, + ) + + mp.spawn( + _reversed_prefetch_order_worker, + args=( + 2, + f"file://{tmp_path / 'prefetch-order-init'}", + str(first), + str(second), + ), + nprocs=2, + join=True, + ) + + +def test_checkpoint_commit_failure_rolls_back_every_rank(tmp_path: Path) -> None: + source_trainer = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + _install_portable_checkpoint(source_trainer) + source_trainer._dynamic_optimizers["student"] = ( + source_trainer._new_dynamic_optimizer("student", AdamParams(learning_rate=3e-4)) + ) + _step_portable_checkpoint(source_trainer, 0.25) + source = tmp_path / "transaction-source" + source_trainer.save_checkpoint(str(source), "student") + + mp.spawn( + _transactional_load_failure_worker, + args=( + 2, + f"file://{tmp_path / 'transaction-init'}", + str(source), + ), + nprocs=2, + join=True, + ) + + +def test_trainer_rank_checkpoint_deduplicates_data_parallel_replicas( + tmp_path: Path, +) -> None: + expected = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + _install_portable_checkpoint(expected) + expected._dynamic_optimizers["student"] = expected._new_dynamic_optimizer( + "student", + AdamParams(learning_rate=3e-4, beta1=0.8, beta2=0.95, weight_decay=0.1), + ) + _step_portable_checkpoint(expected, 0.25) + + output = tmp_path / "replicated" + counts = tmp_path / "replicated-counts.json" + mp.spawn( + _portable_replica_worker, + args=( + 2, + f"file://{tmp_path / 'replica-init'}", + str(output), + str(counts), + None, + ), + nprocs=2, + join=True, + ) + + assert json.loads(counts.read_text()) == [0, 0] + restored = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + load_trainer_checkpoint( + restored, + prepare_checkpoint(str(output)), + "student", + ) + for expected_param, actual_param in zip( + expected._checkpoint_slot_params_by_name["student"], + restored._checkpoint_slot_params_by_name["student"], + strict=True, + ): + torch.testing.assert_close(actual_param, expected_param, atol=0, rtol=0) + expected_optimizer = expected._dynamic_optimizers["student"] + actual_optimizer = restored._dynamic_optimizers["student"] + for expected_master, actual_master in zip( + expected_optimizer.master_params, + actual_optimizer.master_params, + strict=True, + ): + torch.testing.assert_close(actual_master, expected_master, atol=0, rtol=0) + for key in ("step", "exp_avg", "exp_avg_sq"): + torch.testing.assert_close( + actual_optimizer.optimizer.state[actual_master][key], + expected_optimizer.optimizer.state[expected_master][key], + atol=0, + rtol=0, + ) + + +@pytest.mark.parametrize("diverge", ["lora", "optimizer"]) +def test_trainer_rank_checkpoint_rejects_divergent_data_parallel_replicas( + tmp_path: Path, + diverge: Literal["lora", "optimizer"], +) -> None: + output = tmp_path / "replica-mismatch" + with pytest.raises( + ProcessRaisedException, match="Inconsistent replicated tensor contents" + ): + mp.spawn( + _portable_replica_worker, + args=( + 2, + f"file://{tmp_path / 'replica-mismatch-init'}", + str(output), + str(tmp_path / "unused-counts.json"), + diverge, + ), + nprocs=2, + join=True, + ) + assert not output.exists() + assert not list(tmp_path.glob(".replica-mismatch.snapshot-*")) + + +def test_trainer_rank_checkpoint_restores_1_to_2_to_1( + tmp_path: Path, +) -> None: + original = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + _install_portable_checkpoint(original) + original._dynamic_optimizers["student"] = original._new_dynamic_optimizer( + "student", + AdamParams(learning_rate=3e-4, beta1=0.8, beta2=0.95, weight_decay=0.1), + ) + _step_portable_checkpoint(original, 0.25) + one_rank = tmp_path / "one-rank" + original.save_checkpoint(str(one_rank), "student") + _step_portable_checkpoint(original, -0.125) + + two_rank = tmp_path / "two-rank" + mp.spawn( + _portable_topology_worker, + args=( + 2, + f"file://{tmp_path / 'topology-init'}", + str(one_rank), + str(two_rank), + ), + nprocs=2, + join=True, + ) + + restored = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + load_trainer_checkpoint( + restored, + prepare_checkpoint(str(two_rank)), + "student", + ) + for expected, actual in zip( + original._checkpoint_slot_params_by_name["student"], + restored._checkpoint_slot_params_by_name["student"], + strict=True, + ): + torch.testing.assert_close(actual, expected, atol=0, rtol=0) + expected_optimizer = original._dynamic_optimizers["student"] + actual_optimizer = restored._dynamic_optimizers["student"] + for expected, actual in zip( + expected_optimizer.master_params, + actual_optimizer.master_params, + strict=True, + ): + torch.testing.assert_close(actual, expected, atol=0, rtol=0) + expected_state = expected_optimizer.optimizer.state[expected] + actual_state = actual_optimizer.optimizer.state[actual] + for key in ("step", "exp_avg", "exp_avg_sq"): + torch.testing.assert_close( + actual_state[key], expected_state[key], atol=0, rtol=0 + ) + + +def test_checkpoint_restores_pipeline_parallel_next_step(tmp_path: Path) -> None: + original = _portable_trainer(torch.nn.Sequential(*_pipeline_loras(0, 1))) + _install_checkpoint(original, _PIPELINE_ADAPTER, _PORTABLE_CONFIG) + source = tmp_path / "pipeline-source" + _save_before_next_step(original, source) + output = tmp_path / "pipeline-output" + + mp.spawn( + _advanced_topology_worker, + args=( + 2, + f"file://{tmp_path / 'pipeline-init'}", + str(source), + str(output), + "pp", + None, + ), + nprocs=2, + join=True, + ) + + restored = _portable_trainer(torch.nn.Sequential(*_pipeline_loras(0, 1))) + load_trainer_checkpoint(restored, prepare_checkpoint(str(output)), "student") + _assert_checkpoint_state_equal(original, restored) + + +@pytest.mark.parametrize("topology", ["ep", "etp"]) +def test_checkpoint_restores_expert_parallel_next_step( + tmp_path: Path, topology: Literal["ep", "etp"] +) -> None: + original = _portable_trainer( + _moe_lora(num_local_experts=2), + model_identifier=str(_MOE_CONFIG["base_model_name_or_path"]), + ) + _install_checkpoint(original, _MOE_ADAPTER, _MOE_CONFIG) + source = tmp_path / f"{topology}-source" + _save_before_next_step(original, source) + output = tmp_path / f"{topology}-output" + + mp.spawn( + _advanced_topology_worker, + args=( + 2, + f"file://{tmp_path / f'{topology}-init'}", + str(source), + str(output), + topology, + None, + ), + nprocs=2, + join=True, + ) + + restored = _portable_trainer( + _moe_lora(num_local_experts=2), + model_identifier=str(_MOE_CONFIG["base_model_name_or_path"]), + ) + load_trainer_checkpoint(restored, prepare_checkpoint(str(output)), "student") + _assert_checkpoint_state_equal(original, restored) + + +def test_checkpoint_deduplicates_context_parallel_replicas( + tmp_path: Path, +) -> None: + original = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + _install_portable_checkpoint(original) + source = tmp_path / "context-source" + _save_before_next_step(original, source) + output = tmp_path / "context-output" + counts = tmp_path / "context-counts.json" + + mp.spawn( + _advanced_topology_worker, + args=( + 2, + f"file://{tmp_path / 'context-init'}", + str(source), + str(output), + "cp", + str(counts), + ), + nprocs=2, + join=True, + ) + + restored = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + load_trainer_checkpoint(restored, prepare_checkpoint(str(output)), "student") + _assert_checkpoint_state_equal(original, restored) + assert json.loads(counts.read_text()) == [0, 0] @pytest.mark.parametrize( @@ -1678,12 +3242,10 @@ def test_direct_3d_packed_expert_publish_matches_handler_vllm_exactly( down_lora.B_T.data[expert].copy_(tensors["down_proj.lora_B.weight"].T) offset += 1000 - slot_ref = LoRASlotRef("checkpoint", "student") if dynamic_slot else None + slot_ref = LoRASlotRef("student") if dynamic_slot else None if slot_ref is not None: - assert gate_up_lora.load_lora_slot( - slot_ref, full, alpha=rank, requires_grad=True - ) - assert down_lora.load_lora_slot(slot_ref, full, alpha=rank, requires_grad=True) + assert gate_up_lora.load_lora_slot(slot_ref, full, alpha=rank) + assert down_lora.load_lora_slot(slot_ref, full, alpha=rank) adapter_config = _config(base_model, rank=rank, alpha=rank) old_dir = tmp_path / "old" @@ -1826,7 +3388,7 @@ def test_qwen35_megatron_shards_can_merge_to_separate_vllm_checkpoint( publish_dir = tmp_path / "published" adapter_config = _config("Qwen/Qwen3.5-35B-A3B", rank=1, alpha=1) entries_by_key = { - key: [({"sharded": False, "shard_world_size": 1, "shard_rank": 0}, tensor)] + key: [(_manifest(sharded=False, shard_world_size=1, shard_rank=0), tensor)] for key, tensor in full.items() } merged = merge_sharded_adapter_entries(entries_by_key) diff --git a/tests/integration/megatron/model_support/oracle_worker.py b/tests/integration/megatron/model_support/oracle_worker.py index 014871667..5529c3fac 100644 --- a/tests/integration/megatron/model_support/oracle_worker.py +++ b/tests/integration/megatron/model_support/oracle_worker.py @@ -195,9 +195,10 @@ def provider_topology_env(topology: Topology): def _merge_sharded_dicts(shards_by_rank: list[dict[str, Any]]) -> dict[str, Any]: """Merges rank-sharded LoRA tensors into a full state dict on rank 0.""" + from art.megatron.lora import LoraShardManifest from art.megatron.weights.lora_publish import merge_sharded_adapter_entries - entries_by_key: dict[str, list[tuple[dict[str, Any], torch.Tensor]]] = {} + entries_by_key: dict[str, list[tuple[LoraShardManifest, torch.Tensor]]] = {} for rank_entry in shards_by_rank: rank_state = rank_entry["state"] rank_manifest = rank_entry["manifest"] @@ -205,7 +206,7 @@ def _merge_sharded_dicts(shards_by_rank: list[dict[str, Any]]) -> dict[str, Any] if key not in rank_manifest: raise RuntimeError(f"Missing manifest entry for sharded key '{key}'") entries_by_key.setdefault(key, []).append( - (rank_manifest[key], tensor.detach().cpu()) + (cast(LoraShardManifest, rank_manifest[key]), tensor.detach().cpu()) ) return merge_sharded_adapter_entries(entries_by_key) diff --git a/tests/integration/megatron/model_support/test_provider_support.py b/tests/integration/megatron/model_support/test_provider_support.py index 43a641452..3c0aeeba9 100644 --- a/tests/integration/megatron/model_support/test_provider_support.py +++ b/tests/integration/megatron/model_support/test_provider_support.py @@ -1,13 +1,16 @@ from __future__ import annotations +from pathlib import Path from types import SimpleNamespace -from typing import Any, cast +from typing import Any, Protocol, cast import pytest import torch pytest.importorskip("megatron.bridge") +from megatron.bridge import AutoBridge +from megatron.bridge.models.gpt_provider import GPTModelProvider from megatron.core.transformer.enums import AttnBackend from art.megatron.context_parallel.core_attention import ArtContextParallelCoreAttention @@ -24,9 +27,29 @@ from art.megatron.runtime.bridge_runtime import load_unique_hf_keys_once -class _FakeProvider: +class _CoreAttentionSubmodules(Protocol): + core_attention: object + + +class _SelfAttentionSpec(Protocol): + submodules: _CoreAttentionSubmodules + + +class _TransformerLayerSubmodules(Protocol): + self_attention: _SelfAttentionSpec + + +class _TransformerLayerSpec(Protocol): + submodules: _TransformerLayerSubmodules + + +class _TransformerLayerSpecFactory(Protocol): + def __call__(self, provider: object, *, vp_stage: int) -> _TransformerLayerSpec: ... + + +class _FakeProvider(GPTModelProvider): def __init__(self) -> None: - self.transformer_layer_spec = self._base_layer_spec + cast(Any, self).transformer_layer_spec = self._base_layer_spec self.finalized = False self.overlap_moe_expert_parallel_comm = False self.moe_shared_expert_overlap = False @@ -39,7 +62,7 @@ def __init__(self) -> None: self.window_size: int | tuple[int, int] = (128, 0) self.moe_hybridep_num_sms = 16 self.moe_flex_dispatcher_backend = "hybridep" - self.moe_token_dispatcher_type = "" + cast(Any, self).moe_token_dispatcher_type = "" self.recompute_granularity: str | None = None self.recompute_method: str | None = None self.recompute_num_layers: int | None = None @@ -101,13 +124,20 @@ def finalize(self) -> None: self.finalized = True -class _FakeBridge: +class _FakeBridge(AutoBridge): def __init__(self, *, model_bridge: object, provider: _FakeProvider) -> None: - self._model_bridge = model_bridge + self._fake_model_bridge = model_bridge self._provider = provider - self.hf_pretrained = SimpleNamespace(model_name_or_path="unused") + cast(Any, self).hf_pretrained = SimpleNamespace(model_name_or_path="unused") - def to_megatron_provider(self) -> _FakeProvider: + @property + def _model_bridge(self) -> Any: + return self._fake_model_bridge + + def to_megatron_provider( + self, load_weights: bool = True, hf_path: str | Path | None = None + ) -> _FakeProvider: + del load_weights, hf_path return self._provider @@ -431,7 +461,9 @@ def test_get_provider_bundle_honors_single_gpu_env_topology( assert resolved.recompute_method == "uniform" assert resolved.recompute_num_layers == 1 - layer_spec = resolved.transformer_layer_spec(resolved, vp_stage=0) + layer_spec = cast(_TransformerLayerSpecFactory, resolved.transformer_layer_spec)( + resolved, vp_stage=0 + ) assert ( layer_spec.submodules.self_attention.submodules.core_attention is FlexDotProductAttention @@ -464,7 +496,9 @@ def test_get_provider_bundle_honors_context_parallel_env_topology( assert resolved.context_parallel_size == 2 assert resolved.expert_model_parallel_size == 1 assert resolved.expert_tensor_parallel_size == 1 - layer_spec = resolved.transformer_layer_spec(resolved, vp_stage=0) + layer_spec = cast(_TransformerLayerSpecFactory, resolved.transformer_layer_spec)( + resolved, vp_stage=0 + ) assert ( layer_spec.submodules.self_attention.submodules.core_attention is ArtContextParallelCoreAttention diff --git a/tests/integration/megatron/train_inf_mismatch/output_parity.py b/tests/integration/megatron/train_inf_mismatch/output_parity.py index d4c6d70bd..97df2a45a 100644 --- a/tests/integration/megatron/train_inf_mismatch/output_parity.py +++ b/tests/integration/megatron/train_inf_mismatch/output_parity.py @@ -899,14 +899,22 @@ def _build_deterministic_nonzero_lora( def _merge_sharded_lora(shards_by_rank: list[dict[str, Any]]) -> dict[str, Any]: + import torch + + from art.megatron.lora import LoraShardManifest from art.megatron.weights.lora_publish import merge_sharded_adapter_entries - entries_by_key: dict[str, list[tuple[dict[str, Any], Any]]] = {} + entries_by_key: dict[str, list[tuple[LoraShardManifest, torch.Tensor]]] = {} for rank_entry in shards_by_rank: state = rank_entry["state"] manifest = rank_entry["manifest"] for key, tensor in state.items(): - entries_by_key.setdefault(key, []).append((manifest[key], tensor)) + entries_by_key.setdefault(key, []).append( + ( + cast(LoraShardManifest, manifest[key]), + cast(torch.Tensor, tensor), + ) + ) return merge_sharded_adapter_entries(entries_by_key) diff --git a/tests/unit/test_megatron_reference_logprobs.py b/tests/unit/test_megatron_reference_logprobs.py index 1ef1c5f1f..40cf4b22d 100644 --- a/tests/unit/test_megatron_reference_logprobs.py +++ b/tests/unit/test_megatron_reference_logprobs.py @@ -8,6 +8,7 @@ from art import types from art.megatron import train as megatron_train +from art.megatron.model_support.spec import ModelSupportHandler from art.megatron.training import microbatches as megatron_microbatches from art.preprocessing.pack import PackedTensors @@ -173,7 +174,7 @@ def test_calculate_megatron_logprobs_replays_routes(monkeypatch) -> None: logprobs = megatron_train._calculate_megatron_logprobs( model_chunks=cast(megatron_train.ModelChunks, [chunk]), provider=object(), - model_support_handler=_Handler(), + model_support_handler=cast(ModelSupportHandler, _Handler()), inputs=_packed_inputs(), moe_routing_replay_controller=cast( megatron_train.MoeRoutingReplayController, controller diff --git a/tests/unit/test_trainer_rank_validation.py b/tests/unit/test_trainer_rank_validation.py index 042dad3c1..bf4c5596b 100644 --- a/tests/unit/test_trainer_rank_validation.py +++ b/tests/unit/test_trainer_rank_validation.py @@ -1,10 +1,12 @@ from __future__ import annotations +import asyncio from collections.abc import Iterable from dataclasses import dataclass import gc from importlib.util import find_spec import inspect +from pathlib import Path from types import SimpleNamespace from typing import TYPE_CHECKING, Any, cast @@ -20,11 +22,14 @@ TopK, TrainerRank, TrainerRankMemoryError, - TrainerRankOptimizerLayout, - TrainerRankOptimizerState, TrainerRankSlotStateError, Unset, ) +from art.trainer_rank._checkpoint import ( + AdamWRecord, + LocalOptimizerState, + _validate_save_state, +) from art.trainer_rank._impl import ( _anchor_disconnected_outputs, _MemoryCheck, @@ -46,8 +51,6 @@ def test_public_types_have_canonical_module_paths() -> None: assert { "AdapterSelection", - "TrainerRankOptimizerLayout", - "TrainerRankOptimizerState", "Unset", } <= set(art.trainer_rank.__all__) for public_type in ( @@ -57,8 +60,6 @@ def test_public_types_have_canonical_module_paths() -> None: TopK, TrainerRank, TrainerRankMemoryError, - TrainerRankOptimizerLayout, - TrainerRankOptimizerState, TrainerRankSlotStateError, ): assert public_type.__module__ == "art.trainer_rank" @@ -99,7 +100,6 @@ def zero_grad(self) -> None: @dataclass(frozen=True) class _SlotRef: - kind: str name: str | None @@ -122,14 +122,21 @@ def _runtime( model_support_handler=SimpleNamespace( build_gdn_execution_spec=True, canonicalize_loaded_lora_state=lambda state, _model: state, + from_vllm_lora_tensors=lambda state, **_kwargs: state, + to_vllm_lora_tensors=lambda state, **kwargs: ( + state, + kwargs["adapter_config"], + ), zero_internal_padding_grads=lambda _model: None, zero_internal_padding_params=lambda _model: None, ), + rank=0, + world_size=1, ) # type: ignore -def _slot_ref(kind: str, name: str | None) -> "LoRASlotRef": - return _SlotRef(kind, name) # type: ignore +def _slot_ref(name: str | None) -> "LoRASlotRef": + return _SlotRef(name) # type: ignore def _target_request(token: int) -> ForwardInput[torch.Tensor, None, None, None]: @@ -210,8 +217,7 @@ def _tracked_targets( def test_forward_input_validation() -> None: with pytest.raises(ValueError, match="top_k must be >= 1"): ForwardInput(input_tokens=torch.tensor([1]), top_k=0) - with pytest.raises(ValueError, match="cannot set both checkpoint and lora"): - ForwardInput(input_tokens=torch.tensor([1]), checkpoint="a", lora="b") + assert "lora" not in ForwardInput.__dataclass_fields__ with pytest.raises(ValueError, match="top_k=9 exceeds vocabulary size 8"): _validate_top_k(9, _Model()) @@ -223,7 +229,40 @@ def test_forward_input_distinguishes_unset_and_base_checkpoint( request = ForwardInput(input_tokens=torch.tensor([1]), checkpoint=checkpoint) assert request.checkpoint is expected - assert request.lora is Unset + + +def test_dp_rank_forward_rejects_unloaded_explicit_checkpoint( + monkeypatch: pytest.MonkeyPatch, +) -> None: + trainer = TrainerRank(_runtime()) + _stub_forward(monkeypatch, trainer) + request = ForwardInput( + input_tokens=torch.tensor([1]), + target_tokens=torch.tensor([1]), + checkpoint="typo", + ) + + with pytest.raises(TrainerRankSlotStateError, match="unloaded.*'typo'"): + trainer.dp_rank_forward([request]) + + +@pytest.mark.parametrize("checkpoint", (None, "student")) +def test_dp_rank_forward_accepts_base_or_loaded_explicit_checkpoint( + monkeypatch: pytest.MonkeyPatch, + checkpoint: str | None, +) -> None: + trainer = TrainerRank(_runtime()) + trainer._checkpoint_slot_params_by_name["student"] = () + _stub_forward(monkeypatch, trainer) + request = ForwardInput( + input_tokens=torch.tensor([1]), + target_tokens=torch.tensor([1]), + checkpoint=checkpoint, + ) + + output = trainer.dp_rank_forward([request]) + + assert isinstance(output[0], ForwardOutput) def test_forward_input_preserves_public_runtime_shape() -> None: @@ -388,15 +427,251 @@ def test_hybridep_rejects_buffer_growth_with_live_graph( ) -def test_trainer_rank_adapter_stack_errors() -> None: +async def test_trainer_rank_checkpoint_stack_errors() -> None: trainer = TrainerRank(_runtime()) - with pytest.raises(RuntimeError, match="No pushed LoRA or checkpoint"): - trainer.pop_pushed_lora_or_checkpoint() + with pytest.raises(RuntimeError, match="No pushed checkpoint"): + trainer.pop_checkpoint() trainer._slot_stack.append(object()) # type: ignore - for load in (trainer.load_checkpoint_slot, trainer.load_lora_slot): - with pytest.raises(RuntimeError, match="Cannot load a LoRA/checkpoint"): - load("teacher", {}) + with pytest.raises(RuntimeError, match="Cannot load a checkpoint"): + await trainer.load_checkpoint("teacher") + + +async def test_checkpoint_tasks_and_async_context( + monkeypatch: pytest.MonkeyPatch, +) -> None: + trainer = TrainerRank(_runtime()) + trainer._checkpoint_sources = {} + fetched: list[str] = [] + + async def prefetch(path: str) -> object: + fetched.append(path) + return object() + + def install(trainer: TrainerRank, _source: object, path: str) -> None: + trainer._checkpoint_slot_params_by_name[path] = () + trainer._checkpoint_revisions[path] = 0 + + monkeypatch.setattr(trainer, "_prefetch_checkpoint", prefetch) + monkeypatch.setattr("art.trainer_rank._checkpoint.load_checkpoint", install) + + task = trainer.load_checkpoint("student") + assert isinstance(task, asyncio.Task) + await task + assert fetched == ["student"] + assert trainer._default_slot_ref == trainer._slot_ref("student") + + task = trainer.prefetch_checkpoints("teacher", "reference") + assert isinstance(task, asyncio.Task) + await task + assert fetched[-2:] == ["teacher", "reference"] + + pushed = trainer.push_checkpoint("student") + await pushed + assert trainer._slot_stack == [trainer._slot_ref("student")] + trainer.pop_checkpoint() + async with trainer.push_checkpoint("student"): + assert trainer._slot_stack == [trainer._slot_ref("student")] + async with trainer.push_checkpoint("missing"): + assert trainer._slot_stack == [ + trainer._slot_ref("student"), + trainer._slot_ref("missing"), + ] + assert trainer._slot_stack == [trainer._slot_ref("student")] + assert trainer._slot_stack == [] + assert fetched[-1] == "missing" + + +async def test_checkpoint_context_cancellation_after_successful_push_cleans_stack( + monkeypatch: pytest.MonkeyPatch, +) -> None: + trainer = TrainerRank(_runtime()) + trainer._checkpoint_slot_params_by_name["student"] = () + parent = asyncio.current_task() + assert parent is not None + original_slot_ref = trainer._slot_ref + cancellation_scheduled = False + + def cancel_parent_after_resolving(path: str | None): + nonlocal cancellation_scheduled + ref = original_slot_ref(path) + if not cancellation_scheduled: + cancellation_scheduled = True + asyncio.get_running_loop().call_soon(parent.cancel) + return ref + + monkeypatch.setattr(trainer, "_slot_ref", cancel_parent_after_resolving) + pushed = trainer.push_checkpoint("student") + entered = False + + with pytest.raises(asyncio.CancelledError): + async with pushed: + entered = True + + assert cancellation_scheduled + assert pushed.task.done() and not pushed.task.cancelled() + assert not entered + assert trainer._slot_stack == [] + + +async def test_shared_checkpoint_prefetch_survives_waiter_cancellation( + monkeypatch: pytest.MonkeyPatch, +) -> None: + trainer = TrainerRank(_runtime()) + started = asyncio.Event() + release = asyncio.Event() + source = object() + + async def delayed_to_thread(_function: object, *_args: object) -> object: + started.set() + await release.wait() + return source + + monkeypatch.setattr(asyncio, "to_thread", delayed_to_thread) + first = asyncio.create_task(trainer._prefetch_checkpoint("student")) + second = asyncio.create_task(trainer._prefetch_checkpoint("student")) + await started.wait() + first.cancel() + with pytest.raises(asyncio.CancelledError): + await first + release.set() + + assert await second is source + assert ( + trainer._checkpoint_sources[trainer._checkpoint_source_key("student")] is source + ) + assert trainer._checkpoint_prefetch_tasks == {} + + +async def test_materialized_sources_keep_logical_checkpoint_identities( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path +) -> None: + trainer = TrainerRank(_runtime()) + prepared: list[str] = [] + installed: list[tuple[str, object]] = [] + + def prepare(source_path: str) -> object: + prepared.append(source_path) + return object() + + def install(trainer: TrainerRank, source: object, logical_path: str) -> None: + installed.append((logical_path, source)) + trainer._checkpoint_slot_params_by_name[logical_path] = () + trainer._checkpoint_revisions[logical_path] = 0 + + monkeypatch.setattr("art.trainer_rank._checkpoint.prepare_checkpoint", prepare) + monkeypatch.setattr("art.trainer_rank._checkpoint.load_checkpoint", install) + root_a = str(tmp_path / "immutable-a") + root_b = str(tmp_path / "immutable-b") + logical_a = "wandb-artifact:///entity/project/run:step1" + logical_b = "wandb-artifact:///entity/project/run-teacher:step1" + logical_c = "wandb-artifact:///entity/project/run-reference:step1" + + await trainer._prefetch_checkpoints_from_sources( + (logical_a, root_a), (logical_b, root_a), (logical_c, root_b) + ) + assert sorted(prepared) == sorted( + (trainer._checkpoint_source_key(root_a), trainer._checkpoint_source_key(root_b)) + ) + + await asyncio.gather( + trainer._load_checkpoint_from_source(logical_a, root_a), + trainer._load_checkpoint_from_source(logical_b, root_a), + trainer._load_checkpoint_from_source(logical_c, root_b), + ) + assert [logical_path for logical_path, _source in installed] == [ + logical_a, + logical_b, + logical_c, + ] + assert set(trainer._checkpoint_slot_params_by_name) == { + logical_a, + logical_b, + logical_c, + } + assert trainer._default_slot_ref == trainer._slot_ref(logical_c) + for logical_path in (logical_a, logical_b, logical_c): + request = ForwardInput(input_tokens=torch.tensor([1]), checkpoint=logical_path) + assert trainer._resolve_slot_ref(request) == trainer._slot_ref(logical_path) + + refreshed_root = str(tmp_path / "immutable-new") + await trainer._load_checkpoint_from_source(logical_a, refreshed_root) + assert installed[-1][0] == logical_a + assert prepared[-1] == trainer._checkpoint_source_key(refreshed_root) + + prepared_before_push = tuple(prepared) + pushed = trainer._push_checkpoint_from_source( + logical_a, str(tmp_path / "unused-while-loaded") + ) + await pushed + assert tuple(prepared) == prepared_before_push + assert trainer._slot_stack == [trainer._slot_ref(logical_a)] + trainer.pop_checkpoint() + + +async def test_checkpoint_mutations_follow_call_order_and_recover_from_failure( + monkeypatch: pytest.MonkeyPatch, +) -> None: + trainer = TrainerRank(_runtime()) + trainer._checkpoint_sources = {} + started = {name: asyncio.Event() for name in ("first", "second")} + ready = {name: asyncio.Event() for name in ("first", "second")} + installed: list[str] = [] + + async def prefetch(path: str) -> object: + event = started.get(path) + if event is not None: + event.set() + await ready[path].wait() + return object() + + def install(trainer: TrainerRank, _source: object, path: str) -> None: + installed.append(path) + if path == "bad": + raise RuntimeError("injected load failure") + trainer._checkpoint_slot_params_by_name[path] = () + trainer._checkpoint_revisions[path] = 0 + + monkeypatch.setattr(trainer, "_prefetch_checkpoint", prefetch) + monkeypatch.setattr("art.trainer_rank._checkpoint.load_checkpoint", install) + first = trainer.load_checkpoint("first") + second = trainer.load_checkpoint("second") + await asyncio.gather(*(event.wait() for event in started.values())) + ready["second"].set() + await asyncio.sleep(0) + assert installed == [] + ready["first"].set() + await asyncio.gather(first, second) + assert installed == ["first", "second"] + + bad = trainer.load_checkpoint("bad") + after = trainer.load_checkpoint("after") + with pytest.raises(RuntimeError, match="injected load failure"): + await bad + await after + assert installed[-2:] == ["bad", "after"] + + +@pytest.mark.parametrize("local_failure", [False, True]) +async def test_checkpoint_prefetch_failures_are_coordinated( + monkeypatch: pytest.MonkeyPatch, local_failure: bool +) -> None: + trainer = TrainerRank(_runtime()) + + async def prefetch(_path: str) -> object: + if local_failure: + raise OSError("rank-local prefetch failed") + return object() + + def coordinated(error: BaseException | None, phase: str) -> None: + assert phase == "prepare checkpoint" + assert isinstance(error, OSError) is local_failure + raise RuntimeError("a rank failed to prepare checkpoint") + + monkeypatch.setattr(trainer, "_prefetch_checkpoint", prefetch) + monkeypatch.setattr("art.trainer_rank._checkpoint._raise_distributed", coordinated) + with pytest.raises(RuntimeError, match="a rank failed"): + await trainer.load_checkpoint("student") def test_trainer_rank_rejects_adapter_keys_without_installed_lora_site() -> None: @@ -405,11 +680,10 @@ def test_trainer_rank_rejects_adapter_keys_without_installed_lora_site() -> None "base.layer.lora_A.weight": torch.empty(1), "base.layer.lora_B.weight": torch.empty(1), } - trainer._prepare_adapter_model("checkpoint", "student", valid) + trainer._prepare_adapter_model("student", valid) with pytest.raises(ValueError, match="matching LoRA target modules"): trainer._prepare_adapter_model( - "checkpoint", "student", {**valid, "base.other.lora_A.weight": torch.empty(1)}, ) @@ -423,7 +697,7 @@ def test_trainer_rank_normalizes_adapter_tensors_to_installed_site() -> None: "base.layer.lora_B.weight": torch.ones(5, 3, dtype=torch.float32), } - normalized = trainer._prepare_adapter_model("checkpoint", "student", adapter) + normalized = trainer._prepare_adapter_model("student", adapter) assert all(tensor.device == site.A_T.device for tensor in normalized.values()) assert all(tensor.dtype == torch.bfloat16 for tensor in normalized.values()) @@ -438,20 +712,16 @@ def test_checkpoint_slot_adapter_config_is_validated_and_copied() -> None: "target_modules": ["q_proj"], } - retained = trainer._validate_checkpoint_slot_adapter_config( - "student", config, alpha=16 - ) + retained = trainer._validate_checkpoint_adapter_config("student", config, alpha=16) assert retained == config config["target_modules"].append("v_proj") # type: ignore[union-attr] assert retained is not None assert retained["target_modules"] == ["q_proj"] with pytest.raises(ValueError, match="conflicts"): - trainer._validate_checkpoint_slot_adapter_config("student", config, alpha=32) + trainer._validate_checkpoint_adapter_config("student", config, alpha=32) with pytest.raises(ValueError, match="missing"): - trainer._validate_checkpoint_slot_adapter_config( - "student", {"r": 8}, alpha=None - ) + trainer._validate_checkpoint_adapter_config("student", {"r": 8}, alpha=None) @pytest.mark.parametrize( @@ -477,7 +747,7 @@ def test_checkpoint_slot_adapter_config_rejects_invalid_field_types( config[field] = value with pytest.raises(TypeError, match=field): - trainer._validate_checkpoint_slot_adapter_config("student", config, alpha=None) + trainer._validate_checkpoint_adapter_config("student", config, alpha=None) def test_checkpoint_slot_adapter_config_rejects_cross_rank_mismatch( @@ -493,37 +763,7 @@ def gather(output: list[object], value: object) -> None: monkeypatch.setattr("art.trainer_rank.dist.all_gather_object", gather) with pytest.raises(ValueError, match="differs across ranks"): - trainer._validate_checkpoint_slot_adapter_config("student", None, alpha=None) - - -def test_load_checkpoint_slot_retains_config_and_uses_its_alpha( - monkeypatch: pytest.MonkeyPatch, -) -> None: - trainer = TrainerRank(_runtime()) - seen: dict[str, object] = {} - monkeypatch.setattr( - trainer, - "_load_slot", - lambda *_args, **kwargs: seen.update(kwargs) or 1, - ) - monkeypatch.setattr(trainer, "_validate_dynamic_slot_consistency", lambda *_: ()) - monkeypatch.setattr( - trainer, "_validate_loaded_checkpoint_slot_config", lambda *_: None - ) - config = { - "base_model_name_or_path": "Qwen/Qwen3-8B", - "r": 8, - "lora_alpha": 16, - "target_modules": ["q_proj"], - } - - trainer.load_checkpoint_slot("student", {}, adapter_config=config) - - assert seen["alpha"] == 16 - assert trainer._checkpoint_slot_adapter_configs["student"] == config - trainer.load_checkpoint_slot("student", {}, alpha=7) - assert seen["alpha"] == 7 - assert "student" not in trainer._checkpoint_slot_adapter_configs + trainer._validate_checkpoint_adapter_config("student", None, alpha=None) @pytest.mark.skipif(find_spec("megatron") is None, reason="requires Megatron") @@ -533,12 +773,20 @@ def test_slot_load_canonicalizes_only_local_incoming_adapter( calls: list[tuple[dict[str, torch.Tensor], object]] = [] loaded_state: dict[str, torch.Tensor] = {} runtime = _runtime() - runtime.model_support_handler.canonicalize_loaded_lora_state = lambda state, model: ( - calls.append((state, model)) - or {key: torch.zeros_like(value) for key, value in state.items()} + monkeypatch.setattr( + runtime.model_support_handler, + "canonicalize_loaded_lora_state", + lambda state, model: ( + calls.append((state, model)) + or {key: torch.zeros_like(value) for key, value in state.items()} + ), ) - runtime.model_support_handler.zero_internal_padding_params = lambda _model: ( - pytest.fail("slot load must not mutate unrelated slot parameters") + monkeypatch.setattr( + runtime.model_support_handler, + "zero_internal_padding_params", + lambda _model: pytest.fail( + "slot load must not mutate unrelated slot parameters" + ), ) trainer = TrainerRank(runtime) monkeypatch.setattr( @@ -569,7 +817,7 @@ def load_slot( ) adapter = {"weight": torch.ones(1), "remote_weight": torch.ones(1)} - trainer._load_slot("checkpoint", "student", adapter, trainable=True, alpha=None) + trainer._load_checkpoint_slot("student", adapter, alpha=1.0) assert calls == [({"weight": adapter["weight"]}, runtime.model)] torch.testing.assert_close(loaded_state["weight"], torch.zeros(1)) @@ -577,15 +825,30 @@ def load_slot( torch.testing.assert_close(adapter["weight"], torch.ones(1)) -def test_checkpoint_slot_publish_requires_retained_adapter_config() -> None: +def test_checkpoint_export_requires_retained_adapter_config() -> None: trainer = TrainerRank(_runtime()) - with pytest.raises(ValueError, match="Unknown checkpoint slot"): - trainer.save_checkpoint_slot_lora("missing", "/unused") + with pytest.raises(ValueError, match="Unknown checkpoint"): + trainer.export_lora("/unused", "missing") trainer._checkpoint_slot_params_by_name["student"] = () - with pytest.raises(TrainerRankSlotStateError, match="adapter_config"): - trainer.save_checkpoint_slot_lora("student", "/unused") + trainer.export_lora("/unused", "student") + + +def test_checkpoint_save_rejects_accumulated_gradients() -> None: + trainer = TrainerRank(_runtime()) + parameter = torch.nn.Parameter(torch.ones(2)) + parameter.grad = torch.ones_like(parameter) + trainer._checkpoint_slot_params_by_name["student"] = (parameter,) + trainer._checkpoint_slot_adapter_configs["student"] = { + "base_model_name_or_path": "test", + "r": 1, + "lora_alpha": 1, + "target_modules": [], + } + + with pytest.raises(TrainerRankSlotStateError, match="accumulated gradients"): + _validate_save_state(trainer, "student") def test_trainer_rank_default_forward_uses_explicit_base_slot() -> None: @@ -596,7 +859,6 @@ def test_trainer_rank_default_forward_uses_explicit_base_slot() -> None: assert len(plan.groups) == 1 slot = plan.groups[0].slot_ref assert slot is not None - assert getattr(slot, "kind") == "checkpoint" assert getattr(slot, "name") is None @@ -687,9 +949,17 @@ def zero_padding_grads(_model: object) -> None: assert param.grad is not None param.grad[-1] = 0.0 - runtime.model_support_handler.zero_internal_padding_grads = zero_padding_grads - runtime.model_support_handler.zero_internal_padding_params = lambda _model: ( - pytest.fail("slot step must not mutate unrelated slot parameters") + monkeypatch.setattr( + runtime.model_support_handler, + "zero_internal_padding_grads", + zero_padding_grads, + ) + monkeypatch.setattr( + runtime.model_support_handler, + "zero_internal_padding_params", + lambda _model: pytest.fail( + "slot step must not mutate unrelated slot parameters" + ), ) trainer = TrainerRank(runtime) trainer._checkpoint_slot_params_by_name["student"] = (param,) @@ -711,7 +981,7 @@ def zero_padding_grads(_model: object) -> None: assert param[-1].item() == 0.0 -def test_checkpoint_slot_optimizer_state_reproduces_exact_next_step( +def test_canonical_optimizer_state_reproduces_exact_next_step( monkeypatch: pytest.MonkeyPatch, ) -> None: adam = AdamParams( @@ -721,19 +991,34 @@ def test_checkpoint_slot_optimizer_state_reproduces_exact_next_step( weight_decay=0.1, grad_clip_norm=10.0, ) - original, original_param = _trainer_with_checkpoint( monkeypatch, torch.tensor([0.5, -0.25], dtype=torch.bfloat16) ) + original._checkpoint_revisions["student"] = 0 original_param.grad = torch.tensor([0.2, -0.4], dtype=torch.bfloat16) original.optim_step(params=adam) - state = original.checkpoint_slot_optimizer_state("student") - assert state is not None + dynamic = original._dynamic_optimizers["student"] + optimizer_state = dynamic.optimizer.state[dynamic.master_params[0]] + group = dynamic.optimizer.param_groups[0] + beta1, beta2 = cast(tuple[float, float], group["betas"]) + state = LocalOptimizerState( + masters=tuple(param.detach().clone() for param in dynamic.master_params), + exp_avgs=(cast(torch.Tensor, optimizer_state["exp_avg"]).clone(),), + exp_avg_sqs=(cast(torch.Tensor, optimizer_state["exp_avg_sq"]).clone(),), + steps=(float(cast(torch.Tensor, optimizer_state["step"]).item()),), + config=AdamWRecord( + learning_rate=float(group["lr"]), + beta1=beta1, + beta2=beta2, + eps=float(group["eps"]), + weight_decay=float(group["weight_decay"]), + ), + ) restored, restored_param = _trainer_with_checkpoint( monkeypatch, original_param.detach() ) - restored._dynamic_optimizers["student"] = restored._restore_dynamic_optimizer( + restored._dynamic_optimizers["student"] = restored._restore_canonical_optimizer( "student", state ) for param in (original_param, restored_param): @@ -742,10 +1027,11 @@ def test_checkpoint_slot_optimizer_state_reproduces_exact_next_step( restored.optim_step(params=adam) torch.testing.assert_close(restored_param, original_param, atol=0, rtol=0) - original_state = original.checkpoint_slot_optimizer_state("student") - restored_state = restored.checkpoint_slot_optimizer_state("student") - assert original_state is not None and restored_state is not None - _assert_nested_tensors_equal(restored_state, original_state) + _assert_nested_tensors_equal( + restored._dynamic_optimizers["student"].optimizer.state_dict(), + original._dynamic_optimizers["student"].optimizer.state_dict(), + ) + assert original._checkpoint_revisions["student"] == 2 def test_dynamic_optimizer_keeps_fp32_master_weight_and_moments( @@ -773,35 +1059,26 @@ def test_dynamic_optimizer_keeps_fp32_master_weight_and_moments( assert state["exp_avg_sq"].dtype == torch.float32 -@pytest.mark.parametrize( - ("corruption", "error"), - ( - ("layout", "topology or parameter layout"), - ("missing_master", "master parameters"), - ("shape", "topology or parameter layout"), - ), -) -def test_checkpoint_slot_optimizer_state_rejects_incompatible_state( - corruption: str, - error: str, +def test_canonical_optimizer_rejects_incompatible_local_shape( monkeypatch: pytest.MonkeyPatch, ) -> None: - trainer, param = _trainer_with_checkpoint(monkeypatch, torch.ones(2)) - param.grad = torch.ones_like(param) - trainer.optim_step( - params=AdamParams(learning_rate=1e-2, weight_decay=0.0, grad_clip_norm=10.0) - ) - state = trainer.checkpoint_slot_optimizer_state("student") - assert state is not None - if corruption == "layout": - cast(dict[str, object], state)["layout"] = {"different": True} - elif corruption == "missing_master": - state["master_params"] = () - restored, _ = _trainer_with_checkpoint( - monkeypatch, torch.ones(3 if corruption == "shape" else 2) + trainer, _ = _trainer_with_checkpoint(monkeypatch, torch.ones(2)) + state = LocalOptimizerState( + masters=(torch.ones(3),), + exp_avgs=(torch.zeros(3),), + exp_avg_sqs=(torch.zeros(3),), + steps=(1.0,), + config=AdamWRecord( + learning_rate=1e-3, + beta1=0.9, + beta2=0.99, + eps=1e-8, + weight_decay=0.0, + ), ) - with pytest.raises(TrainerRankSlotStateError, match=error): - restored._restore_dynamic_optimizer("student", state) + + with pytest.raises(TrainerRankSlotStateError, match="master parameter shape"): + trainer._restore_canonical_optimizer("student", state) @pytest.mark.parametrize("operation", ("load", "step")) @@ -810,7 +1087,7 @@ def test_trainer_rank_rejects_mutating_slot_with_pending_graph( monkeypatch: pytest.MonkeyPatch, ) -> None: trainer = TrainerRank(_runtime()) - ref = _slot_ref("checkpoint", "teacher") + ref = _slot_ref("teacher") monkeypatch.setattr(trainer, "_slot_ref", _slot_ref) target = _tracked_targets(trainer, ref, 2)[0] guard = ( @@ -837,7 +1114,7 @@ def test_trainer_rank_step_allows_missing_slot_graph_bookkeeping( def test_trainer_rank_zero_grad_does_not_clear_live_slot_graphs() -> None: trainer = TrainerRank(_runtime()) - ref = _slot_ref("lora", "teacher") + ref = _slot_ref("teacher") output = ForwardOutput( None, TopK( @@ -858,7 +1135,7 @@ def test_trainer_rank_zero_grad_does_not_clear_live_slot_graphs() -> None: def test_trainer_rank_retained_backward_keeps_slot_graph_guard() -> None: trainer = TrainerRank(_runtime()) - ref = _slot_ref("checkpoint", "teacher") + ref = _slot_ref("teacher") target = _tracked_targets(trainer, ref, 2)[0] target.sum().backward(retain_graph=True) @@ -871,7 +1148,7 @@ def test_trainer_rank_retained_backward_keeps_slot_graph_guard() -> None: def test_trainer_rank_tracks_each_independent_output_graph() -> None: trainer = TrainerRank(_runtime()) - ref = _slot_ref("checkpoint", "teacher") + ref = _slot_ref("teacher") first, second = _tracked_targets(trainer, ref, 2, 3) first.sum().backward() @@ -884,7 +1161,7 @@ def test_trainer_rank_tracks_each_independent_output_graph() -> None: def test_trainer_rank_tracks_graph_after_output_is_replaced_by_loss() -> None: trainer = TrainerRank(_runtime()) - ref = _slot_ref("checkpoint", "teacher") + ref = _slot_ref("teacher") target = _tracked_targets(trainer, ref, 2)[0] loss = target.sum() del target @@ -899,7 +1176,7 @@ def test_trainer_rank_tracks_graph_after_output_is_replaced_by_loss() -> None: def test_trainer_rank_releases_abandoned_output_graph() -> None: trainer = TrainerRank(_runtime()) - ref = _slot_ref("checkpoint", "teacher") + ref = _slot_ref("teacher") target = _tracked_targets(trainer, ref, 2)[0] del target gc.collect() @@ -1102,6 +1379,7 @@ def test_forward_micro_batches_rejects_mismatched_replicated_counts( monkeypatch.setattr(trainer_rank.dist, "is_available", lambda: True) monkeypatch.setattr(trainer_rank.dist, "is_initialized", lambda: True) monkeypatch.setattr(trainer_rank.dist, "get_world_size", lambda: 2) + monkeypatch.setattr(trainer_rank.dist, "all_reduce", lambda *_args, **_kwargs: None) def gather(output, value): output[:] = [value, value + 1] diff --git a/tests/unit/test_trainer_rank_weird_shapes.py b/tests/unit/test_trainer_rank_weird_shapes.py index 509e102d5..1cbc33d2c 100644 --- a/tests/unit/test_trainer_rank_weird_shapes.py +++ b/tests/unit/test_trainer_rank_weird_shapes.py @@ -68,7 +68,6 @@ def _target_request( logits: bool = False, hidden_states: bool = False, checkpoint: AdapterSelection = Unset, - lora: AdapterSelection = Unset, ) -> ForwardInput: labels = ( tokens @@ -85,7 +84,6 @@ def _target_request( logits=logits, hidden_states=hidden_states, checkpoint=checkpoint, - lora=lora, ) @@ -442,13 +440,16 @@ def test_heterogeneous_slots_split_packing_without_losing_output_estimates( monkeypatch.setattr( TrainerRank, "_slot_ref", - staticmethod(lambda kind, name: (kind, name)), + staticmethod(lambda name: name), + ) + rank._default_slot_ref = rank._slot_ref("student") + rank._checkpoint_slot_params_by_name.update( + {"student": (), "teacher": (), "critic": ()} ) - rank.set_checkpoint("student") requests = [ _target_request(_tokens(1, 2, 3), top_k=3), _target_request(_tokens(1, 2, 4), checkpoint=None, logits=True), - _target_request(_tokens(1, 2, 5), lora="teacher", hidden_states=True), + _target_request(_tokens(1, 2, 5), checkpoint="teacher", hidden_states=True), _target_request(_tokens(1, 2, 6), checkpoint="critic", target_count=4), ] @@ -462,10 +463,10 @@ def test_heterogeneous_slots_split_packing_without_losing_output_estimates( assert signature == plan.signature assert plan.signature.slot_group_count == 4 assert {group.slot_ref for group in plan.groups} == { - ("checkpoint", "student"), - ("checkpoint", None), - ("lora", "teacher"), - ("checkpoint", "critic"), + "student", + None, + "teacher", + "critic", } From 118e34b4897acfd11072f9ad950840b5e11afa8a Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 11 Aug 2026 05:44:27 +0000 Subject: [PATCH 2/8] Test concurrent checkpoint prefetches --- tests/unit/test_trainer_rank_validation.py | 43 ++++++++++++++++++---- 1 file changed, 35 insertions(+), 8 deletions(-) diff --git a/tests/unit/test_trainer_rank_validation.py b/tests/unit/test_trainer_rank_validation.py index bf4c5596b..457d1f074 100644 --- a/tests/unit/test_trainer_rank_validation.py +++ b/tests/unit/test_trainer_rank_validation.py @@ -543,6 +543,34 @@ async def delayed_to_thread(_function: object, *_args: object) -> object: assert trainer._checkpoint_prefetch_tasks == {} +async def test_shared_checkpoint_prefetch_serves_successful_waiters( + monkeypatch: pytest.MonkeyPatch, +) -> None: + trainer = TrainerRank(_runtime()) + started = asyncio.Event() + release = asyncio.Event() + source = object() + calls = 0 + + async def delayed_to_thread(_function: object, *_args: object) -> object: + nonlocal calls + calls += 1 + started.set() + await release.wait() + return source + + monkeypatch.setattr(asyncio, "to_thread", delayed_to_thread) + first = asyncio.create_task(trainer._prefetch_checkpoint("student")) + second = asyncio.create_task(trainer._prefetch_checkpoint("student")) + await started.wait() + await asyncio.sleep(0) + release.set() + + assert await asyncio.gather(first, second) == [source, source] + assert calls == 1 + assert trainer._checkpoint_prefetch_tasks == {} + + async def test_materialized_sources_keep_logical_checkpoint_identities( monkeypatch: pytest.MonkeyPatch, tmp_path: Path ) -> None: @@ -567,17 +595,16 @@ def install(trainer: TrainerRank, source: object, logical_path: str) -> None: logical_b = "wandb-artifact:///entity/project/run-teacher:step1" logical_c = "wandb-artifact:///entity/project/run-reference:step1" - await trainer._prefetch_checkpoints_from_sources( - (logical_a, root_a), (logical_b, root_a), (logical_c, root_b) - ) - assert sorted(prepared) == sorted( - (trainer._checkpoint_source_key(root_a), trainer._checkpoint_source_key(root_b)) - ) - await asyncio.gather( trainer._load_checkpoint_from_source(logical_a, root_a), trainer._load_checkpoint_from_source(logical_b, root_a), - trainer._load_checkpoint_from_source(logical_c, root_b), + ) + assert prepared == [trainer._checkpoint_source_key(root_a)] + + await trainer._prefetch_checkpoints_from_sources((logical_c, root_b)) + await trainer._load_checkpoint_from_source(logical_c, root_b) + assert sorted(prepared) == sorted( + (trainer._checkpoint_source_key(root_a), trainer._checkpoint_source_key(root_b)) ) assert [logical_path for logical_path, _source in installed] == [ logical_a, From 98eecd8585c23135690849a59d7b1bdb9b481925 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 11 Aug 2026 06:05:10 +0000 Subject: [PATCH 3/8] Harden checkpoint restoration invariants --- src/art/trainer_rank/_checkpoint.py | 23 +++-- src/art/trainer_rank/_impl.py | 39 ++++---- .../megatron/lora/test_dynamic_lora_slots.py | 2 +- .../megatron/lora/test_lora_disk_codecs.py | 88 ++++++++++++++++++- 4 files changed, 117 insertions(+), 35 deletions(-) diff --git a/src/art/trainer_rank/_checkpoint.py b/src/art/trainer_rank/_checkpoint.py index 9d46e26b8..9e294b0df 100644 --- a/src/art/trainer_rank/_checkpoint.py +++ b/src/art/trainer_rank/_checkpoint.py @@ -508,10 +508,10 @@ def _export_lora_files( def _load_stage_lora( - trainer: TrainerRank, source: PreparedCheckpoint + trainer: TrainerRank, source: PreparedCheckpoint, adapter_rank: int ) -> dict[str, torch.Tensor]: if _native_checkpoint(source): - return _load_native_lora(trainer, source) + return _load_native_lora(trainer, source, adapter_rank) local_layers = { match.group("layer") for key in trainer._local_adapter_keys() @@ -546,6 +546,7 @@ def _native_checkpoint(source: PreparedCheckpoint) -> bool: def _local_tensor_plan( trainer: TrainerRank, + adapter_rank: int, ) -> dict[str, tuple[LoraShardManifest, tuple[int, ...]]]: from art.megatron.lora import LoRA @@ -559,7 +560,9 @@ def _local_tensor_plan( keys = module._expected_weight_keys(suffix) for expert, key in enumerate(keys): local = parameter[expert] if parameter.ndim == 3 else parameter - plan[key] = (manifest, tuple(reversed(local.shape))) + shape = list(reversed(local.shape)) + shape[0 if suffix == "lora_A" else -1] = adapter_rank + plan[key] = (manifest, tuple(shape)) return plan @@ -616,11 +619,11 @@ def _read_local_slice( def _load_native_lora( - trainer: TrainerRank, source: PreparedCheckpoint + trainer: TrainerRank, source: PreparedCheckpoint, adapter_rank: int ) -> dict[str, torch.Tensor]: from art.megatron.model_support.lora_disk import safe_open - plan = _local_tensor_plan(trainer) + plan = _local_tensor_plan(trainer, adapter_rank) tensors: dict[str, torch.Tensor] = {} with safe_open(source.path / "adapter_model.safetensors", framework="pt") as file: keys = set(file.keys()) @@ -645,7 +648,7 @@ def load_checkpoint( adapter_model: dict[str, torch.Tensor] = {} read_error: BaseException | None = None try: - adapter_model = _load_stage_lora(trainer, source) + adapter_model = _load_stage_lora(trainer, source, int(config["r"])) except BaseException as exc: read_error = exc _raise_distributed(read_error, "read checkpoint") @@ -698,7 +701,9 @@ def load_checkpoint( staged_params = tuple(trainer._iter_slot_parameters(temporary_ref)) trainer._checkpoint_slot_params_by_name[temporary_name] = staged_params if source.manifest is not None and source.manifest.optimizer is not None: - local_optimizer = _load_local_optimizer(trainer, source, temporary_name) + local_optimizer = _load_local_optimizer( + trainer, source, temporary_name, int(config["r"]) + ) dynamic = trainer._restore_canonical_optimizer( temporary_name, local_optimizer ) @@ -1325,10 +1330,11 @@ def _load_local_optimizer( trainer: TrainerRank, source: PreparedCheckpoint, name: str, + adapter_rank: int, ) -> LocalOptimizerState: manifest = source.manifest assert manifest is not None and manifest.optimizer is not None - plan = _local_tensor_plan(trainer) + plan = _local_tensor_plan(trainer, adapter_rank) localized: dict[str, tuple[torch.Tensor, ...]] = {} for component in ("master", "exp_avg", "exp_avg_sq"): records = { @@ -1484,6 +1490,7 @@ def _commit_output(temporary: Path, destination: Path, digest: str) -> None: existing_manifest.read_text() ) if existing.digest == digest: + _checkpoint_metadata(destination) return if any(destination.iterdir()): raise FileExistsError(f"Checkpoint output is not empty: {destination}") diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 0d7663661..df41c315b 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -2087,17 +2087,21 @@ def _has_live_slot_graph(self, ref: "LoRASlotRef") -> bool: return bool(self._slot_graphs().get(ref)) def _guard_slot_can_load(self, ref: "LoRASlotRef") -> None: - if not self._has_live_slot_graph(ref): - return - raise TrainerRankSlotStateError( - f"Cannot load checkpoint {ref.name!r} while outputs from an " - "earlier forward using that slot still have a live backward graph. " - "Activation checkpoint recompute resolves slots by name, so replacing " - "the slot before backward can compute gradients with different LoRA " - "weights than the original forward. Finish backward first; if the " - "forward was abandoned, release all references to its outputs; or load " - "the new weights under a different slot name." - ) + if self._has_live_slot_graph(ref): + raise TrainerRankSlotStateError( + f"Cannot load checkpoint {ref.name!r} while outputs from an " + "earlier forward using that slot still have a live backward graph. " + "Activation checkpoint recompute resolves slots by name, so replacing " + "the slot before backward can compute gradients with different LoRA " + "weights than the original forward. Finish backward first; if the " + "forward was abandoned, release all references to its outputs; or load " + "the new weights under a different slot name." + ) + if any(param.grad is not None for param in self._iter_slot_parameters(ref)): + raise TrainerRankSlotStateError( + f"Cannot load checkpoint {ref.name!r} with accumulated gradients. " + "Call optim_step() or zero_grad() before replacing the checkpoint." + ) def _guard_checkpoint_can_step(self, name: str) -> None: ref = self._slot_ref(name) @@ -3366,16 +3370,3 @@ def _nested_forward_children(inputs: ForwardInputs) -> Iterator[ForwardInputs]: "TrainerRank forward inputs must be ForwardInput objects or nested " "iterables of ForwardInput objects" ) from exc - - -__all__ = [ - "AdamParams", - "ForwardInput", - "ForwardOutput", - "MicroBatch", - "MicroBatchStats", - "TopK", - "TrainerRank", - "TrainerRankMemoryError", - "TrainerRankSlotStateError", -] diff --git a/tests/integration/megatron/lora/test_dynamic_lora_slots.py b/tests/integration/megatron/lora/test_dynamic_lora_slots.py index 4d128b235..1809aebd7 100644 --- a/tests/integration/megatron/lora/test_dynamic_lora_slots.py +++ b/tests/integration/megatron/lora/test_dynamic_lora_slots.py @@ -232,7 +232,7 @@ def _tp_head_backward_worker(rank: int, world: int, init_method: str) -> None: from megatron.core import tensor_parallel - local_hidden = torch.randn(2, 1, 3, device=device) + local_hidden = torch.randn(2, 1, 3, device=device, requires_grad=True) gathered_hidden = tensor_parallel.gather_from_sequence_parallel_region( local_hidden, tensor_parallel_output_grad=False, diff --git a/tests/integration/megatron/lora/test_lora_disk_codecs.py b/tests/integration/megatron/lora/test_lora_disk_codecs.py index fdab49760..eac73ade0 100644 --- a/tests/integration/megatron/lora/test_lora_disk_codecs.py +++ b/tests/integration/megatron/lora/test_lora_disk_codecs.py @@ -51,7 +51,7 @@ merge_sharded_adapter_entries, save_vllm_lora_from_model, ) -from art.trainer_rank import AdamParams, TrainerRank +from art.trainer_rank import AdamParams, TrainerRank, TrainerRankSlotStateError from art.trainer_rank._checkpoint import ( _PreparedSave, materialize_lora, @@ -2002,7 +2002,10 @@ def test_prepared_checkpoint_pins_a_symlink_target(tmp_path: Path) -> None: checkpoint_module = importlib.import_module("art.trainer_rank._checkpoint") assert prepared.path == first.resolve() _assert_tensors_equal( - checkpoint_module._load_stage_lora(trainer, prepared), _PORTABLE_ADAPTER + checkpoint_module._load_stage_lora( + trainer, prepared, int(_PORTABLE_CONFIG["r"]) + ), + _PORTABLE_ADAPTER, ) @@ -2180,6 +2183,29 @@ def test_checkpoint_prepare_captures_immutable_state_and_finish_is_idempotent( assert not list(tmp_path.glob(".immutable.snapshot-*")) +def test_checkpoint_idempotence_rejects_corrupt_existing_payload( + tmp_path: Path, +) -> None: + trainer = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + _install_portable_checkpoint(trainer) + output = tmp_path / "corrupt-existing" + trainer.save_checkpoint(str(output), "student") + tensors = load_file(output / "adapter_model.safetensors") + key = next(iter(tensors)) + tensors[key].add_(1) + save_file(tensors, output / "adapter_model.safetensors") + + with pytest.raises(RuntimeError, match="Checkpoint digest mismatch"): + trainer.save_checkpoint(str(output), "student") + + with pytest.raises(RuntimeError, match="Checkpoint digest mismatch"): + validate_checkpoint(output) + assert not list(tmp_path.glob(".corrupt-existing.snapshot-*")) + assert not list(tmp_path.glob(".corrupt-existing.tmp-*")) + + def test_checkpoint_finalizers_run_fifo_and_clean_snapshots( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, @@ -2924,6 +2950,44 @@ def test_checkpoint_commit_failure_rolls_back_every_rank(tmp_path: Path) -> None ) +def test_checkpoint_reload_rejects_accumulated_gradients(tmp_path: Path) -> None: + source_trainer = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + _install_portable_checkpoint(source_trainer) + source = tmp_path / "reload-source" + source_trainer.save_checkpoint(str(source), "student") + + trainer = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + _install_portable_checkpoint(trainer) + trainer._dynamic_optimizers["student"] = trainer._new_dynamic_optimizer( + "student", AdamParams(learning_rate=3e-4) + ) + params = trainer._checkpoint_slot_params_by_name["student"] + params[0].grad = torch.ones_like(params[0]) + values = tuple(param.detach().clone() for param in params) + dynamic = trainer._dynamic_optimizers["student"] + + with pytest.raises(TrainerRankSlotStateError, match="accumulated gradients"): + load_trainer_checkpoint(trainer, prepare_checkpoint(str(source)), "student") + + assert trainer._checkpoint_slot_params_by_name["student"] is params + assert trainer._dynamic_optimizers["student"] is dynamic + assert params[0].grad is not None + for expected, actual in zip(values, params, strict=True): + torch.testing.assert_close(actual, expected, atol=0, rtol=0) + assert all( + not (ref.name or "").startswith("__art_loading_") + for ref in next( + module + for module in trainer.runtime.model[0].modules() + if isinstance(module, LoRA) + )._slot_keys + ) + + def test_trainer_rank_checkpoint_deduplicates_data_parallel_replicas( tmp_path: Path, ) -> None: @@ -3068,6 +3132,26 @@ def test_trainer_rank_checkpoint_restores_1_to_2_to_1( ) +def test_checkpoint_restores_different_adapter_rank(tmp_path: Path) -> None: + original = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + _install_portable_checkpoint(original) + output = tmp_path / "rank-two" + _save_before_next_step(original, output) + + ambient = LoRA(_PORTABLE_PREFIX, 3, 4, 4, 4, torch.float32, torch.device("cpu")) + restored = _portable_trainer(ambient) + load_trainer_checkpoint(restored, prepare_checkpoint(str(output)), "student") + assert ambient.A_T.shape[-1] == 4 + assert all( + parameter.shape[-1] == 2 or parameter.shape[-2] == 2 + for parameter in restored._checkpoint_slot_params_by_name["student"] + ) + _step_portable_checkpoint(restored, -0.125) + _assert_checkpoint_state_equal(original, restored) + + def test_checkpoint_restores_pipeline_parallel_next_step(tmp_path: Path) -> None: original = _portable_trainer(torch.nn.Sequential(*_pipeline_loras(0, 1))) _install_checkpoint(original, _PIPELINE_ADAPTER, _PORTABLE_CONFIG) From 330b16f0c120848f62a161787990942cd777bb86 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 11 Aug 2026 06:59:17 +0000 Subject: [PATCH 4/8] Bound checkpoint save idempotency cache --- src/art/trainer_rank/_checkpoint.py | 8 +++--- src/art/trainer_rank/_impl.py | 3 ++- .../megatron/lora/test_lora_disk_codecs.py | 25 +++++++++++++++++++ 3 files changed, 32 insertions(+), 4 deletions(-) diff --git a/src/art/trainer_rank/_checkpoint.py b/src/art/trainer_rank/_checkpoint.py index 9e294b0df..7df2085aa 100644 --- a/src/art/trainer_rank/_checkpoint.py +++ b/src/art/trainer_rank/_checkpoint.py @@ -2,6 +2,7 @@ from __future__ import annotations +from collections import deque from collections.abc import Iterable, Mapping, Sequence from copy import deepcopy from dataclasses import dataclass @@ -286,7 +287,8 @@ def prepare_checkpoint_save( if output_dir in trainer._prepared_checkpoint_saves: raise RuntimeError(f"Checkpoint save is already pending: {output_dir}") with trainer._checkpoint_save_condition: - trainer._completed_checkpoint_saves.discard(output_dir) + if output_dir in trainer._completed_checkpoint_saves: + trainer._completed_checkpoint_saves.remove(output_dir) _ensure_checkpoint_group(trainer) _validate_save_state(trainer, checkpoint_name) adapter_config = deepcopy(dict(_checkpoint_config(trainer, checkpoint_name))) @@ -448,7 +450,7 @@ def finish_checkpoint_save(trainer: TrainerRank, output_dir: str) -> None: with condition: trainer._checkpoint_finishing_saves.discard(output_dir) if error is None and cleanup_error is None: - trainer._completed_checkpoint_saves.add(output_dir) + trainer._completed_checkpoint_saves.append(output_dir) trainer._checkpoint_finish_sequence += 1 condition.notify_all() if error is not None and cleanup_error is not None: @@ -853,7 +855,7 @@ def _ensure_checkpoint_save_state(trainer: TrainerRank) -> None: trainer._checkpoint_finish_sequence = 0 trainer._prepared_checkpoint_saves = {} trainer._checkpoint_finishing_saves = set() - trainer._completed_checkpoint_saves = set() + trainer._completed_checkpoint_saves = deque(maxlen=128) def _ensure_checkpoint_group(trainer: TrainerRank) -> dist.ProcessGroup | None: diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index df41c315b..3aba628ea 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -3,6 +3,7 @@ from __future__ import annotations import asyncio +from collections import deque from collections.abc import ( Awaitable, Callable, @@ -572,7 +573,7 @@ def __init__( self._checkpoint_finish_sequence = 0 self._prepared_checkpoint_saves: dict[str, _PreparedSave] = {} self._checkpoint_finishing_saves: set[str] = set() - self._completed_checkpoint_saves: set[str] = set() + self._completed_checkpoint_saves: deque[str] = deque(maxlen=128) self._pending_slot_graphs: dict[ LoRASlotRef, list[weakref.ReferenceType[torch.Tensor]] ] = {} diff --git a/tests/integration/megatron/lora/test_lora_disk_codecs.py b/tests/integration/megatron/lora/test_lora_disk_codecs.py index eac73ade0..3d28030e5 100644 --- a/tests/integration/megatron/lora/test_lora_disk_codecs.py +++ b/tests/integration/megatron/lora/test_lora_disk_codecs.py @@ -1,4 +1,5 @@ import asyncio +from collections import deque from collections.abc import Sequence import importlib.util import json @@ -2183,6 +2184,30 @@ def test_checkpoint_prepare_captures_immutable_state_and_finish_is_idempotent( assert not list(tmp_path.glob(".immutable.snapshot-*")) +def test_checkpoint_completed_save_idempotency_cache_is_bounded( + tmp_path: Path, +) -> None: + trainer = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + ) + _install_portable_checkpoint(trainer) + checkpoint_module = importlib.import_module("art.trainer_rank._checkpoint") + checkpoint_module._ensure_checkpoint_save_state(trainer) + trainer._completed_checkpoint_saves = deque(maxlen=2) + + outputs = [tmp_path / f"completed-{index}" for index in range(3)] + for output in outputs: + trainer.save_checkpoint(str(output), "student") + + assert list(trainer._completed_checkpoint_saves) == [ + str(outputs[1]), + str(outputs[2]), + ] + trainer._finish_checkpoint_save(str(outputs[1])) + with pytest.raises(RuntimeError, match="Checkpoint save was not prepared"): + trainer._finish_checkpoint_save(str(outputs[0])) + + def test_checkpoint_idempotence_rejects_corrupt_existing_payload( tmp_path: Path, ) -> None: From cd92267ae0ac364ef4bef92bcda4b06009a61352 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 11 Aug 2026 07:31:45 +0000 Subject: [PATCH 5/8] Pin trainer base model revisions --- src/art/megatron/provider.py | 10 ++++ src/art/megatron/train.py | 6 +++ src/art/trainer_rank/_impl.py | 29 +++++++++-- .../megatron/lora/test_lora_disk_codecs.py | 51 +++++++++++++++++++ .../model_support/test_provider_support.py | 26 ++++++++++ 5 files changed, 119 insertions(+), 3 deletions(-) diff --git a/src/art/megatron/provider.py b/src/art/megatron/provider.py index 657f66742..5d99e39c5 100644 --- a/src/art/megatron/provider.py +++ b/src/art/megatron/provider.py @@ -81,6 +81,7 @@ class ProviderBundle(BaseModel): model_config = ConfigDict(arbitrary_types_allowed=True) model_identifier: str + model_revision: str | None = None provider: GPTModelProvider bridge: AutoBridge handler: ModelSupportHandler @@ -612,6 +613,7 @@ def _register_art_flex_attention_mapping_types() -> None: def _build_provider_bundle( model: str, *, + model_revision: str | None = None, torch_dtype: torch.dtype, allow_unvalidated_arch: bool = False, ) -> ProviderBundle: @@ -623,6 +625,7 @@ def _build_provider_bundle( handler = get_model_support_handler_for_spec(spec) bridge = AutoBridge.from_hf_pretrained( model, + revision=model_revision, dtype=torch_dtype, trust_remote_code=True, ) @@ -630,6 +633,7 @@ def _build_provider_bundle( handler.patch_bridge(bridge) return ProviderBundle( model_identifier=model, + model_revision=model_revision, provider=provider, bridge=bridge, handler=handler, @@ -640,12 +644,14 @@ def _build_provider_bundle( def prepare_provider_bundle( model: str, *, + model_revision: str | None = None, torch_dtype: torch.dtype = torch.bfloat16, allow_unvalidated_arch: bool = False, ) -> ProviderBundle: runtime_env = _ProviderRuntimeEnv.from_environ() bundle = _build_provider_bundle( model, + model_revision=model_revision, torch_dtype=torch_dtype, allow_unvalidated_arch=allow_unvalidated_arch, ) @@ -765,12 +771,14 @@ def _validate_art_gdn_context_parallel_provider(provider: GPTModelProvider) -> N def get_provider_bundle( model: str, *, + model_revision: str | None = None, torch_dtype: torch.dtype = torch.bfloat16, allow_unvalidated_arch: bool = False, ) -> ProviderBundle: return finalize_provider_bundle( prepare_provider_bundle( model, + model_revision=model_revision, torch_dtype=torch_dtype, allow_unvalidated_arch=allow_unvalidated_arch, ) @@ -780,11 +788,13 @@ def get_provider_bundle( def get_provider( model: str, *, + model_revision: str | None = None, torch_dtype: torch.dtype = torch.bfloat16, allow_unvalidated_arch: bool = False, ) -> GPTModelProvider: return get_provider_bundle( model, + model_revision=model_revision, torch_dtype=torch_dtype, allow_unvalidated_arch=allow_unvalidated_arch, ).provider diff --git a/src/art/megatron/train.py b/src/art/megatron/train.py index 5519a60a1..683b6eca8 100644 --- a/src/art/megatron/train.py +++ b/src/art/megatron/train.py @@ -188,6 +188,10 @@ def _validate_model(cls, value: ModelChunks) -> ModelChunks: def model_identifier(self) -> str: return self.provider_bundle.model_identifier + @property + def model_revision(self) -> str | None: + return self.provider_bundle.model_revision + @property def bridge(self) -> Any: return self.provider_bundle.bridge @@ -374,6 +378,7 @@ def _enable_native_moe_routing_replay(provider: Any) -> None: def build_training_runtime( *, model_identifier: str | None = None, + model_revision: str | None = None, provider_torch_dtype: torch.dtype = torch.bfloat16, provider_bundle_configure: Callable[[ProviderBundle], None] | None = None, provider_configure: Callable[[Any], None] | None = None, @@ -396,6 +401,7 @@ def build_training_runtime( provider_bundle = prepare_provider_bundle( model_identifier or os.environ.get("MODEL_IDENTIFIER", DEFAULT_MODEL_IDENTIFIER), + model_revision=model_revision, torch_dtype=provider_torch_dtype, allow_unvalidated_arch=( os.environ.get("ART_MEGATRON_ALLOW_UNVALIDATED_ARCH", "").strip().lower() diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 3aba628ea..f444058cd 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -21,6 +21,7 @@ TYPE_CHECKING, Generic, Literal, + NotRequired, Self, TypedDict, TypeVar, @@ -85,6 +86,7 @@ class TopK: class _AdapterConfig(TypedDict): base_model_name_or_path: str + revision: NotRequired[str | None] r: int lora_alpha: float target_modules: str | list[str] @@ -819,15 +821,36 @@ def _validate_checkpoint_adapter_config( alpha: float | None, ) -> _AdapterConfig | None: config = None if adapter_config is None else deepcopy(dict(adapter_config)) + runtime_revision = getattr(self.runtime, "model_revision", None) if dist.is_available() and dist.is_initialized(): - gathered: list[dict[str, object] | None] = [None] * dist.get_world_size() - dist.all_gather_object(gathered, config) - if any(value != config for value in gathered): + gathered: list[tuple[dict[str, object] | None, object] | None] = [ + None + ] * dist.get_world_size() + dist.all_gather_object(gathered, (config, runtime_revision)) + if any(value is None or value[0] != config for value in gathered): raise ValueError( f"Adapter config for checkpoint slot {name!r} differs across ranks" ) + if any(value is None or value[1] != runtime_revision for value in gathered): + raise ValueError("Runtime model revision differs across ranks") if config is None: return None + source_revision = config.get("revision") + if source_revision is not None and not isinstance(source_revision, str): + raise TypeError("adapter_config['revision'] must be a string or null") + if runtime_revision is not None and not isinstance(runtime_revision, str): + raise TypeError("runtime model_revision must be a string or null") + if ( + source_revision is not None + and runtime_revision is not None + and source_revision != runtime_revision + ): + raise ValueError( + f"Checkpoint {name!r} base-model revision {source_revision!r} " + f"does not match runtime revision {runtime_revision!r}" + ) + if runtime_revision is not None: + config["revision"] = runtime_revision required = {"base_model_name_or_path", "r", "lora_alpha", "target_modules"} if missing := sorted(required - config.keys()): raise ValueError( diff --git a/tests/integration/megatron/lora/test_lora_disk_codecs.py b/tests/integration/megatron/lora/test_lora_disk_codecs.py index 3d28030e5..d0fa8bbc6 100644 --- a/tests/integration/megatron/lora/test_lora_disk_codecs.py +++ b/tests/integration/megatron/lora/test_lora_disk_codecs.py @@ -1800,6 +1800,7 @@ def _portable_trainer( rank: int = 0, world_size: int = 1, model_identifier: str = str(_PORTABLE_CONFIG["base_model_name_or_path"]), + model_revision: str | None = None, model_names: tuple[str, ...] = (), ) -> TrainerRank: trainer = TrainerRank.__new__(TrainerRank) @@ -1809,6 +1810,7 @@ def _portable_trainer( rank=rank, world_size=world_size, model_identifier=model_identifier, + model_revision=model_revision, model_support_spec=SimpleNamespace(model_names=model_names), ) trainer.device = torch.device("cpu") @@ -1891,6 +1893,55 @@ def _save_before_next_step(trainer: TrainerRank, path: Path) -> None: _step_portable_checkpoint(trainer, -0.125) +@pytest.mark.parametrize("source_revision", [None, "a" * 40]) +def test_checkpoint_load_pins_runtime_revision( + tmp_path: Path, source_revision: str | None +) -> None: + revision = "a" * 40 + source = tmp_path / "source" + config = dict(_PORTABLE_CONFIG) + if source_revision is not None: + config["revision"] = source_revision + save_vllm_lora_tensors(source, _PORTABLE_ADAPTER, config) + trainer = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")), + model_revision=revision, + ) + + load_trainer_checkpoint(trainer, prepare_checkpoint(str(source)), "student") + + assert trainer._checkpoint_slot_adapter_configs["student"]["revision"] == revision + + +def test_checkpoint_load_rejects_conflicting_runtime_revision_transactionally( + tmp_path: Path, +) -> None: + source = tmp_path / "source" + save_vllm_lora_tensors( + source, + _PORTABLE_ADAPTER, + {**_PORTABLE_CONFIG, "revision": "a" * 40}, + ) + trainer = _portable_trainer( + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")), + model_revision="b" * 40, + ) + _install_portable_checkpoint(trainer) + previous_params = trainer._checkpoint_slot_params_by_name["student"] + previous_values = tuple(param.detach().clone() for param in previous_params) + previous_config = trainer._checkpoint_slot_adapter_configs["student"] + + with pytest.raises(ValueError, match="base-model revision"): + load_trainer_checkpoint(trainer, prepare_checkpoint(str(source)), "student") + + restored_params = trainer._checkpoint_slot_params_by_name["student"] + assert tuple(map(id, restored_params)) == tuple(map(id, previous_params)) + assert trainer._checkpoint_slot_adapter_configs["student"] is previous_config + assert trainer._checkpoint_revisions["student"] == 0 + for expected, actual in zip(previous_values, restored_params, strict=True): + torch.testing.assert_close(actual, expected, atol=0, rtol=0) + + def test_single_local_expert_uses_global_expert_keys( monkeypatch: pytest.MonkeyPatch, ) -> None: diff --git a/tests/integration/megatron/model_support/test_provider_support.py b/tests/integration/megatron/model_support/test_provider_support.py index 3c0aeeba9..e19aba873 100644 --- a/tests/integration/megatron/model_support/test_provider_support.py +++ b/tests/integration/megatron/model_support/test_provider_support.py @@ -429,6 +429,32 @@ def test_finalize_provider_bundle_uses_post_prepare_topology( assert dispatcher_calls == [] +def test_get_provider_bundle_pins_hf_revision( + monkeypatch: pytest.MonkeyPatch, +) -> None: + provider = _FakeProvider() + fake_bridge = _FakeBridge(model_bridge=object(), provider=provider) + call: dict[str, object] = {} + + def from_hf_pretrained(model: str, **kwargs: object) -> _FakeBridge: + call["model"] = model + call.update(kwargs) + return fake_bridge + + monkeypatch.setattr( + provider_module.AutoBridge, "from_hf_pretrained", from_hf_pretrained + ) + monkeypatch.setattr(provider_module.torch.cuda, "device_count", lambda: 1) + revision = "a" * 40 + + bundle = provider_module.get_provider_bundle( + "Qwen/Qwen3-30B-A3B-Instruct-2507", model_revision=revision + ) + + assert call["revision"] == revision + assert bundle.model_revision == revision + + def test_get_provider_bundle_honors_single_gpu_env_topology( monkeypatch: pytest.MonkeyPatch, ) -> None: From a9304edb5032e6966f6b1eeed53efff1aa3bf3d0 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 11 Aug 2026 07:53:50 +0000 Subject: [PATCH 6/8] Validate checkpoint collective payloads --- src/art/trainer_rank/_impl.py | 12 ++++++++++-- tests/unit/test_trainer_rank_validation.py | 3 ++- 2 files changed, 12 insertions(+), 3 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index f444058cd..79ac0aee2 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -827,11 +827,19 @@ def _validate_checkpoint_adapter_config( None ] * dist.get_world_size() dist.all_gather_object(gathered, (config, runtime_revision)) - if any(value is None or value[0] != config for value in gathered): + if any( + not isinstance(value, tuple) or len(value) != 2 or value[0] != config + for value in gathered + ): raise ValueError( f"Adapter config for checkpoint slot {name!r} differs across ranks" ) - if any(value is None or value[1] != runtime_revision for value in gathered): + if any( + not isinstance(value, tuple) + or len(value) != 2 + or value[1] != runtime_revision + for value in gathered + ): raise ValueError("Runtime model revision differs across ranks") if config is None: return None diff --git a/tests/unit/test_trainer_rank_validation.py b/tests/unit/test_trainer_rank_validation.py index 457d1f074..9b8241fa5 100644 --- a/tests/unit/test_trainer_rank_validation.py +++ b/tests/unit/test_trainer_rank_validation.py @@ -785,7 +785,8 @@ def test_checkpoint_slot_adapter_config_rejects_cross_rank_mismatch( monkeypatch.setattr("art.trainer_rank.dist.get_world_size", lambda: 2) def gather(output: list[object], value: object) -> None: - output[:] = [value, {"different": True}] + revision = value[1] if isinstance(value, tuple) and len(value) == 2 else None + output[:] = [value, ({"different": True}, revision)] monkeypatch.setattr("art.trainer_rank.dist.all_gather_object", gather) From 3b18695ddd1f941fae68704c1ddc3dca3034febe Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 11 Aug 2026 09:09:44 +0000 Subject: [PATCH 7/8] Fix mixed precision checkpoint exchange --- dev/trainer_rank_checkpoint_benchmark.py | 16 ++++----- src/art/megatron/weights/lora_publish.py | 11 ++++++ src/art/trainer_rank/_checkpoint.py | 6 +++- .../megatron/lora/test_lora_disk_codecs.py | 36 ++++++++++++++++--- 4 files changed, 56 insertions(+), 13 deletions(-) diff --git a/dev/trainer_rank_checkpoint_benchmark.py b/dev/trainer_rank_checkpoint_benchmark.py index f17621a1a..10eeed684 100644 --- a/dev/trainer_rank_checkpoint_benchmark.py +++ b/dev/trainer_rank_checkpoint_benchmark.py @@ -219,14 +219,6 @@ def _parse_args() -> argparse.Namespace: def main() -> None: args = _parse_args() - if not torch.cuda.is_available(): - raise RuntimeError("checkpoint benchmark requires CUDA") - device = int(os.environ["LOCAL_RANK"]) - torch.cuda.set_device(device) - dist.init_process_group("nccl") - rank = dist.get_rank() - world_size = dist.get_world_size() - topology = _topology(args, world_size) for key, value in ( ("TENSOR_MODEL", args.tp), ("PIPELINE_MODEL", args.pp), @@ -235,6 +227,14 @@ def main() -> None: ("EXPERT_TENSOR", args.etp), ): os.environ[f"ART_MEGATRON_{key}_PARALLEL_SIZE"] = str(value) + if not torch.cuda.is_available(): + raise RuntimeError("checkpoint benchmark requires CUDA") + device = int(os.environ["LOCAL_RANK"]) + torch.cuda.set_device(device) + dist.init_process_group("nccl") + rank = dist.get_rank() + world_size = dist.get_world_size() + topology = _topology(args, world_size) try: from art.megatron import train as megatron_train diff --git a/src/art/megatron/weights/lora_publish.py b/src/art/megatron/weights/lora_publish.py index 8cb2e76ac..1704c32e6 100644 --- a/src/art/megatron/weights/lora_publish.py +++ b/src/art/megatron/weights/lora_publish.py @@ -392,6 +392,17 @@ def _prepare_exchange_buffers( identity = (meta.owner_rank, meta.key) if rank == meta.owner_rank: tensor = local_tensors[meta.key].detach().contiguous() + if tuple(tensor.shape) != meta.shape: + raise RuntimeError( + f"Tensor {meta.key!r} shape {tuple(tensor.shape)} does not match " + f"exchange metadata {meta.shape}" + ) + dtype_name = _dtype_name(tensor.dtype) + if dtype_name != meta.dtype_name: + raise RuntimeError( + f"Tensor {meta.key!r} dtype {dtype_name!r} does not match " + f"exchange metadata {meta.dtype_name!r}" + ) if rank == 0: received[identity] = tensor.cpu().contiguous() else: diff --git a/src/art/trainer_rank/_checkpoint.py b/src/art/trainer_rank/_checkpoint.py index 7df2085aa..52f21b986 100644 --- a/src/art/trainer_rank/_checkpoint.py +++ b/src/art/trainer_rank/_checkpoint.py @@ -1132,7 +1132,11 @@ def _serialize_snapshot_optimizer( for component in ("master", "exp_avg", "exp_avg_sq"): component_records: dict[str, TensorRecord] = {} for index, block in enumerate(blocks): - block_metadata = [item for item in metadata if item.block == block] + block_metadata = [ + item._replace(dtype_name=_dtype_name(torch.float32)) + for item in metadata + if item.block == block + ] local_tensors: dict[str, torch.Tensor] = {} error: BaseException | None = None try: diff --git a/tests/integration/megatron/lora/test_lora_disk_codecs.py b/tests/integration/megatron/lora/test_lora_disk_codecs.py index d0fa8bbc6..644859231 100644 --- a/tests/integration/megatron/lora/test_lora_disk_codecs.py +++ b/tests/integration/megatron/lora/test_lora_disk_codecs.py @@ -2602,6 +2602,7 @@ def _portable_topology_worker( init_method: str, source_path: str, output_path: str, + dtype_name: str = "float32", ) -> None: torch.distributed.init_process_group( "gloo", rank=rank, world_size=world_size, init_method=init_method @@ -2627,7 +2628,7 @@ def _portable_topology_worker( 4 // world_size, 2, 2, - torch.float32, + getattr(torch, dtype_name), torch.device("cpu"), b_parallel_spec=LoRAParallelSpec(sharded=True, shard_dim=-1), ) @@ -2979,6 +2980,25 @@ def test_lora_exchange_coordinates_asymmetric_preparation_failure( ) +def test_lora_exchange_rejects_mismatched_metadata() -> None: + key = "base_model.model.model.layers.0.self_attn.q_proj.lora_A.weight" + metadata = LoraShardMeta( + key=key, + owner_rank=0, + shape=(2,), + dtype_name="bfloat16", + manifest={"sharded": False, "shard_world_size": 1, "shard_rank": 0}, + block="base_model.model.model.layers.0", + ) + with pytest.raises(RuntimeError, match="dtype 'float32'.*metadata 'bfloat16'"): + lora_publish._prepare_exchange_buffers( + (metadata,), + local_tensors={key: torch.ones(2)}, + rank=0, + device=torch.device("cpu"), + ) + + def test_checkpoint_loads_follow_call_order_across_ranks(tmp_path: Path) -> None: first = tmp_path / "prefetch-first" second = tmp_path / "prefetch-second" @@ -3153,7 +3173,7 @@ def test_trainer_rank_checkpoint_restores_1_to_2_to_1( tmp_path: Path, ) -> None: original = _portable_trainer( - LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.bfloat16, torch.device("cpu")) ) _install_portable_checkpoint(original) original._dynamic_optimizers["student"] = original._new_dynamic_optimizer( @@ -3173,14 +3193,22 @@ def test_trainer_rank_checkpoint_restores_1_to_2_to_1( f"file://{tmp_path / 'topology-init'}", str(one_rank), str(two_rank), + "bfloat16", ), nprocs=2, join=True, ) restored = _portable_trainer( - LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.float32, torch.device("cpu")) - ) + LoRA(_PORTABLE_PREFIX, 3, 4, 2, 2, torch.bfloat16, torch.device("cpu")) + ) + manifest = validate_checkpoint(two_rank, require_optimizer=True) + assert manifest is not None + assert { + record.dtype + for parameter in manifest.parameters.values() + for record in (parameter.master, parameter.exp_avg, parameter.exp_avg_sq) + } == {"float32"} load_trainer_checkpoint( restored, prepare_checkpoint(str(two_rank)), From 26b4a40eb0d2ef29d56d7c6715f9813f6e56a75e Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 11 Aug 2026 09:22:18 +0000 Subject: [PATCH 8/8] Honor checkpoint model revision offline --- .../model_support/handlers/qwen3_5.py | 8 +++-- .../megatron/lora/test_lora_disk_codecs.py | 36 +++++++++++++++++++ 2 files changed, 42 insertions(+), 2 deletions(-) diff --git a/src/art/megatron/model_support/handlers/qwen3_5.py b/src/art/megatron/model_support/handlers/qwen3_5.py index db72e0687..1a0ad5fd5 100644 --- a/src/art/megatron/model_support/handlers/qwen3_5.py +++ b/src/art/megatron/model_support/handlers/qwen3_5.py @@ -508,11 +508,12 @@ def _is_self_attn_q_proj_lora_b(key: str) -> bool: @lru_cache(maxsize=8) -def _qwen35_text_config(base_model_name_or_path: str) -> Any: +def _qwen35_text_config(base_model_name_or_path: str, revision: str | None) -> Any: from transformers import AutoConfig config = AutoConfig.from_pretrained( base_model_name_or_path, + revision=revision, local_files_only=True, trust_remote_code=True, ) @@ -528,7 +529,10 @@ def _qwen35_attention_dims(adapter_config: dict[str, Any]) -> tuple[int, int, in base_model = adapter_config.get("base_model_name_or_path") if not base_model: raise RuntimeError("Qwen3.5 LoRA adapter config is missing base model path") - config = _qwen35_text_config(str(base_model)) + revision = adapter_config.get("revision") + if revision is not None and not isinstance(revision, str): + raise RuntimeError("Qwen3.5 LoRA adapter revision must be a string") + config = _qwen35_text_config(str(base_model), revision) num_heads = getattr(config, "num_attention_heads") num_groups = getattr(config, "num_key_value_heads", num_heads) head_dim = getattr(config, "head_dim", None) diff --git a/tests/integration/megatron/lora/test_lora_disk_codecs.py b/tests/integration/megatron/lora/test_lora_disk_codecs.py index 644859231..ccdd395ca 100644 --- a/tests/integration/megatron/lora/test_lora_disk_codecs.py +++ b/tests/integration/megatron/lora/test_lora_disk_codecs.py @@ -34,6 +34,7 @@ QWEN3_5_MOE_HANDLER, QWEN3_MOE_HANDLER, ) +from art.megatron.model_support.handlers import qwen3_5 as qwen35_module from art.megatron.model_support.handlers.dsv4 import DSV4_HANDLER from art.megatron.model_support.handlers.gemma4 import GEMMA4_MOE_HANDLER from art.megatron.model_support.lora_disk import ( @@ -766,6 +767,41 @@ def test_qwen3_target_parameter_identity_normalizes_to_per_expert_vllm_layout( ] +def test_qwen35_config_lookup_uses_checkpoint_revision( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import transformers + + calls: list[tuple[str, dict[str, object]]] = [] + + def from_pretrained(base_model: str, **kwargs: object) -> object: + calls.append((base_model, kwargs)) + return SimpleNamespace( + num_attention_heads=4, + num_key_value_heads=2, + head_dim=3, + ) + + monkeypatch.setattr(transformers.AutoConfig, "from_pretrained", from_pretrained) + qwen35_module._qwen35_text_config.cache_clear() + try: + assert qwen35_module._qwen35_attention_dims( + {"base_model_name_or_path": "Qwen/test", "revision": "abc123"} + ) == (4, 2, 3) + finally: + qwen35_module._qwen35_text_config.cache_clear() + assert calls == [ + ( + "Qwen/test", + { + "revision": "abc123", + "local_files_only": True, + "trust_remote_code": True, + }, + ) + ] + + def test_qwen35_and_qwen36_vllm_canonical_roundtrip_and_stock_loader(tmp_path: Path): art_prefix = "base_model.model.model.layers.0" original = _qwen35_moe_art_tensors(art_prefix)