From 6c3a0faddd42abffaf9e9d02f899b624992be419 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 22:00:58 +0000 Subject: [PATCH 1/2] Move TrainerRank checkpoint-slot bookkeeping into trainer_rank/_slots Extract the 31 checkpoint-slot bookkeeping methods (prefetch registry, slot loading and validation, the slot stack push/pop, slot-graph liveness guards) from TrainerRank into a new _slots module. Bodies are moved verbatim as module functions taking the rank as self; TrainerRank binds them in its class body so self._x(...) dispatch, class-level access and per-instance monkeypatches keep working, and module globals (torch, dist, sibling helpers) are still resolved through _impl at call time so tests that patch _impl.torch and friends keep intercepting them. _slots is imported immediately before the class body (alongside _memory and _micro_batch_planner) because _resolve_slot_ref and _ensure_checkpoint_slots_for read _impl.Unset as a definition-time default. Behavior-preserving code motion only. --- src/art/trainer_rank/_impl.py | 637 ++----------------------------- src/art/trainer_rank/_slots.py | 663 +++++++++++++++++++++++++++++++++ 2 files changed, 702 insertions(+), 598 deletions(-) create mode 100644 src/art/trainer_rank/_slots.py diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index d21313eaa..7c4965d61 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -14,7 +14,6 @@ from concurrent.futures import Future, ThreadPoolExecutor from contextlib import contextmanager from contextvars import ContextVar -from copy import deepcopy from dataclasses import dataclass, replace from dataclasses import field as dataclass_field from functools import partial @@ -1844,10 +1843,10 @@ def _moe_output_bytes_per_token( return coefficient -# ``_memory`` and ``_micro_batch_planner`` hold TrainerRank method bodies and -# read ``_impl`` names such as ``Unset`` at definition time, so they are imported -# once those exist and before the class body binds their functions. -from art.trainer_rank import _memory, _micro_batch_planner +# ``_memory``, ``_micro_batch_planner`` and ``_slots`` hold TrainerRank method +# bodies and read ``_impl`` names such as ``Unset`` at definition time, so they are +# imported once those exist and before the class body binds their functions. +from art.trainer_rank import _memory, _micro_batch_planner, _slots class TrainerRank: @@ -2235,21 +2234,6 @@ def _custom_object( tracker.active = True return custom.value - def _resolve_custom_checkpoint(self, checkpoint: AdapterSelection) -> str: - if checkpoint is Unset: - ref = self._slot_stack[-1] if self._slot_stack else self._default_slot_ref - name = None if ref is None else ref.name - else: - name = cast(str | None, checkpoint) - if name is None: - raise TrainerRankSlotStateError( - "Custom checkpoint objects require a loaded named checkpoint" - ) - self._ensure_checkpoint_slots((name,)) - if name not in self._checkpoint_slots: - raise TrainerRankSlotStateError(f"Unknown checkpoint: {name!r}") - return name - def _initialize_custom_object( self, checkpoint: str, @@ -2359,148 +2343,12 @@ def _extend_dynamic_optimizer( } return _DynamicOptimizer(optimizer, all_masters) - def prefetch_checkpoints( - self, *checkpoints: str | MaterializedCheckpoint - ) -> asyncio.Task[None]: - futures = [] - for checkpoint in checkpoints: - logical, source = self._checkpoint_source(checkpoint) - assert logical is not None and source is not None - futures.append(self._register_checkpoint_prefetch(logical, source)) - - async def prefetch() -> None: - await asyncio.gather( - *(self._await_checkpoint_prefetch(future) for future in futures) - ) - - return asyncio.create_task(prefetch()) - @staticmethod async def _await_checkpoint_prefetch( future: Future[PreparedCheckpoint], ) -> PreparedCheckpoint: return await asyncio.shield(asyncio.wrap_future(future)) - def _register_checkpoint_prefetch( - self, - checkpoint: str, - source: str, - prepare: Callable[[], PreparedCheckpoint] | None = None, - ) -> Future[PreparedCheckpoint]: - key = self._checkpoint_source_key(source) - with self._checkpoint_prefetch_lock: - previous = self._checkpoint_prefetch_sources.get(checkpoint) - self._checkpoint_prefetch_sources[checkpoint] = key - if ( - previous is not None - and previous != key - and previous not in self._checkpoint_prefetch_sources.values() - ): - self._checkpoint_prefetches.pop(previous, None) - future = self._checkpoint_prefetches.get(key) - if future is None or ( - future.done() and (future.cancelled() or future.exception() is not None) - ): - if prepare is None: - from ._checkpoint import prepare_checkpoint - - prepare = lambda: prepare_checkpoint(key) - future = _checkpoint_prefetch_executor().submit(prepare) - self._checkpoint_prefetches[key] = future - return future - - def _checkpoint_prefetch_waiter(self, *checkpoints: str) -> asyncio.Task[None]: - with self._checkpoint_prefetch_lock: - futures = [ - self._checkpoint_prefetches[self._checkpoint_prefetch_sources[name]] - for name in checkpoints - ] - - async def wait() -> None: - await asyncio.gather( - *(self._await_checkpoint_prefetch(future) for future in futures) - ) - - return asyncio.create_task(wait()) - - def _prefetched_checkpoint(self, checkpoint: str) -> PreparedCheckpoint: - with self._checkpoint_prefetch_lock: - key = self._checkpoint_prefetch_sources.get(checkpoint) - future = None if key is None else self._checkpoint_prefetches.get(key) - if future is None: - raise TrainerRankSlotStateError( - f"Checkpoint {checkpoint!r} has not been prefetched" - ) - return future.result() - - def _load_registered_checkpoint(self, checkpoint: str) -> None: - from . import _checkpoint - - source: PreparedCheckpoint | None = None - error: BaseException | None = None - try: - source = self._prefetched_checkpoint(checkpoint) - except BaseException as exc: - error = exc - group = _checkpoint._ensure_group(self) - _checkpoint.raise_distributed(error, "prepare checkpoint", group) - assert source is not None - if checkpoint in self._snapshot_checkpoint_names: - _checkpoint.load_checkpoint(self, source, checkpoint, forward_only=True) - else: - _checkpoint.load_checkpoint(self, source, checkpoint) - - def _ensure_checkpoint_slots(self, checkpoints: Iterable[str]) -> None: - from . import _checkpoint - - requested = tuple(dict.fromkeys(checkpoints)) - with self._checkpoint_mutation_lock: - group = _checkpoint._ensure_group(self) - names = sorted( - { - name - for rank_names in _checkpoint._gather(requested, group) - for name in rank_names - } - ) - for name in names: - with self._checkpoint_prefetch_lock: - state = ( - name in self._checkpoint_slots, - name in self._checkpoint_prefetch_sources, - ) - states = _checkpoint._gather(state, group) - if all(loaded for loaded, _prefetched in states): - continue - if any(loaded for loaded, _prefetched in states): - raise TrainerRankSlotStateError( - f"Checkpoint {name!r} is not loaded consistently across ranks" - ) - if not all(prefetched for _loaded, prefetched in states): - raise TrainerRankSlotStateError( - f"Explicit selection references unloaded checkpoint {name!r}; " - "it has not been prefetched on every rank" - ) - self._load_registered_checkpoint(name) - - def load_checkpoint(self, checkpoint: str | MaterializedCheckpoint | None) -> None: - self._guard_forward_collective("load_checkpoint") - logical, source = self._checkpoint_source(checkpoint) - with self._checkpoint_mutation_lock: - if self._slot_stack: - raise RuntimeError("Cannot load a checkpoint while one is pushed") - if logical is None: - self._set_default_slot(self._slot_ref(None)) - return - assert source is not None - if ( - isinstance(checkpoint, MaterializedCheckpoint) - or logical not in self._checkpoint_prefetch_sources - ): - self._register_checkpoint_prefetch(logical, source) - self._load_registered_checkpoint(logical) - self._set_default_slot(self._slot_ref(logical)) - def snapshot_checkpoint(self, source: str, destination: str) -> bool: """Clone a loaded checkpoint into a forward-only resident snapshot.""" self._guard_forward_collective("snapshot_checkpoint") @@ -2509,43 +2357,12 @@ def snapshot_checkpoint(self, source: str, destination: str) -> bool: self._ensure_checkpoint_slots((source,)) return _checkpoint.snapshot_checkpoint(self, source, destination) - def _discard_snapshot_checkpoint(self, checkpoint: str) -> None: - """Discard a forward-only resident snapshot.""" - from . import _checkpoint - - _checkpoint.discard_snapshot_checkpoint(self, checkpoint) - def push_checkpoint( self, checkpoint: str | MaterializedCheckpoint | None ) -> PushedCheckpoint: logical, directory = self._checkpoint_source(checkpoint) return PushedCheckpoint(self, logical, directory) - def _push_checkpoint(self, checkpoint: str | MaterializedCheckpoint | None) -> None: - logical, source = self._checkpoint_source(checkpoint) - self._push_checkpoint_sync(logical, source) - - def _push_checkpoint_sync( - self, logical_path: str | None, source_path: str | None - ) -> None: - self._guard_forward_collective("push_checkpoint") - with self._checkpoint_mutation_lock: - if source_path is not None: - assert logical_path is not None - if ( - logical_path not in self._checkpoint_slots - and logical_path not in self._checkpoint_prefetch_sources - ): - self._register_checkpoint_prefetch(logical_path, source_path) - self._ensure_checkpoint_slots((logical_path,)) - self._slot_stack.append(self._slot_ref(logical_path)) - - def pop_checkpoint(self) -> None: - with self._checkpoint_mutation_lock: - if not self._slot_stack: - raise RuntimeError("No pushed checkpoint to pop") - self._slot_stack.pop() - def save_checkpoint( self, output_dir: str, @@ -2632,103 +2449,10 @@ def _checkpoint_source( return checkpoint.path, checkpoint.directory return checkpoint, checkpoint - def _resolve_checkpoint_name(self, checkpoint_path: str | Literal["active"]) -> str: - if checkpoint_path != "active": - self._ensure_checkpoint_slots((checkpoint_path,)) - 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 - @staticmethod def _slot_state_error(message: str) -> TrainerRankSlotStateError: return TrainerRankSlotStateError(message) - def _checkpoint_group(self) -> dist.ProcessGroup | None: - from ._checkpoint import _ensure_group - - return _ensure_group(self) - - def _validate_checkpoint_adapter_config( - self, - name: str, - adapter_config: Mapping[str, object] | None, - *, - alpha: float | None, - ) -> _AdapterConfig | None: - config = None if adapter_config is None else deepcopy(dict(adapter_config)) - 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, group=self._checkpoint_group()) - if any(value != config for value in gathered): - raise ValueError( - f"Adapter config for checkpoint slot {name!r} differs across ranks" - ) - if config is None: - return None - required = {"base_model_name_or_path", "r", "lora_alpha", "target_modules"} - if missing := sorted(required - config.keys()): - raise ValueError( - f"Adapter config for checkpoint slot {name!r} is missing {missing}" - ) - base_model = config["base_model_name_or_path"] - rank = config["r"] - config_alpha_value = config["lora_alpha"] - target_modules = config["target_modules"] - if not isinstance(base_model, str): - raise TypeError( - "adapter_config['base_model_name_or_path'] must be a string" - ) - if base_model.startswith(("Qwen/Qwen3.5-", "Qwen/Qwen3.6-", "Qwen/Qwen3.8-")): - from art.megatron.model_support.lora_disk import ( - model_attention_dimensions, - ) - - config.update(model_attention_dimensions(self.runtime.provider)) - if not isinstance(rank, int) or isinstance(rank, bool): - raise TypeError("adapter_config['r'] must be an integer") - if not isinstance(config_alpha_value, int | float) or isinstance( - config_alpha_value, bool - ): - raise TypeError("adapter_config['lora_alpha'] must be numeric") - if not isinstance(target_modules, str | list) or ( - isinstance(target_modules, list) - and not all(isinstance(module, str) for module in target_modules) - ): - raise TypeError( - "adapter_config['target_modules'] must be a string or list of strings" - ) - if rank < 1: - raise ValueError("adapter_config['r'] must be >= 1") - config_alpha = float(config_alpha_value) - if alpha is not None and float(alpha) != config_alpha: - raise ValueError( - f"alpha={alpha} conflicts with adapter_config lora_alpha={config_alpha}" - ) - return cast(_AdapterConfig, config) - - def _validate_loaded_checkpoint_config( - self, name: str, config: _AdapterConfig - ) -> None: - from art.megatron.lora import LoRA - - ref = self._slot_ref(name) - slots = [ - slot - for chunk in self.runtime.model - for module in chunk.modules() - if isinstance(module, LoRA) - if (slot := module._slot(ref)) is not None - ] - expected = (int(config["r"]), float(config["lora_alpha"])) - actual = {(slot.rank, slot.alpha) for slot in slots} - if actual != {expected}: - raise ValueError( - f"Adapter config for checkpoint slot {name!r} declares " - f"rank/alpha={expected}, loaded weights use {sorted(actual)}" - ) - @overload def forward_micro_batches( self, @@ -3477,32 +3201,6 @@ def adam_values( "Optimizer checkpoint selection or AdamParams differ across ranks" ) - def _load_checkpoint_slot( - self, - name: str, - adapter_model: Mapping[str, torch.Tensor], - *, - alpha: float, - _prepared: bool = False, - ) -> int: - if self._slot_stack: - raise RuntimeError("Cannot load a checkpoint while one is pushed") - adapter_model = self._prepare_adapter_model( - name, adapter_model, canonicalized=_prepared - ) - from art.megatron.lora import load_lora_slot_into_model - - ref = self._slot_ref(name) - self._guard_slot_can_load(ref) - self._compact_lora_slot_keys() - return load_lora_slot_into_model( - self.runtime.model, - ref, - adapter_model, - alpha=alpha, - requires_grad=True, - ) - def _compact_lora_slot_keys(self) -> None: from art.megatron.lora import LoRA @@ -3586,11 +3284,6 @@ def _local_lora_adapter_templates(self) -> dict[str, torch.Tensor]: ) return templates - def _iter_slot_parameters(self, ref: "LoRASlotRef") -> Iterator[torch.nn.Parameter]: - from art.megatron.lora import iter_lora_slot_parameters - - return iter_lora_slot_parameters(self.runtime.model, ref) - def _local_parameter_key_groups(self, name: str) -> tuple[tuple[str, ...], ...]: ref = self._slot_ref(name) return tuple( @@ -3602,32 +3295,6 @@ def _local_parameter_key_groups(self, name: str) -> tuple[tuple[str, ...], ...]: for suffix, _param in lora_params(ref) ) - def _validate_checkpoint_consistency( - self, name: str, loaded_sites: int, expected_keys: set[str] - ) -> tuple[torch.nn.Parameter, ...]: - params = tuple(self._iter_slot_parameters(self._slot_ref(name))) - local_keys = { - key for group in self._local_parameter_key_groups(name) for key in group - } - gathered = ( - [local_keys] - if not (dist.is_available() and dist.is_initialized()) - else [None] * dist.get_world_size() - ) - if dist.is_available() and dist.is_initialized(): - dist.all_gather_object(gathered, local_keys, group=self._checkpoint_group()) - covered = set().union(*(keys for keys in gathered if keys is not None)) - if loaded_sites < 1 or covered != expected_keys: - raise TrainerRankSlotStateError( - f"Checkpoint {name!r} has inconsistent distributed coverage" - ) - return params - - def _set_default_slot(self, ref: "LoRASlotRef") -> None: - if self._slot_stack: - raise RuntimeError("Cannot select a checkpoint while one is pushed") - self._default_slot_ref = ref - @staticmethod def _slot_ref(name: str | None) -> "LoRASlotRef": try: @@ -3640,109 +3307,6 @@ def _slot_ref(name: str | None) -> "LoRASlotRef": return LoRASlotRef(kind="checkpoint", name=name) - def _resolve_slot_ref( - self, - request: AnyForwardInput, - *, - checkpoint: AdapterSelection = Unset, - ) -> "LoRASlotRef | None": - selection = ( - request.checkpoint if request.checkpoint is not Unset else checkpoint - ) - if selection is not Unset: - name = cast(str | None, selection) - if name is not None and name not in self._checkpoint_slots: - raise TrainerRankSlotStateError( - f"Forward selects unloaded checkpoint {name!r}" - ) - 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(None) - - def _selected_dynamic_checkpoints( - self, - checkpoints: Sequence[str] | None, - ) -> tuple[str, ...]: - if checkpoints is not None: - self._ensure_checkpoint_slots(checkpoints) - loaded = set(self._checkpoint_slots) - if not loaded: - raise TrainerRankSlotStateError( - "TrainerRank.optim_step requires a loaded checkpoint slot. Call " - "load_checkpoint(...) and run backward on outputs produced by " - "that slot before stepping." - ) - requested = ( - tuple( - sorted( - name for name in loaded if not self._checkpoint_slots[name].snapshot - ) - ) - if checkpoints is None - else tuple(dict.fromkeys(checkpoints)) - ) - if not requested: - if checkpoints is None: - raise TrainerRankSlotStateError( - "TrainerRank.optim_step requires a loaded trainable checkpoint " - "slot. Call load_checkpoint(...) and run backward on outputs " - "produced by that slot before stepping." - ) - raise TrainerRankSlotStateError( - "TrainerRank.optim_step(checkpoints=...) received no checkpoint " - "names. Pass at least one loaded checkpoint slot." - ) - if unknown := set(requested) - loaded: - raise ValueError(f"Unknown checkpoint slots: {sorted(unknown)}") - if snapshots := [ - name for name in requested if self._checkpoint_slots[name].snapshot - ]: - raise TrainerRankSlotStateError( - "Snapshot checkpoints are forward-only and cannot be stepped: " - f"{snapshots}" - ) - flags = self._checkpoint_grad_flags(requested) - selected = tuple( - name for name, has_grad in zip(requested, flags, strict=True) if has_grad - ) - if checkpoints is None: - if selected: - return selected - raise TrainerRankSlotStateError( - "TrainerRank.optim_step found loaded checkpoint slots, but none " - "have gradients on any rank. Call loss.backward() first." - ) - if missing := [ - name - for name, has_grad in zip(requested, flags, strict=True) - if not has_grad - ]: - raise TrainerRankSlotStateError( - "TrainerRank.optim_step was asked to step checkpoint slots with no " - f"gradients on any rank: {missing}. Call loss.backward() for those " - "slots first, or omit them from checkpoints=[...]." - ) - return selected - - def _checkpoint_grad_flags(self, names: Sequence[str]) -> tuple[bool, ...]: - flags = torch.tensor( - [ - any( - param.grad is not None - for param in self._checkpoint_slots[name].params - ) - for name in names - ], - device=self.device, - dtype=torch.int32, - ) - if dist.is_available() and dist.is_initialized(): - dist.all_reduce(flags, op=dist.ReduceOp.MAX) - return tuple(bool(flag) for flag in flags.tolist()) - def _dynamic_optim_step( self, checkpoint_names: Sequence[str], @@ -4157,32 +3721,6 @@ def _forced_test_anchor(self) -> str | None: return None return os.environ.get(_TEST_ANCHOR_ENV) or None - def _ensure_checkpoint_slots_for( - self, - requests: Sequence[AnyForwardInput], - *, - checkpoint: AdapterSelection, - ) -> None: - self._ensure_checkpoint_slots( - cast(str, selection) - for request in requests - if ( - request.target_tokens is not None - or request.logits - or request.top_k is not None - or request.hidden_states - ) - if ( - selection := ( - request.checkpoint - if request.checkpoint is not Unset - else checkpoint - ) - ) - is not Unset - and selection is not None - ) - def _group_active_request_indices( self, requests: Sequence[AnyForwardInput], @@ -4409,52 +3947,6 @@ def _execute_flat_plan(self, plan: _FlatForwardPlan) -> list[AnyForwardOutput]: self._set_hybridep_rows(hybridep[1]) return outputs - def _track_slot_graph_outputs( - self, - ref: "LoRASlotRef | None", - outputs: Sequence[AnyForwardOutput], - ) -> list[AnyForwardOutput]: - track_slot = ref is not None and ref.name is not None - track_hybridep = bool(getattr(self, "_hybridep_graph_tracking", False)) - if not track_slot and not track_hybridep: - return list(outputs) - - marker: torch.Tensor | None = None - - def track(tensor: torch.Tensor | None) -> torch.Tensor | None: - nonlocal marker - if tensor is None or not tensor.requires_grad: - return tensor - if marker is None: - marker = tensor.new_empty(0) - return cast(torch.Tensor, _SlotGraphSentinel.apply(tensor, marker)) - - tracked_outputs = [ - ForwardOutput( - target_logprobs=track(output.target_logprobs), - top_k=( - None - if output.top_k is None - else TopK( - logprobs=cast(torch.Tensor, track(output.top_k.logprobs)), - tokens=output.top_k.tokens, - ) - ), - logits=track(output.logits), - hidden_states=track(output.hidden_states), - checkpoint=output.checkpoint, - no_grad=output.no_grad, - ) - for output in outputs - ] - if marker is not None: - marker_ref = weakref.ref(marker) - if track_slot: - self._slot_graphs().setdefault(ref, []).append(marker_ref) - if track_hybridep: - self._hybridep_graphs().append(marker_ref) - return tracked_outputs - def _forward_output_metadata( self, request: AnyForwardInput, @@ -4486,92 +3978,6 @@ def _has_live_hybridep_graphs(self) -> bool: graphs[:] = [marker for marker in graphs if marker() is not None] return bool(graphs) - def _slot_graphs( - self, - ) -> dict["LoRASlotRef", list[weakref.ReferenceType[torch.Tensor]]]: - graphs = getattr(self, "_pending_slot_graphs", None) - if graphs is None: - graphs = {} - self._pending_slot_graphs = graphs - return graphs - - def _prune_slot_graphs(self, ref: "LoRASlotRef | None" = None) -> None: - graphs = self._slot_graphs() - refs = tuple(graphs) if ref is None else (ref,) - for current in refs: - live = [ - marker - for marker in graphs.get(current, ()) - if _graph_marker_is_live(marker) - ] - if live: - graphs[current] = live - else: - graphs.pop(current, None) - - def _has_live_slot_graph(self, ref: "LoRASlotRef") -> bool: - self._prune_slot_graphs(ref) - return bool(self._slot_graphs().get(ref)) - - def _guard_slot_can_load(self, ref: "LoRASlotRef") -> None: - slot = None if ref.name is None else self._checkpoint_slots.get(ref.name) - if slot is not None and slot.snapshot: - raise TrainerRankSlotStateError( - f"Cannot load over forward-only snapshot checkpoint {ref.name!r}" - ) - if slot is not None and any(param.grad is not None for param in slot.params): - raise TrainerRankSlotStateError( - f"Cannot load checkpoint {ref.name!r} while it has accumulated " - "gradients. Call optim_step() or zero_grad() before replacing it." - ) - 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." - ) - - def _guard_checkpoint_can_step(self, name: str) -> None: - if not self._has_live_slot_graph(self._slot_ref(name)): - return - raise TrainerRankSlotStateError( - f"Cannot optim_step checkpoint slot {name!r} while outputs from an " - "earlier forward using that slot have not been backpropagated. Call " - "loss.backward() without retaining the graph before optim_step(); if " - "the forward was abandoned, release all references to its outputs." - ) - - def _guard_checkpoints_can_step(self, names: Sequence[str]) -> None: - local_live = [self._has_live_slot_graph(self._slot_ref(name)) for name in names] - if dist.is_available() and dist.is_initialized(): - live = torch.tensor( - local_live, - device=self.device, - dtype=torch.int32, - ) - dist.all_reduce(live, op=dist.ReduceOp.MAX) - live_flags = live.tolist() - else: - live_flags = local_live - blocked = [ - name for name, is_live in zip(names, live_flags, strict=True) if is_live - ] - if not blocked: - return - raise TrainerRankSlotStateError( - f"Cannot optim_step checkpoint slots {blocked!r} while outputs from an " - "earlier forward using those slots have a live backward graph on at " - "least one rank. Call loss.backward() without retaining the graph " - "before optim_step(); if the forward was abandoned, release all " - "references to its outputs; or pass on_live_graphs='allow' to accept " - "responsibility for any retained graphs." - ) - def _topology_key(self) -> tuple[int, int, int, int]: try: topology = self._topology() @@ -5882,6 +5288,41 @@ def _gather_tensor_parallel_logits(self, logits: torch.Tensor) -> torch.Tensor: _plan_retained_tokens = _micro_batch_planner._plan_retained_tokens _planning_status = _micro_batch_planner._planning_status + # Checkpoint-slot bookkeeping lives in ``_slots``; binding the + # functions here keeps ``self._x(...)`` dispatch and per-instance + # overrides (tests monkeypatch these) behaving as before. + _resolve_custom_checkpoint = _slots._resolve_custom_checkpoint + prefetch_checkpoints = _slots.prefetch_checkpoints + _register_checkpoint_prefetch = _slots._register_checkpoint_prefetch + _checkpoint_prefetch_waiter = _slots._checkpoint_prefetch_waiter + _prefetched_checkpoint = _slots._prefetched_checkpoint + _load_registered_checkpoint = _slots._load_registered_checkpoint + _ensure_checkpoint_slots = _slots._ensure_checkpoint_slots + load_checkpoint = _slots.load_checkpoint + _discard_snapshot_checkpoint = _slots._discard_snapshot_checkpoint + _push_checkpoint = _slots._push_checkpoint + _push_checkpoint_sync = _slots._push_checkpoint_sync + pop_checkpoint = _slots.pop_checkpoint + _resolve_checkpoint_name = _slots._resolve_checkpoint_name + _checkpoint_group = _slots._checkpoint_group + _validate_checkpoint_adapter_config = _slots._validate_checkpoint_adapter_config + _validate_loaded_checkpoint_config = _slots._validate_loaded_checkpoint_config + _load_checkpoint_slot = _slots._load_checkpoint_slot + _iter_slot_parameters = _slots._iter_slot_parameters + _validate_checkpoint_consistency = _slots._validate_checkpoint_consistency + _set_default_slot = _slots._set_default_slot + _resolve_slot_ref = _slots._resolve_slot_ref + _selected_dynamic_checkpoints = _slots._selected_dynamic_checkpoints + _checkpoint_grad_flags = _slots._checkpoint_grad_flags + _ensure_checkpoint_slots_for = _slots._ensure_checkpoint_slots_for + _track_slot_graph_outputs = _slots._track_slot_graph_outputs + _slot_graphs = _slots._slot_graphs + _prune_slot_graphs = _slots._prune_slot_graphs + _has_live_slot_graph = _slots._has_live_slot_graph + _guard_slot_can_load = _slots._guard_slot_can_load + _guard_checkpoint_can_step = _slots._guard_checkpoint_can_step + _guard_checkpoints_can_step = _slots._guard_checkpoints_can_step + def _validate_top_k(top_k: int, model: object) -> None: vocab_size = _padded_vocab_size(model) diff --git a/src/art/trainer_rank/_slots.py b/src/art/trainer_rank/_slots.py new file mode 100644 index 000000000..7e867632d --- /dev/null +++ b/src/art/trainer_rank/_slots.py @@ -0,0 +1,663 @@ +"""TrainerRank checkpoint-slot bookkeeping: prefetch registry, slot +loading and validation, the slot stack, and slot-graph liveness guards. + +These are ``TrainerRank`` method bodies moved out of ``_impl`` verbatim: each +function takes the owning rank as ``self`` and ``TrainerRank`` binds them as +methods, so ``self._x(...)`` dispatch and per-instance overrides keep working. + +Module globals the bodies used to read from ``_impl`` (``torch``, ``dist``, +sibling helpers) are still resolved through ``_impl`` at call time, so tests +that patch ``_impl.torch`` and friends keep intercepting them; only pure +stdlib helpers are imported here directly. Referencing ``_impl`` as a module +also lets the circular import resolve lazily. +""" + +from __future__ import annotations + +import asyncio +from collections.abc import Callable, Iterable, Iterator, Mapping, Sequence +from concurrent.futures import Future +from copy import deepcopy +from typing import TYPE_CHECKING, Literal, cast +import weakref + +from art.trainer_rank import _impl + +if TYPE_CHECKING: + import torch + import torch.distributed as dist + + from art.megatron.lora import LoRASlotRef + from art.trainer_rank._checkpoint import PreparedCheckpoint + from art.trainer_rank._impl import ( + AdapterSelection, + AnyForwardInput, + AnyForwardOutput, + MaterializedCheckpoint, + TrainerRank, + _AdapterConfig, + ) + + +def _resolve_custom_checkpoint(self: TrainerRank, checkpoint: AdapterSelection) -> str: + if checkpoint is _impl.Unset: + ref = self._slot_stack[-1] if self._slot_stack else self._default_slot_ref + name = None if ref is None else ref.name + else: + name = cast(str | None, checkpoint) + if name is None: + raise _impl.TrainerRankSlotStateError( + "Custom checkpoint objects require a loaded named checkpoint" + ) + self._ensure_checkpoint_slots((name,)) + if name not in self._checkpoint_slots: + raise _impl.TrainerRankSlotStateError(f"Unknown checkpoint: {name!r}") + return name + + +def prefetch_checkpoints( + self: TrainerRank, *checkpoints: str | MaterializedCheckpoint +) -> asyncio.Task[None]: + futures = [] + for checkpoint in checkpoints: + logical, source = self._checkpoint_source(checkpoint) + assert logical is not None and source is not None + futures.append(self._register_checkpoint_prefetch(logical, source)) + + async def prefetch() -> None: + await asyncio.gather( + *(self._await_checkpoint_prefetch(future) for future in futures) + ) + + return asyncio.create_task(prefetch()) + + +def _register_checkpoint_prefetch( + self: TrainerRank, + checkpoint: str, + source: str, + prepare: Callable[[], PreparedCheckpoint] | None = None, +) -> Future[PreparedCheckpoint]: + key = self._checkpoint_source_key(source) + with self._checkpoint_prefetch_lock: + previous = self._checkpoint_prefetch_sources.get(checkpoint) + self._checkpoint_prefetch_sources[checkpoint] = key + if ( + previous is not None + and previous != key + and previous not in self._checkpoint_prefetch_sources.values() + ): + self._checkpoint_prefetches.pop(previous, None) + future = self._checkpoint_prefetches.get(key) + if future is None or ( + future.done() and (future.cancelled() or future.exception() is not None) + ): + if prepare is None: + from ._checkpoint import prepare_checkpoint + + prepare = lambda: prepare_checkpoint(key) + future = _impl._checkpoint_prefetch_executor().submit(prepare) + self._checkpoint_prefetches[key] = future + return future + + +def _checkpoint_prefetch_waiter( + self: TrainerRank, *checkpoints: str +) -> asyncio.Task[None]: + with self._checkpoint_prefetch_lock: + futures = [ + self._checkpoint_prefetches[self._checkpoint_prefetch_sources[name]] + for name in checkpoints + ] + + async def wait() -> None: + await asyncio.gather( + *(self._await_checkpoint_prefetch(future) for future in futures) + ) + + return asyncio.create_task(wait()) + + +def _prefetched_checkpoint(self: TrainerRank, checkpoint: str) -> PreparedCheckpoint: + with self._checkpoint_prefetch_lock: + key = self._checkpoint_prefetch_sources.get(checkpoint) + future = None if key is None else self._checkpoint_prefetches.get(key) + if future is None: + raise _impl.TrainerRankSlotStateError( + f"Checkpoint {checkpoint!r} has not been prefetched" + ) + return future.result() + + +def _load_registered_checkpoint(self: TrainerRank, checkpoint: str) -> None: + from . import _checkpoint + + source: PreparedCheckpoint | None = None + error: BaseException | None = None + try: + source = self._prefetched_checkpoint(checkpoint) + except BaseException as exc: + error = exc + group = _checkpoint._ensure_group(self) + _checkpoint.raise_distributed(error, "prepare checkpoint", group) + assert source is not None + if checkpoint in self._snapshot_checkpoint_names: + _checkpoint.load_checkpoint(self, source, checkpoint, forward_only=True) + else: + _checkpoint.load_checkpoint(self, source, checkpoint) + + +def _ensure_checkpoint_slots(self: TrainerRank, checkpoints: Iterable[str]) -> None: + from . import _checkpoint + + requested = tuple(dict.fromkeys(checkpoints)) + with self._checkpoint_mutation_lock: + group = _checkpoint._ensure_group(self) + names = sorted( + { + name + for rank_names in _checkpoint._gather(requested, group) + for name in rank_names + } + ) + for name in names: + with self._checkpoint_prefetch_lock: + state = ( + name in self._checkpoint_slots, + name in self._checkpoint_prefetch_sources, + ) + states = _checkpoint._gather(state, group) + if all(loaded for loaded, _prefetched in states): + continue + if any(loaded for loaded, _prefetched in states): + raise _impl.TrainerRankSlotStateError( + f"Checkpoint {name!r} is not loaded consistently across ranks" + ) + if not all(prefetched for _loaded, prefetched in states): + raise _impl.TrainerRankSlotStateError( + f"Explicit selection references unloaded checkpoint {name!r}; " + "it has not been prefetched on every rank" + ) + self._load_registered_checkpoint(name) + + +def load_checkpoint( + self: TrainerRank, checkpoint: str | MaterializedCheckpoint | None +) -> None: + self._guard_forward_collective("load_checkpoint") + logical, source = self._checkpoint_source(checkpoint) + with self._checkpoint_mutation_lock: + if self._slot_stack: + raise RuntimeError("Cannot load a checkpoint while one is pushed") + if logical is None: + self._set_default_slot(self._slot_ref(None)) + return + assert source is not None + if ( + isinstance(checkpoint, _impl.MaterializedCheckpoint) + or logical not in self._checkpoint_prefetch_sources + ): + self._register_checkpoint_prefetch(logical, source) + self._load_registered_checkpoint(logical) + self._set_default_slot(self._slot_ref(logical)) + + +def _discard_snapshot_checkpoint(self: TrainerRank, checkpoint: str) -> None: + """Discard a forward-only resident snapshot.""" + from . import _checkpoint + + _checkpoint.discard_snapshot_checkpoint(self, checkpoint) + + +def _push_checkpoint( + self: TrainerRank, checkpoint: str | MaterializedCheckpoint | None +) -> None: + logical, source = self._checkpoint_source(checkpoint) + self._push_checkpoint_sync(logical, source) + + +def _push_checkpoint_sync( + self: TrainerRank, logical_path: str | None, source_path: str | None +) -> None: + self._guard_forward_collective("push_checkpoint") + with self._checkpoint_mutation_lock: + if source_path is not None: + assert logical_path is not None + if ( + logical_path not in self._checkpoint_slots + and logical_path not in self._checkpoint_prefetch_sources + ): + self._register_checkpoint_prefetch(logical_path, source_path) + self._ensure_checkpoint_slots((logical_path,)) + self._slot_stack.append(self._slot_ref(logical_path)) + + +def pop_checkpoint(self: TrainerRank) -> None: + with self._checkpoint_mutation_lock: + if not self._slot_stack: + raise RuntimeError("No pushed checkpoint to pop") + self._slot_stack.pop() + + +def _resolve_checkpoint_name( + self: TrainerRank, checkpoint_path: str | Literal["active"] +) -> str: + if checkpoint_path != "active": + self._ensure_checkpoint_slots((checkpoint_path,)) + 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 _impl.TrainerRankSlotStateError("No active trainable checkpoint") + return ref.name + + +def _checkpoint_group(self: TrainerRank) -> dist.ProcessGroup | None: + from ._checkpoint import _ensure_group + + return _ensure_group(self) + + +def _validate_checkpoint_adapter_config( + self: TrainerRank, + name: str, + adapter_config: Mapping[str, object] | None, + *, + alpha: float | None, +) -> _AdapterConfig | None: + config = None if adapter_config is None else deepcopy(dict(adapter_config)) + if _impl.dist.is_available() and _impl.dist.is_initialized(): + gathered: list[dict[str, object] | None] = [None] * _impl.dist.get_world_size() + _impl.dist.all_gather_object(gathered, config, group=self._checkpoint_group()) + if any(value != config for value in gathered): + raise ValueError( + f"Adapter config for checkpoint slot {name!r} differs across ranks" + ) + if config is None: + return None + required = {"base_model_name_or_path", "r", "lora_alpha", "target_modules"} + if missing := sorted(required - config.keys()): + raise ValueError( + f"Adapter config for checkpoint slot {name!r} is missing {missing}" + ) + base_model = config["base_model_name_or_path"] + rank = config["r"] + config_alpha_value = config["lora_alpha"] + target_modules = config["target_modules"] + if not isinstance(base_model, str): + raise TypeError("adapter_config['base_model_name_or_path'] must be a string") + if base_model.startswith(("Qwen/Qwen3.5-", "Qwen/Qwen3.6-", "Qwen/Qwen3.8-")): + from art.megatron.model_support.lora_disk import ( + model_attention_dimensions, + ) + + config.update(model_attention_dimensions(self.runtime.provider)) + if not isinstance(rank, int) or isinstance(rank, bool): + raise TypeError("adapter_config['r'] must be an integer") + if not isinstance(config_alpha_value, int | float) or isinstance( + config_alpha_value, bool + ): + raise TypeError("adapter_config['lora_alpha'] must be numeric") + if not isinstance(target_modules, str | list) or ( + isinstance(target_modules, list) + and not all(isinstance(module, str) for module in target_modules) + ): + raise TypeError( + "adapter_config['target_modules'] must be a string or list of strings" + ) + if rank < 1: + raise ValueError("adapter_config['r'] must be >= 1") + config_alpha = float(config_alpha_value) + if alpha is not None and float(alpha) != config_alpha: + raise ValueError( + f"alpha={alpha} conflicts with adapter_config lora_alpha={config_alpha}" + ) + return cast(_impl._AdapterConfig, config) + + +def _validate_loaded_checkpoint_config( + self: TrainerRank, name: str, config: _AdapterConfig +) -> None: + from art.megatron.lora import LoRA + + ref = self._slot_ref(name) + slots = [ + slot + for chunk in self.runtime.model + for module in chunk.modules() + if isinstance(module, LoRA) + if (slot := module._slot(ref)) is not None + ] + expected = (int(config["r"]), float(config["lora_alpha"])) + actual = {(slot.rank, slot.alpha) for slot in slots} + if actual != {expected}: + raise ValueError( + f"Adapter config for checkpoint slot {name!r} declares " + f"rank/alpha={expected}, loaded weights use {sorted(actual)}" + ) + + +def _load_checkpoint_slot( + self: TrainerRank, + name: str, + adapter_model: Mapping[str, torch.Tensor], + *, + alpha: float, + _prepared: bool = False, +) -> int: + if self._slot_stack: + raise RuntimeError("Cannot load a checkpoint while one is pushed") + adapter_model = self._prepare_adapter_model( + name, adapter_model, canonicalized=_prepared + ) + from art.megatron.lora import load_lora_slot_into_model + + ref = self._slot_ref(name) + self._guard_slot_can_load(ref) + self._compact_lora_slot_keys() + return load_lora_slot_into_model( + self.runtime.model, + ref, + adapter_model, + alpha=alpha, + requires_grad=True, + ) + + +def _iter_slot_parameters( + self: TrainerRank, ref: "LoRASlotRef" +) -> Iterator[torch.nn.Parameter]: + from art.megatron.lora import iter_lora_slot_parameters + + return iter_lora_slot_parameters(self.runtime.model, ref) + + +def _validate_checkpoint_consistency( + self: TrainerRank, name: str, loaded_sites: int, expected_keys: set[str] +) -> tuple[torch.nn.Parameter, ...]: + params = tuple(self._iter_slot_parameters(self._slot_ref(name))) + local_keys = { + key for group in self._local_parameter_key_groups(name) for key in group + } + gathered = ( + [local_keys] + if not (_impl.dist.is_available() and _impl.dist.is_initialized()) + else [None] * _impl.dist.get_world_size() + ) + if _impl.dist.is_available() and _impl.dist.is_initialized(): + _impl.dist.all_gather_object( + gathered, local_keys, group=self._checkpoint_group() + ) + covered = set().union(*(keys for keys in gathered if keys is not None)) + if loaded_sites < 1 or covered != expected_keys: + raise _impl.TrainerRankSlotStateError( + f"Checkpoint {name!r} has inconsistent distributed coverage" + ) + return params + + +def _set_default_slot(self: TrainerRank, ref: "LoRASlotRef") -> None: + if self._slot_stack: + raise RuntimeError("Cannot select a checkpoint while one is pushed") + self._default_slot_ref = ref + + +def _resolve_slot_ref( + self: TrainerRank, + request: AnyForwardInput, + *, + checkpoint: AdapterSelection = _impl.Unset, +) -> "LoRASlotRef | None": + selection = ( + request.checkpoint if request.checkpoint is not _impl.Unset else checkpoint + ) + if selection is not _impl.Unset: + name = cast(str | None, selection) + if name is not None and name not in self._checkpoint_slots: + raise _impl.TrainerRankSlotStateError( + f"Forward selects unloaded checkpoint {name!r}" + ) + 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(None) + + +def _selected_dynamic_checkpoints( + self: TrainerRank, + checkpoints: Sequence[str] | None, +) -> tuple[str, ...]: + if checkpoints is not None: + self._ensure_checkpoint_slots(checkpoints) + loaded = set(self._checkpoint_slots) + if not loaded: + raise _impl.TrainerRankSlotStateError( + "TrainerRank.optim_step requires a loaded checkpoint slot. Call " + "load_checkpoint(...) and run backward on outputs produced by " + "that slot before stepping." + ) + requested = ( + tuple( + sorted(name for name in loaded if not self._checkpoint_slots[name].snapshot) + ) + if checkpoints is None + else tuple(dict.fromkeys(checkpoints)) + ) + if not requested: + if checkpoints is None: + raise _impl.TrainerRankSlotStateError( + "TrainerRank.optim_step requires a loaded trainable checkpoint " + "slot. Call load_checkpoint(...) and run backward on outputs " + "produced by that slot before stepping." + ) + raise _impl.TrainerRankSlotStateError( + "TrainerRank.optim_step(checkpoints=...) received no checkpoint " + "names. Pass at least one loaded checkpoint slot." + ) + if unknown := set(requested) - loaded: + raise ValueError(f"Unknown checkpoint slots: {sorted(unknown)}") + if snapshots := [ + name for name in requested if self._checkpoint_slots[name].snapshot + ]: + raise _impl.TrainerRankSlotStateError( + f"Snapshot checkpoints are forward-only and cannot be stepped: {snapshots}" + ) + flags = self._checkpoint_grad_flags(requested) + selected = tuple( + name for name, has_grad in zip(requested, flags, strict=True) if has_grad + ) + if checkpoints is None: + if selected: + return selected + raise _impl.TrainerRankSlotStateError( + "TrainerRank.optim_step found loaded checkpoint slots, but none " + "have gradients on any rank. Call loss.backward() first." + ) + if missing := [ + name for name, has_grad in zip(requested, flags, strict=True) if not has_grad + ]: + raise _impl.TrainerRankSlotStateError( + "TrainerRank.optim_step was asked to step checkpoint slots with no " + f"gradients on any rank: {missing}. Call loss.backward() for those " + "slots first, or omit them from checkpoints=[...]." + ) + return selected + + +def _checkpoint_grad_flags(self: TrainerRank, names: Sequence[str]) -> tuple[bool, ...]: + flags = _impl.torch.tensor( + [ + any(param.grad is not None for param in self._checkpoint_slots[name].params) + for name in names + ], + device=self.device, + dtype=_impl.torch.int32, + ) + if _impl.dist.is_available() and _impl.dist.is_initialized(): + _impl.dist.all_reduce(flags, op=_impl.dist.ReduceOp.MAX) + return tuple(bool(flag) for flag in flags.tolist()) + + +def _ensure_checkpoint_slots_for( + self: TrainerRank, + requests: Sequence[AnyForwardInput], + *, + checkpoint: AdapterSelection, +) -> None: + self._ensure_checkpoint_slots( + cast(str, selection) + for request in requests + if ( + request.target_tokens is not None + or request.logits + or request.top_k is not None + or request.hidden_states + ) + if ( + selection := ( + request.checkpoint + if request.checkpoint is not _impl.Unset + else checkpoint + ) + ) + is not _impl.Unset + and selection is not None + ) + + +def _track_slot_graph_outputs( + self: TrainerRank, + ref: "LoRASlotRef | None", + outputs: Sequence[AnyForwardOutput], +) -> list[AnyForwardOutput]: + track_slot = ref is not None and ref.name is not None + track_hybridep = bool(getattr(self, "_hybridep_graph_tracking", False)) + if not track_slot and not track_hybridep: + return list(outputs) + + marker: torch.Tensor | None = None + + def track(tensor: torch.Tensor | None) -> torch.Tensor | None: + nonlocal marker + if tensor is None or not tensor.requires_grad: + return tensor + if marker is None: + marker = tensor.new_empty(0) + return cast(_impl.torch.Tensor, _impl._SlotGraphSentinel.apply(tensor, marker)) + + tracked_outputs = [ + _impl.ForwardOutput( + target_logprobs=track(output.target_logprobs), + top_k=( + None + if output.top_k is None + else _impl.TopK( + logprobs=cast(_impl.torch.Tensor, track(output.top_k.logprobs)), + tokens=output.top_k.tokens, + ) + ), + logits=track(output.logits), + hidden_states=track(output.hidden_states), + checkpoint=output.checkpoint, + no_grad=output.no_grad, + ) + for output in outputs + ] + if marker is not None: + marker_ref = weakref.ref(marker) + if track_slot: + self._slot_graphs().setdefault(ref, []).append(marker_ref) + if track_hybridep: + self._hybridep_graphs().append(marker_ref) + return tracked_outputs + + +def _slot_graphs( + self: TrainerRank, +) -> dict["LoRASlotRef", list[weakref.ReferenceType[torch.Tensor]]]: + graphs = getattr(self, "_pending_slot_graphs", None) + if graphs is None: + graphs = {} + self._pending_slot_graphs = graphs + return graphs + + +def _prune_slot_graphs(self: TrainerRank, ref: "LoRASlotRef | None" = None) -> None: + graphs = self._slot_graphs() + refs = tuple(graphs) if ref is None else (ref,) + for current in refs: + live = [ + marker + for marker in graphs.get(current, ()) + if _impl._graph_marker_is_live(marker) + ] + if live: + graphs[current] = live + else: + graphs.pop(current, None) + + +def _has_live_slot_graph(self: TrainerRank, ref: "LoRASlotRef") -> bool: + self._prune_slot_graphs(ref) + return bool(self._slot_graphs().get(ref)) + + +def _guard_slot_can_load(self: TrainerRank, ref: "LoRASlotRef") -> None: + slot = None if ref.name is None else self._checkpoint_slots.get(ref.name) + if slot is not None and slot.snapshot: + raise _impl.TrainerRankSlotStateError( + f"Cannot load over forward-only snapshot checkpoint {ref.name!r}" + ) + if slot is not None and any(param.grad is not None for param in slot.params): + raise _impl.TrainerRankSlotStateError( + f"Cannot load checkpoint {ref.name!r} while it has accumulated " + "gradients. Call optim_step() or zero_grad() before replacing it." + ) + if not self._has_live_slot_graph(ref): + return + raise _impl.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." + ) + + +def _guard_checkpoint_can_step(self: TrainerRank, name: str) -> None: + if not self._has_live_slot_graph(self._slot_ref(name)): + return + raise _impl.TrainerRankSlotStateError( + f"Cannot optim_step checkpoint slot {name!r} while outputs from an " + "earlier forward using that slot have not been backpropagated. Call " + "loss.backward() without retaining the graph before optim_step(); if " + "the forward was abandoned, release all references to its outputs." + ) + + +def _guard_checkpoints_can_step(self: TrainerRank, names: Sequence[str]) -> None: + local_live = [self._has_live_slot_graph(self._slot_ref(name)) for name in names] + if _impl.dist.is_available() and _impl.dist.is_initialized(): + live = _impl.torch.tensor( + local_live, + device=self.device, + dtype=_impl.torch.int32, + ) + _impl.dist.all_reduce(live, op=_impl.dist.ReduceOp.MAX) + live_flags = live.tolist() + else: + live_flags = local_live + blocked = [name for name, is_live in zip(names, live_flags, strict=True) if is_live] + if not blocked: + return + raise _impl.TrainerRankSlotStateError( + f"Cannot optim_step checkpoint slots {blocked!r} while outputs from an " + "earlier forward using those slots have a live backward graph on at " + "least one rank. Call loss.backward() without retaining the graph " + "before optim_step(); if the forward was abandoned, release all " + "references to its outputs; or pass on_live_graphs='allow' to accept " + "responsibility for any retained graphs." + ) From e2c832fcf511b7e7014f623321ed5282000e3dbd Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 22:02:47 +0000 Subject: [PATCH 2/2] Move TrainerRank dynamic-optimizer management into trainer_rank/_optimizer Extract the 10 dynamic (per-checkpoint) optimizer methods (optim_step and its configuration guard, dynamic optimizer creation, extension and restore, padding masks and step flags) from TrainerRank into a new _optimizer module, using the same mechanics as _slots: bodies moved verbatim as module functions taking the rank as self, bound in the TrainerRank class body, with module globals (torch, dist, _telemetry_phase and other _impl helpers) resolved through _impl at call time so tests that patch _impl.torch and friends keep intercepting them. Drops the now-unused hashlib import and LocalOptimizerState type import from _impl. Behavior-preserving code motion only. --- src/art/trainer_rank/_impl.py | 553 +--------------------------- src/art/trainer_rank/_optimizer.py | 561 +++++++++++++++++++++++++++++ 2 files changed, 580 insertions(+), 534 deletions(-) create mode 100644 src/art/trainer_rank/_optimizer.py diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 7c4965d61..7b1a72425 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -17,7 +17,6 @@ from dataclasses import dataclass, replace from dataclasses import field as dataclass_field from functools import partial -import hashlib import logging import math import os @@ -87,7 +86,6 @@ from art.megatron.train import TrainingRuntime from art.trainer_rank._checkpoint import ( CustomOptimizerState, - LocalOptimizerState, PreparedCheckpoint, PreparedCustomPayload, _FinalizedSave, @@ -1843,10 +1841,11 @@ def _moe_output_bytes_per_token( return coefficient -# ``_memory``, ``_micro_batch_planner`` and ``_slots`` hold TrainerRank method -# bodies and read ``_impl`` names such as ``Unset`` at definition time, so they are -# imported once those exist and before the class body binds their functions. -from art.trainer_rank import _memory, _micro_batch_planner, _slots +# ``_memory``, ``_micro_batch_planner``, ``_optimizer`` and ``_slots`` hold +# TrainerRank method bodies and read ``_impl`` names such as ``Unset`` at +# definition time, so they are imported once those exist and before the class +# body binds their functions. +from art.trainer_rank import _memory, _micro_batch_planner, _optimizer, _slots class TrainerRank: @@ -2281,68 +2280,6 @@ def _tag_custom_parameters(params: Sequence[torch.nn.Parameter]) -> None: setattr(param, "grad_sync_domain", "tp_default") setattr(param, "grad_sync_op", "avg") - def _extend_dynamic_optimizer( - self, - name: str, - params: Sequence[tuple[str, torch.nn.Parameter]], - ) -> _DynamicOptimizer: - from . import _checkpoint - - slot = self._checkpoint_slots[name] - dynamic = slot.optimizer - assert dynamic is not None - restored = _checkpoint.load_custom_optimizer( - slot.custom_payload, tuple(key for key, _param in params) - ) - masters = [] - for key, param in params: - state = restored.get(key) - if state is not None: - _validate_custom_optimizer_state(name, key, param, state) - source = ( - param.detach().float() - if state is None - else state.master.to(param.device) - ) - masters.append(torch.nn.Parameter(source.clone())) - master_params = tuple(masters) - all_masters = dynamic.master_params + master_params - optimizer = torch.optim.AdamW( - all_masters, - **{ - key: dynamic.optimizer.defaults[key] - for key in ( - "lr", - "betas", - "eps", - "weight_decay", - "amsgrad", - "maximize", - "foreach", - "capturable", - "differentiable", - "fused", - ) - }, - ) - optimizer.param_groups[0].update( - { - key: value - for key, value in dynamic.optimizer.param_groups[0].items() - if key != "params" - } - ) - optimizer.state.update(dynamic.optimizer.state) - for (key, _param), master in zip(params, master_params, strict=True): - state = restored.get(key) - if state is not None: - optimizer.state[master] = { - "step": torch.tensor(state.step, device=master.device), - "exp_avg": state.exp_avg.to(master.device).clone(), - "exp_avg_sq": state.exp_avg_sq.to(master.device).clone(), - } - return _DynamicOptimizer(optimizer, all_masters) - @staticmethod async def _await_checkpoint_prefetch( future: Future[PreparedCheckpoint], @@ -3061,146 +2998,6 @@ def dp_reduce( group=ps.get_data_parallel_group(with_context_parallel=False), ) - def optim_step( - self, - *, - params: AdamParams | Mapping[str, AdamParams], - scale_grads: float | Mapping[str, float] = 1.0, - checkpoints: Sequence[str] | None = None, - on_live_graphs: Literal["allow", "error"] = "allow", - ) -> dict[str, float]: - self._guard_forward_collective("optim_step") - if on_live_graphs not in ("allow", "error"): - raise ValueError( - "on_live_graphs must be either 'allow' or 'error', got " - f"{on_live_graphs!r}" - ) - params_by_checkpoint = dict(params) if isinstance(params, Mapping) else None - if params_by_checkpoint is not None: - if not params_by_checkpoint: - raise ValueError("params mapping must select at least one checkpoint") - if any(not isinstance(name, str) for name in params_by_checkpoint): - raise TypeError("params keys must be checkpoint names") - if any( - not isinstance(value, AdamParams) - for value in params_by_checkpoint.values() - ): - raise TypeError("params values must be AdamParams") - elif not isinstance(params, AdamParams): - raise TypeError( - "params must be AdamParams or a mapping of checkpoint names" - ) - if isinstance(scale_grads, Mapping): - raw_scales = cast(Mapping[object, object], scale_grads) - if not raw_scales: - raise ValueError( - "scale_grads mapping must select at least one checkpoint" - ) - if any(not isinstance(name, str) for name in raw_scales): - raise TypeError("scale_grads keys must be checkpoint names") - try: - scales_by_checkpoint = { - cast(str, name): float(cast(Any, value)) - for name, value in raw_scales.items() - } - except (TypeError, ValueError) as error: - raise TypeError("scale_grads values must be floats") from error - scale_grads_value = None - else: - scales_by_checkpoint = None - scale_grads_value = float(scale_grads) - configured = [ - tuple(value) - for value in (params_by_checkpoint, scales_by_checkpoint) - if value is not None - ] - if checkpoints is not None: - configured.append(tuple(dict.fromkeys(checkpoints))) - if configured and any(set(names) != set(configured[0]) for names in configured): - raise ValueError( - "params, scale_grads, and checkpoints must select the same " - "checkpoint names" - ) - checkpoint_selection = ( - checkpoints - if checkpoints is not None - else sorted(configured[0]) - if configured - else None - ) - self._guard_optim_step_configuration( - checkpoint_selection, params, on_live_graphs - ) - selected_checkpoints = self._selected_dynamic_checkpoints(checkpoint_selection) - params_by_checkpoint = ( - params_by_checkpoint - if params_by_checkpoint is not None - else dict.fromkeys(selected_checkpoints, cast(AdamParams, params)) - ) - scales_by_checkpoint = ( - scales_by_checkpoint - if scales_by_checkpoint is not None - else dict.fromkeys(selected_checkpoints, cast(float, scale_grads_value)) - ) - if on_live_graphs == "error": - self._guard_checkpoints_can_step(selected_checkpoints) - with _telemetry_phase( - "optim", - {"checkpoint_count": len(selected_checkpoints)}, - ): - return self._dynamic_optim_step( - selected_checkpoints, - params=params_by_checkpoint, - scale_grads=scales_by_checkpoint, - ) - - def _guard_optim_step_configuration( - self, - checkpoints: Sequence[str] | None, - params: AdamParams | Mapping[str, AdamParams], - on_live_graphs: Literal["allow", "error"], - ) -> None: - if not (dist.is_available() and dist.is_initialized()): - return - - def adam_values( - value: AdamParams, - ) -> tuple[float, float, float, float, float]: - return ( - value.learning_rate, - value.beta1, - value.beta2, - value.weight_decay, - value.grad_clip_norm, - ) - - config_values = ( - tuple( - (name, *adam_values(value)) - for name, value in sorted( - cast(Mapping[str, AdamParams], params).items() - ) - ) - if isinstance(params, Mapping) - else adam_values(params) - ) - digest = hashlib.sha256( - repr( - ( - None if checkpoints is None else tuple(checkpoints), - config_values, - on_live_graphs, - ) - ).encode() - ).digest() - local = torch.tensor(tuple(digest), device=self.device, dtype=torch.uint8) - gathered = [torch.empty_like(local) for _ in range(dist.get_world_size())] - dist.all_gather(gathered, local) - if any(not torch.equal(value, local) for value in gathered): - raise TrainerRankSlotStateError( - "Optimizer checkpoint selection or AdamParams differ across ranks" - ) - def _compact_lora_slot_keys(self) -> None: from art.megatron.lora import LoRA @@ -3307,332 +3104,6 @@ def _slot_ref(name: str | None) -> "LoRASlotRef": return LoRASlotRef(kind="checkpoint", name=name) - def _dynamic_optim_step( - self, - checkpoint_names: Sequence[str], - *, - params: Mapping[str, AdamParams], - scale_grads: Mapping[str, float], - ) -> dict[str, float]: - self.runtime.model_support_handler.zero_internal_padding_grads( - self.runtime.model - ) - selected = [] - for name in checkpoint_names: - slot_params = self._checkpoint_slots[name].params - step_flags = self._dynamic_param_step_flags(slot_params) - slot_grads = self._reduce_dynamic_grads( - slot_params, scale_grads=scale_grads[name] - ) - selected.append((name, slot_params, slot_grads, step_flags)) - - grad_norms = dict( - zip( - checkpoint_names, - _distributed_grad_norms( - [(model_params, grads) for _, model_params, grads, _ in selected] - ), - strict=True, - ) - ) - grad_norm = math.sqrt(sum(value**2 for value in grad_norms.values())) - metrics = { - "grad_norm": float(grad_norm), - "update_successful": float(math.isfinite(grad_norm)), - "num_zeros_in_grad": 0.0, - } - for name in checkpoint_names: - metrics[f"learning_rate/{name}"] = float(params[name].learning_rate) - metrics[f"grad_norm/{name}"] = float(grad_norms[name]) - learning_rates = {params[name].learning_rate for name in checkpoint_names} - if len(learning_rates) == 1: - metrics["learning_rate"] = float(params[checkpoint_names[0]].learning_rate) - if not math.isfinite(grad_norm): - for name in checkpoint_names: - for param in self._checkpoint_slots[name].params: - param.grad = None - self._prune_slot_graphs(self._slot_ref(name)) - return metrics - previous = { - name: ( - slot.optimizer, - None - if slot.optimizer is None - else [ - {key: group[key] for key in ("lr", "betas", "weight_decay")} - for group in slot.optimizer.optimizer.param_groups - ], - ) - for name in checkpoint_names - for slot in (self._checkpoint_slots[name],) - } - try: - dynamics = { - name: self._dynamic_optimizer(name, params[name]) - for name in checkpoint_names - } - except BaseException: - for name, (optimizer, groups) in previous.items(): - self._checkpoint_slots[name].optimizer = optimizer - if optimizer is not None and groups is not None: - for group, values in zip( - optimizer.optimizer.param_groups, groups, strict=True - ): - group.update(values) - raise - for name, model_params, grads, step_flags in selected: - checkpoint_params = params[name] - checkpoint_grad_norm = grad_norms[name] - clip = ( - min( - 1.0, - checkpoint_params.grad_clip_norm / (checkpoint_grad_norm + 1.0e-6), - ) - if checkpoint_params.grad_clip_norm > 0.0 - else 1.0 - ) - dynamic = dynamics[name] - for master, grad, should_step in zip( - dynamic.master_params, grads, step_flags, strict=True - ): - master.grad = grad.mul(clip) if should_step else None - dynamic.optimizer.step() - dynamic.optimizer.zero_grad(set_to_none=True) - with torch.no_grad(): - for model, master in zip( - model_params, dynamic.master_params, strict=True - ): - model.copy_(master) - model.grad = None - self._prune_slot_graphs(self._slot_ref(name)) - self._checkpoint_slots[name].revision += 1 - return metrics - - def _dynamic_param_step_flags( - self, params: Sequence[torch.nn.Parameter] - ) -> tuple[bool, ...]: - custom = [ - (index, param) - for index, param in enumerate(params) - if bool(getattr(param, "_art_custom_checkpoint_param", False)) - ] - if not custom: - return (True,) * len(params) - flags = torch.tensor( - [param.grad is not None for _, param in custom], - device=self.device, - dtype=torch.int32, - ) - if dist.is_available() and dist.is_initialized(): - dist.all_reduce(flags, op=dist.ReduceOp.MAX) - result = [True] * len(params) - for (index, _param), flag in zip(custom, flags.tolist(), strict=True): - result[index] = bool(flag) - return tuple(result) - - def _dynamic_optimizer( - self, - name: str, - params: AdamParams, - ) -> _DynamicOptimizer: - slot = self._checkpoint_slots[name] - dynamic = slot.optimizer - if dynamic is None: - dynamic = self._new_dynamic_optimizer(name, params) - slot.optimizer = dynamic - return dynamic - for group in dynamic.optimizer.param_groups: - group["lr"] = params.learning_rate - group["betas"] = (params.beta1, params.beta2) - group["weight_decay"] = params.weight_decay - return dynamic - - def _new_dynamic_optimizer( - self, - name: str, - params: AdamParams, - *, - master_params: Sequence[torch.Tensor] | None = None, - ) -> _DynamicOptimizer: - model_params = self._checkpoint_slots[name].params - sources = model_params if master_params is None else tuple(master_params) - if len(sources) != len(model_params) or any( - not isinstance(source, torch.Tensor) for source in sources - ): - raise TrainerRankSlotStateError( - 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() - ) - for model, source in zip( - model_params, - sources, - strict=True, - ) - ) - optimizer = torch.optim.AdamW( - masters, - lr=params.learning_rate, - betas=(params.beta1, params.beta2), - weight_decay=params.weight_decay, - ) - dynamic = _DynamicOptimizer(optimizer, masters) - slot = self._checkpoint_slots[name] - if slot.custom: - from ._checkpoint import load_custom_optimizer - - named = tuple( - pair - for custom_name, custom in slot.custom.items() - for pair in _custom_named_parameters(custom_name, custom) - if pair[1].requires_grad - ) - lora_count = len(model_params) - len(named) - restored = load_custom_optimizer( - slot.custom_payload, tuple(key for key, _param in named) - ) - for (key, _param), master in zip(named, masters[lora_count:], strict=True): - state = restored.get(key) - if state is not None: - _validate_custom_optimizer_state(name, key, _param, state) - with torch.no_grad(): - master.copy_(state.master.to(master.device)) - optimizer.state[master] = { - "step": torch.tensor(state.step, device=master.device), - "exp_avg": state.exp_avg.to(master.device).clone(), - "exp_avg_sq": state.exp_avg_sq.to(master.device).clone(), - } - return dynamic - - def _restore_canonical_optimizer( - self, - name: str, - state: "LocalOptimizerState", - ) -> _DynamicOptimizer: - dynamic = self._new_dynamic_optimizer( - name, - 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, - ) - 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(master.device, torch.float32).clone(), - "exp_avg_sq": exp_avg_sq.to(master.device, torch.float32).clone(), - } - 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_slots[name].params - 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(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 " - 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) - elif len(keys) == 1: - key = str(keys[0]) - exported[key] = torch.ones_like(param.T) - owners[key] = (index, None) - else: - raise TrainerRankSlotStateError( - f"Cannot map optimizer padding for checkpoint {name!r}: " - f"expected one adapter key, got {len(keys)}." - ) - - if mapped_indices and ( - missing := sorted( - index - for index, param in enumerate(params) - if index not in mapped_indices - and not bool(getattr(param, "_art_custom_checkpoint_param", False)) - ) - ): - raise TrainerRankSlotStateError( - f"Cannot map optimizer padding for checkpoint {name!r}: parameter " - f"indices {missing} do not belong to installed LoRA sites." - ) - - 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 - def _reduce_dynamic_grads( self, params: Sequence[torch.nn.Parameter], @@ -5323,6 +4794,20 @@ def _gather_tensor_parallel_logits(self, logits: torch.Tensor) -> torch.Tensor: _guard_checkpoint_can_step = _slots._guard_checkpoint_can_step _guard_checkpoints_can_step = _slots._guard_checkpoints_can_step + # Dynamic-optimizer management lives in ``_optimizer``; binding the + # functions here keeps ``self._x(...)`` dispatch and per-instance + # overrides (tests monkeypatch these) behaving as before. + _extend_dynamic_optimizer = _optimizer._extend_dynamic_optimizer + optim_step = _optimizer.optim_step + _guard_optim_step_configuration = _optimizer._guard_optim_step_configuration + _dynamic_optim_step = _optimizer._dynamic_optim_step + _dynamic_param_step_flags = _optimizer._dynamic_param_step_flags + _dynamic_optimizer = _optimizer._dynamic_optimizer + _new_dynamic_optimizer = _optimizer._new_dynamic_optimizer + _restore_canonical_optimizer = _optimizer._restore_canonical_optimizer + _zero_dynamic_optimizer_padding = _optimizer._zero_dynamic_optimizer_padding + _dynamic_optimizer_padding_masks = _optimizer._dynamic_optimizer_padding_masks + def _validate_top_k(top_k: int, model: object) -> None: vocab_size = _padded_vocab_size(model) diff --git a/src/art/trainer_rank/_optimizer.py b/src/art/trainer_rank/_optimizer.py new file mode 100644 index 000000000..0285582af --- /dev/null +++ b/src/art/trainer_rank/_optimizer.py @@ -0,0 +1,561 @@ +"""TrainerRank dynamic (per-checkpoint) optimizer management: optim_step, +its configuration guard, and dynamic optimizer creation, extension, +restore, padding masks and step flags. + +These are ``TrainerRank`` method bodies moved out of ``_impl`` verbatim: each +function takes the owning rank as ``self`` and ``TrainerRank`` binds them as +methods, so ``self._x(...)`` dispatch and per-instance overrides keep working. + +Module globals the bodies used to read from ``_impl`` (``torch``, ``dist``, +sibling helpers) are still resolved through ``_impl`` at call time, so tests +that patch ``_impl.torch`` and friends keep intercepting them; only pure +stdlib helpers are imported here directly. Referencing ``_impl`` as a module +also lets the circular import resolve lazily. +""" + +from __future__ import annotations + +from collections.abc import Mapping, Sequence +import hashlib +import math +from typing import TYPE_CHECKING, Any, Literal, cast + +from art.trainer_rank import _impl + +if TYPE_CHECKING: + import torch + + from art.trainer_rank._checkpoint import LocalOptimizerState + from art.trainer_rank._impl import AdamParams, TrainerRank, _DynamicOptimizer + + +def _extend_dynamic_optimizer( + self: TrainerRank, + name: str, + params: Sequence[tuple[str, torch.nn.Parameter]], +) -> _DynamicOptimizer: + from . import _checkpoint + + slot = self._checkpoint_slots[name] + dynamic = slot.optimizer + assert dynamic is not None + restored = _checkpoint.load_custom_optimizer( + slot.custom_payload, tuple(key for key, _param in params) + ) + masters = [] + for key, param in params: + state = restored.get(key) + if state is not None: + _impl._validate_custom_optimizer_state(name, key, param, state) + source = ( + param.detach().float() if state is None else state.master.to(param.device) + ) + masters.append(_impl.torch.nn.Parameter(source.clone())) + master_params = tuple(masters) + all_masters = dynamic.master_params + master_params + optimizer = _impl.torch.optim.AdamW( + all_masters, + **{ + key: dynamic.optimizer.defaults[key] + for key in ( + "lr", + "betas", + "eps", + "weight_decay", + "amsgrad", + "maximize", + "foreach", + "capturable", + "differentiable", + "fused", + ) + }, + ) + optimizer.param_groups[0].update( + { + key: value + for key, value in dynamic.optimizer.param_groups[0].items() + if key != "params" + } + ) + optimizer.state.update(dynamic.optimizer.state) + for (key, _param), master in zip(params, master_params, strict=True): + state = restored.get(key) + if state is not None: + optimizer.state[master] = { + "step": _impl.torch.tensor(state.step, device=master.device), + "exp_avg": state.exp_avg.to(master.device).clone(), + "exp_avg_sq": state.exp_avg_sq.to(master.device).clone(), + } + return _impl._DynamicOptimizer(optimizer, all_masters) + + +def optim_step( + self: TrainerRank, + *, + params: AdamParams | Mapping[str, AdamParams], + scale_grads: float | Mapping[str, float] = 1.0, + checkpoints: Sequence[str] | None = None, + on_live_graphs: Literal["allow", "error"] = "allow", +) -> dict[str, float]: + self._guard_forward_collective("optim_step") + if on_live_graphs not in ("allow", "error"): + raise ValueError( + f"on_live_graphs must be either 'allow' or 'error', got {on_live_graphs!r}" + ) + params_by_checkpoint = dict(params) if isinstance(params, Mapping) else None + if params_by_checkpoint is not None: + if not params_by_checkpoint: + raise ValueError("params mapping must select at least one checkpoint") + if any(not isinstance(name, str) for name in params_by_checkpoint): + raise TypeError("params keys must be checkpoint names") + if any( + not isinstance(value, _impl.AdamParams) + for value in params_by_checkpoint.values() + ): + raise TypeError("params values must be AdamParams") + elif not isinstance(params, _impl.AdamParams): + raise TypeError("params must be AdamParams or a mapping of checkpoint names") + if isinstance(scale_grads, Mapping): + raw_scales = cast(Mapping[object, object], scale_grads) + if not raw_scales: + raise ValueError("scale_grads mapping must select at least one checkpoint") + if any(not isinstance(name, str) for name in raw_scales): + raise TypeError("scale_grads keys must be checkpoint names") + try: + scales_by_checkpoint = { + cast(str, name): float(cast(Any, value)) + for name, value in raw_scales.items() + } + except (TypeError, ValueError) as error: + raise TypeError("scale_grads values must be floats") from error + scale_grads_value = None + else: + scales_by_checkpoint = None + scale_grads_value = float(scale_grads) + configured = [ + tuple(value) + for value in (params_by_checkpoint, scales_by_checkpoint) + if value is not None + ] + if checkpoints is not None: + configured.append(tuple(dict.fromkeys(checkpoints))) + if configured and any(set(names) != set(configured[0]) for names in configured): + raise ValueError( + "params, scale_grads, and checkpoints must select the same checkpoint names" + ) + checkpoint_selection = ( + checkpoints + if checkpoints is not None + else sorted(configured[0]) + if configured + else None + ) + self._guard_optim_step_configuration(checkpoint_selection, params, on_live_graphs) + selected_checkpoints = self._selected_dynamic_checkpoints(checkpoint_selection) + params_by_checkpoint = ( + params_by_checkpoint + if params_by_checkpoint is not None + else dict.fromkeys(selected_checkpoints, cast(_impl.AdamParams, params)) + ) + scales_by_checkpoint = ( + scales_by_checkpoint + if scales_by_checkpoint is not None + else dict.fromkeys(selected_checkpoints, cast(float, scale_grads_value)) + ) + if on_live_graphs == "error": + self._guard_checkpoints_can_step(selected_checkpoints) + with _impl._telemetry_phase( + "optim", + {"checkpoint_count": len(selected_checkpoints)}, + ): + return self._dynamic_optim_step( + selected_checkpoints, + params=params_by_checkpoint, + scale_grads=scales_by_checkpoint, + ) + + +def _guard_optim_step_configuration( + self: TrainerRank, + checkpoints: Sequence[str] | None, + params: AdamParams | Mapping[str, AdamParams], + on_live_graphs: Literal["allow", "error"], +) -> None: + if not (_impl.dist.is_available() and _impl.dist.is_initialized()): + return + + def adam_values( + value: AdamParams, + ) -> tuple[float, float, float, float, float]: + return ( + value.learning_rate, + value.beta1, + value.beta2, + value.weight_decay, + value.grad_clip_norm, + ) + + config_values = ( + tuple( + (name, *adam_values(value)) + for name, value in sorted( + cast(Mapping[str, _impl.AdamParams], params).items() + ) + ) + if isinstance(params, Mapping) + else adam_values(params) + ) + digest = hashlib.sha256( + repr( + ( + None if checkpoints is None else tuple(checkpoints), + config_values, + on_live_graphs, + ) + ).encode() + ).digest() + local = _impl.torch.tensor( + tuple(digest), device=self.device, dtype=_impl.torch.uint8 + ) + gathered = [ + _impl.torch.empty_like(local) for _ in range(_impl.dist.get_world_size()) + ] + _impl.dist.all_gather(gathered, local) + if any(not _impl.torch.equal(value, local) for value in gathered): + raise _impl.TrainerRankSlotStateError( + "Optimizer checkpoint selection or AdamParams differ across ranks" + ) + + +def _dynamic_optim_step( + self: TrainerRank, + checkpoint_names: Sequence[str], + *, + params: Mapping[str, AdamParams], + scale_grads: Mapping[str, float], +) -> dict[str, float]: + self.runtime.model_support_handler.zero_internal_padding_grads(self.runtime.model) + selected = [] + for name in checkpoint_names: + slot_params = self._checkpoint_slots[name].params + step_flags = self._dynamic_param_step_flags(slot_params) + slot_grads = self._reduce_dynamic_grads( + slot_params, scale_grads=scale_grads[name] + ) + selected.append((name, slot_params, slot_grads, step_flags)) + + grad_norms = dict( + zip( + checkpoint_names, + _impl._distributed_grad_norms( + [(model_params, grads) for _, model_params, grads, _ in selected] + ), + strict=True, + ) + ) + grad_norm = math.sqrt(sum(value**2 for value in grad_norms.values())) + metrics = { + "grad_norm": float(grad_norm), + "update_successful": float(math.isfinite(grad_norm)), + "num_zeros_in_grad": 0.0, + } + for name in checkpoint_names: + metrics[f"learning_rate/{name}"] = float(params[name].learning_rate) + metrics[f"grad_norm/{name}"] = float(grad_norms[name]) + learning_rates = {params[name].learning_rate for name in checkpoint_names} + if len(learning_rates) == 1: + metrics["learning_rate"] = float(params[checkpoint_names[0]].learning_rate) + if not math.isfinite(grad_norm): + for name in checkpoint_names: + for param in self._checkpoint_slots[name].params: + param.grad = None + self._prune_slot_graphs(self._slot_ref(name)) + return metrics + previous = { + name: ( + slot.optimizer, + None + if slot.optimizer is None + else [ + {key: group[key] for key in ("lr", "betas", "weight_decay")} + for group in slot.optimizer.optimizer.param_groups + ], + ) + for name in checkpoint_names + for slot in (self._checkpoint_slots[name],) + } + try: + dynamics = { + name: self._dynamic_optimizer(name, params[name]) + for name in checkpoint_names + } + except BaseException: + for name, (optimizer, groups) in previous.items(): + self._checkpoint_slots[name].optimizer = optimizer + if optimizer is not None and groups is not None: + for group, values in zip( + optimizer.optimizer.param_groups, groups, strict=True + ): + group.update(values) + raise + for name, model_params, grads, step_flags in selected: + checkpoint_params = params[name] + checkpoint_grad_norm = grad_norms[name] + clip = ( + min( + 1.0, + checkpoint_params.grad_clip_norm / (checkpoint_grad_norm + 1.0e-6), + ) + if checkpoint_params.grad_clip_norm > 0.0 + else 1.0 + ) + dynamic = dynamics[name] + for master, grad, should_step in zip( + dynamic.master_params, grads, step_flags, strict=True + ): + master.grad = grad.mul(clip) if should_step else None + dynamic.optimizer.step() + dynamic.optimizer.zero_grad(set_to_none=True) + with _impl.torch.no_grad(): + for model, master in zip(model_params, dynamic.master_params, strict=True): + model.copy_(master) + model.grad = None + self._prune_slot_graphs(self._slot_ref(name)) + self._checkpoint_slots[name].revision += 1 + return metrics + + +def _dynamic_param_step_flags( + self: TrainerRank, params: Sequence[torch.nn.Parameter] +) -> tuple[bool, ...]: + custom = [ + (index, param) + for index, param in enumerate(params) + if bool(getattr(param, "_art_custom_checkpoint_param", False)) + ] + if not custom: + return (True,) * len(params) + flags = _impl.torch.tensor( + [param.grad is not None for _, param in custom], + device=self.device, + dtype=_impl.torch.int32, + ) + if _impl.dist.is_available() and _impl.dist.is_initialized(): + _impl.dist.all_reduce(flags, op=_impl.dist.ReduceOp.MAX) + result = [True] * len(params) + for (index, _param), flag in zip(custom, flags.tolist(), strict=True): + result[index] = bool(flag) + return tuple(result) + + +def _dynamic_optimizer( + self: TrainerRank, + name: str, + params: AdamParams, +) -> _DynamicOptimizer: + slot = self._checkpoint_slots[name] + dynamic = slot.optimizer + if dynamic is None: + dynamic = self._new_dynamic_optimizer(name, params) + slot.optimizer = dynamic + return dynamic + for group in dynamic.optimizer.param_groups: + group["lr"] = params.learning_rate + group["betas"] = (params.beta1, params.beta2) + group["weight_decay"] = params.weight_decay + return dynamic + + +def _new_dynamic_optimizer( + self: TrainerRank, + name: str, + params: AdamParams, + *, + master_params: Sequence[torch.Tensor] | None = None, +) -> _DynamicOptimizer: + model_params = self._checkpoint_slots[name].params + sources = model_params if master_params is None else tuple(master_params) + if len(sources) != len(model_params) or any( + not isinstance(source, _impl.torch.Tensor) for source in sources + ): + raise _impl.TrainerRankSlotStateError( + 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 _impl.TrainerRankSlotStateError( + f"Optimizer master parameter shape does not match checkpoint {name!r}" + ) + masters = tuple( + _impl.torch.nn.Parameter( + source.detach().to(device=model.device, dtype=_impl.torch.float32).clone() + ) + for model, source in zip( + model_params, + sources, + strict=True, + ) + ) + optimizer = _impl.torch.optim.AdamW( + masters, + lr=params.learning_rate, + betas=(params.beta1, params.beta2), + weight_decay=params.weight_decay, + ) + dynamic = _impl._DynamicOptimizer(optimizer, masters) + slot = self._checkpoint_slots[name] + if slot.custom: + from ._checkpoint import load_custom_optimizer + + named = tuple( + pair + for custom_name, custom in slot.custom.items() + for pair in _impl._custom_named_parameters(custom_name, custom) + if pair[1].requires_grad + ) + lora_count = len(model_params) - len(named) + restored = load_custom_optimizer( + slot.custom_payload, tuple(key for key, _param in named) + ) + for (key, _param), master in zip(named, masters[lora_count:], strict=True): + state = restored.get(key) + if state is not None: + _impl._validate_custom_optimizer_state(name, key, _param, state) + with _impl.torch.no_grad(): + master.copy_(state.master.to(master.device)) + optimizer.state[master] = { + "step": _impl.torch.tensor(state.step, device=master.device), + "exp_avg": state.exp_avg.to(master.device).clone(), + "exp_avg_sq": state.exp_avg_sq.to(master.device).clone(), + } + return dynamic + + +def _restore_canonical_optimizer( + self: TrainerRank, + name: str, + state: "LocalOptimizerState", +) -> _DynamicOptimizer: + dynamic = self._new_dynamic_optimizer( + name, + _impl.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, + ) + 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 _impl.TrainerRankSlotStateError( + f"Canonical optimizer moment shape does not match {name!r}" + ) + dynamic.optimizer.state[master] = { + "step": _impl.torch.tensor(step, dtype=_impl.torch.float32), + "exp_avg": exp_avg.to(master.device, _impl.torch.float32).clone(), + "exp_avg_sq": exp_avg_sq.to(master.device, _impl.torch.float32).clone(), + } + self._zero_dynamic_optimizer_padding(name, dynamic) + return dynamic + + +def _zero_dynamic_optimizer_padding( + self: TrainerRank, + name: str, + dynamic: _DynamicOptimizer, +) -> None: + masks = self._dynamic_optimizer_padding_masks(name) + with _impl.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, _impl.torch.Tensor) and value.shape == param.shape: + value.masked_fill_(mask, 0) + + +def _dynamic_optimizer_padding_masks( + self: TrainerRank, name: str +) -> tuple[torch.Tensor, ...]: + params = self._checkpoint_slots[name].params + masks = tuple( + _impl.torch.zeros_like(param, dtype=_impl.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(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 _impl.TrainerRankSlotStateError( + f"Cannot map optimizer padding for checkpoint " + f"{name!r}: {len(keys)} adapter keys describe " + f"{int(param.shape[0])} local experts." + ) + for expert, key in enumerate(keys): + exported[str(key)] = _impl.torch.ones_like(param[expert].T) + owners[str(key)] = (index, expert) + elif len(keys) == 1: + key = str(keys[0]) + exported[key] = _impl.torch.ones_like(param.T) + owners[key] = (index, None) + else: + raise _impl.TrainerRankSlotStateError( + f"Cannot map optimizer padding for checkpoint {name!r}: " + f"expected one adapter key, got {len(keys)}." + ) + + if mapped_indices and ( + missing := sorted( + index + for index, param in enumerate(params) + if index not in mapped_indices + and not bool(getattr(param, "_art_custom_checkpoint_param", False)) + ) + ): + raise _impl.TrainerRankSlotStateError( + f"Cannot map optimizer padding for checkpoint {name!r}: parameter " + f"indices {missing} do not belong to installed LoRA sites." + ) + + 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, _impl.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