From f89213adbb1e10825a1205251bd50725e8b28f41 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 18 Sep 2026 00:26:25 +0000 Subject: [PATCH 001/150] Add versioned trainer state and logical rank execution --- .github/workflows/prek.yml | 6 +- dev/trainer_rank.py | 8 +- dev/trainer_rank_check.py | 16 +- dev/trainer_rank_collective_trace.py | 4 +- dev/trainer_rank_landing_acceptance.py | 79 +- dev/trainer_rank_landing_acceptance_README.md | 6 +- ...r_rank_landing_acceptance_dp2_tp2.sky.yaml | 2 +- ...ainer_rank_landing_acceptance_tp2.sky.yaml | 2 +- dev/trainer_rank_planner_design.md | 14 +- dev/trainer_rank_recompute_memory.py | 2 +- dev/trainer_v1_acceptance.py | 509 ++++++ dev/trainer_v1_benchmark.py | 330 ++++ dev/trainer_v1_checkpoint.py | 42 + dev/trainer_v1_facade_benchmark.py | 235 +++ dev/trainer_v1_head_memory.py | 181 +++ dev/trainer_v1_memory.py | 383 +++++ dev/trainer_v1_validation.sky.yaml | 34 + pyproject.toml | 1 + scripts/ci/trainer-rank-gpu-tests.sh | 20 + src/art/_tensor_residency.py | 28 + src/art/megatron/compile_workarounds.py | 111 ++ src/art/megatron/context_parallel/executor.py | 37 +- src/art/megatron/lora.py | 43 +- src/art/megatron/prefix_tree_packing.py | 2 +- src/art/megatron/runtime/compile_cache.py | 21 +- src/art/megatron/training/compile.py | 12 +- src/art/trainer_rank/__init__.py | 121 +- src/art/trainer_rank/_checkpoint.py | 26 + src/art/trainer_rank/_commands.py | 1144 +++++++++++++ src/art/trainer_rank/_corrections.py | 332 ++++ src/art/trainer_rank/_graphs.py | 704 ++++++++ src/art/trainer_rank/_heads.py | 1363 ++++++++++++++++ src/art/trainer_rank/_impl.py | 1401 ++++++++++++++-- src/art/trainer_rank/_memory_policy.py | 352 ++++ src/art/trainer_rank/_operations.py | 150 ++ src/art/trainer_rank/_options.py | 162 ++ src/art/trainer_rank/_parameter_hooks.py | 141 ++ src/art/trainer_rank/_tensors.py | 491 ++++++ src/art/trainer_rank/_versions.py | 336 ++++ .../test_public_contract.py | 22 +- .../cp_attn/test_cpu_offload_residency.py | 253 +++ .../cp_attn/test_retained_backward.py | 206 +++ .../megatron/lora/test_dynamic_lora_slots.py | 5 +- .../megatron/lora/test_lora_versions.py | 206 +++ .../lora/test_trainer_v1_graph_cache.py | 219 +++ .../megatron/lora/test_trainer_v1_versions.py | 171 ++ .../model_support/test_compile_flags.py | 31 +- .../unit/test_trainer_batch_input_capture.py | 168 ++ tests/unit/test_trainer_command_transport.py | 238 +++ tests/unit/test_trainer_driver_transport.py | 255 +++ .../unit/test_trainer_live_parameter_roots.py | 137 ++ tests/unit/test_trainer_operations.py | 325 ++++ tests/unit/test_trainer_rank_active_memory.py | 4 +- .../unit/test_trainer_rank_cache_recovery.py | 59 +- .../test_trainer_rank_calibration_harness.py | 19 +- .../test_trainer_rank_callback_cleanup.py | 179 ++ .../test_trainer_rank_callback_failures.py | 234 +++ .../test_trainer_rank_callback_lifecycle.py | 182 +++ tests/unit/test_trainer_rank_commands.py | 809 +++++++++ tests/unit/test_trainer_rank_corrections.py | 230 +++ .../unit/test_trainer_rank_custom_tensors.py | 50 +- tests/unit/test_trainer_rank_graph_order.py | 116 ++ tests/unit/test_trainer_rank_graphs.py | 525 ++++++ tests/unit/test_trainer_rank_graphs_cuda.py | 596 +++++++ tests/unit/test_trainer_rank_head_memory.py | 198 +++ .../test_trainer_rank_head_memory_cuda.py | 129 ++ .../unit/test_trainer_rank_head_recompute.py | 4 +- tests/unit/test_trainer_rank_live_heads.py | 1448 +++++++++++++++++ .../test_trainer_rank_memory_admission.py | 564 +++++++ tests/unit/test_trainer_rank_memory_policy.py | 306 ++++ .../test_trainer_rank_memory_policy_cuda.py | 213 +++ ...rainer_rank_memory_recovery_distributed.py | 194 +++ tests/unit/test_trainer_rank_options.py | 252 +++ tests/unit/test_trainer_rank_output_memory.py | 104 ++ .../test_trainer_rank_output_memory_cuda.py | 86 + .../unit/test_trainer_rank_parameter_hooks.py | 290 ++++ .../test_trainer_rank_parameter_no_grad.py | 100 ++ .../test_trainer_rank_recompute_memory.py | 2 +- .../unit/test_trainer_rank_recovery_slots.py | 5 +- ...trainer_rank_recovery_slots_distributed.py | 4 +- .../test_trainer_rank_release_completion.py | 88 + .../test_trainer_rank_release_lifetime.py | 122 ++ .../unit/test_trainer_rank_resident_memory.py | 49 + .../test_trainer_rank_resident_memory_cuda.py | 60 + .../test_trainer_rank_slot_graph_lifetime.py | 168 ++ tests/unit/test_trainer_rank_split.py | 77 +- tests/unit/test_trainer_rank_split_peak.py | 10 +- tests/unit/test_trainer_rank_tensors.py | 1009 ++++++++++++ tests/unit/test_trainer_rank_validation.py | 221 +-- tests/unit/test_trainer_rank_versions.py | 539 ++++++ tests/unit/test_trainer_rank_weird_shapes.py | 48 +- uv.lock | 2 + 92 files changed, 19926 insertions(+), 456 deletions(-) create mode 100644 dev/trainer_v1_acceptance.py create mode 100644 dev/trainer_v1_benchmark.py create mode 100644 dev/trainer_v1_checkpoint.py create mode 100644 dev/trainer_v1_facade_benchmark.py create mode 100644 dev/trainer_v1_head_memory.py create mode 100644 dev/trainer_v1_memory.py create mode 100644 dev/trainer_v1_validation.sky.yaml create mode 100644 src/art/_tensor_residency.py create mode 100644 src/art/trainer_rank/_commands.py create mode 100644 src/art/trainer_rank/_corrections.py create mode 100644 src/art/trainer_rank/_graphs.py create mode 100644 src/art/trainer_rank/_heads.py create mode 100644 src/art/trainer_rank/_memory_policy.py create mode 100644 src/art/trainer_rank/_operations.py create mode 100644 src/art/trainer_rank/_options.py create mode 100644 src/art/trainer_rank/_parameter_hooks.py create mode 100644 src/art/trainer_rank/_tensors.py create mode 100644 src/art/trainer_rank/_versions.py create mode 100644 tests/integration/megatron/cp_attn/test_cpu_offload_residency.py create mode 100644 tests/integration/megatron/cp_attn/test_retained_backward.py create mode 100644 tests/integration/megatron/lora/test_lora_versions.py create mode 100644 tests/integration/megatron/lora/test_trainer_v1_graph_cache.py create mode 100644 tests/integration/megatron/lora/test_trainer_v1_versions.py create mode 100644 tests/unit/test_trainer_batch_input_capture.py create mode 100644 tests/unit/test_trainer_command_transport.py create mode 100644 tests/unit/test_trainer_driver_transport.py create mode 100644 tests/unit/test_trainer_live_parameter_roots.py create mode 100644 tests/unit/test_trainer_operations.py create mode 100644 tests/unit/test_trainer_rank_callback_cleanup.py create mode 100644 tests/unit/test_trainer_rank_callback_failures.py create mode 100644 tests/unit/test_trainer_rank_callback_lifecycle.py create mode 100644 tests/unit/test_trainer_rank_commands.py create mode 100644 tests/unit/test_trainer_rank_corrections.py create mode 100644 tests/unit/test_trainer_rank_graph_order.py create mode 100644 tests/unit/test_trainer_rank_graphs.py create mode 100644 tests/unit/test_trainer_rank_graphs_cuda.py create mode 100644 tests/unit/test_trainer_rank_head_memory.py create mode 100644 tests/unit/test_trainer_rank_head_memory_cuda.py create mode 100644 tests/unit/test_trainer_rank_live_heads.py create mode 100644 tests/unit/test_trainer_rank_memory_admission.py create mode 100644 tests/unit/test_trainer_rank_memory_policy.py create mode 100644 tests/unit/test_trainer_rank_memory_policy_cuda.py create mode 100644 tests/unit/test_trainer_rank_memory_recovery_distributed.py create mode 100644 tests/unit/test_trainer_rank_options.py create mode 100644 tests/unit/test_trainer_rank_output_memory.py create mode 100644 tests/unit/test_trainer_rank_output_memory_cuda.py create mode 100644 tests/unit/test_trainer_rank_parameter_hooks.py create mode 100644 tests/unit/test_trainer_rank_parameter_no_grad.py create mode 100644 tests/unit/test_trainer_rank_release_completion.py create mode 100644 tests/unit/test_trainer_rank_release_lifetime.py create mode 100644 tests/unit/test_trainer_rank_resident_memory.py create mode 100644 tests/unit/test_trainer_rank_resident_memory_cuda.py create mode 100644 tests/unit/test_trainer_rank_slot_graph_lifetime.py create mode 100644 tests/unit/test_trainer_rank_tensors.py create mode 100644 tests/unit/test_trainer_rank_versions.py diff --git a/.github/workflows/prek.yml b/.github/workflows/prek.yml index 1f8f27e54..63007cbbc 100644 --- a/.github/workflows/prek.yml +++ b/.github/workflows/prek.yml @@ -225,6 +225,8 @@ jobs: tests/unit/test_prefix_tree_packing.py \ tests/unit/test_trainer_rank_validation.py \ tests/unit/test_trainer_rank_weird_shapes.py \ + tests/unit/test_trainer_rank_slot_graph_lifetime.py \ + tests/unit/test_trainer_rank_commands.py::test_gloo_dp2_tp2_participation_and_gradients \ tests/unit/test_trainer_rank_split.py \ tests/unit/test_trainer_rank_cache_recovery.py::test_dense_cp_exact_demand_fits_after_recovery \ tests/unit/test_trainer_rank_cache_recovery.py::test_dense_cp_exact_demand_refuses_after_recovery \ @@ -244,6 +246,7 @@ jobs: uv run --no-sync pytest --nbval --current-env --tb=short tests/unit \ --deselect=tests/unit/test_trainer_rank_cache_recovery.py::test_dense_cp_exact_demand_fits_after_recovery \ --deselect=tests/unit/test_trainer_rank_cache_recovery.py::test_dense_cp_exact_demand_refuses_after_recovery \ + --deselect=tests/unit/test_trainer_rank_commands.py::test_gloo_dp2_tp2_participation_and_gradients \ --ignore=tests/unit/test_megatron_reference_logprobs.py \ --ignore=tests/unit/test_moe_routing_replay.py \ --ignore=tests/unit/test_moe_routing_real_path.py \ @@ -253,4 +256,5 @@ jobs: --ignore=tests/unit/test_prefix_tree_grad_parity.py \ --ignore=tests/unit/test_prefix_tree_packing.py \ --ignore=tests/unit/test_trainer_rank_validation.py \ - --ignore=tests/unit/test_trainer_rank_weird_shapes.py + --ignore=tests/unit/test_trainer_rank_weird_shapes.py \ + --ignore=tests/unit/test_trainer_rank_slot_graph_lifetime.py diff --git a/dev/trainer_rank.py b/dev/trainer_rank.py index 65aa2c6ff..99478890c 100644 --- a/dev/trainer_rank.py +++ b/dev/trainer_rank.py @@ -70,17 +70,17 @@ def main( for step in range(steps): loss_sum = torch.tensor(0.0, device=rank.device) token_count = torch.tensor(0.0, device=rank.device) - for micro in rank.forward_micro_batches(inputs, checkpoint=slot): + for micro in rank.forward_batches(inputs, checkpoint=slot): loss = torch.tensor(0.0, device=rank.device) for output in micro.outputs: assert output.target_logprobs is not None loss = loss - output.target_logprobs.sum() token_count += output.target_logprobs.numel() - loss.backward() + rank.backward(loss) loss_sum += loss.detach() - rank.dp_reduce(loss_sum) - rank.dp_reduce(token_count) + rank.reduce(loss_sum) + rank.reduce(token_count) scale = 1.0 / max(float(token_count.item()), 1.0) metrics = rank.optim_step( params=AdamParams(learning_rate=lr), diff --git a/dev/trainer_rank_check.py b/dev/trainer_rank_check.py index b0acf6c42..afcbac6fd 100644 --- a/dev/trainer_rank_check.py +++ b/dev/trainer_rank_check.py @@ -226,7 +226,7 @@ def _local_outputs( rank: TrainerRank, indexed_requests: Sequence[tuple[int, ForwardInput]], ) -> list[dict[str, object]]: - outputs = rank.dp_rank_forward([request for _, request in indexed_requests]) + outputs = rank.forward([request for _, request in indexed_requests]) return [ _output_record(index, torch.arange(request.input_tokens.numel()), output) for (index, request), output in zip(indexed_requests, outputs, strict=True) @@ -502,7 +502,7 @@ def _head_backward_chunk_parity( outputs = rank._project_head(items, prepared, candidate) loss = _output_loss(outputs) logprob_sums.append(loss.detach()) - loss.backward() + rank.backward(loss) assert candidate.grad is not None gradients.append(candidate.grad) finally: @@ -584,7 +584,7 @@ def _slot_gradients( slots: Sequence[str], ) -> dict[str, list[torch.Tensor]]: rank.zero_grad() - _output_loss(rank.dp_rank_forward(requests)).backward() + rank.backward(_output_loss(rank.forward(requests))) return { slot: [ torch.zeros_like(parameter, dtype=torch.float32, device="cpu") @@ -665,14 +665,16 @@ def step() -> list[MicroBatchStats]: rank.zero_grad() stats: list[MicroBatchStats] = [] if adaptive: - for micro in rank.forward_micro_batches(requests): - _output_loss(cast(Sequence[ForwardOutput], micro.outputs)).backward() + for micro in rank.forward_batches(requests): + rank.backward( + _output_loss(cast(Sequence[ForwardOutput], micro.outputs)) + ) stats.append(micro.stats) else: - outputs = rank.dp_rank_forward(requests[dp_rank::dp_size]) + outputs = rank.forward(requests[dp_rank::dp_size]) if workload == "unequal_slots": _trace_unequal_slots("forward_ready") - _output_loss(outputs).backward() + rank.backward(_output_loss(outputs)) if workload == "unequal_slots": _trace_unequal_slots("backward_ready") if optimizer_step: diff --git a/dev/trainer_rank_collective_trace.py b/dev/trainer_rank_collective_trace.py index 6c1a05a16..94af149a2 100644 --- a/dev/trainer_rank_collective_trace.py +++ b/dev/trainer_rank_collective_trace.py @@ -141,7 +141,7 @@ def logged(*args, **kwargs): # gdn.layout imported the function by name; rebind it to the logged one. gdn_layout.all_to_all_single = dist.all_to_all_single - original_forward = TrainerRank.dp_rank_forward + original_forward = TrainerRank.forward @functools.wraps(original_forward) def logged_forward(self, *args, **kwargs): @@ -166,7 +166,7 @@ def logged_forward(self, *args, **kwargs): ) return outputs - TrainerRank.dp_rank_forward = logged_forward # type: ignore[method-assign] + TrainerRank.forward = logged_forward # type: ignore[method-assign] def _start_watchdog(logger: _Logger, log_dir: Path, stall_seconds: float) -> None: diff --git a/dev/trainer_rank_landing_acceptance.py b/dev/trainer_rank_landing_acceptance.py index dadd80c97..a4269581c 100644 --- a/dev/trainer_rank_landing_acceptance.py +++ b/dev/trainer_rank_landing_acceptance.py @@ -137,17 +137,35 @@ def phase_contract() -> None: problems: list[str] = [] constructor = _public_parameters(trainer_rank.TrainerRank.__init__) - if list(constructor) != ["runtime"]: + if ( + list(constructor) != ["runtime", "options"] + or constructor["options"].kind is not inspect.Parameter.KEYWORD_ONLY + or constructor["options"].default is not None + ): problems.append( - f"TrainerRank must accept exactly (runtime); found {sorted(constructor)}" + f"TrainerRank must accept (runtime, *, options=None); found {constructor}" ) - for method_name in ("forward_micro_batches", "dp_rank_forward"): + for method_name in ("forward_batches", "forward"): parameters = _public_parameters(getattr(trainer_rank.TrainerRank, method_name)) # ``yield_empty`` (PR #864) is a keyword-only flag defaulting to False: # off, the contract the acceptance suite pins is unchanged. - extra = set(parameters) - {"inputs", "checkpoint", "no_grad", "yield_empty"} + extra = set(parameters) - { + "inputs", + "checkpoint", + "no_grad", + "yield_empty", + "options", + } if extra: problems.append(f"{method_name} has extra parameters {sorted(extra)}") + options = parameters.get("options") + if options is None or ( + options.kind is not inspect.Parameter.KEYWORD_ONLY + or options.default is not None + ): + problems.append( + f"{method_name}: options must be keyword-only and default to None" + ) flag = parameters.get("yield_empty") if flag is not None and ( flag.kind is not inspect.Parameter.KEYWORD_ONLY or flag.default is not False @@ -374,8 +392,8 @@ def phase_measure(cell: str, arm: str, output_jsonl: str, repeat: int) -> None: # Behavior smoke from the contract: empty inputs are valid zero-work # calls that must not disturb subsequent planning. - empty = rank.dp_rank_forward([]) - assert len(list(empty)) == 0, "dp_rank_forward([]) must return no outputs" + empty = rank.forward([]) + assert len(list(empty)) == 0, "forward([]) must return no outputs" rows: list[dict[str, object]] = [] for sample in range(repeat + 4): # 1 cold + 3 warmup + repeat measured @@ -387,8 +405,8 @@ def phase_measure(cell: str, arm: str, output_jsonl: str, repeat: int) -> None: start.record() admission_failed = False try: - outputs = rank.dp_rank_forward(requests) - _output_loss(outputs).backward() + outputs = rank.forward(requests) + rank.backward(_output_loss(outputs)) except TrainerRankMemoryError as error: admission_failed = True rows.append( @@ -629,7 +647,7 @@ def forward( ) -> tuple[list[torch.Tensor], dict[str, object]]: torch.cuda.synchronize() torch.cuda.reset_peak_memory_stats() - outputs = rank.dp_rank_forward(requests(), no_grad=no_grad) + outputs = rank.forward(requests(), no_grad=no_grad) telemetry = rank.last_forward_telemetry() logprobs = [ output.target_logprobs.detach().float().clone() for output in outputs @@ -638,9 +656,11 @@ def forward( def combined_backward(info: dict[str, object]) -> None: outputs = cast(list, info["outputs"]) - torch.stack( - [output.target_logprobs.float().sum() for output in outputs] - ).sum().backward() + rank.backward( + torch.stack( + [output.target_logprobs.float().sum() for output in outputs] + ).sum() + ) torch.cuda.synchronize() def unsplit_requirement() -> int: @@ -684,7 +704,7 @@ def expect_decline(arm: str) -> None: torch.cuda.synchronize() before = int(torch.cuda.memory_allocated()) try: - rank.dp_rank_forward(requests()) + rank.forward(requests()) except TrainerRankMemoryError as error: message = str(error).lower() if "unable to find a feasible split" not in message: @@ -773,9 +793,14 @@ def pressured_cap() -> int: if len(partition) < 2: _fail("reverse-order arm expected a split plan") for indices in reversed(partition): - torch.stack( - [outputs[index].target_logprobs.float().sum() for index in indices] - ).sum().backward() + rank.backward( + torch.stack( + [ + outputs[index].target_logprobs.float().sum() + for index in indices + ] + ).sum() + ) rank.zero_grad() del info, outputs @@ -972,7 +997,7 @@ def conversion( # backward through an active LoRA slot, compared against the depth-one arm. # ``dp2-tp2-waves`` (4 ranks, DP2 x TP2) exercises the global wave planner's # collectives (world scope) composed with TP execution collectives (pair scope) -# through public ``forward_micro_batches``, including an empty DP slot. +# through public ``forward_batches``, including an empty DP slot. TP_GATES: dict[str, float] = { # bf16 kernels reorder reductions across packings (same metric/tolerance as @@ -1337,7 +1362,7 @@ def _write_rows(evidence: str | None, rows: list[dict[str, object]], name: str) def phase_tp2_public( evidence: str | None, repeat: int, *, tp: int, dump_dir: str | None ) -> None: - """DP1 x TP{tp} x CP1 public ``dp_rank_forward`` cell (Qwen3.5-4B full model). + """DP1 x TP{tp} x CP1 public ``forward`` cell (Qwen3.5-4B full model). Run at ``--tp 2`` (the gate) and at ``--tp 1`` (the control: the identical cell on one GPU). Structural gates run here; the numerics gates compare @@ -1407,9 +1432,9 @@ def run(arm: str) -> dict[str, Any]: start = torch.cuda.Event(enable_timing=True) end = torch.cuda.Event(enable_timing=True) start.record() - outputs = rank.dp_rank_forward(requests, checkpoint=slot) + outputs = rank.forward(requests, checkpoint=slot) loss = _output_loss(outputs) - loss.backward() + rank.backward(loss) end.record() torch.cuda.synchronize() telemetry = rank.last_forward_telemetry() @@ -1759,7 +1784,7 @@ def structured(label: str, profile: dict[str, float]) -> None: def phase_dp2_tp2_waves(evidence: str | None) -> None: - """DP2 x TP2 public ``forward_micro_batches`` gate (4 ranks, Qwen3.5-4B). + """DP2 x TP2 public ``forward_batches`` gate (4 ranks, Qwen3.5-4B). Arm A (branchy): six hierarchical GRPO groups as top-level items, the test-only memory cap sized so the stream needs at least two waves, forward @@ -1816,7 +1841,7 @@ def stream(arm: str, cap_bytes: int | None) -> dict[str, Any]: logprobs: dict[int, list[torch.Tensor]] = {} loss_total = 0.0 depths: list[int] = [] - for batch in rank.forward_micro_batches(items, checkpoint=slot): + for batch in rank.forward_batches(items, checkpoint=slot): seen.extend(int(index) for index in batch.indices) telemetry = rank.last_forward_telemetry() depths.append(int(telemetry["selected_max_depth"])) @@ -1840,7 +1865,7 @@ def stream(arm: str, cap_bytes: int | None) -> dict[str, Any]: if flat: loss = _output_loss(flat) loss_total += float(loss.detach().float().item()) - loss.backward() + rank.backward(loss) torch.cuda.synchronize() result = { "arm": arm, @@ -1955,12 +1980,12 @@ def stream(arm: str, cap_bytes: int | None) -> dict[str, Any]: rank.zero_grad() waves = 0 local_outputs = 0 - for batch in rank.forward_micro_batches([single], checkpoint=slot): + for batch in rank.forward_batches([single], checkpoint=slot): waves += 1 flat = [output for group in batch.outputs for output in group] local_outputs += len(flat) if flat: - _output_loss(flat).backward() + rank.backward(_output_loss(flat)) torch.cuda.synchronize() counts = _gather_objects((dp_rank, waves, local_outputs)) rows.append({"arm": "dp2-tp2-empty-slot", "per_rank": counts}) @@ -2625,9 +2650,9 @@ def run(label: str) -> dict[str, object]: start.record() failed = 0 try: - outputs = rank.dp_rank_forward(requests, checkpoint=slot) + outputs = rank.forward(requests, checkpoint=slot) loss = _output_loss(outputs) - loss.backward() + rank.backward(loss) except TrainerRankMemoryError as error: failed = 1 message = str(error) diff --git a/dev/trainer_rank_landing_acceptance_README.md b/dev/trainer_rank_landing_acceptance_README.md index 342f7a3a2..9fca2b264 100644 --- a/dev/trainer_rank_landing_acceptance_README.md +++ b/dev/trainer_rank_landing_acceptance_README.md @@ -18,8 +18,8 @@ every gate below now passes on the landed implementation. | `tests/unit/test_trainer_rank_split.py` | CPU | best-effort splitting contract: bounded ladder (failed rungs rejected by cheap bounds; planner runs only for the executing rung), cumulative live-graph admission, retained profile trusted only near its observed scale and max-merged once observed, caller-order reconstruction, honest refusal wording, minimum-wave splitting, deterministic partitions, one slot-ensure collective per call, independent slot-graph sentinels per subforward, `TrainerRankPartialExecutionError` on execution-time failure, `subforward_count` telemetry | pass | | `--phase split-conversion --pressure cap` | 1x H200 | sealed cell shape (Qwen3.5-4B, 4 layers, 4 inputs) under the test-only cap: unlimited runs unsplit; cap converts (>=2 subforwards) with output parity, combined and reverse-order backward; sub-request cap refuses before execution | see evidence log | | `tests/unit/test_trainer_rank_topology.py` | CPU | TP>1 runtimes construct; PP>1 and multi-chunk runtimes still refuse | pass | -| `--phase tp2-public --tp 2` (`dev/trainer_rank_landing_acceptance_tp2.sky.yaml`) + `--tp 1` control (1x H200) + `--phase tp-compare` (CPU) | 2x H200 k8s | Qwen3.5-4B full model, DP1×TP2×CP1, public `dp_rank_forward`, active LoRA: TP peers plan identical physical layouts; automatic shares deeper than depth-one (group lengths 9,199 vs 37,871 for 53,216 logical tokens; physical 9,200 / 37,872 after per-group TP padding); odd group lengths exercise SP padding; measured rows compile-free and plan-cache-stable; numerics gated relative to the TP1 control's cross-layout divergence with source/workload fingerprints (same-layout TP2-vs-TP1 ratios 1.06/1.05, cross-layout ratios 1.05/1.06, losses within 0.06%, unstructured differences). Also runs the CI check script at TP=2 (0.0 divergence) | pass: automatic 720 ms vs depth-one 1,921 ms (62.5% paired gain at TP2; 71.9% at TP1) | -| `--phase dp2-tp2-waves` (`dev/trainer_rank_landing_acceptance_dp2_tp2.sky.yaml`) | 4x H200 k8s | DP2×TP2 public `forward_micro_batches`: ≥2 waves, distinct DP payloads, identical wave shapes within each TP pair, every input returned once in order, forward+backward per wave, automatic vs depth-one parity, empty-DP-slot arm, no hang | see evidence log | +| `--phase tp2-public --tp 2` (`dev/trainer_rank_landing_acceptance_tp2.sky.yaml`) + `--tp 1` control (1x H200) + `--phase tp-compare` (CPU) | 2x H200 k8s | Qwen3.5-4B full model, DP1×TP2×CP1, public `forward`, active LoRA: TP peers plan identical physical layouts; automatic shares deeper than depth-one (group lengths 9,199 vs 37,871 for 53,216 logical tokens; physical 9,200 / 37,872 after per-group TP padding); odd group lengths exercise SP padding; measured rows compile-free and plan-cache-stable; numerics gated relative to the TP1 control's cross-layout divergence with source/workload fingerprints (same-layout TP2-vs-TP1 ratios 1.06/1.05, cross-layout ratios 1.05/1.06, losses within 0.06%, unstructured differences). Also runs the CI check script at TP=2 (0.0 divergence) | pass: automatic 720 ms vs depth-one 1,921 ms (62.5% paired gain at TP2; 71.9% at TP1) | +| `--phase dp2-tp2-waves` (`dev/trainer_rank_landing_acceptance_dp2_tp2.sky.yaml`) | 4x H200 k8s | DP2×TP2 public `forward_batches`: ≥2 waves, distinct DP payloads, identical wave shapes within each TP pair, every input returned once in order, forward+backward per wave, automatic vs depth-one parity, empty-DP-slot arm, no hang | see evidence log | | `--phase cost-calibrate` (`dev/trainer_rank_cost_calibration_{cp4,2gpu}.sky.yaml`, `dev/trainer_rank_cost_calibration_local.sh`) | 1x/2x/4x H200 | every mandatory candidate layout of a cell timed through the public API (forward+backward, active LoRA, compile-free, max-rank) plus the production selection; layout features and topology/model facts to JSONL; `--planner-ab` times every layout under the current and the legacy CP planner in alternating rounds on the same node (rows carry `planner_variant`; the fitter uses only current-planner rows). `--gdn-ab` brackets the GDN planner's chain decision (production, never chain, chain every legal segment) and `--gdn-legacy-ab` pairs the production GDN planner against the one before the 2026-09-09 recalibration, the same way. Cells: GRPO g8/g16/g4x4 (Qwen3.5-4B GDN and Qwen3-4B attention, 2-layer and full height), three heterogeneous controls, Ellavox groups, at TP1/TP2 × CP1/CP2/CP4 | 58 cells, 3,849 within-cell pairs | | `dev/trainer_rank_cost_fit.py` | CPU | paired within-cell deltas, non-negative least squares over the production term functions, regret-minimizing refinement; gates on whole held-out cells: pairwise ordering ≥90% on pairs separated >3%, median regret ≤2%, p95 ≤5%, none >10%, clear winners selected within 5%; `--selector-check` runs the shipped table through the real selector | final table (fit on 45 cells, evaluated on all 58; 3,849 pairs; 13 held-out cells): 98.1% pairwise, p95 regret 2.9%, max 4.2%, no clear misses; pre-registered odd-Ellavox holdout passes; the two Ellavox CP4 cells re-measured after the issue #840 harness fix are held out too (shipped-table regret 0% and 2.8%); profile narrowed to the measured envelope (H200-class SM 9.0, bf16, hidden 2,560, dense) and bound to the certificate by test; expected-cell manifest validated; TP2, CP2 and attention-model ablations pass (all-heterogeneous ablation: one CP4 cell at 9.5%); campaign-1 table on the 18 later cells within 4.2%; timed production selection on the later CP2/TP2 cells: median −0.2%, max 0.4%; hand-set score: 78.6% pairwise, max regret 67%. Re-certified 2026-09-08 under the recalibrated CP planner (issue #854; every CP > 1 cell re-measured with the paired A/B, `--planner-ab`): 58 cells, 98.7% pairwise, p95 1.9%, max 2.8%, held-out pass | | `dev/trainer_rank_cost_calibration_lattice.sky.yaml` (`MODEL`, `SHAPES` = tp,cp,ep[,etp], `CELLS`) | 4x H200 | one model's calibration cells over an explicit parallel-shape lattice, one torchrun per shape with EP/ETP pinned (builds HybridEP in setup); used for the Qwen3.5-35B-A3B class (EP1 at CP1/2/4/TP2, EP2 at CP2/CP4/TP2, EP4 at CP4) and the dense controls | 35B class: 73 cells / 3,658 pairs certified as `gdn-moe-h2048-h200-bf16` (97.1% pairwise, p95 regret 2.4%, max 4.9%, gates pass per shape); 15 lattice cells excluded with reasons (#848, #851) | @@ -40,7 +40,7 @@ deliberately uses a 2-layer model whose ~130 ms execution would make any fraction meaningless. The screen gates planning absolutely instead. Every measured sample uses fresh tokens, so these planning numbers are all cache-miss (cold) costs; steady-state identical-content calls are a content -hash plus dictionary hit, and `forward_micro_batches` additionally pre-plans +hash plus dictionary hit, and `forward_batches` additionally pre-plans the predicted next wave in the background during the caller's GPU time (measured benefit is marginal — about 1–2 ms/step on a 2-wave GPU benchmark — because sharing-aware width pricing already plans the accepted width; it is diff --git a/dev/trainer_rank_landing_acceptance_dp2_tp2.sky.yaml b/dev/trainer_rank_landing_acceptance_dp2_tp2.sky.yaml index bb84ae466..a5fec7f8d 100644 --- a/dev/trainer_rank_landing_acceptance_dp2_tp2.sky.yaml +++ b/dev/trainer_rank_landing_acceptance_dp2_tp2.sky.yaml @@ -1,4 +1,4 @@ -# TP support gate: DP2 x TP2 public forward_micro_batches cell (4x H200, Kubernetes). +# TP support gate: DP2 x TP2 public forward_batches cell (4x H200, Kubernetes). # # The global wave planner's world-scope collectives composed with TP execution # collectives (pair scope): at least two waves under the test-only memory cap, diff --git a/dev/trainer_rank_landing_acceptance_tp2.sky.yaml b/dev/trainer_rank_landing_acceptance_tp2.sky.yaml index 3f3de9eb3..6f28d3272 100644 --- a/dev/trainer_rank_landing_acceptance_tp2.sky.yaml +++ b/dev/trainer_rank_landing_acceptance_tp2.sky.yaml @@ -1,4 +1,4 @@ -# TP support gate: DP1 x TP2 x CP1 public dp_rank_forward cell (2x H200, Kubernetes). +# TP support gate: DP1 x TP2 x CP1 public forward cell (2x H200, Kubernetes). # # First public-API execution of the automatic planner under tensor # parallelism on the full Qwen3.5-4B: planner-selected prefix sharing -> GDN diff --git a/dev/trainer_rank_planner_design.md b/dev/trainer_rank_planner_design.md index 6363f3602..e03105163 100644 --- a/dev/trainer_rank_planner_design.md +++ b/dev/trainer_rank_planner_design.md @@ -31,7 +31,7 @@ document records the verified facts the acceptance suite pins). sharing-aware (a no-sharing token count accepts a width, and the planner's actual layouts are priced only when that bound would reject one); head chunking and memory margins are internal calibrated constants, not planner - decisions; `dp_rank_forward` plans once and raises + decisions; `forward` plans once and raises `TrainerRankMemoryError(predicted_peak_bytes, usable_limit_bytes, suggestion)` when the unsplit plan cannot be admitted (best-effort internal splitting is a follow-up PR); `TrainerRankRuntimeSupportError` at PP>1 @@ -632,7 +632,7 @@ memory-minimal (full-sharing) layout fits": full sharing minimizes packed tokens and its count is monotone in width by construction. Admission then executes the cost-optimal layout when it fits and the memory-minimal layout otherwise; the chosen mode is recorded per width so materialization builds -exactly the layouts that were priced. `dp_rank_forward` applies the same +exactly the layouts that were priced. `forward` applies the same fallback before refusing. Both bounds are cheap O(tokens) walks of the packing primitive (no-sharing and unlimited-depth sharing); planner pricing runs only inside the band where they disagree. @@ -658,7 +658,7 @@ keeping the memory-to-throughput crossover. ## Overlapped (speculative) next-wave planning -``forward_micro_batches`` pre-plans the predicted next wave (exactly the +``forward_batches`` pre-plans the predicted next wave (exactly the width the search will seed with — the largest width so far — over this DP rank's strided slice) on a single background thread while the generator is suspended at the yield — i.e. during the caller's forward/backward GPU time. @@ -681,7 +681,7 @@ if finding out is too expensive or fragile, refuse — worded as "unable to find a feasible split", never as a claim that none exists. Mechanism: -- `dp_rank_forward` (and the minimum wave of `forward_micro_batches`) plans +- `forward` (and the minimum wave of `forward_batches`) plans unsplit first (cost-optimal, then memory-minimal). If neither is admitted, a bounded, deterministic ladder tries 2, 4, ... subforwards (at most one request each), cutting the requests in prefix-local depth-first order into @@ -719,7 +719,7 @@ Mechanism: so a cold call that cannot fit unsplit refuses until a profile exists. Limitation: the observation is taken at forward return and says nothing about backward; TrainerRank cannot see the caller's backward peak for - `dp_rank_forward` (the micro-batch path folds the post-yield peak into + `forward` (the micro-batch path folds the post-yield peak into `bytes_per_token`, not into the retained fraction). - Collectives. Ensuring checkpoint slots is a world collective; the ladder's length depends on this rank's DP-local inputs, so slots are ensured exactly @@ -812,7 +812,7 @@ Gates (test-first; all failed on the refusing tree): - `tests/unit/test_trainer_rank_topology.py`: TP>1 constructs, PP>1 and multi-chunk runtimes still refuse. - `--phase tp2-public` (2× H200, Qwen3.5-4B full model, DP1×TP2×CP1, public - `dp_rank_forward`, active LoRA slot), plus the identical cell at `--tp 1` + `forward`, active LoRA slot), plus the identical cell at `--tp 1` as the control: both TP peers plan the same physical layout on every call; the automatic planner shares more deeply than depth-one on the hierarchical GRPO shape; odd packed lengths exercise sequence-parallel padding with @@ -831,7 +831,7 @@ Gates (test-first; all failed on the refusing tree): the body, no bias). Measured: same-layout TP2-vs-TP1 1.39% vs the 1.31% reference (ratio 1.06), cross-layout ratios 1.05 (outputs) and 1.06 (gradients), losses within 0.06%, flat per-request profile, tail = body. -- `--phase dp2-tp2-waves` (4× H200, DP2×TP2, public `forward_micro_batches`): +- `--phase dp2-tp2-waves` (4× H200, DP2×TP2, public `forward_batches`): at least two waves under the test-only cap, DP replicas with different payloads, identical wave shapes within each TP pair, every input returned exactly once in order, forward and backward per wave, automatic vs diff --git a/dev/trainer_rank_recompute_memory.py b/dev/trainer_rank_recompute_memory.py index ba9fbc76e..62f78a271 100644 --- a/dev/trainer_rank_recompute_memory.py +++ b/dev/trainer_rank_recompute_memory.py @@ -274,7 +274,7 @@ def record(module, inputs, output): ] assert len(terms) == len(requests) loss = torch.stack(terms).sum() - loss.backward() + rank.backward(loss) torch.cuda.synchronize() backward_peak = torch.cuda.max_memory_allocated() backward_seconds = time.monotonic() - started diff --git a/dev/trainer_v1_acceptance.py b/dev/trainer_v1_acceptance.py new file mode 100644 index 000000000..cc0604637 --- /dev/null +++ b/dev/trainer_v1_acceptance.py @@ -0,0 +1,509 @@ +"""Native model acceptance against a global analytical-loss cotangent oracle. + +Run ``oracle`` at DP=TP=CP=1, then ``zero`` on each target topology, passing the +oracle's .pt file as --reference. The oracle bypasses the callback tensor bridge +and differentiates native model outputs with hand-derived global cotangents. +""" + +from __future__ import annotations + +import argparse +import asyncio +from collections import defaultdict +from dataclasses import asdict +import json +import os +from pathlib import Path +import resource +import subprocess +import sys +import time +import traceback + +import torch +import torch.distributed as dist +from trainer_rank_diag import rank0_checked +from trainer_rank_support import load_random_checkpoints + +from art.trainer_rank import ForwardInput, TrainerRank + + +def _requests(checkpoint, offset, lengths, options=None): + leaves = [] + for index, length in enumerate(lengths): + tokens = (torch.arange(length) * 17 + offset + index * 103) % 30_000 + leaves.append( + ForwardInput( + input_tokens=tokens, + target_tokens=(tokens * 7 + 3) % 30_000, + checkpoint=checkpoint, + hidden_states=True, + **({"options": options} if options is not None else {}), + ) + ) + # Nested complete roots, odd lengths, and an unused differentiable output. + return [[leaves[0], *leaves[1:2]], leaves[2:]] if len(leaves) > 1 else leaves + + +def _leaves(tree): + if isinstance(tree, (list, tuple)): + return [leaf for item in tree for leaf in _leaves(item)] + return [tree] + + +def _loss_and_cotangents(first, second): + a = torch.cat([leaf.target_logprobs for leaf in _leaves(first)]) + b = torch.cat([leaf.target_logprobs for leaf in _leaves(second)]) + difference = a.mean() - 0.4 * b.mean() + loss = difference.square() + 0.03 * a.square().mean() + 0.02 * b.square().mean() + da = (2 * difference + 0.06 * a.detach()) / a.numel() + db = (-0.8 * difference + 0.04 * b.detach()) / b.numel() + tensors = [leaf.target_logprobs for leaf in _leaves(first) + _leaves(second)] + gradients = list( + da.detach().split([x.target_logprobs.numel() for x in _leaves(first)]) + ) + gradients += list( + db.detach().split([x.target_logprobs.numel() for x in _leaves(second)]) + ) + return loss, tensors, gradients + + +def _cached_backward(rank, loss): + with rank._gradient_transaction(): + packets = rank._forward_cotangent_collector().backward(loss) + rank._forward_graph_cache().backward_many( + [(packet.handle, packet.gradients) for packet in packets] + ) + + +def _canonical_gradients(rank, checkpoint): + """Reduce once, then gather canonical LoRA shards without altering weights.""" + from art.megatron.lora import LoRA + from art.megatron.weights.lora_publish import _merge_manifest_entries + + parameters = rank._checkpoint_slots[checkpoint].params + reduced = rank._reduce_dynamic_grads(parameters, scale_grads=1.0) + by_id = { + id(parameter): gradient + for parameter, gradient in zip(parameters, reduced, strict=True) + } + local = {} + for chunk in rank.runtime.model: + for module in chunk.modules(): + if not isinstance(module, LoRA): + continue + for key, parameter, expert in module._export_items( + rank._slot_ref(checkpoint) + ): + value = by_id[id(parameter)] + value = value if expert is None else value[expert] + local[key] = ( + module._manifest_for_param(parameter), + value.T.float().cpu(), + ) + if dist.get_rank() == 0: + for name, custom in rank._checkpoint_slots[checkpoint].custom.items(): + for key, parameter in custom.value.named_parameters(): + local[f"custom.{name}.{key}"] = ( + {"sharded": False, "shard_world_size": 1}, + by_id[id(parameter)].float().cpu(), + ) + gathered = [None] * dist.get_world_size() + dist.all_gather_object(gathered, local) + if dist.get_rank() != 0: + return None + groups = defaultdict(list) + for shard in gathered: + for key, entry in shard.items(): + groups[key].append(entry) + return { + key: _merge_manifest_entries(key, entries) for key, entries in groups.items() + } + + +def _compare(actual, reference): + if actual["gradients"].keys() != reference["gradients"].keys(): + raise AssertionError("Canonical gradient keys differ") + torch.testing.assert_close( + torch.tensor(actual["loss"]), + torch.tensor(reference["loss"]), + atol=0.03, + rtol=0.003, + ) + for output, expected in zip(actual["outputs"], reference["outputs"], strict=True): + torch.testing.assert_close(output, expected, atol=0.03, rtol=0.003) + rows, numerator, denominator = [], 0.0, 0.0 + for key, value in actual["gradients"].items(): + expected = reference["gradients"][key] + error = (value.double() - expected.double()).square().sum().item() + scale = expected.double().square().sum().item() + relative = (error / max(scale, 1e-24)) ** 0.5 + rows.append({"key": key, "relative_l2": relative}) + numerator += error + denominator += scale + if scale > 1e-12 and relative > 0.06: + raise AssertionError( + f"{key}: gradient relative L2 {relative:.6f} exceeds 0.06" + ) + relative = (numerator / max(denominator, 1e-24)) ** 0.5 + if denominator <= 1e-12 or relative > 0.03: + raise AssertionError( + f"global gradient relative L2 {relative:.6f}, reference norm² {denominator}" + ) + return {"gradient_relative_l2": relative, "per_parameter": rows} + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument( + "mode", choices=["oracle", "control", "zero", "rank", "retained"] + ) + parser.add_argument("--model", default="Qwen/Qwen3-0.6B") + parser.add_argument("--layers", type=int, default=1) + parser.add_argument("--output", type=Path, required=True) + parser.add_argument("--reference", type=Path) + parser.add_argument("--retention", choices=["gpu", "cpu", "replay"], default="gpu") + parser.add_argument("--output-device", choices=["model", "cpu"], default="model") + parser.add_argument( + "--batched", + action="store_true", + help="Physical control uses replicated-root batching", + ) + parser.add_argument( + "--head", + action="store_true", + help="Include a registered linear-head analytical oracle", + ) + parser.add_argument("--cuda-values", action="store_true") + parser.add_argument("--reverse-devices", action="store_true") + parser.add_argument("--split-head-backward", action="store_true") + args = parser.parse_args() + if args.split_head_backward and ( + not args.head or args.mode not in ("zero", "rank") + ): + parser.error("Split head backward requires --head and zero/rank mode") + for axis in ("TENSOR_MODEL", "CONTEXT", "PIPELINE_MODEL"): + os.environ.setdefault(f"ART_MEGATRON_{axis}_PARALLEL_SIZE", "1") + os.environ["ART_TRAINER_RANK_TEST_HOOKS"] = "1" + os.environ["ART_TRAINER_RANK_TEST_ANCHOR"] = "no_sharing" + device = int(os.environ["LOCAL_RANK"]) + if args.reverse_devices: + device = int(os.environ["LOCAL_WORLD_SIZE"]) - 1 - device + torch.cuda.set_device(device) + dist.init_process_group("nccl") + try: + from megatron.core import parallel_state as ps + + from art.megatron.train import build_training_runtime + + if args.mode in ("oracle", "retained") and dist.get_world_size() != 1: + raise ValueError( + "The mathematical global-loss oracle requires one physical rank" + ) + torch.manual_seed(90217) + runtime = build_training_runtime( + model_identifier=args.model, + provider_configure=(lambda p: setattr(p, "num_layers", args.layers)) + if args.layers + else None, + print_env=dist.get_rank() == 0, + ) + for chunk in runtime.model: + chunk.eval() + physical = TrainerRank(runtime) + if args.mode == "control" and ps.get_data_parallel_world_size() != 1: + raise ValueError("A physical topology control requires DP=1") + (checkpoint,) = load_random_checkpoints( + runtime, physical, 1, base_model=args.model, lora_rank=2 + ) + options = None + if args.mode == "retained": + from art.trainer_rank import ForwardOptions + + options = ForwardOptions( + backward_state=args.retention, + output_device=args.output_device, + stale_gradient_corrections=(), + ) + first = _requests(checkpoint, 29, [17, 23, 31], options) + # One complete root forces an empty DP partition when DP > 1. + second = _requests(checkpoint, 113, [19], options) + callbacks = 0 + cache_before_backward = None + hidden_size = runtime.provider.hidden_size + + def head_factory(): + head = torch.nn.Linear( + hidden_size, + 1, + bias=False, + device=physical.device if args.cuda_values else None, + ) + with torch.no_grad(): + head.weight.copy_( + torch.linspace(-1, 1, hidden_size).reshape(1, -1) / hidden_size**0.5 + ) + return head + + def callback(rank): + nonlocal callbacks, cache_before_backward + callbacks += 1 + if args.cuda_values: + for request in _leaves(first) + _leaves(second): + request.input_tokens = request.input_tokens.to(rank.device) + assert request.target_tokens is not None + request.target_tokens = request.target_tokens.to(rank.device) + head = ( + rank.module("validation_head", head_factory, checkpoint=checkpoint) + if args.head + else None + ) + # The original main branch names the physical operation differently; + # keeping it usable as an oracle permits before/after comparisons. + forward = getattr(rank, "forward", None) or rank.dp_rank_forward + if args.batched: + batches = ( + getattr(rank, "forward_batches", None) or rank.forward_micro_batches + ) + + def forward(roots): + return [ + root + for batch in batches(roots, yield_empty=True) + for root in batch.outputs + ] + + one, two = forward(first), forward(second) + loss, outputs, cotangents = _loss_and_cotangents(one, two) + model_loss = loss + split_losses = None + report_outputs = list(outputs) + if head is not None: + hidden = [leaf.hidden_states for leaf in _leaves(one) + _leaves(two)] + scale = 0.01 / len(hidden) + if args.mode in ("oracle", "control"): + # Explicit dH and dW avoid the live-head/snapshot autograd + # machinery in the unchanged-source mathematical reference. + weight = ( + rank._checkpoint_slots[checkpoint] + .custom["validation_head"] + .value.weight + ) + scores = [value.float() @ weight.detach().T for value in hidden] + derivatives = [ + 2 * scale * score.detach() / score.numel() for score in scores + ] + outputs.extend(hidden) + cotangents.extend( + [ + (gradient @ weight.detach()).to(value.dtype) + for value, gradient in zip(hidden, derivatives, strict=True) + ] + ) + outputs.append(weight) + cotangents.append( + sum( + gradient.T @ value.detach().float() + for value, gradient in zip(hidden, derivatives, strict=True) + ) + ) + else: + scores = [head(value.float()) for value in hidden] + loss = loss + scale * sum(score.square().mean() for score in scores) + if args.split_head_backward: + # Distinct head captures publish twice to the same targets; + # their sum retains the independent one-copy dH/dW oracle. + repeated = [head(value.float()) for value in hidden] + split_losses = ( + loss / 2, + (model_loss + scale * sum(x.square().mean() for x in repeated)) + / 2, + ) + report_outputs.extend(scores) + if args.mode in ("oracle", "control"): + # No callback packet/autograd bridge or loss autograd contributes + # to the reference gradients. + torch.autograd.backward(outputs, cotangents) + elif args.mode == "retained": + from art.trainer_rank import AdamParams + + fresh = forward(first) + update_loss = -torch.cat( + [leaf.target_logprobs for leaf in _leaves(fresh)] + ).mean() + _cached_backward(rank, update_loss) + metrics = rank.optim_step( + params=AdamParams(learning_rate=0.01, grad_clip_norm=0), + checkpoints=[checkpoint], + ) + if metrics["update_successful"] != 1: + raise AssertionError( + f"Intervening optimizer update failed: {metrics}" + ) + # Replay must use captured tokens and historical adapter tensors. + for request in _leaves(first) + _leaves(second): + request.input_tokens.fill_(999) + cache = rank._forward_graph_cache() + cache_before_backward = [ + asdict(cache.state(handle)) for handle in cache.handles() + ] + rng = torch.cuda.get_rng_state() + _cached_backward(rank, loss) + if not torch.equal(rng, torch.cuda.get_rng_state()): + raise AssertionError( + "Retained/replayed backward changed ambient CUDA RNG" + ) + if cache.handles(): + raise AssertionError("Consumed model graphs remained resident") + elif args.mode == "rank": + # Each logical DP rank owns the same complete local workload. + # Sum its scaled loss/gradients to recover the one-copy oracle; + # TP/CP replicas must not multiply either reduction. + size = ps.get_data_parallel_world_size() + if split_losses is None: + rank.backward(loss / size) + else: + rank.backward(split_losses[0] / size, retain_graph=True) + rank.backward(split_losses[1] / size) + loss = loss.detach().clone() / size + rank.reduce(loss) + else: + if split_losses is None: + rank.backward(loss) + else: + rank.backward(split_losses[0], retain_graph=True) + rank.backward(split_losses[1]) + return { + "loss": loss.item(), + "outputs": [value.detach().float().cpu() for value in report_outputs], + } + + torch.cuda.synchronize() + torch.cuda.reset_peak_memory_stats() + started = time.perf_counter() + if args.mode in ("oracle", "control", "retained"): + result = callback(physical) + else: + from art.trainer_rank import run_rank_callback + + callback_error = None + try: + wrapped = asyncio.run( + run_rank_callback( + physical, + callback, + mode="rank" if args.mode == "rank" else "zero", + ) + ) + except BaseException: + callback_error = traceback.format_exc() + print(callback_error, file=sys.stderr, flush=True) + errors = [None] * dist.get_world_size() + dist.all_gather_object(errors, callback_error) + if any(errors): + raise RuntimeError( + "Callback oracle failed before result collection:\n" + + "\n".join(error for error in errors if error is not None) + ) + result = wrapped.value + torch.cuda.synchronize() + elapsed = time.perf_counter() - started + counts = [None] * dist.get_world_size() + dist.all_gather_object(counts, callbacks) + devices = [None] * dist.get_world_size() + dist.all_gather_object(devices, str(physical.device)) + if args.reverse_devices: + local_size = int(os.environ["LOCAL_WORLD_SIZE"]) + expected_devices = [ + f"cuda:{local_size - 1 - rank % local_size}" + for rank in range(dist.get_world_size()) + ] + if devices != expected_devices: + raise AssertionError(f"Device mapping {devices} != {expected_devices}") + expected_counts = ( + [1] * dist.get_world_size() + if args.mode == "control" + else [1] + [0] * (dist.get_world_size() - 1) + ) + if args.mode == "rank": + is_leader = int( + ps.get_tensor_model_parallel_rank() == 0 + and ps.get_context_parallel_rank() == 0 + and ps.get_pipeline_model_parallel_rank() == 0 + ) + dist.all_gather_object(expected_counts, is_leader) + if counts != expected_counts: + raise AssertionError(f"User callback counts are {counts}") + gradients = _canonical_gradients(physical, checkpoint) + measurements = [None] * dist.get_world_size() + dist.all_gather_object( + measurements, + { + "elapsed_seconds": elapsed, + "peak_gpu_allocated_bytes": torch.cuda.max_memory_allocated(), + "peak_gpu_reserved_bytes": torch.cuda.max_memory_reserved(), + "process_peak_rss_bytes": resource.getrusage( + resource.RUSAGE_SELF + ).ru_maxrss + * 1024, + }, + ) + + def finish(): + result["gradients"] = gradients + args.output.parent.mkdir(parents=True, exist_ok=True) + torch.save(result, args.output) + comparison = None + if args.reference: + comparison = _compare( + result, torch.load(args.reference, weights_only=True) + ) + metadata = { + "mode": args.mode, + "model": args.model, + "layers": args.layers, + "callback_counts": counts, + "physical_devices": devices, + "loss": result["loss"], + "topology": { + "dp": ps.get_data_parallel_world_size(), + "tp": ps.get_tensor_model_parallel_world_size(), + "cp": ps.get_context_parallel_world_size(), + }, + "device": torch.cuda.get_device_name(), + "torch": torch.__version__, + "source_commit": subprocess.check_output( + ["git", "rev-parse", "HEAD"], + cwd=Path(sys.modules["art.trainer_rank"].__file__) + .resolve() + .parents[3], + text=True, + ).strip(), + "harness_commit": subprocess.check_output( + ["git", "rev-parse", "HEAD"], + cwd=Path(__file__).resolve().parent.parent, + text=True, + ).strip(), + "measurements": measurements, + "comparison": comparison, + "retention": args.retention if args.mode == "retained" else None, + "output_device": args.output_device, + "cache_before_backward": cache_before_backward, + "batched_control": args.batched, + "registered_head": args.head, + "cuda_values": args.cuda_values, + "reverse_devices": args.reverse_devices, + "split_head_backward": args.split_head_backward, + } + args.output.with_suffix(".json").write_text( + json.dumps(metadata, indent=2) + "\n" + ) + print(json.dumps(metadata), flush=True) + + rank0_checked("trainer v1 global-loss acceptance", finish) + finally: + dist.destroy_process_group() + + +if __name__ == "__main__": + main() diff --git a/dev/trainer_v1_benchmark.py b/dev/trainer_v1_benchmark.py new file mode 100644 index 000000000..f716d45d2 --- /dev/null +++ b/dev/trainer_v1_benchmark.py @@ -0,0 +1,330 @@ +"""Fixed-workload native throughput and learning canary, including delayed graphs. + +Run baseline on unchanged source, then gpu/cpu/replay with its tensor artifact as +--reference. Delayed cpu/replay instead use delayed gpu as their schedule oracle. +Use identical model, seed, lengths and compiler settings for all compared arms. +""" + +from __future__ import annotations + +import argparse +from dataclasses import asdict +import gc +import json +import math +import os +from pathlib import Path +import resource +import statistics +import subprocess +import sys +import time +import weakref + +from dotenv import load_dotenv +import torch +import torch.distributed as dist +from trainer_rank_support import load_random_checkpoints +from trainer_v1_acceptance import _cached_backward, _leaves + +from art.trainer_rank import AdamParams, ForwardInput, TrainerRank + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("mode", choices=["baseline", "gpu", "cpu", "replay"]) + parser.add_argument("--model", default="Qwen/Qwen3-0.6B") + parser.add_argument("--layers", type=int, default=0, help="0 keeps full model") + parser.add_argument("--tokens", type=int, default=1024) + parser.add_argument("--leaves", type=int, default=4) + parser.add_argument("--rounds", type=int, default=10) + parser.add_argument("--warmups", type=int, default=2) + parser.add_argument("--delay-ms", type=float, default=0) + parser.add_argument("--learning-rate", type=float, default=1e-4) + parser.add_argument("--explicit-cotangents", action="store_true") + parser.add_argument("--deterministic", action="store_true") + parser.add_argument("--no-update", action="store_true") + parser.add_argument("--delayed", action="store_true") + parser.add_argument("--output-device", choices=["model", "cpu"], default="model") + parser.add_argument("--output", type=Path, required=True) + parser.add_argument("--reference", type=Path) + args = parser.parse_args() + if args.delayed and args.mode == "baseline": + parser.error("Delayed graphs require v1; use gpu as the delayed oracle") + load_dotenv(".env") + if args.deterministic: + os.environ["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8" + torch.use_deterministic_algorithms(True) + import art.megatron.flex_attn.compiled as flex + + flex._FORCED_FLEX_BACKEND = "TRITON" + flex._FORCED_FLEX_KERNEL_OPTIONS = {"BACKEND": "TRITON"} + flex.dense_compiled_flex_attention = flex.triton_dense_compiled_flex_attention + flex.sparse_compiled_flex_attention = flex.triton_sparse_compiled_flex_attention + for axis in ("TENSOR_MODEL", "CONTEXT", "DATA", "PIPELINE_MODEL"): + os.environ[f"ART_MEGATRON_{axis}_PARALLEL_SIZE"] = "1" + os.environ["ART_TRAINER_RANK_TEST_HOOKS"] = "1" + os.environ["ART_TRAINER_RANK_TEST_ANCHOR"] = "no_sharing" + torch.cuda.set_device(int(os.environ["LOCAL_RANK"])) + dist.init_process_group("nccl") + try: + if dist.get_world_size() != 1: + raise ValueError("This paired benchmark requires one physical rank") + from art.megatron.train import build_training_runtime + + torch.manual_seed(90217) + runtime = build_training_runtime( + model_identifier=args.model, + provider_configure=(lambda p: setattr(p, "num_layers", args.layers)) + if args.layers + else None, + print_env=False, + ) + for chunk in runtime.model: + chunk.eval() + rank = TrainerRank(runtime) + (checkpoint,) = load_random_checkpoints( + runtime, rank, 1, base_model=args.model, lora_rank=2 + ) + forward = getattr(rank, "forward", None) or rank.dp_rank_forward + options = None + if args.mode != "baseline": + from art.trainer_rank import ForwardOptions + + options = ForwardOptions( + backward_state=args.mode, + output_device=args.output_device, + stale_gradient_corrections=(), + ) + + def inputs(offset=0, no_grad=False): + leaves = [] + for index in range(args.leaves): + tokens = (torch.arange(args.tokens) * 17 + index * 103 + offset) % 30000 + leaves.append( + ForwardInput( + input_tokens=tokens, + target_tokens=(tokens * 7 + 3) % 30000, + hidden_states=True, + checkpoint=checkpoint, + no_grad=no_grad, + **({"options": options} if options is not None else {}), + ) + ) + return [leaves[:2], leaves[2:]] + + boundary_outputs = {} + + def loss(outputs): + tensors = [value.target_logprobs for value in _leaves(outputs)] + value = -torch.cat(tensors).mean() + if args.explicit_cotangents and value.requires_grad: + boundary_outputs[id(value)] = tensors + return value + + def backward(value): + if args.mode == "baseline": + if args.explicit_cotangents: + tensors = boundary_outputs.pop(id(value)) + cotangents = torch.autograd.grad(value, tensors, retain_graph=True) + torch.autograd.backward(tensors, cotangents) + else: + value.backward() + else: + _cached_backward(rank, value) + + parameters = rank._checkpoint_slots[checkpoint].params + counters = {"physical_forwards": 0} + record_refs, version_refs = [], [] + + native_forward = rank._forward_packed + + def count_forward(*call_args, **kwargs): + counters["physical_forwards"] += 1 + return native_forward(*call_args, **kwargs) + + rank._forward_packed = count_forward + + def cache_metrics(): + if args.mode == "baseline": + return {"states": [], "historical_lora_bytes": 0} + cache = rank._forward_graph_cache() + record_refs.extend( + weakref.ref(record) for record in cache._records.values() + ) + storages = {} + for version in rank._version_state().lora.values(): + version_refs.append(weakref.ref(version)) + for slot in version.slots.values(): + for parameter in slot.parameters(): + storage = parameter.untyped_storage() + storages[(parameter.device, storage.data_ptr())] = ( + storage.nbytes() + ) + return { + "states": [asdict(cache.state(handle)) for handle in cache.handles()], + "historical_lora_bytes": sum(storages.values()), + } + + optimizer = AdamParams(learning_rate=args.learning_rate, grad_clip_norm=1) + first_step_gradients = None + optimizer_metrics = [] + + def step(value): + nonlocal first_step_gradients + backward(value) + if first_step_gradients is None: + first_step_gradients = [ + None if p.grad is None else p.grad.detach().float().cpu() + for p in parameters + ] + if args.no_update: + rank.zero_grad() + return + metrics = rank.optim_step(params=optimizer, checkpoints=[checkpoint]) + optimizer_metrics.append(metrics) + if metrics["update_successful"] != 1: + raise AssertionError(f"Failed optimizer update: {metrics}") + + # Warm every selected path without altering the learning initial state. + for _ in range(args.warmups): + value = loss(forward(inputs())) + backward(value) + rank.zero_grad() + with torch.no_grad(): + initial_eval = loss(forward(inputs(no_grad=True))).item() + rows, losses = [], [] + for iteration in range(args.rounds): + rank.zero_grad() + torch.cuda.synchronize() + torch.cuda.reset_peak_memory_stats() + counters["physical_forwards"] = 0 + started = time.perf_counter() + old = loss(forward(inputs())) + torch.cuda.synchronize() + forward_seconds = time.perf_counter() - started + diagnostic_started = time.perf_counter() + capture = cache_metrics() + losses.append(old.detach().item()) + diagnostic_seconds = time.perf_counter() - diagnostic_started + started = time.perf_counter() + if args.delayed: + fresh = loss(forward(inputs(offset=19))) + step(fresh) + if args.delay_ms: + time.sleep(args.delay_ms / 1000) + step(old) + torch.cuda.synchronize() + elapsed = forward_seconds + time.perf_counter() - started + tokens = args.tokens * args.leaves * (2 if args.delayed else 1) + row = { + "iteration": iteration, + "loss": losses[-1], + "elapsed_seconds": elapsed, + "excluded_diagnostic_seconds": diagnostic_seconds, + "logical_tokens": tokens, + "tokens_per_second": tokens / elapsed, + "physical_forwards": counters["physical_forwards"], + "active_graphs_after_backward": len( + rank._forward_graph_cache().handles() + ) + if args.mode != "baseline" + else 0, + "gpu_peak_allocated_bytes": torch.cuda.max_memory_allocated(), + "gpu_peak_reserved_bytes": torch.cuda.max_memory_reserved(), + "gpu_resident_bytes": torch.cuda.memory_allocated(), + "process_peak_rss_bytes": resource.getrusage( + resource.RUSAGE_SELF + ).ru_maxrss + * 1024, + **capture, + } + rows.append(row) + print("BENCHMARK=" + json.dumps(row), flush=True) + with torch.no_grad(): + final_eval = loss(forward(inputs(no_grad=True))).item() + rank._forward_packed = native_forward + del old, value + if args.delayed: + del fresh + gc.collect() + after_gc = cache_metrics() + after_gc["gpu_resident_bytes"] = torch.cuda.memory_allocated() + after_gc["live_captured_records"] = sum( + ref() is not None for ref in record_refs + ) + after_gc["live_captured_versions"] = sum( + ref() is not None for ref in version_refs + ) + artifact = { + "losses": losses, + "weights": [p.detach().float().cpu() for p in parameters], + "first_step_gradients": first_step_gradients, + } + args.output.parent.mkdir(parents=True, exist_ok=True) + torch.save(artifact, args.output) + comparison = None + if args.reference: + reference = torch.load(args.reference, weights_only=True) + numerator = denominator = 0.0 + for value, expected in zip( + artifact["weights"], reference["weights"], strict=True + ): + numerator += (value.double() - expected.double()).square().sum().item() + denominator += expected.double().square().sum().item() + comparison = { + "weight_relative_l2": math.sqrt(numerator / max(denominator, 1e-24)) + } + metadata = { + "arguments": { + k: str(v) if isinstance(v, Path) else v for k, v in vars(args).items() + }, + "source_commit": subprocess.check_output( + ["git", "rev-parse", "HEAD"], + cwd=Path(sys.modules["art.trainer_rank"].__file__).resolve().parents[3], + text=True, + ).strip(), + "harness_commit": subprocess.check_output( + ["git", "rev-parse", "HEAD"], + cwd=Path(__file__).resolve().parent.parent, + text=True, + ).strip(), + "torch": torch.__version__, + "device": torch.cuda.get_device_name(), + "layers": runtime.provider.num_layers, + "dtype": str(next(runtime.model[0].parameters()).dtype), + "initial_fixed_objective": initial_eval, + "final_fixed_objective": final_eval, + "median_tokens_per_second": statistics.median( + row["tokens_per_second"] for row in rows + ), + "comparison": comparison, + "optimizer_metrics": optimizer_metrics, + "after_gc": after_gc, + "rows": rows, + } + args.output.with_suffix(".json").write_text( + json.dumps(metadata, indent=2) + "\n" + ) + print(json.dumps(metadata), flush=True) + if args.reference: + torch.testing.assert_close( + torch.tensor(losses), + torch.tensor(reference["losses"]), + atol=0.03, + rtol=0.003, + ) + if comparison["weight_relative_l2"] > 0.005: + raise AssertionError(f"Learning trajectory changed: {comparison}") + if not math.isfinite(final_eval) or ( + not args.no_update and final_eval >= initial_eval + ): + raise AssertionError( + f"Fixed objective did not improve: {initial_eval} -> {final_eval}" + ) + finally: + dist.destroy_process_group() + + +if __name__ == "__main__": + main() diff --git a/dev/trainer_v1_checkpoint.py b/dev/trainer_v1_checkpoint.py new file mode 100644 index 000000000..528d6239b --- /dev/null +++ b/dev/trainer_v1_checkpoint.py @@ -0,0 +1,42 @@ +"""Create a native checkpoint fixture for the remote public API canary.""" + +import argparse +import os +from pathlib import Path + +from dotenv import load_dotenv +import torch +import torch.distributed as dist + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("output", type=Path) + parser.add_argument("--model", default="Qwen/Qwen3-0.6B") + args = parser.parse_args() + load_dotenv(".env") + for axis in ("TENSOR_MODEL", "CONTEXT", "DATA", "PIPELINE_MODEL"): + os.environ[f"ART_MEGATRON_{axis}_PARALLEL_SIZE"] = "1" + torch.cuda.set_device(int(os.environ["LOCAL_RANK"])) + dist.init_process_group("nccl") + try: + from trainer_rank_support import load_random_checkpoints + + from art.megatron.train import build_training_runtime + from art.trainer_rank import TrainerRank, validate_checkpoint + + torch.manual_seed(90217) + runtime = build_training_runtime(model_identifier=args.model, print_env=False) + rank = TrainerRank(runtime) + (checkpoint,) = load_random_checkpoints( + runtime, rank, 1, base_model=args.model, lora_rank=2 + ) + rank.save_checkpoint(str(args.output), checkpoint) + assert validate_checkpoint(args.output) is not None + print(f"native checkpoint saved: {args.output}", flush=True) + finally: + dist.destroy_process_group() + + +if __name__ == "__main__": + main() diff --git a/dev/trainer_v1_facade_benchmark.py b/dev/trainer_v1_facade_benchmark.py new file mode 100644 index 000000000..282c39192 --- /dev/null +++ b/dev/trainer_v1_facade_benchmark.py @@ -0,0 +1,235 @@ +"""Paired raw/logical facade forward-backward timings on identical frozen weights.""" + +import argparse +import asyncio +from collections import defaultdict +import cProfile +from functools import wraps +import json +import os +from pathlib import Path +import pstats +import statistics +import time + +from dotenv import load_dotenv +import torch +import torch.distributed as dist +from trainer_rank_support import load_random_checkpoints + +from art.trainer_rank import ForwardInput, ForwardOptions, TrainerRank, _commands + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--layers", type=int, default=28) + parser.add_argument("--tokens", type=int, default=1024) + parser.add_argument("--rounds", type=int, default=6) + parser.add_argument("--output", type=Path, required=True) + parser.add_argument("--profile", action="store_true") + args = parser.parse_args() + load_dotenv(".env") + torch.cuda.set_device(int(os.environ["LOCAL_RANK"])) + dist.init_process_group("nccl") + try: + world = dist.get_world_size() + for axis in ("TENSOR_MODEL", "CONTEXT", "DATA", "PIPELINE_MODEL"): + os.environ[f"ART_MEGATRON_{axis}_PARALLEL_SIZE"] = str( + world if axis == "DATA" else 1 + ) + os.environ["ART_TRAINER_RANK_TEST_HOOKS"] = "1" + os.environ["ART_TRAINER_RANK_TEST_ANCHOR"] = "no_sharing" + from art.megatron.train import build_training_runtime + + torch.manual_seed(90217) + runtime = build_training_runtime( + model_identifier="Qwen/Qwen3-0.6B", + provider_configure=lambda p: setattr(p, "num_layers", args.layers), + print_env=False, + ) + for chunk in runtime.model: + chunk.eval() + physical = TrainerRank(runtime) + (checkpoint,) = load_random_checkpoints( + runtime, physical, 1, base_model="Qwen/Qwen3-0.6B", lora_rank=2 + ) + options = ForwardOptions( + backward_state="gpu", output_device="model", stale_gradient_corrections=() + ) + metrics = defaultdict(float) + profiles = defaultdict(cProfile.Profile) + for cls, name in ( + (_commands._Executor, "_packet"), + (_commands._Executor, "_gather_outputs"), + (_commands._RankView, "_place_outputs"), + (_commands._RankView, "_attach"), + (type(physical), "_capture_forward_options"), + (type(physical), "_plan_admissible_forward"), + (type(physical), "_execute_graph_group"), + (type(physical), "_run_flat_plan_with_memory_tracking"), + ): + original = getattr(cls, name) + + @wraps(original) + def measured(*a, _original=original, _name=name, **kw): + start = time.perf_counter() + try: + return _original(*a, **kw) + finally: + metrics[_name] += time.perf_counter() - start + + setattr(cls, name, measured) + + rows, references = [], {} + for kind in ("logprobs", "hidden"): + requests = [] + for index in range(4): + tokens = (torch.arange(args.tokens) * 17 + index * 103) % 30000 + requests.append( + ForwardInput( + input_tokens=tokens, + target_tokens=(tokens * 7 + 3) % 30000 + if kind == "logprobs" + else None, + hidden_states=kind == "hidden", + checkpoint=checkpoint, + options=options, + ) + ) + inputs = [requests[:2], requests[2:]] + + local_inputs = inputs if world == 1 else [inputs[dist.get_rank()]] + + def run(view, mode, repetitions): + for repetition in repetitions: + view.zero_grad() + metrics.clear() + torch.cuda.synchronize() + started = time.perf_counter() + profile = ( + profiles[kind, mode] + if args.profile and repetition >= 2 + else None + ) + if profile is not None: + profile.enable() + try: + outputs = view.forward( + inputs if mode == "zero" else local_inputs + ) + finally: + if profile is not None: + profile.disable() + torch.cuda.synchronize() + forward_seconds = time.perf_counter() - started + values = [ + output.target_logprobs + if kind == "logprobs" + else output.hidden_states + for root in outputs + for output in root + ] + loss = sum( + value.float().square().mean() + if kind == "hidden" + else value.float().mean() + for value in values + ) + started = time.perf_counter() + view.backward(loss) + torch.cuda.synchronize() + backward_seconds = time.perf_counter() - started + row = dict( + physical_rank=dist.get_rank(), + mode=mode, + output=kind, + iteration=repetition - 2, + forward_seconds=forward_seconds, + backward_seconds=backward_seconds, + total_seconds=forward_seconds + backward_seconds, + output_bytes=sum(v.numel() * v.element_size() for v in values), + loss=loss.item(), + stages=dict(metrics), + telemetry=physical.last_forward_telemetry(), + ) + if repetition == 1: + gradient = torch.cat( + [ + p.grad.detach().float().reshape(-1) + for p in physical._checkpoint_slots[checkpoint].params + if p.grad is not None + ] + ) + if mode == "native": + references[kind] = gradient.clone() + row["gradient_relative_l2"] = ( + (gradient - references[kind]).norm() + / references[kind].norm().clamp_min(1e-20) + ).item() + if repetition >= 1: + rows.append(row) + print("FACADE=" + json.dumps(row), flush=True) + del outputs, values, loss + if physical._forward_graph_cache().handles(): + raise AssertionError("Unconsumed physical graph") + + def measure(mode, repetitions): + if mode == "native": + run(physical, mode, repetitions) + else: + asyncio.run( + _commands.run_rank_callback( + physical, + lambda view: run(view, mode, repetitions), + mode=mode, + ) + ) + dist.barrier() + + modes = ("native", "rank", "zero") + for mode in modes: + measure(mode, range(2)) + for repetition in range(args.rounds): + offset = repetition % len(modes) + for mode in modes[offset:] + modes[:offset]: + measure(mode, (repetition + 2,)) + all_rows = [None] * world + dist.all_gather_object(all_rows, rows) + if dist.get_rank() != 0: + return + rows = [row for peer in all_rows for row in peer] + summary = { + f"{kind}/{mode}": { + key: statistics.median( + max( + row[key] + for row in rows + if row["output"] == kind + and row["mode"] == mode + and row["iteration"] == iteration + ) + for iteration in range(args.rounds) + ) + for key in ("forward_seconds", "backward_seconds", "total_seconds") + } + for kind in ("logprobs", "hidden") + for mode in ("native", "rank", "zero") + } + args.output.parent.mkdir(parents=True, exist_ok=True) + args.output.write_text( + json.dumps(dict(dp=world, summary=summary, rows=rows), indent=2) + ) + for (kind, mode), profile in profiles.items(): + prefix = args.output.with_suffix(f".{kind}-{mode}") + profile.dump_stats(str(prefix) + ".pstats") + with open(str(prefix) + ".txt", "w") as stream: + pstats.Stats(profile, stream=stream).strip_dirs().sort_stats( + "cumulative" + ).print_stats(100) + print("FACADE_SUMMARY=" + json.dumps(summary), flush=True) + finally: + dist.destroy_process_group() + + +if __name__ == "__main__": + main() diff --git a/dev/trainer_v1_head_memory.py b/dev/trainer_v1_head_memory.py new file mode 100644 index 000000000..3f20d52b2 --- /dev/null +++ b/dev/trainer_v1_head_memory.py @@ -0,0 +1,181 @@ +"""Reserved-GPU oracle for registered-head admission and streamed cotangents.""" + +import argparse +import asyncio +import gc +import json +import os +from pathlib import Path + +from dotenv import load_dotenv +import torch +import torch.distributed as dist +from torch.multiprocessing.reductions import StorageWeakRef +from trainer_rank_support import load_random_checkpoints + +from art.trainer_rank import ( + ForwardInput, + ForwardOptions, + TrainerRank, + TrainerRankMemoryError, + run_rank_callback, +) + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--output", type=Path, required=True) + args = parser.parse_args() + load_dotenv(".env") + torch.set_num_threads(2) + torch.cuda.set_device(int(os.environ.get("LOCAL_RANK", "0"))) + dist.init_process_group("nccl") + try: + from art.megatron.train import build_training_runtime + + os.environ["ART_TRAINER_RANK_TEST_HOOKS"] = "1" + torch.manual_seed(716) + runtime = build_training_runtime( + model_identifier="Qwen/Qwen3-0.6B", + model_initialization="random", + provider_configure=lambda provider: setattr(provider, "num_layers", 2), + print_env=False, + ) + rank = TrainerRank(runtime) + (checkpoint,) = load_random_checkpoints( + runtime, rank, 1, base_model="Qwen/Qwen3-0.6B", lora_rank=2 + ) + request = ForwardInput( + input_tokens=torch.arange(32), + hidden_states=True, + checkpoint=checkpoint, + options=ForwardOptions( + backward_state="replay", + output_device="cpu", + stale_gradient_corrections=(), + ), + ) + # Warm native kernels before allocator measurements. + rank.backward(rank.forward(request).hidden_states.float().sum()) + rank.zero_grad() + incoming = [] + commit = rank._commit_versioned_gradients + + def record(gradients): + incoming.extend( + StorageWeakRef(gradient.untyped_storage()) + for _, _, _, gradient in gradients + ) + return commit(gradients) + + rank._commit_versioned_gradients = record + state = rank._version_state() + publish = state._publish + publications = [] + + def check(prepared): + assert all(reference.expired() for reference in incoming), ( + "GPU head cotangent survived until publication" + ) + publications.append(len(incoming)) + publish(prepared) + + state._publish = check + rows = [] + + def callback(view): + head = view.module( + "head", + lambda: torch.nn.Linear(rank.hidden_size, 4096, bias=False), + checkpoint=checkpoint, + ) + head.cpu() # Match driver/client head placement; isolate worker staging. + parameter = rank._checkpoint_slots[checkpoint].custom["head"].value.weight + target_bytes = parameter.numel() * parameter.element_size() + gc.collect() + torch.cuda.empty_cache() + torch.cuda.synchronize() + os.environ["ART_TRAINER_RANK_TEST_MEMORY_LIMIT_BYTES"] = str( + torch.cuda.memory_allocated() + 2 * target_bytes + ) + try: + try: + view.forward(request) + except TrainerRankMemoryError: + pass + else: + raise AssertionError( + "Known registered head staging was not reserved: " + + json.dumps(rank.last_forward_telemetry()) + ) + finally: + os.environ.pop("ART_TRAINER_RANK_TEST_MEMORY_LIMIT_BYTES", None) + # Two complete outstanding model roots, each using the same head + # eight times. Packet source tensors stay on the caller's CPU. + outputs = [view.forward(request) for _ in range(2)] + losses = [ + sum( + head(output.hidden_states.float().mean(0)).sum() * scale + for scale in range(1, 9) + ) + for output in outputs + ] + reserved = rank._lora_gradient_staging_bytes(rank._slot_ref(checkpoint)) + workspace = max( + rank._forward_graph_cache().state(handle).restore_workspace_bytes + for handle in rank._forward_graph_cache().handles() + ) + torch.cuda.synchronize() + baseline = torch.cuda.memory_allocated() + torch.cuda.reset_peak_memory_stats() + for loss in losses: + view.backward(loss) + torch.cuda.synchronize() + peak = torch.cuda.max_memory_allocated() - baseline + print( + "HEAD_PEAK=" + + json.dumps( + dict( + peak=peak, + reserved=reserved, + workspace=workspace, + publications=publications, + ) + ), + flush=True, + ) + assert publications == [8, 16] + assert peak <= reserved + workspace + expected = 36 * sum( + output.hidden_states.float().mean(0).detach().cpu() + for output in outputs + ) + torch.testing.assert_close( + parameter.grad.cpu(), + expected.expand(parameter.shape), + rtol=1e-5, + atol=1e-5, + ) + rows.append( + dict( + target_bytes=target_bytes, + staging_reserved_bytes=reserved, + restore_workspace_bytes=workspace, + backward_peak_bytes=peak, + publications=publications, + outstanding_roots=2, + captures_per_root=8, + ) + ) + + asyncio.run(run_rank_callback(rank, callback, mode="zero")) + assert not rank._forward_graph_cache().handles() + args.output.parent.mkdir(parents=True, exist_ok=True) + args.output.write_text(json.dumps(rows, indent=2)) + print("HEAD_MEMORY=" + json.dumps(rows), flush=True) + finally: + dist.destroy_process_group() + + +if __name__ == "__main__": + main() diff --git a/dev/trainer_v1_memory.py b/dev/trainer_v1_memory.py new file mode 100644 index 000000000..db0065a1c --- /dev/null +++ b/dev/trainer_v1_memory.py @@ -0,0 +1,383 @@ +"""Native complete-root admission, measured peaks and paired headroom timings.""" + +from __future__ import annotations + +import argparse +from dataclasses import asdict +import gc +import json +import os +from pathlib import Path +import time + +from dotenv import load_dotenv +import torch +import torch.distributed as dist +from trainer_rank_support import load_random_checkpoints + +from art.trainer_rank import ( + ForwardInput, + ForwardOptions, + TrainerRank, + TrainerRankMemoryError, +) +from art.trainer_rank._impl import _SplitForwardPlan +from art.trainer_rank._memory_policy import host_memory_budget, placement_cost + + +def _leaves(tree): + if isinstance(tree, (list, tuple)): + return [leaf for child in tree for leaf in _leaves(child)] + return [tree] + + +def _backward(rank, outputs): + loss = sum( + output.hidden_states.float().mean() for output in _leaves(outputs) + ).square() + with rank._gradient_transaction(): + packets = rank._forward_cotangent_collector().backward(loss) + rank._forward_graph_cache().backward_many( + [(p.handle, p.gradients) for p in packets] + ) + return float(loss.detach().cpu()) + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--model", default="Qwen/Qwen3-0.6B") + parser.add_argument("--layers", type=int, default=2) + parser.add_argument("--tokens", type=int, default=512) + parser.add_argument("--rounds", type=int, default=10) + parser.add_argument("--deterministic", action="store_true") + parser.add_argument("--output", type=Path, required=True) + args = parser.parse_args() + load_dotenv(".env") + if args.deterministic: + os.environ["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8" + torch.use_deterministic_algorithms(True) + import art.megatron.flex_attn.compiled as flex + + flex._FORCED_FLEX_BACKEND = "TRITON" + flex._FORCED_FLEX_KERNEL_OPTIONS = {"BACKEND": "TRITON"} + flex.dense_compiled_flex_attention = flex.triton_dense_compiled_flex_attention + flex.sparse_compiled_flex_attention = flex.triton_sparse_compiled_flex_attention + torch.cuda.set_device(int(os.environ.get("LOCAL_RANK", "0"))) + dist.init_process_group("nccl") + try: + if dist.get_world_size() != 1: + raise ValueError("This admission oracle uses one physical rank") + os.environ["ART_TRAINER_RANK_TEST_HOOKS"] = "1" + os.environ["ART_TRAINER_RANK_TEST_ANCHOR"] = "no_sharing" + from art.megatron.train import build_training_runtime + + def configure(provider): + provider.num_layers = args.layers + provider.recompute_granularity = None + provider.recompute_method = None + provider.recompute_num_layers = None + provider.recompute_modules = [] + + torch.manual_seed(913) + runtime = build_training_runtime( + model_identifier=args.model, + model_initialization="random", + provider_configure=configure, + print_env=False, + ) + for chunk in runtime.model: + chunk.eval() + rank = TrainerRank(runtime) + (checkpoint,) = load_random_checkpoints( + runtime, rank, 1, base_model=args.model, lora_rank=2 + ) + forward = getattr(rank, "forward", None) or rank.dp_rank_forward + batches = getattr(rank, "forward_batches", None) or rank.forward_micro_batches + options = lambda state, device: ForwardOptions( + backward_state=state, output_device=device, stale_gradient_corrections=() + ) + + def root(state, device): + leaves = [ + ForwardInput( + input_tokens=(torch.arange(args.tokens) * 17 + 103 * i) % 30000, + hidden_states=True, + checkpoint=checkpoint, + options=options(state, device), + ) + for i in range(4) + ] + return [leaves[:2], leaves[2:]] + + # Learn both forward retention and caller backward peak with the existing + # profiler, at the same physical child size later used by the split root. + for _ in range(2): + rank.zero_grad() + for batch in batches([root("gpu", "model")[0][0]]): + _backward(rank, batch.outputs) + del batch + + rows = [] + gradients_by_label = {} + child_peaks = [] + run_child = rank._run_flat_plan_with_memory_tracking + + def track(*call_args, **kwargs): + result = run_child(*call_args, **kwargs) + child_peaks.append(torch.cuda.max_memory_allocated()) + return result + + rank._run_flat_plan_with_memory_tracking = track + + def run( + state, + device, + label, + *, + cap=None, + chunks=None, + compare=None, + legacy_admission=False, + ): + rank.zero_grad() + gc.collect() + torch.cuda.empty_cache() + torch.cuda.synchronize() + baseline = torch.cuda.memory_allocated() + if cap is None: + os.environ.pop("ART_TRAINER_RANK_TEST_MEMORY_LIMIT_BYTES", None) + else: + os.environ["ART_TRAINER_RANK_TEST_MEMORY_LIMIT_BYTES"] = str( + baseline + cap + ) + child_peaks.clear() + torch.cuda.reset_peak_memory_stats() + started = time.perf_counter() + find = rank._find_admissible_forward + enabled = rank._graph_memory_policy_enabled + + def fixed_split(requests, *, checkpoint, **_kwargs): + plan, check = rank._admit_split_rung( + chunks, + requests, + [request.input_tokens for request in requests], + checkpoint=checkpoint, + ) + assert plan is not None and check.fits + return plan, check + + if chunks is not None: + rank._find_admissible_forward = fixed_split + if legacy_admission: + # Isolate admission overhead while retaining identical graph + # execution, kernels and the existing calibrated GPU profiler. + rank._graph_memory_policy_enabled = lambda: False + try: + outputs = forward(root(state, device)) + finally: + rank._find_admissible_forward = find + rank._graph_memory_policy_enabled = enabled + torch.cuda.synchronize() + retained = torch.cuda.memory_allocated() - baseline + assert [len(part) for part in outputs] == [2, 2] + assert all( + tuple(value.hidden_states.shape) == (args.tokens, rank._hidden_size) + for value in _leaves(outputs) + ) + cache = rank._forward_graph_cache() + states = [asdict(cache.state(handle)) for handle in cache.handles()] + loss = _backward(rank, outputs) + torch.cuda.synchronize() + elapsed = time.perf_counter() - started + peak = max([torch.cuda.max_memory_allocated(), *child_peaks]) - baseline + gradients = [ + None if p.grad is None else p.grad.detach().float().cpu() + for p in rank._checkpoint_slots[checkpoint].params + ] + gradients_by_label[label] = gradients + if cache.handles(): + raise AssertionError("Consumed forward graph remained cached") + row = dict( + label=label, + state=state, + output_device=device, + legacy_admission=legacy_admission, + loss=loss, + elapsed_seconds=elapsed, + logical_tokens=4 * args.tokens, + gpu_retained_bytes=retained, + gpu_peak_bytes=peak, + budget_bytes=cap, + graph_states=states, + telemetry=rank.last_forward_telemetry(), + transfer_stats=asdict(cache.transfer_stats) + if hasattr(cache, "transfer_stats") + else None, + ) + rows.append(row) + print("NATIVE_MEMORY=" + json.dumps(row, default=str), flush=True) + # Preserve measured evidence even when a correctness gate fails. + args.output.parent.mkdir(parents=True, exist_ok=True) + args.output.write_text(json.dumps({"rows": rows}, indent=2, default=str)) + torch.save(gradients_by_label, args.output.with_suffix(".gradients.pt")) + if compare is not None: + for actual, expected in zip( + gradients, gradients_by_label[compare], strict=True + ): + if (actual is None) != (expected is None): + raise AssertionError( + "Used parameter set changed across policies" + ) + if actual is not None: + torch.testing.assert_close( + actual, expected, atol=5e-5, rtol=0.02 + ) + return row + + run("gpu", "model", "reference") + # The cap is derived from the same measured/calibrated child costs used + # by admission. It limits available memory, never replaces a model cost. + leaves = _leaves(root("replay", "cpu")) + children = tuple(rank._plan_flat_forward([leaf]) for leaf in leaves) + split = _SplitForwardPlan(children, tuple((i,) for i in range(4)), 4) + placements = [ + placement_cost((cost,), backward_state="replay", output_device="cpu") + for _, _, cost, _ in rank._graph_memory_units(split) + ] + cap = sum( + p.gpu_retained_bytes + p.gpu_backward_bytes for p in placements + ) + max( + p.gpu_required_bytes - p.gpu_retained_bytes - p.gpu_backward_bytes + for p in placements + ) + cap = int(cap * 1.05) + try: + run("gpu", "model", "must_refuse", cap=cap) + except TrainerRankMemoryError: + pass + else: + raise AssertionError("Constrained retained-GPU root was not refused") + accepted = run("replay", "cpu", "constrained_replay", cap=cap) + if accepted["telemetry"]["subforward_count"] <= 1: + raise AssertionError( + "Constrained complete root did not execute as split physical forwards" + ) + if accepted["gpu_peak_bytes"] > cap: + raise AssertionError("Native measured peak exceeded admission cap") + # BF16 packed and split matmuls can differ numerically. Compare replay + # against the identical admitted physical partition, with CPU reduction + # on both arms; retain the packed reference as a separate diagnostic. + run( + "gpu", + "cpu", + "same_split_gpu", + chunks=accepted["telemetry"]["subforward_request_indices"], + compare="constrained_replay", + ) + run( + "cpu", + "cpu", + "same_split_cpu_offload", + chunks=accepted["telemetry"]["subforward_request_indices"], + compare="constrained_replay", + ) + adaptive = run( + "auto", + "cpu", + "constrained_measured_auto", + cap=cap, + chunks=accepted["telemetry"]["subforward_request_indices"], + compare="constrained_replay", + ) + evidence = adaptive["telemetry"]["fallback_costs"] + if evidence["source"] != "measured_forward_and_transfers": + raise AssertionError("Matching measured fallback evidence was not used") + if {state["retention"] for state in adaptive["graph_states"]} != { + evidence["preferred"] + }: + raise AssertionError("Admitted fallback did not follow measured costs") + if adaptive["gpu_peak_bytes"] > cap: + raise AssertionError( + "Measured adaptive fallback peak exceeded admission cap" + ) + # Paired alternating warmed rounds compare the default choice with the + # forced retained path under headroom, using identical computation. + for repetition in range(args.rounds): + order = ("auto", "gpu", "legacy") + for state in order if repetition % 2 == 0 else tuple(reversed(order)): + run( + "gpu" if state == "legacy" else state, + "model", + f"headroom_{state}", + compare="reference", + legacy_admission=state == "legacy", + ) + # An older, larger replay must retain its restore reservation even when + # the next root is small enough to fit by itself. + os.environ.pop("ART_TRAINER_RANK_TEST_MEMORY_LIMIT_BYTES", None) + rank.zero_grad() + old_outputs = forward(root("auto", "cpu")) + cache = rank._forward_graph_cache() + for handle in cache.handles(): + cache.evict(handle) + gc.collect() + torch.cuda.empty_cache() + workspace = max( + cache.state(handle).restore_workspace_bytes for handle in cache.handles() + ) + small = root("replay", "cpu")[0][0] + small_cost = next(rank._graph_memory_units(rank._plan_flat_forward([small])))[2] + small_required = placement_cost( + (small_cost,), backward_state="replay", output_device="cpu" + ).gpu_required_bytes + cap = workspace - 1024**2 + assert 0 < small_required < cap + baseline = torch.cuda.memory_allocated() + os.environ["ART_TRAINER_RANK_TEST_MEMORY_LIMIT_BYTES"] = str(baseline + cap) + try: + try: + forward([small]) + except TrainerRankMemoryError: + pass + else: + raise AssertionError("A new root consumed the older replay reservation") + finally: + os.environ.pop("ART_TRAINER_RANK_TEST_MEMORY_LIMIT_BYTES", None) + torch.cuda.reset_peak_memory_stats() + _backward(rank, old_outputs) + torch.cuda.synchronize() + restored_peak = torch.cuda.max_memory_allocated() - baseline + assert restored_peak <= workspace + assert not cache.handles() + row = dict( + label="prior_replay_reservation", + old_restore_workspace_bytes=workspace, + small_root_required_bytes=small_required, + budget_bytes=cap, + measured_restore_peak_bytes=restored_peak, + ) + rows.append(row) + print("NATIVE_MEMORY=" + json.dumps(row), flush=True) + args.output.parent.mkdir(parents=True, exist_ok=True) + args.output.write_text( + json.dumps( + { + "model": args.model, + "layers": args.layers, + "tokens_per_child": args.tokens, + "torch": torch.__version__, + "deterministic": args.deterministic, + "device": torch.cuda.get_device_name(), + "host_budget": asdict(host_memory_budget(local_world_size=1)), + "rows": rows, + }, + indent=2, + default=str, + ) + ) + finally: + dist.destroy_process_group() + + +if __name__ == "__main__": + main() diff --git a/dev/trainer_v1_validation.sky.yaml b/dev/trainer_v1_validation.sky.yaml new file mode 100644 index 000000000..cf02b3aee --- /dev/null +++ b/dev/trainer_v1_validation.sky.yaml @@ -0,0 +1,34 @@ +name: trainer-v1-validation + +workdir: . + +resources: + infra: k8s/cks-wb3 + accelerators: H200:4 + cpus: 32+ + memory: 256+ + image_id: docker:docker.io/bradhiltonnw/art-gpu:latest + +setup: | + INSTALL_VLLM_RUNTIME=false bash src/art/megatron/setup.sh + uv sync --project megatron_runtime --extra cuda12 --group test \ + --frozen --no-install-project --inexact + +run: | + set -euo pipefail + nvidia-smi + timeout --signal=TERM --kill-after=30s 30m \ + bash scripts/ci/trainer-rank-gpu-tests.sh + +config: + kubernetes: + pod_config: + spec: + schedulerName: binpack-scheduler + activeDeadlineSeconds: 108000 + containers: + - name: ray-node + imagePullPolicy: Always + env: + - name: UV_LINK_MODE + value: copy diff --git a/pyproject.toml b/pyproject.toml index 82a527972..a10727649 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -5,6 +5,7 @@ description = "The OpenPipe Agent Reinforcement Training (ART) library" readme = "README.md" requires-python = ">=3.12" dependencies = [ + "cloudpickle>=3.1.1", "aiohttp>=3.10.0", "anthropic>=0.77.0", "openai>=2.14.0,<3", diff --git a/scripts/ci/trainer-rank-gpu-tests.sh b/scripts/ci/trainer-rank-gpu-tests.sh index 690d99db3..d6729a6b4 100755 --- a/scripts/ci/trainer-rank-gpu-tests.sh +++ b/scripts/ci/trainer-rank-gpu-tests.sh @@ -3,6 +3,7 @@ set -euo pipefail export CUDA_VISIBLE_DEVICES=0,1 export PYTHONUNBUFFERED=1 +export ART_GRAPH_GPU_TEST=1 runtime_python="$( .venv/bin/python -c 'from art.megatron.runtime.managed import ensure_megatron_runtime; print(ensure_megatron_runtime(art_build_sha256="trainer-rank-ci").python)' )" @@ -11,6 +12,19 @@ test -x "${runtime_python}" "${runtime_python}" -m pytest --tb=short \ tests/unit/test_trainer_rank_head_recompute.py \ tests/unit/test_trainer_rank_custom_tensors.py \ + tests/unit/test_trainer_rank_tensors.py \ + tests/unit/test_trainer_rank_graphs_cuda.py \ + tests/unit/test_trainer_rank_memory_policy_cuda.py \ + tests/unit/test_trainer_rank_head_memory_cuda.py \ + tests/unit/test_trainer_rank_output_memory_cuda.py \ + tests/unit/test_trainer_rank_live_heads.py \ + tests/unit/test_trainer_rank_commands.py \ + tests/unit/test_trainer_command_transport.py \ + tests/unit/test_trainer_driver_transport.py \ + tests/unit/test_trainer_rank_versions.py \ + tests/integration/megatron/lora/test_lora_versions.py \ + tests/integration/megatron/lora/test_trainer_v1_versions.py \ + tests/integration/megatron/lora/test_trainer_v1_graph_cache.py \ tests/integration/megatron/cp_attn/test_attention_packed_vs_flattened.py \ 'tests/integration/megatron/gdn_shared_prefix/test_gdn_cp_packed_correctness.py::test_gdn_cp_packed_sibling_order_matches_cp1_oracle[2]' \ 'tests/integration/megatron/gdn_shared_prefix/test_gdn_cp_packed_correctness.py::test_gdn_cp_tree_chain_matches_cp1_oracle[2]' \ @@ -21,6 +35,12 @@ test -x "${runtime_python}" tests/integration/megatron/lora/test_dynamic_lora_slots.py::test_trainer_rank_custom_parameter_reduction_oracle \ 'tests/integration/megatron/lora/test_dynamic_lora_slots.py::test_trainer_rank_tp_head_backward_matches_unsharded_oracle[2]' +# Bound distributed retained-backward and residency regressions independently. +timeout --signal=TERM --kill-after=30s 10m "${runtime_python}" -m pytest --tb=short \ + tests/unit/test_trainer_rank_resident_memory_cuda.py \ + tests/integration/megatron/cp_attn/test_retained_backward.py \ + tests/integration/megatron/cp_attn/test_cpu_offload_residency.py + # Keep SFT distributed state and compiler workarounds in separate test processes. "${runtime_python}" -m pytest --tb=short \ tests/integration/megatron/test_sft_packing.py::test_sft_packing_loss_and_gradients diff --git a/src/art/_tensor_residency.py b/src/art/_tensor_residency.py new file mode 100644 index 000000000..226dfb1cc --- /dev/null +++ b/src/art/_tensor_residency.py @@ -0,0 +1,28 @@ +"""Observe tensors held outside autograd saved-variable hooks without owning them.""" + +from collections.abc import Callable, Iterator, Sequence +from contextlib import contextmanager +from contextvars import ContextVar + +import torch + +_observer: ContextVar[Callable[[Sequence[torch.Tensor]], None] | None] = ContextVar( + "art_tensor_residency_observer", default=None +) + + +@torch.compiler.disable +def record_resident_tensors(tensors: Sequence[torch.Tensor]) -> None: + if observer := _observer.get(): + observer(tensors) + + +@contextmanager +def observe_resident_tensors( + observer: Callable[[Sequence[torch.Tensor]], None], +) -> Iterator[None]: + token = _observer.set(observer) + try: + yield + finally: + _observer.reset(token) diff --git a/src/art/megatron/compile_workarounds.py b/src/art/megatron/compile_workarounds.py index 0759ace4b..07c4bef7f 100644 --- a/src/art/megatron/compile_workarounds.py +++ b/src/art/megatron/compile_workarounds.py @@ -1,9 +1,14 @@ from __future__ import annotations +from copy import copy +from functools import wraps import os from typing import Any +import weakref import torch +from torch._C import _current_graph_task_id +from torch._C._autograd import _get_current_graph_task_keep_graph from art.megatron.model_support.spec import CompileWorkaroundConfig @@ -13,6 +18,112 @@ ) +def install_te_reusable_backward() -> None: + """Preserve TE's saved-tensor metadata until the final backward.""" + from transformer_engine.pytorch.module.layernorm_linear import _LayerNormLinear + from transformer_engine.pytorch.module.layernorm_mlp import _LayerNormMLP + from transformer_engine.pytorch.module.linear import _Linear + from transformer_engine.pytorch.ops.fuser import _OperationFuserAutogradFunction + + for function in ( + _OperationFuserAutogradFunction, + _Linear, + _LayerNormLinear, + _LayerNormMLP, + ): + _preserve_te_backward_metadata(function) + + +def _preserve_te_backward_metadata(function) -> None: + original = function.backward + if getattr(original, "__art_reusable_backward__", False): + return + + @wraps(original) + def backward(ctx, *gradients): + if not _get_current_graph_task_keep_graph(): + return original(ctx, *gradients) + tensor_objects = ctx.tensor_objects + if any(value is not None for value in tensor_objects): + raise RuntimeError( + "Retained backward is not supported for Transformer Engine quantized saved tensors" + ) + contexts = getattr(ctx, "basic_op_ctxs", ()) + ranges = [op_ctx._saved_tensors_range for op_ctx in contexts] + try: + return original(ctx, *gradients) + finally: + ctx.tensor_objects = tensor_objects + for op_ctx, tensor_range in zip(contexts, ranges, strict=True): + op_ctx._saved_tensors_range = tensor_range + # Do not keep unpacked tensors alive between backward calls, + # including when an operation raises before TE's own cleanup. + op_ctx.saved_tensors = None + + setattr(backward, "__art_reusable_backward__", True) + setattr(function, "backward", staticmethod(backward)) + + +def install_reusable_checkpoint_backward() -> None: + """Rebuild selective checkpoint outputs separately for each backward.""" + from megatron.core.tensor_parallel.random import CheckpointWithoutOutput + + original = CheckpointWithoutOutput._recompute + if getattr(original, "__art_reusable_backward__", False): + return + fields = ("run_function", "rng_states", "outputs", "ctx") + original_discard = CheckpointWithoutOutput.discard_output_and_register_recompute + + @wraps(original_discard) + def discard(self, hook_tensor): + # TransformerLayer retains the controller on the module. Transfer its + # recipe to the graph hook so eviction can release the physical graph. + if self.ctx is None: + owned_ref = getattr(self, "_art_recompute_owner", None) + if owned_ref is None: + return original_discard(self, hook_tensor) + owned = owned_ref() + if owned is None: + return + else: + owned = copy(self) + self._art_recompute_owner = weakref.ref(owned) + try: + return original_discard(owned, hook_tensor) + except BaseException: + for field in fields: + setattr(owned, field, None) + raise + finally: + for field in fields: + setattr(self, field, None) + + @wraps(original) + def recompute(self, gradient): + if not _get_current_graph_task_keep_graph(): + return original(self, gradient) + task = _current_graph_task_id() + if getattr(self, "_art_recompute_task", None) == task: + return + # The inner autograd graph is consumed normally. Preserve the forward + # recipe so the next outer backward recomputes a fresh inner graph. + state = tuple(getattr(self, field) for field in fields) + try: + result = original(self, gradient) + self._art_recompute_task = task + return result + except BaseException: + state = (None,) * len(fields) + raise + finally: + for field, value in zip(fields, state, strict=True): + setattr(self, field, value) + + setattr(recompute, "__art_reusable_backward__", True) + setattr(CheckpointWithoutOutput, "_recompute", recompute) + setattr(CheckpointWithoutOutput, "discard_output_and_register_recompute", discard) + + def _require_attr(obj: Any, name: str) -> Any: value = getattr(obj, name, None) if value is None: diff --git a/src/art/megatron/context_parallel/executor.py b/src/art/megatron/context_parallel/executor.py index 4015011fe..a0cc6f556 100644 --- a/src/art/megatron/context_parallel/executor.py +++ b/src/art/megatron/context_parallel/executor.py @@ -3,12 +3,14 @@ from typing import Any, cast import torch +from torch._C._autograd import _get_current_graph_task_keep_graph from torch._dynamo import config as dynamo_config import torch.distributed as dist from torch.nn.attention.flex_attention import BlockMask import triton import triton.language as tl +from art._tensor_residency import record_resident_tensors from art.megatron.flex_attn.compiled import ( SparseBlockSize, flash_sparse_block_size_for_head_dim, @@ -672,6 +674,16 @@ def run( head_dim_v=int(v.shape[-1]), device=q.device, ) + if ( + backend == "FLASH" + and q.device.type == "cuda" + and int(q.shape[-1]) <= 64 + and torch.cuda.get_device_capability(q.device)[0] == 9 + ): + # SM90 sparse FLASH dQ is incorrect at these head widths. Both + # backends use 128x128 mask blocks here; select Triton's distinct + # compiled kernel and LSE convention together. + backend = "TRITON" if compile_key is None: _q_len, _k_len, compile_key = select_sparse_execution_family( is_local_stage=bool(is_local_stage), @@ -2063,6 +2075,7 @@ def _run_context_parallel_backward( replay_records: list[dict[str, Any]] | None = None, replay_accum_out: torch.Tensor | None = None, replay_accum_lse: torch.Tensor | None = None, + retain_graph: bool = False, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor | None]: kernel = FlexAttentionKernel( compile_enabled=compile_enabled, @@ -2240,6 +2253,7 @@ def _run_context_parallel_backward( inputs=inputs, grad_outputs=tuple(stage_output_grads), allow_unused=True, + retain_graph=retain_graph, ) grad_map: dict[str, torch.Tensor | None] = { name: grad for name, grad in zip(input_names, input_grads, strict=True) @@ -2376,6 +2390,14 @@ def forward( tensors_to_save.extend((replay_accum_out, replay_accum_lse)) ctx.save_for_backward(*tensors_to_save) ctx.replay_records = replay_records + record_resident_tensors( + tuple( + value + for record in replay_records + for value in record.values() + if isinstance(value, torch.Tensor) + ) + ) return output.detach() @staticmethod @@ -2390,6 +2412,14 @@ def backward(ctx, *grad_outputs: Any): softmax_offset = None replay_accum_out = None replay_accum_lse = None + retain_graph = _get_current_graph_task_keep_graph() + replay_records = cast(list[dict[str, Any]], ctx.replay_records) + # Stage backward consumes its dictionaries and merge tape. A retained + # outer graph needs both that metadata and the inner attention graphs + # again; copying dictionaries preserves them without copying tensors. + if retain_graph: + replay_records = [record.copy() for record in replay_records] + succeeded = False try: dq, dk, dv, grad_softmax_offset = _run_context_parallel_backward( grad_output=grad_output, @@ -2403,12 +2433,15 @@ def backward(ctx, *grad_outputs: Any): sliding_window=ctx.sliding_window, triton_num_stages_2_head_dims=ctx.triton_num_stages_2_head_dims, softmax_offset=softmax_offset, - replay_records=cast(list[dict[str, Any]], ctx.replay_records), + replay_records=replay_records, replay_accum_out=replay_accum_out, replay_accum_lse=replay_accum_lse, + retain_graph=retain_graph, ) + succeeded = True finally: - ctx.replay_records = None + if not retain_graph or not succeeded: + ctx.replay_records = None return dq, dk, dv, grad_softmax_offset, None, None, None, None, None, None diff --git a/src/art/megatron/lora.py b/src/art/megatron/lora.py index a7fdfb3c0..626f32666 100644 --- a/src/art/megatron/lora.py +++ b/src/art/megatron/lora.py @@ -1,4 +1,6 @@ -from collections.abc import Iterator, Sequence +from __future__ import annotations + +from collections.abc import Iterator, Mapping, Sequence from contextlib import contextmanager import contextvars from dataclasses import dataclass, replace @@ -66,27 +68,51 @@ class LoRASlotRef: _CURRENT_LORA_SLOT: contextvars.ContextVar[LoRASlotRef | None] = contextvars.ContextVar( "art_megatron_current_lora_slot", default=None ) +_CURRENT_LORA_VERSION: contextvars.ContextVar[LoRAVersion | None] = ( + contextvars.ContextVar("art_megatron_current_lora_version", default=None) +) + + +@dataclass(frozen=True) +class LoRAVersion: + ref: LoRASlotRef + version: Any + slots: Mapping[int, LoRASlot] + validate: Callable[[], None] + weight_version: Any = None + + @property + def nbytes(self) -> int: + return sum( + param.numel() * param.element_size() + for slot in self.slots.values() + for param in (slot.A_T, slot.B_T) + ) @contextmanager -def use_lora_slot(ref: LoRASlotRef | None) -> Iterator[None]: +def use_lora_slot( + ref: LoRASlotRef | None, *, version: LoRAVersion | None = None +) -> Iterator[None]: + if version is not None and version.ref != ref: + raise ValueError("LoRA version belongs to a different slot") token = _CURRENT_LORA_SLOT.set(ref) + version_token = _CURRENT_LORA_VERSION.set(version) try: yield finally: _CURRENT_LORA_SLOT.reset(token) + _CURRENT_LORA_VERSION.reset(version_token) def _with_captured_lora_slot(function: _F) -> _F: context = _CURRENT_LORA_SLOT.get() + version = _CURRENT_LORA_VERSION.get() @functools.wraps(function) def wrapped(*args: Any, **kwargs: Any) -> Any: - token = _CURRENT_LORA_SLOT.set(context) - try: + with use_lora_slot(context, version=version): return function(*args, **kwargs) - finally: - _CURRENT_LORA_SLOT.reset(token) return cast(_F, wrapped) @@ -101,7 +127,7 @@ def _patch_function_once(module: Any, name: str, wrapper: Callable[[_F], _F]) -> def install_lora_checkpoint_context_hooks() -> None: - """Preserve the selected dynamic LoRA slot across activation recompute.""" + """Preserve the selected slot and immutable tensors across recompute.""" def wrap_checkpoint(original: _F, function_index: int) -> _F: @functools.wraps(original) @@ -910,7 +936,8 @@ def active_lora_tensors( return self.A_T, self.B_T, self.scale if ref.name is None: return None - slot = self._slot(ref) + version = _CURRENT_LORA_VERSION.get() + slot = self._slot(ref) if version is None else version.slots.get(id(self)) if slot is None: return None return slot.A_T, slot.B_T, slot.scale diff --git a/src/art/megatron/prefix_tree_packing.py b/src/art/megatron/prefix_tree_packing.py index 381388292..0530d5e0b 100644 --- a/src/art/megatron/prefix_tree_packing.py +++ b/src/art/megatron/prefix_tree_packing.py @@ -50,7 +50,7 @@ def prefix_tree_pack( ) -> PrefixTreePack: """Pack token sequences by storing prefix trees once. - This is the small packing step that lets `TrainerRank.dp_rank_forward()` run one + This is the small packing step that lets `TrainerRank.forward()` run one model pass over a compact prefix tree instead of replaying the same prompt tokens for every request. Think of each input sequence as a path through a tree: when several paths start with the same tokens, this function writes diff --git a/src/art/megatron/runtime/compile_cache.py b/src/art/megatron/runtime/compile_cache.py index 3b4aa1fba..2a4661e35 100644 --- a/src/art/megatron/runtime/compile_cache.py +++ b/src/art/megatron/runtime/compile_cache.py @@ -7,7 +7,7 @@ from pathlib import Path import sys import time -from typing import Any, Literal +from typing import Any, Literal, cast import uuid from pydantic import BaseModel, ConfigDict, Field @@ -17,6 +17,22 @@ _PACKAGES = ("megatron-core", "torchmonarch", "transformer-engine", "transformers") +def configure_reusable_backward() -> None: + """Set native compile policy before loading artifacts or compiling forwards.""" + import torch + from torch._functorch import config as functorch_config + + # A backward-only toggle would bypass AOT's donated-buffer safety check. + # This cannot repair graphs already compiled by an external runtime. + cast(Any, functorch_config).donated_buffer = False + # AOT hashes the flag, but Inductor's lower FX cache does not. Separate + # donating kernels there too, preserving any caller-provided cache tag. + suffix = "|art-retained-backward-v1" + tag = torch.compiler.config.cache_key_tag + if not tag.endswith(suffix): + torch.compiler.config.cache_key_tag = tag + suffix + + class CompileCacheEvent(BaseModel): model_config = ConfigDict(extra="forbid", frozen=True) @@ -79,6 +95,8 @@ def _compile_cache_key(spec: TrainerRuntimeSpec, rank: int) -> str: "compile_workarounds": os.environ.get( "ART_MEGATRON_COMPILE_WORKAROUNDS", "1" ), + "donated_buffer": False, + "cache_key_tag": torch.compiler.config.cache_key_tag, }, } encoded = json.dumps(payload, sort_keys=True, separators=(",", ":")).encode() @@ -91,6 +109,7 @@ class TrainerCompileCache: def __init__( self, spec: TrainerRuntimeSpec, *, rank: int, cache_root: Path ) -> None: + configure_reusable_backward() self.key = _compile_cache_key(spec, rank) self.path = cache_root / "megatron" / "compile_cache" / "v1" / self.key self.path.parent.mkdir(parents=True, exist_ok=True) diff --git a/src/art/megatron/training/compile.py b/src/art/megatron/training/compile.py index 7d824c976..99dfae1da 100644 --- a/src/art/megatron/training/compile.py +++ b/src/art/megatron/training/compile.py @@ -7,8 +7,13 @@ import torch from torch._dynamo import config as dynamo_config -from art.megatron.compile_workarounds import install_torch_compile_workarounds +from art.megatron.compile_workarounds import ( + install_reusable_checkpoint_backward, + install_te_reusable_backward, + install_torch_compile_workarounds, +) from art.megatron.provider import ProviderBundle +from art.megatron.runtime.compile_cache import configure_reusable_backward from art.megatron.training.model_chunks import ModelChunks _DYNAMO_CONFIG = cast(Any, dynamo_config) @@ -16,6 +21,7 @@ def _configure_dynamo() -> None: """Set the process-wide Dynamo policy required by dynamic LoRA slots.""" + configure_reusable_backward() # Dynamic checkpoint slots register differently shaped projection parameters # behind one LoRA.forward code object. Let automatic dynamic shapes generalize # those parameter dimensions instead of compiling once per projection site. @@ -64,6 +70,10 @@ def configure_training_compile( provider: Any, provider_bundle: ProviderBundle, ) -> bool: + # Flex-attention suboperators may compile even with layer compilation off. + configure_reusable_backward() + install_te_reusable_backward() + install_reusable_checkpoint_backward() compile_workaround_config = provider_bundle.handler.compile_workaround_config( provider ) diff --git a/src/art/trainer_rank/__init__.py b/src/art/trainer_rank/__init__.py index fd321a820..49b573815 100644 --- a/src/art/trainer_rank/__init__.py +++ b/src/art/trainer_rank/__init__.py @@ -9,6 +9,14 @@ from . import _impl from ._checkpoint import CheckpointManifest, materialize_lora, validate_checkpoint +from ._heads import ModuleHandle +from ._options import ( + ForwardOptions, + ImportanceSamplingGradientCorrection, + ResolvedForwardOptions, + resolve_forward_options, +) +from ._options import _Unset as _Unset AdapterSelection = _impl.AdapterSelection AdamParams = _impl.AdamParams @@ -42,7 +50,10 @@ for _public_type in ( AdamParams, ForwardInput, + ForwardOptions, ForwardOutput, + ImportanceSamplingGradientCorrection, + ResolvedForwardOptions, MicroBatch, MicroBatchStats, TopK, @@ -60,8 +71,8 @@ class TrainerRank(_impl.TrainerRank): """Execute TrainerRank forwards using automatic, data-dependent planning. - The constructor intentionally accepts only the training runtime. Prefix - sharing and microbatch width are data-dependent planner decisions; + The constructor accepts the training runtime and optional forward policy. + Prefix sharing and microbatch width are data-dependent planner decisions; output-head chunking and memory margins are internal calibrated policy. None are user tuning parameters. Requires PP=1 (TrainerRank does not use the MCore pipeline schedule); PP>1 raises ``TrainerRankRuntimeSupportError`` @@ -69,8 +80,10 @@ class TrainerRank(_impl.TrainerRank): profile is keyed by topology and calibrates itself online. """ - def __init__(self, runtime: TrainingRuntime) -> None: - super().__init__(runtime) + def __init__( + self, runtime: TrainingRuntime, *, options: ForwardOptions | None = None + ) -> None: + super().__init__(runtime, options=options) @property def hidden_size(self) -> int: @@ -86,8 +99,8 @@ def module( factory: Callable[[], ModuleT], *, checkpoint: AdapterSelection = Unset, - ) -> ModuleT: - """Register or retrieve a checkpoint-owned PyTorch module.""" + ) -> ModuleHandle: + """Retrieve a live module whose calls capture immutable checkpoint weights.""" return super().module(name, factory, checkpoint=checkpoint) def parameter( @@ -159,10 +172,11 @@ def export_lora( return super().export_lora(output_dir, checkpoint_path) @overload - def forward_micro_batches( + def forward_batches( self, inputs: Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]], *, + options: ForwardOptions | None = None, checkpoint: AdapterSelection = Unset, no_grad: bool | None = None, yield_empty: bool = False, @@ -174,12 +188,13 @@ def forward_micro_batches( ]: ... @overload - def forward_micro_batches( + def forward_batches( self, inputs: Iterable[ Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]] ], *, + options: ForwardOptions | None = None, checkpoint: AdapterSelection = Unset, no_grad: bool | None = None, yield_empty: bool = False, @@ -191,12 +206,13 @@ def forward_micro_batches( ]: ... @overload - def forward_micro_batches( + def forward_batches( self, inputs: Iterable[ Iterable[Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]]] ], *, + options: ForwardOptions | None = None, checkpoint: AdapterSelection = Unset, no_grad: bool | None = None, yield_empty: bool = False, @@ -208,7 +224,7 @@ def forward_micro_batches( ]: ... @overload - def forward_micro_batches( + def forward_batches( self, inputs: Iterable[ Iterable[ @@ -218,6 +234,7 @@ def forward_micro_batches( ] ], *, + options: ForwardOptions | None = None, checkpoint: AdapterSelection = Unset, no_grad: bool | None = None, yield_empty: bool = False, @@ -236,10 +253,11 @@ def forward_micro_batches( ] ]: ... - def forward_micro_batches( + def forward_batches( self, inputs: Iterable[ForwardInputs], *, + options: ForwardOptions | None = None, checkpoint: AdapterSelection = Unset, no_grad: bool | None = None, yield_empty: bool = False, @@ -254,10 +272,11 @@ def forward_micro_batches( the caller-owned `ForwardInput` objects. Per-position outputs contain the full flattened input sequence in source - order, including with context parallelism. Callers must compute identical - losses on every TP/CP replica; ART routes gradients to owning rows - without multiplying them by the number of replicas. `dp_reduce` combines - only distinct data-parallel batches. + order, including with context parallelism. Logical callbacks execute + once per DP rank; use `backward(loss)` to route cotangents to internal + TP/CP participants. Direct physical callers must invoke matching + forwards and backwards on their TP/CP peers. `reduce` combines only + distinct data-parallel batches. Empty local microbatches are skipped unless `yield_empty=True`. Every rank must use the same setting. When a wave skips ranks, TrainerRank @@ -270,28 +289,44 @@ def forward_micro_batches( """ forward = cast( Callable[..., Iterator[MicroBatch[ForwardInputs, ForwardOutputs]]], - super().forward_micro_batches, + super().forward_batches, ) return forward( - inputs, checkpoint=checkpoint, no_grad=no_grad, yield_empty=yield_empty + inputs, + options=options, + checkpoint=checkpoint, + no_grad=no_grad, + yield_empty=yield_empty, ) @overload - def dp_rank_forward( + def forward( + self, + inputs: ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT], + *, + options: ForwardOptions | None = None, + checkpoint: AdapterSelection = Unset, + no_grad: bool | None = None, + ) -> ForwardOutput[LogprobsT, TopKT, LogitsT, HiddenStatesT]: ... + + @overload + def forward( self, inputs: Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]], *, + options: ForwardOptions | None = None, checkpoint: AdapterSelection = Unset, no_grad: bool | None = None, ) -> Sequence[ForwardOutput[LogprobsT, TopKT, LogitsT, HiddenStatesT]]: ... @overload - def dp_rank_forward( + def forward( self, inputs: Iterable[ Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]] ], *, + options: ForwardOptions | None = None, checkpoint: AdapterSelection = Unset, no_grad: bool | None = None, ) -> Sequence[ @@ -299,12 +334,13 @@ def dp_rank_forward( ]: ... @overload - def dp_rank_forward( + def forward( self, inputs: Iterable[ Iterable[Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]]] ], *, + options: ForwardOptions | None = None, checkpoint: AdapterSelection = Unset, no_grad: bool | None = None, ) -> Sequence[ @@ -312,7 +348,7 @@ def dp_rank_forward( ]: ... @overload - def dp_rank_forward( + def forward( self, inputs: Iterable[ Iterable[ @@ -322,6 +358,7 @@ def dp_rank_forward( ] ], *, + options: ForwardOptions | None = None, checkpoint: AdapterSelection = Unset, no_grad: bool | None = None, ) -> Sequence[ @@ -330,17 +367,18 @@ def dp_rank_forward( ] ]: ... - def dp_rank_forward( + def forward( self, inputs: ForwardInputs, *, + options: ForwardOptions | None = None, checkpoint: AdapterSelection = Unset, no_grad: bool | None = None, ) -> ForwardOutputs: """Forward inputs already local to this data-parallel rank. Outputs contain full sequences in source order on every TP/CP rank, - with the same loss and reduction contract as `forward_micro_batches`. + with the same loss and reduction contract as `forward_batches`. Per-input checkpoints and `no_grad` values override the method defaults. `no_grad=None` inherits the ambient PyTorch grad mode; `True` disables @@ -351,18 +389,18 @@ def dp_rank_forward( """ forward = cast( Callable[..., ForwardOutputs], - super().dp_rank_forward, + super().forward, ) - return forward(inputs, checkpoint=checkpoint, no_grad=no_grad) + return forward(inputs, options=options, checkpoint=checkpoint, no_grad=no_grad) - def dp_reduce( + def reduce( self, tensor: torch.Tensor, *, op: dist.ReduceOp.RedOpType = dist.ReduceOp.SUM, ) -> None: """Reduce in place over data-parallel batches, excluding TP/CP replicas.""" - super().dp_reduce(tensor, op=op) + super().reduce(tensor, op=op) def optim_step( self, @@ -381,11 +419,10 @@ def optim_step( norm is clipped independently. If any selected norm is nonfinite, no selected checkpoint is updated. - By default, caller-retained forward graphs do not block the step. ART does - not detach or free those graphs, and backward through one after the step is - unsafe: it may fail PyTorch's version checks or recompute against updated - checkpoint-slot weights. Pass `on_live_graphs="error"` to raise before - mutating any selected slot when a live graph remains on any rank. + Retained forwards use immutable checkpoint versions and may be consumed + after this step within their captured `max_gradient_staleness` policy. + Pass `on_live_graphs="error"` to additionally refuse updates while a + selected checkpoint still has a live forward graph on any rank. """ return super().optim_step( params=params, @@ -395,15 +432,35 @@ def optim_step( ) +from ._commands import ( + RankCallbackResult, + TrainerRankZero, + get_rank_callback_metadata, + rank_callback_leader, + run_rank_callback, + run_rank_callback_stream, +) + __all__ = [ + "RankCallbackResult", + "TrainerRankZero", + "get_rank_callback_metadata", + "rank_callback_leader", + "run_rank_callback", + "run_rank_callback_stream", "AdapterSelection", "AdamParams", "CheckpointManifest", "ForwardInput", + "ForwardOptions", + "ImportanceSamplingGradientCorrection", + "ResolvedForwardOptions", + "resolve_forward_options", "ForwardOutput", "MicroBatch", "MicroBatchStats", "MaterializedCheckpoint", + "ModuleHandle", "materialize_lora", "TopK", "TrainerRank", diff --git a/src/art/trainer_rank/_checkpoint.py b/src/art/trainer_rank/_checkpoint.py index 1c3b33d97..97fd4805a 100644 --- a/src/art/trainer_rank/_checkpoint.py +++ b/src/art/trainer_rank/_checkpoint.py @@ -885,6 +885,9 @@ def prepare_checkpoint_save( raise trainer._slot_state_error( f"Unknown checkpoint on at least one rank: {checkpoint_name!r}" ) + from ._heads import synchronize_head_buffers + + synchronize_head_buffers(trainer, (checkpoint_name,)) config = deepcopy(_validate_save_state(trainer, checkpoint_name)) if any(value != config for value in _gather(config, group)): raise trainer._slot_state_error( @@ -1540,6 +1543,22 @@ def _forward_custom_payload( return PreparedCustomPayload(payload.records, payload.tensors, {}) +def _reserve_generation(trainer: TrainerRank, group: dist.ProcessGroup | None) -> int: + state = trainer._version_state() + generation = ( + max( + ( + state.generation, + *(slot.generation for slot in trainer._checkpoint_slots.values()), + ) + ) + + 1 + ) + # Keep the high-water mark even if creation rolls back or a snapshot is discarded. + state.generation = max(_gather(generation, group)) + return state.generation + + def snapshot_checkpoint(trainer: TrainerRank, source: str, destination: str) -> bool: """Clone one loaded checkpoint into a forward-only resident slot.""" from art.trainer_rank._impl import ( @@ -1574,6 +1593,7 @@ def snapshot_checkpoint(trainer: TrainerRank, source: str, destination: str) -> source, destination, dict(source_slot.config), + source_slot.generation, source_slot.revision, destination_slot is not None, ) @@ -1588,6 +1608,7 @@ def snapshot_checkpoint(trainer: TrainerRank, source: str, destination: str) -> destination_ref = trainer._slot_ref(destination) custom: dict[str, _CustomObject] = {} trackers: list[_CustomTensorTracker] = [] + generation = _reserve_generation(trainer, group) try: for chunk in trainer.runtime.model: for module in chunk.modules(): @@ -1628,6 +1649,7 @@ def snapshot_checkpoint(trainer: TrainerRank, source: str, destination: str) -> custom=custom, custom_payload=_forward_custom_payload(source_slot.custom_payload), snapshot=True, + generation=generation, ) for tracker in trackers: tracker.active = True @@ -1891,6 +1913,7 @@ def load_checkpoint( temporary = f"__art_loading_{uuid.uuid4().hex}" snapshot = _slot_snapshot(trainer) previous = trainer._checkpoint_slots.get(name) + generation = _reserve_generation(trainer, group) try: loaded = _phase( lambda: trainer._load_checkpoint_slot( @@ -1950,6 +1973,9 @@ def load_checkpoint( def commit() -> None: _commit_slot(trainer, temporary, name) staged = trainer._checkpoint_slots.pop(temporary) + staged.generation = generation + # Reload invalidates old graphs, but publication ordering still uses + # this slot's revision independently of the graph generation. staged.revision = 0 if previous is None else previous.revision + 1 trainer._checkpoint_slots[name] = staged diff --git a/src/art/trainer_rank/_commands.py b/src/art/trainer_rank/_commands.py new file mode 100644 index 000000000..315adf5b8 --- /dev/null +++ b/src/art/trainer_rank/_commands.py @@ -0,0 +1,1144 @@ +"""Logical callback leaders and ordered physical-rank participation.""" + +from __future__ import annotations + +import asyncio +from collections.abc import AsyncGenerator, Callable, Generator, Iterator, Sequence +from contextlib import asynccontextmanager, nullcontext +from dataclasses import dataclass, field, replace +from functools import partial +import inspect +from io import BytesIO +from typing import TYPE_CHECKING, Any, Literal, TypeVar, cast +import weakref + +import cloudpickle +import torch +import torch.distributed as dist + +from . import _impl + +if TYPE_CHECKING: + from . import TrainerRank + from ._heads import ModuleHandle + +Mode = Literal["rank", "zero"] +T = TypeVar("T") + + +def _coordinate_call(call: Callable[[], T], *, group: dist.ProcessGroup | None) -> T: + result, error = None, None + try: + result = call() + except Exception as exc: + error = exc + failures = [None if error is None else f"{type(error).__name__}: {error}"] + if dist.is_initialized(): + local = failures[0] + failures = [None] * dist.get_world_size(group) + dist.all_gather_object(failures, local, group=group) + if any(failures): + if error is not None: + raise error + raise RuntimeError(f"Physical trainer preflight failed: {failures}") + return cast(T, result) + + +@dataclass(frozen=True) +class RankCallbackResult: + logical_rank: int | None + value: Any = None + + +@dataclass(frozen=True) +class _Command: + sequence: int + operation: str + args: tuple[Any, ...] + kwargs: dict[str, Any] + grad_enabled: bool + + +def _encode_command(command: _Command) -> bytes: + # Torch owns storage serialization so nested modules and tensor views keep + # their shared storages. Cloudpickle still supports callback-local objects. + stream = BytesIO() + torch.save(command, stream, pickle_module=cloudpickle) + return stream.getvalue() + + +@dataclass(frozen=True) +class _OutputPacket: + packet: Any + cpu: tuple[bool, ...] + managed: bool + + +@dataclass +class _Release: + completed: asyncio.Future[None] + gathered: list[Any] + + +@dataclass +class _State: + sequence: int = 0 + graphs: dict[str, tuple[torch.Tensor, ...]] = field(default_factory=dict) + exports: dict[str, tuple[torch.Tensor, ...]] = field(default_factory=dict) + collector: Any = None + released: set[str] = field(default_factory=set) + iterators: dict[str, Iterator[Any]] = field(default_factory=dict) + batch_inputs: dict[str, Any] = field(default_factory=dict) + control_groups: dict[tuple[int, ...], dist.ProcessGroup] = field( + default_factory=dict + ) + release_group: dist.ProcessGroup | None = None + pending_release: _Release | None = None + release_error: str | None = None + + +def get_rank_callback_metadata(rank: TrainerRank) -> int | None: + """Logical DP index on its leader; None on internal TP/CP participants.""" + if not dist.is_initialized(): + return 0 + from megatron.core import parallel_state as ps + + if ps.get_tensor_model_parallel_rank() or ps.get_context_parallel_rank(): + return None + return rank._dp_rank_and_size()[0] + + +def rank_callback_leader(rank: TrainerRank, *, mode: Mode = "rank") -> bool: + return ( + not dist.is_initialized() or dist.get_rank() == 0 + if mode == "zero" + else get_rank_callback_metadata(rank) is not None + ) + + +class _Executor: + def __init__(self, rank: TrainerRank, mode: Mode) -> None: + from ._tensors import CotangentCollector + + if mode not in ("rank", "zero"): + raise ValueError(f"Unknown callback mode {mode!r}") + self.rank, self.mode = rank, mode + self.group = None + self.distributed = dist.is_initialized() + if self.distributed and mode == "rank": + from megatron.core import parallel_state as ps + + self.group = ps.get_tensor_and_context_parallel_group() + self.members = ( + dist.get_process_group_ranks(self.group) + if self.group is not None + else list(range(dist.get_world_size())) + if self.distributed + else [0] + ) + self.leader = self.members[0] + self.is_leader = not self.distributed or dist.get_rank() == self.leader + self.dp_rank = rank._dp_rank_and_size()[0] + state = getattr(rank, "_rank_command_state", None) + if state is None: + state = _State(collector=CotangentCollector()) + setattr(rank, "_rank_command_state", state) + self.state: _State = state + if self.distributed and state.release_group is None: + state.release_group = dist.new_group(backend="gloo") + if self.distributed and dist.get_backend(self.group) != "gloo": + key = tuple(self.members) + if key not in state.control_groups: + state.control_groups[key] = dist.new_group( + self.members, backend="gloo", use_local_synchronization=True + ) + self.group = state.control_groups[key] + self.iterators: dict[int, Generator[Any, None, None]] = {} + self.stopped = False + + def _start_release(self) -> None: + if self.state.release_error is not None: + raise RuntimeError(self.state.release_error) + if self.state.pending_release is not None: + raise RuntimeError("Previous callback release is still pending") + pending = tuple(self.state.released) + gathered: list[Any] = [pending] + loop = asyncio.get_running_loop() + if self.distributed: + gathered = [None] * dist.get_world_size() + # A finished DP callback must leave its actor loop available while + # another DP callback awaits unrelated async work. Only Gloo runs + # off-thread; dropping graph ownership stays on the actor thread. + completed = loop.run_in_executor( + None, + partial( + dist.all_gather_object, + gathered, + pending, + group=self.state.release_group, + ), + ) + else: + completed = loop.create_future() + completed.set_result(None) + release = self.state.pending_release = _Release(completed, gathered) + # The callback's exception can reach its controller before every DP + # sibling exits. Keep ownership until their matching cleanup completes. + completed.add_done_callback(lambda _: self._finish_release(release)) + + def _finish_release(self, release: _Release) -> None: + if self.state.pending_release is not release or not release.completed.done(): + return + self.state.pending_release = None + try: + release.completed.result() + handles = {handle for values in release.gathered for handle in values} + for handle in handles: + self.state.graphs.pop(handle, None) + self.state.released.difference_update(handles) + except BaseException as error: + self.state.release_error = ( + f"Callback release reconciliation failed: {error}" + ) + release.completed.get_loop().call_exception_handler( + {"message": self.state.release_error, "exception": error} + ) + + async def _join_release(self) -> asyncio.CancelledError | None: + cancelled = None + if release := self.state.pending_release: + while True: + try: + await asyncio.shield(release.completed) + break + except asyncio.CancelledError as error: + if release.completed.cancelled(): + break + cancelled = error + except Exception: + break # Finalization records and reports the terminal error. + # A done callback may still be queued; ownership must be settled + # synchronously before this actor can execute another command. + self._finish_release(release) + if self.state.release_error is not None: + raise RuntimeError(self.state.release_error) + return cancelled + + async def reconcile_releases( + self, *, defer_cancellation: bool = False + ) -> asyncio.CancelledError | None: + """Join prior cleanup before entering another all-rank boundary.""" + cancelled = await self._join_release() + self._start_release() + cancelled = await self._join_release() or cancelled + if cancelled is not None and not defer_cancellation: + raise cancelled + return cancelled + + @asynccontextmanager + async def release_on_exit(self) -> AsyncGenerator[None, None]: + try: + yield + except GeneratorExit: + await self.reconcile_releases() + raise + except BaseException: + # FIRST_EXCEPTION controllers need this error to cancel siblings + # which may still be awaiting user work. Their cleanup joins ours. + self._start_release() + raise + else: + await self.reconcile_releases() + + def _broadcast(self, command: _Command | None) -> _Command: + if not self.distributed or len(self.members) == 1: + assert command is not None + return command + payload = None + if command is not None: + try: + payload = _encode_command(command) + except Exception as exc: + command = _Command( + command.sequence, + "error", + (f"Command serialization failed: {exc}",), + {}, + False, + ) + payload = _encode_command(command) + objects: list[Any] = [payload] + dist.broadcast_object_list(objects, src=self.leader, group=self.group) + return self._decode(objects[0]) + + def _decode(self, payload: Any) -> _Command: + error, decoded = None, None + try: + # A leader's CUDA ordinal is not a peer's local model device. Decode + # storages on CPU; native command handlers place them on their rank. + decoded = torch.load( + BytesIO(payload), map_location="cpu", weights_only=False + ) + except Exception as exc: + error = f"Command deserialization failed: {exc}" + failures = self._gather(error) + if any(failures): + return _Command(self.state.sequence, "error", (str(failures),), {}, False) + assert isinstance(decoded, _Command) + return decoded + + def invoke(self, operation: str, *args: Any, **kwargs: Any) -> Any: + if self.stopped: + raise RuntimeError("Trainer callback session has stopped") + self.state.sequence += 1 + command = _Command( + self.state.sequence, operation, args, kwargs, torch.is_grad_enabled() + ) + command = self._broadcast(command) + return self._execute(command) + + async def serve(self) -> None: + cancelled = None + try: + while True: + objects: list[Any] = [None] + # Only the Gloo CPU receive leaves the actor thread. Decoding, + # model commands and iterator cleanup retain its CUDA context. + # Use a Future, not a child Task that asyncio.run shutdown could + # cancel independently before the serving task joins it. + received = asyncio.get_running_loop().run_in_executor( + None, + partial( + dist.broadcast_object_list, + objects, + src=self.leader, + group=self.group, + ), + ) + while True: + try: + await asyncio.shield(received) + break + except asyncio.CancelledError as error: + # An abandoned receive could consume the next callback's + # command. Drain this session through its leader stop. + cancelled = error + command = self._decode(objects[0]) + self.state.sequence = max(self.state.sequence, command.sequence) + if command.operation == "stop": + if cancelled is not None: + raise cancelled + return + try: + self._execute(command) + except Exception: + # The leader receives the same coordinated error and chooses + # whether to catch it, continue, or stop the callback. + pass + finally: + self.stopped = True + self._close_iterators() + + def stop(self) -> None: + if self.stopped: + return + self.stopped = True + try: + self.state.sequence += 1 + self._broadcast(_Command(self.state.sequence, "stop", (), {}, False)) + finally: + self._close_iterators() + + def _close_iterators(self) -> None: + iterators, self.iterators = self.iterators, {} + for iterator in iterators.values(): + iterator.close() + + def _gather(self, value: Any) -> list[Any]: + if not self.distributed or len(self.members) == 1: + return [value] + gathered: list[Any] = [None] * len(self.members) + dist.all_gather_object(gathered, value, group=self.group) + return gathered + + def _execute(self, command: _Command) -> Any: + result, error = None, None + try: + with torch.set_grad_enabled(command.grad_enabled): + result = self._dispatch(command) + except Exception as exc: + error = exc + errors = self._gather( + None if error is None else f"{type(error).__name__}: {error}" + ) + if any(errors): + self.state.graphs.pop( + f"{self.mode}:{command.sequence}:dp:{self.dp_rank}", None + ) + if error is not None: + raise error + raise RuntimeError( + f"Physical trainer command {command.operation!r} failed: {errors}" + ) + if command.operation in ("forward", "next", "batches_next"): + value = ( + result if get_rank_callback_metadata(self.rank) is not None else None + ) + try: + return self._gather_outputs(value) + except BaseException: + self.state.graphs.pop( + f"{self.mode}:{command.sequence}:dp:{self.dp_rank}", None + ) + raise + return result + + def _gather_outputs(self, value: Any) -> list[Any] | None: + if not self.distributed or len(self.members) == 1: + return [value] + + def admit_serialization() -> None: + from ._tensors import flatten_tensors + + tensors, _ = flatten_tensors(value) + # Pickling creates storage bytes before gather admission can sample + # their size. Reserve tensor storage, a copy, and per-leaf metadata. + required = 2 * sum(t.numel() * t.element_size() for t in tensors) + required += 4096 * (len(tensors) + 1) + if required > self._available_host_memory(): + raise MemoryError( + f"Trainer output serialization requires {required} CPU bytes" + ) + + self._coordinated_preflight(admit_serialization) + payload = self._coordinated_preflight(lambda: cloudpickle.dumps(value)) + sizes = self._gather(len(payload)) + + def admit() -> None: + available = self._available_host_memory() + # Gloo gather_object pads every sender to the largest serialized + # payload. Include receive storage and unpickling copies on leader. + padded = max(sizes) + 1024 + required = 2 * padded + if self.is_leader: + required += 2 * len(sizes) * padded + sum(sizes) + if required > available: + raise MemoryError( + f"Trainer output transfer requires {required} CPU bytes, " + f"but the per-process shared-host budget has {available}" + ) + + self._coordinated_preflight(admit) + values: list[Any] | None = ( + [None] * len(self.members) if self.is_leader else None + ) + dist.gather_object(payload, values, dst=self.leader, group=self.group) + + def decode() -> list[Any] | None: + if values is None: + return None + decoded = [cloudpickle.loads(item) for item in values] + return [item for item in decoded if item is not None] + + return self._coordinated_preflight(decode) + + def _available_host_memory(self) -> int: + from ._memory_policy import host_memory_budget, local_rank_count + + if hasattr(self.rank, "_available_cpu_memory_bytes"): + return self.rank._available_cpu_memory_bytes() + return host_memory_budget( + local_world_size=local_rank_count( + world_size=dist.get_world_size() if self.distributed else 1 + ) + ).available_bytes + + def _packet(self, tree: Any, sequence: int) -> Any: + from ._tensors import ManagedTensor, detach_tree, flatten_tensors + + handle = f"{self.mode}:{sequence}:dp:{self.dp_rank}" + tensors, _ = flatten_tensors(tree) + if any(tensor.requires_grad for tensor in tensors): + self.state.graphs[handle] = tuple(tensors) + if get_rank_callback_metadata(self.rank) is None: + return None + required = sum(tensor.numel() * tensor.element_size() for tensor in tensors) + if required > self._available_host_memory(): + raise MemoryError(f"Trainer output snapshot requires {required} CPU bytes") + return _OutputPacket( + detach_tree(handle, tree, device="cpu"), + tuple(tensor.device.type == "cpu" for tensor in tensors), + any(isinstance(tensor, ManagedTensor) for tensor in tensors), + ) + + def _dispatch(self, command: _Command) -> Any: + op, args, kwargs = command.operation, command.args, command.kwargs + if op == "error": + raise RuntimeError(args[0]) + if op == "forward": + return self._packet(self.rank.forward(*args, **kwargs), command.sequence) + if op == "batches": + self.iterators[command.sequence] = cast( + Generator[Any, None, None], self.rank.forward_batches(*args, **kwargs) + ) + return command.sequence + if op == "batches_open": + handle = f"{self.mode}:batches:{command.sequence}" + self.state.iterators[handle] = self.rank.forward_batches(*args, **kwargs) + return handle + if op in ("next", "batches_next"): + iterator = ( + self.iterators[args[0]] + if op == "next" + else self.state.iterators[args[0]] + ) + batch = next(iterator, None) + if batch is None: + return (None, None) + return replace(batch, inputs=[], outputs=[]), self._packet( + batch.outputs, command.sequence + ) + if op in ("close", "batches_close"): + iterator = ( + self.iterators.pop(args[0], None) + if op == "close" + else self.state.iterators.pop(args[0], None) + ) + if op == "batches_close": + self.state.batch_inputs.pop(args[0], None) + if iterator is not None: + cast(Generator[Any, None, None], iterator).close() + return None + if op == "release": + for handle in args[0]: + self.state.graphs.pop(handle, None) + return None + if op == "backward": + return self._backward(*args, **kwargs) + if op == "head": + from ._heads import execute_head_operation, synchronize_head_buffers + + if self.mode == "zero" and args[0] == "head_export": + synchronize_head_buffers(self.rank) + return execute_head_operation( + self.rank, *args, coordinate=self._coordinated_preflight, **kwargs + ) + if op == "reduce_value": + tensor = args[0].to(self.rank.device) + self.rank.reduce(tensor, **kwargs) + return tensor + if op == "pop_checkpoint_context": + + def validate() -> None: + ref = self.rank._slot_ref(args[0]) + if not self.rank._slot_stack or self.rank._slot_stack[-1] != ref: + raise RuntimeError( + "Pushed checkpoint stack changed before context exit" + ) + + self._coordinated_preflight(validate) + return self.rank.pop_checkpoint() + return getattr(self.rank, op)(*args, **kwargs) + + def _coordinated_preflight(self, validate: Callable[[], T]) -> T: + return _coordinate_call(validate, group=self.group) + + def _backward(self, packets: Sequence[Any], *, retain_graph: bool) -> None: + outputs, gradients, handles, head_gradients = [], [], [], [] + + def prepare() -> None: + for packet in packets: + if packet.handle.startswith("head:"): + from ._heads import head_gradient_targets + + targets = head_gradient_targets( + self.rank, packet, materialize=False + ) + # The global caller's head loss occurs once; DP SUM must + # count it once while TP/CP retain their replicated contract. + if self.mode != "zero" or self.dp_rank == 0: + head_gradients.extend(targets) + continue + tensors = self.state.graphs.get(packet.handle) + if tensors is None: + if packet.handle.endswith(f":dp:{self.dp_rank}"): + raise ValueError( + f"Unknown or released forward {packet.handle!r}" + ) + continue # Other DP rank owns this root. + if len(tensors) != len(packet.gradients): + raise ValueError( + "Cotangent packet does not match its physical forward" + ) + handles.append(packet.handle) + for tensor, gradient in zip(tensors, packet.gradients, strict=True): + if gradient is not None: + if gradient.shape != tensor.shape or not tensor.requires_grad: + raise ValueError( + "Cotangent does not match its physical forward" + ) + outputs.append(tensor) + gradients.append( + gradient.to(device=tensor.device, dtype=tensor.dtype) + ) + + self._coordinated_preflight(prepare) + transaction = ( + self.rank._gradient_transaction(before_commit=self._coordinated_preflight) + if hasattr(self.rank, "_gradient_transaction") + else nullcontext() + ) + try: + with transaction: + + def stage_heads() -> None: + gradient = None + try: + for version, maximum, parameter, source in head_gradients: + gradient = source.to( + device=parameter.device, dtype=parameter.dtype + ) + self.rank._commit_versioned_gradients( + ((version, maximum, parameter, gradient),) + ) + gradient = None + finally: + gradient = None + + # Every peer finishes (or rolls back) copies/staging before any + # participant can enter model backward's TP/CP collectives. + if any(packet.handle.startswith("head:") for packet in packets): + self._coordinated_preflight(stage_heads) + if hasattr(self.rank, "_forward_cotangent_collector"): + inner: list[tuple[str, Any]] = [] + + def collect() -> None: + if outputs: + packets = self.rank._forward_cotangent_collector().backward( + outputs, gradients, retain_graph=retain_graph + ) + inner.extend( + (packet.handle, packet.gradients) for packet in packets + ) + self.rank._forward_graph_cache().validate_many(inner) + + self._coordinated_preflight(collect) + self.rank._forward_graph_cache().backward_many( + inner, + retain_graph=retain_graph, + coordinate=lambda call: _coordinate_call( + call, group=self.rank._forward_memory_group() + ), + ) + elif outputs: + torch.autograd.backward( + outputs, gradients, retain_graph=retain_graph + ) + finally: + if not retain_graph: + for handle in handles: + self.state.graphs.pop(handle, None) + + +_COLLECTIVE_METHODS = frozenset( + { + "zero_grad", + "load_checkpoint", + "snapshot_checkpoint", + "_push_checkpoint_sync", + "prefetch_checkpoints", + "pop_checkpoint", + "save_checkpoint", + "prepare_checkpoint_save", + "finish_checkpoint_save", + "abort_checkpoint_save", + "export_lora", + "optim_step", + } +) + + +class _PushedCheckpoint(_impl.PushedCheckpoint): + def _pop(self) -> None: + if not self._entered: + return + cast("_RankView", self._trainer)._invoke("pop_checkpoint_context", self._path) + self._entered = False + self._closed = True + + +class _RankView: + def __init__(self, executor: _Executor) -> None: + self._executor = executor + self._rank = executor.rank + self._transport_handles: list[str] | None = None + + @property + def device(self) -> torch.device: + return self._rank.device + + @property + def hidden_size(self) -> int: + return self._rank.hidden_size + + def __getattribute__(self, name: str) -> Any: + if name in _COLLECTIVE_METHODS: + return lambda *args, **kwargs: self._invoke(name, *args, **kwargs) + return object.__getattribute__(self, name) + + def _flush_heads(self) -> None: + if hasattr(self._rank, "_logical_head_handles"): + from ._heads import flush_logical_heads + + flush_logical_heads(self) + + def _refresh_heads(self) -> None: + if hasattr(self._rank, "_logical_head_handles"): + from ._heads import refresh_logical_heads + + refresh_logical_heads(self) + + def _invoke(self, operation: str, *args: Any, **kwargs: Any) -> Any: + if self._executor.stopped: + raise RuntimeError("Trainer callback session has stopped") + released = self._executor.state.released + handles = tuple( + handle + for handle in released + if handle.startswith(f"{self._executor.mode}:") + ) + if handles: + self._executor.invoke("release", handles) + released.difference_update(handles) + if operation != "head": + self._flush_heads() + result = self._executor.invoke(operation, *args, **kwargs) + if operation in { + "optim_step", + "load_checkpoint", + "pop_checkpoint", + "pop_checkpoint_context", + "_push_checkpoint_sync", + }: + self._refresh_heads() + return result + + def push_checkpoint(self, checkpoint: Any) -> _impl.PushedCheckpoint: + path, directory = self._rank._checkpoint_source(checkpoint) + return _PushedCheckpoint(cast("TrainerRank", self), path, directory) + + def forward(self, inputs: _impl.ForwardInputs, **kwargs: Any) -> Any: + materialized = _impl._materialize(inputs) + if self._executor.mode == "rank": + if hasattr(self._rank, "_capture_forward_options"): + materialized = self._rank._capture_forward_options( + materialized, kwargs.get("options") + ) + (packet,) = self._invoke("forward", materialized, **kwargs) + return self._attach(self._place_outputs([(packet, materialized)])[0]) + single = isinstance(materialized, _impl.ForwardInput) + roots = [materialized] if single else materialized + outputs: list[Any] = [] + for batch in self.forward_batches(roots, **kwargs): + outputs.extend(batch.outputs) + return ( + outputs[0] if single else _impl._rebuild_forward_tree(materialized, outputs) + ) + + def forward_batches(self, inputs: Any, **kwargs: Any) -> Iterator[_impl.MicroBatch]: + items = self._prepare_batches(inputs, kwargs) + return self._iterate_batches(items, kwargs) + + def _prepare_batches(self, inputs: Any, kwargs: dict[str, Any]) -> Any: + items = [_impl._materialize(item) for item in inputs] + if kwargs.get("no_grad") is None: + kwargs["no_grad"] = not torch.is_grad_enabled() + if hasattr(self._rank, "_capture_forward_options"): + items = self._rank._capture_forward_options(items, kwargs.get("options")) + if self._executor.mode == "zero": + kwargs["yield_empty"] = True + return items + + def open_forward_batches(self, inputs: Any, **kwargs: Any) -> str: + """Bind a lazily pulled iterator that survives callback boundaries.""" + items = self._prepare_batches(inputs, kwargs) + if kwargs.get("checkpoint", _impl.Unset) is _impl.Unset: + stack = getattr(self._rank, "_slot_stack", ()) + selected = ( + stack[-1] if stack else getattr(self._rank, "_default_slot_ref", None) + ) + kwargs["checkpoint"] = None if selected is None else selected.name + handle = self._invoke("batches_open", items, **kwargs) + self._executor.state.batch_inputs[handle] = items + return handle + + def next_forward_batch(self, handle: str) -> _impl.MicroBatch | None: + items = self._executor.state.batch_inputs.get(handle) + if items is None: + return None + try: + batch = self._combine_wave(self._invoke("batches_next", handle), items) + except BaseException: + self.close_forward_batches(handle) + raise + if batch is None: + self.close_forward_batches(handle) + return batch + + def close_forward_batches(self, handle: str) -> None: + self._invoke("batches_close", handle) + + def _iterate_batches( + self, items: Any, kwargs: dict[str, Any] + ) -> Iterator[_impl.MicroBatch]: + identifier = self._invoke("batches", items, **kwargs) + try: + while True: + batch = self._combine_wave(self._invoke("next", identifier), items) + if batch is None: + return + yield batch + finally: + # This iterator belongs to its creating callback, even if a retained + # traceback delays its finalizer until a later callback is serving. + if not self._executor.stopped: + self._invoke("close", identifier) + + def _combine_wave(self, wave: Any, items: Any) -> _impl.MicroBatch | None: + if all(batch is None for batch, _ in wave): + return None + if any(batch is None for batch, _ in wave): + raise RuntimeError("Physical forward iterators ended on different waves") + packets = self._place_outputs( + [ + (packet, [items[index] for index in batch.indices]) + for batch, packet in wave + ] + ) + batches = [ + replace( + batch, + inputs=[items[index] for index in batch.indices], + outputs=self._attach(packet), + ) + for (batch, _), packet in zip(wave, packets, strict=True) + ] + if self._executor.mode == "rank": + return batches[0] + rows = sorted( + (index, output) + for batch in batches + for index, output in zip(batch.indices, batch.outputs, strict=True) + ) + batch = batches[0] + return replace( + batch, + inputs=[items[index] for index, _ in rows], + outputs=[output for _, output in rows], + indices=[index for index, _ in rows], + stats=replace(batch.stats, local_count=len(rows)), + ) + + def _place_outputs( + self, outputs: Sequence[tuple[_OutputPacket, Any]] + ) -> list[_OutputPacket]: + if self._transport_handles is not None: + self._transport_handles.extend( + output.packet.handle for output, _ in outputs + ) + return [ + replace(output, cpu=(True,) * len(output.cpu), managed=True) + for output, _ in outputs + ] + from ._memory_policy import choose_output_placements + from ._options import resolve_forward_options + from ._tensors import flatten_tensors, unflatten_tensors + + costs: list[tuple[int, Literal["auto", "model", "cpu"]]] = [] + for output, inputs in outputs: + policies: dict[int, Literal["auto", "model", "cpu"]] = {} + + def visit(request: Any, result: Any) -> None: + if isinstance(request, _impl.ForwardInput): + policy = resolve_forward_options( + input=request.options + ).output_device + tensors, _ = flatten_tensors(result) + for tensor in tensors: + policies[id(tensor)] = policy + else: + for child, value in zip(request, result, strict=True): + visit(child, value) + + visit(inputs, unflatten_tensors(output.packet.spec, output.packet.tensors)) + costs.extend( + ( + tensor.numel() * tensor.element_size(), + "cpu" if cpu else policies[id(tensor)], + ) + for tensor, cpu in zip(output.packet.tensors, output.cpu, strict=True) + ) + available = ( + self._rank._available_memory_bytes() + if hasattr(self._rank, "_available_memory_bytes") + else 1 << 60 + ) + if hasattr(self._rank, "_pending_backward_memory"): + available -= sum(self._rank._pending_backward_memory()) + try: + placements = iter( + choose_output_placements(costs, gpu_available_bytes=available) + ) + except BaseException: + self._invoke( + "release", tuple(output.packet.handle for output, _ in outputs) + ) + raise + result = [] + for output, _ in outputs: + cpu = tuple(next(placements) == "cpu" for _ in output.cpu) + result.append( + replace(output, cpu=cpu, managed=output.managed or cpu != output.cpu) + ) + return result + + def _attach(self, output: _OutputPacket) -> Any: + packet = replace( + output.packet, + tensors=tuple( + tensor if cpu else tensor.to(self.device) + for tensor, cpu in zip(output.packet.tensors, output.cpu, strict=True) + ), + ) + state = self._executor.state + state_ref, handle = weakref.ref(state), packet.handle + + def release() -> None: + if owner := state_ref(): + owner.released.add(handle) + + with torch.enable_grad(): + return state.collector.attach( + packet, managed=output.managed, on_release=release + ) + + def backward( + self, loss: Any, gradient: Any = None, *, retain_graph: bool = False + ) -> None: + packets = self._executor.state.collector.backward( + loss, gradient, retain_graph=retain_graph + ) + self._submit_backward(packets, retain_graph=retain_graph) + + def _submit_backward(self, packets: Sequence[Any], *, retain_graph: bool) -> None: + self._invoke( + "backward", + tuple( + replace( + packet, + gradients=tuple( + None if value is None else value.cpu() + for value in packet.gradients + ), + ) + for packet in packets + ), + retain_graph=retain_graph, + ) + + def export_forward(self, tree: Any) -> Any: + from ._tensors import detach_tree, flatten_tensors + + state = self._executor.state + state.sequence += 1 + handle = f"client:{state.sequence}:dp:{self._executor.dp_rank}" + tensors, _ = flatten_tensors(tree) + if self._transport_handles is not None: + # Gathered CPU storage remains owned by exports until backward or + # release, and fresh host/cgroup headroom excludes that live use. + # Reserve the additional independent reply snapshot before copying. + required = sum(tensor.numel() * tensor.element_size() for tensor in tensors) + if required > self._executor._available_host_memory(): + raise MemoryError( + f"Trainer export snapshot requires {required} CPU bytes" + ) + packet = detach_tree(handle, tree, device="cpu") + if any(tensor.requires_grad for tensor in tensors): + state.exports[handle] = tuple( + reply + if self._transport_handles is not None and not tensor.requires_grad + else tensor + for tensor, reply in zip(tensors, packet.tensors, strict=True) + ) + return packet + + def release_forward(self, handles: Sequence[str]) -> None: + """Idempotently release exported client graphs after caller collection.""" + for handle in handles: + self._executor.state.exports.pop(handle, None) + self._invoke("release", ()) + + def backward_packets( + self, packets: Sequence[Any], *, retain_graph: bool = False + ) -> None: + handles = tuple( + packet.handle for packet in packets if not packet.handle.startswith("head:") + ) + try: + outputs, gradients, heads = [], [], [] + for packet in packets: + if packet.handle.startswith("head:"): + heads.append(packet) + continue + tensors = self._executor.state.exports[packet.handle] + if len(tensors) != len(packet.gradients): + raise ValueError( + "Client cotangent packet does not match its forward" + ) + for tensor, gradient in zip(tensors, packet.gradients, strict=True): + if gradient is not None: + outputs.append(tensor) + gradients.append( + gradient.to(device=tensor.device, dtype=tensor.dtype) + ) + model = ( + self._executor.state.collector.backward( + outputs, gradients, retain_graph=retain_graph + ) + if outputs + else () + ) + self._submit_backward((*model, *heads), retain_graph=retain_graph) + finally: + if not retain_graph: + for handle in handles: + self._executor.state.exports.pop(handle, None) + + def module( + self, name: str, factory: Callable[[], Any], **kwargs: Any + ) -> ModuleHandle: + return cast("ModuleHandle", self._register("module", name, factory, **kwargs)) + + def parameter( + self, name: str, factory: Callable[[], Any], **kwargs: Any + ) -> torch.nn.Parameter: + return cast( + torch.nn.Parameter, self._register("parameter", name, factory, **kwargs) + ) + + def buffer( + self, name: str, factory: Callable[[], Any], **kwargs: Any + ) -> torch.Tensor: + return cast(torch.Tensor, self._register("buffer", name, factory, **kwargs)) + + def _register( + self, + kind: Literal["module", "parameter", "buffer"], + name: str, + factory: Callable[[], Any], + **kwargs: Any, + ) -> Any: + from ._heads import logical_register_head + + return logical_register_head(self, kind, name, factory, **kwargs) + + def last_forward_telemetry(self) -> dict[str, Any]: + return self._rank.last_forward_telemetry() + + +class TrainerRankZero(_RankView): + """Single callback view over every physical trainer rank, without reduce.""" + + +def _view(executor: _Executor) -> _RankView: + from . import TrainerRank + + class LogicalTrainerRank(_RankView, TrainerRank): + def reduce(self, tensor: torch.Tensor, **kwargs: Any) -> None: + result = self._invoke("reduce_value", tensor.detach().cpu(), **kwargs) + tensor.copy_(result.to(tensor.device)) + + return ( + TrainerRankZero(executor) + if executor.mode == "zero" + else LogicalTrainerRank(executor) + ) + + +async def run_rank_callback( + rank: TrainerRank, callback: Callable[[Any], Any], *, mode: Mode = "rank" +) -> RankCallbackResult: + executor = _Executor(rank, mode) + cancelled = await executor.reconcile_releases(defer_cancellation=True) + async with executor.release_on_exit(): + if not executor.is_leader: + await executor.serve() + if cancelled is not None: + raise cancelled + return RankCallbackResult(None) + view = _view(executor) + try: + if cancelled is not None: + raise cancelled + view._refresh_heads() + result = callback(view) + if inspect.isawaitable(result): + result = await result + if inspect.isgenerator(result) or inspect.isasyncgen(result): + raise TypeError("Use run_rank_callback_stream for generator callbacks") + return RankCallbackResult(0 if mode == "zero" else executor.dp_rank, result) + finally: + try: + view._flush_heads() + finally: + executor.stop() + + +async def run_rank_callback_stream( + rank: TrainerRank, callback: Callable[[Any], Any], *, mode: Mode = "rank" +) -> AsyncGenerator[RankCallbackResult, Any]: + """Drive one user generator on each logical leader, forwarding sends.""" + executor = _Executor(rank, mode) + cancelled = await executor.reconcile_releases(defer_cancellation=True) + if not executor.is_leader: + async with executor.release_on_exit(): + await executor.serve() + if cancelled is not None: + raise cancelled + yield RankCallbackResult(None) + return + iterator = value = None + view = _view(executor) + async with executor.release_on_exit(): + try: + if cancelled is not None: + raise cancelled + view._refresh_heads() + iterator = callback(view) + if inspect.isawaitable(iterator): + iterator = await iterator + if not (inspect.isgenerator(iterator) or inspect.isasyncgen(iterator)): + raise TypeError("Stream callback must return a generator") + sent = None + while True: + try: + value = ( + await iterator.asend(sent) + if inspect.isasyncgen(iterator) + else iterator.send(sent) + ) + except (StopIteration, StopAsyncIteration): + return + sent = yield RankCallbackResult( + 0 if mode == "zero" else executor.dp_rank, value + ) + finally: + try: + if inspect.isasyncgen(iterator): + await iterator.aclose() + elif inspect.isgenerator(iterator): + iterator.close() + view._flush_heads() + finally: + value = None + executor.stop() diff --git a/src/art/trainer_rank/_corrections.py b/src/art/trainer_rank/_corrections.py new file mode 100644 index 000000000..7feacb602 --- /dev/null +++ b/src/art/trainer_rank/_corrections.py @@ -0,0 +1,332 @@ +"""Numerical helpers for explicitly requested selected-token correction.""" + +from __future__ import annotations + +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +import math +from typing import Any, Iterator + +import torch + +from ._options import ImportanceSamplingGradientCorrection, ResolvedForwardOptions + + +def importance_weights( + original_logprobs: torch.Tensor, + current_logprobs: torch.Tensor, + correction: ImportanceSamplingGradientCorrection, +) -> torch.Tensor: + """Detached clipped p_current/p_original, without low-precision overflow. + + The caller must supply probabilities for identical events (token IDs and + contexts). Original probabilities must be positive; current zero probability + is allowed. Missing support and undefined 0/0 ratios raise rather than being + silently replaced. Computation uses at least float32 and promotes to float64 + for float64 inputs or clipping bounds outside the float32 normal range. + """ + if original_logprobs.shape != current_logprobs.shape: + raise ValueError("correction logprob shapes must match exactly") + if original_logprobs.device != current_logprobs.device: + raise ValueError("correction logprobs must be on the same device") + if ( + not original_logprobs.is_floating_point() + or not current_logprobs.is_floating_point() + ): + raise TypeError("correction logprobs must be floating-point tensors") + if not bool(torch.isfinite(original_logprobs).all()): + raise ValueError("original logprobs must be finite (positive sampling support)") + if bool((torch.isnan(current_logprobs) | torch.isposinf(current_logprobs)).any()): + raise ValueError("current logprobs must be finite or negative infinity") + dtype = ( + torch.float64 + if torch.float64 in (original_logprobs.dtype, current_logprobs.dtype) + or correction.clip_high > torch.finfo(torch.float32).max + or any( + 0 < bound < torch.finfo(torch.float32).tiny + for bound in (correction.clip_low, correction.clip_high) + ) + else torch.float32 + ) + log_ratio = current_logprobs.detach().to(dtype) - original_logprobs.detach().to( + dtype + ) + if correction.clip_high == 0: + return torch.zeros_like(log_ratio) + log_low = math.log(correction.clip_low) if correction.clip_low else -math.inf + return ( + log_ratio.clamp(log_low, math.log(correction.clip_high)) + .exp() + .clamp(correction.clip_low, correction.clip_high) + ) + + +def correct_logprob_cotangent( + cotangent: torch.Tensor, + *, + original_logprobs: torch.Tensor, + current_logprobs: torch.Tensor | None, + correction: ImportanceSamplingGradientCorrection, + original_tokens: torch.Tensor | None = None, + current_tokens: torch.Tensor | None = None, +) -> torch.Tensor: + """Apply explicitly requested importance weights to aligned logprob outputs. + + For top-k, both token tensors are required by the caller's output contract. + The IDs must match elementwise: independently recomputing top-k can change + their identity/order. Values are full-vocabulary logprobs, not probabilities + renormalized over top-k. No forward is performed here. + """ + if cotangent.shape != original_logprobs.shape: + raise ValueError("cotangent and correction logprob shapes must match exactly") + active = cotangent != 0 + if not bool(active.any()): + return cotangent + if current_logprobs is None: + if correction.policy == "always": + raise RuntimeError( + "importance sampling correction requires current logprobs" + ) + return cotangent + if (original_tokens is None) != (current_tokens is None): + raise ValueError("correction requires both original and current token IDs") + if original_tokens is not None and current_tokens is not None: + if ( + original_tokens.shape != original_logprobs.shape + or current_tokens.shape != current_logprobs.shape + ): + raise ValueError("correction token IDs must match logprob shapes") + if not torch.equal(original_tokens, current_tokens): + raise ValueError( + "correction must compare the same token IDs in the same order" + ) + if current_logprobs.shape != original_logprobs.shape: + raise ValueError("correction logprob shapes must match exactly") + if current_logprobs.device != original_logprobs.device: + raise ValueError("correction logprobs must be on the same device") + selected = active.to(original_logprobs.device) + weights = importance_weights( + original_logprobs[selected], current_logprobs[selected], correction + ) + corrected = cotangent.clone() + corrected[active] = (cotangent[active] * weights.to(cotangent.device)).to( + cotangent.dtype + ) + return corrected + + +@dataclass(frozen=True) +class _OutputCorrection: + index: int + kind: str + original_logprobs: torch.Tensor + token_index: int | None = None + original_tokens: torch.Tensor | None = None + logits_index: int | None = None + + +@dataclass(frozen=True) +class ForwardCorrectionContext: + """Correction metadata and owned CPU copies, independent of physical graphs. + + Current tensors must come from the original inputs/contexts, with the same + flattened output layout. Correct only stale forwards; exact original-weight + replay itself does not produce current probabilities. Validate physical + current-weight replay even when corrections are disabled. + """ + + output_count: int + correction: ImportanceSamplingGradientCorrection | None + outputs: tuple[_OutputCorrection, ...] + + @property + def tensors(self) -> tuple[torch.Tensor, ...]: + return tuple( + tensor + for output in self.outputs + for tensor in (output.original_logprobs, output.original_tokens) + if tensor is not None + ) + + def requires_current(self, gradients: Sequence[torch.Tensor | None]) -> bool: + """Whether active eligible outputs need current data before mutation.""" + if len(gradients) != self.output_count: + raise ValueError( + "correction cotangents must match the captured output count" + ) + for output in self.outputs: + gradient = gradients[output.index] + if ( + gradient is not None + and gradient.shape != output.original_logprobs.shape + ): + raise ValueError( + "correction cotangent shape must match the original output" + ) + return ( + self.correction is not None + and self.correction.policy == "always" + and any( + (gradient := gradients[output.index]) is not None + and bool((gradient != 0).any()) + for output in self.outputs + ) + ) + + def correct( + self, + gradients: Sequence[torch.Tensor | None], + current_tensors: Sequence[torch.Tensor] | None = None, + ) -> tuple[torch.Tensor | None, ...]: + """Stage corrected cotangents without mutating gradients or model state.""" + self.requires_current(gradients) + if current_tensors is not None and len(current_tensors) != self.output_count: + raise ValueError("current tensors must match the captured output count") + corrected = list(gradients) + if self.correction is None: + return tuple(corrected) + for output in self.outputs: + gradient = gradients[output.index] + original = output.original_logprobs + if gradient is None or not bool((gradient != 0).any()): + continue + current = ( + None + if current_tensors is None + else current_tensors[output.index].detach() + ) + if current is not None and output.original_tokens is not None: + assert current_tensors is not None and output.token_index is not None + tokens = current_tensors[output.token_index] + original_tokens = output.original_tokens.to(tokens.device) + if tokens.shape != original_tokens.shape: + raise ValueError( + "current top-k token shape must match original top-k" + ) + if not torch.equal(tokens, original_tokens): + # A changed top-k ordering can still contain every old ID. + sorted_tokens, order = tokens.sort(dim=-1) + positions = torch.searchsorted( + sorted_tokens.contiguous(), original_tokens.contiguous() + ).clamp_max(tokens.shape[-1] - 1) + matched = sorted_tokens.gather(-1, positions) == original_tokens + if bool((matched | (gradient == 0).to(matched.device)).all()): + current = current.gather(-1, order.gather(-1, positions)) + elif output.logits_index is not None: + logits = current_tensors[output.logits_index].detach() + dtype = ( + torch.float64 + if logits.dtype == torch.float64 + else torch.float32 + ) + logits = logits.to(dtype) + current = logits.gather( + -1, output.original_tokens.to(logits.device) + ) - logits.logsumexp(-1, keepdim=True) + else: + # The new top-k lacks original events: no ratio is available. + current = None + corrected[output.index] = correct_logprob_cotangent( + gradient, + original_logprobs=original + if current is None + else original.to(current.device), + current_logprobs=current, + correction=self.correction, + ) + return tuple(corrected) + + def validate_replay( + self, + gradients: Sequence[torch.Tensor | None], + current_tensors: Sequence[torch.Tensor], + ) -> None: + """Require active top-k cotangents to address the same replayed events. + + Ratio evaluation on a separate current forward can realign top-k IDs. + Physical current-weight replay cannot feed original-position cotangents + to a Jacobian whose selected token at that position has changed. + """ + self.requires_current(gradients) + if len(current_tensors) != self.output_count: + raise ValueError("current tensors must match the captured output count") + for output in self.outputs: + gradient = gradients[output.index] + if output.original_tokens is None or gradient is None: + continue + active = gradient != 0 + if not bool(active.any()): + continue + assert output.token_index is not None + tokens = current_tensors[output.token_index] + original = output.original_tokens.to(tokens.device) + if tokens.shape != original.shape or bool( + ((tokens != original) & active.to(tokens.device)).any() + ): + raise RuntimeError( + "current replay changed active top-k token identities; " + "replay the original weights instead" + ) + + +def capture_forward_corrections( + outputs: Any, + tensors: Sequence[torch.Tensor], + options: ResolvedForwardOptions, +) -> ForwardCorrectionContext: + """Map a ForwardOutput tree to the caller's deduplicated flat tensor layout.""" + from ._impl import ForwardOutput + + def leaves(value: Any) -> Iterator[Any]: + if isinstance(value, ForwardOutput): + yield value + elif isinstance(value, Mapping): + for item in value.values(): + yield from leaves(item) + elif isinstance(value, (tuple, list)): + for item in value: + yield from leaves(item) + else: + raise TypeError("correction capture requires a ForwardOutput tree") + + corrections = options.stale_gradient_corrections + correction = corrections[0] if corrections else None + indices = {id(tensor): index for index, tensor in enumerate(tensors)} + if len(indices) != len(tensors): + raise ValueError("correction capture requires deduplicated flat tensors") + entries: dict[int, _OutputCorrection] = {} + for output in leaves(outputs): + for kind, tensor in ( + ("target_logprobs", output.target_logprobs), + ("top_k", None if output.top_k is None else output.top_k.logprobs), + ): + if ( + tensor is None + or not tensor.requires_grad + or (correction is None and kind != "top_k") + ): + continue + index = indices[id(tensor)] + entry = _OutputCorrection( + index=index, + kind=kind, + original_logprobs=tensor.detach().to("cpu", copy=True), + token_index=indices[id(output.top_k.tokens)] + if kind == "top_k" + else None, + original_tokens=output.top_k.tokens.detach().to("cpu", copy=True) + if kind == "top_k" + else None, + logits_index=indices[id(output.logits)] + if kind == "top_k" and output.logits is not None + else None, + ) + if index in entries and ( + entries[index].kind != kind + or entries[index].token_index != entry.token_index + ): + raise ValueError( + "an aliased output tensor has ambiguous correction semantics" + ) + entries.setdefault(index, entry) + return ForwardCorrectionContext(len(tensors), correction, tuple(entries.values())) diff --git a/src/art/trainer_rank/_graphs.py b/src/art/trainer_rank/_graphs.py new file mode 100644 index 000000000..3294f413b --- /dev/null +++ b/src/art/trainer_rank/_graphs.py @@ -0,0 +1,704 @@ +"""Evictable physical forward graphs, separate from caller autograd graphs.""" + +from __future__ import annotations + +from collections.abc import Callable, Iterable, Sequence +from contextlib import AbstractContextManager, ExitStack, contextmanager, nullcontext +from copy import deepcopy +from dataclasses import dataclass, field, fields, is_dataclass, replace +import random +from time import perf_counter +from typing import Any, Literal +from uuid import uuid4 +import weakref + +import torch +from torch._C._autograd import _get_current_graph_task_keep_graph +from torch.multiprocessing.reductions import StorageWeakRef + +from art._tensor_residency import observe_resident_tensors + +ForwardHandle = str +type Retention = Literal["gpu", "cpu", "replay"] + + +@dataclass(frozen=True) +class _InputTensor: + value: torch.Tensor + device: torch.device + requires_grad: bool + + +def _snapshot(value: Any, *, inputs: bool = False) -> Any: + if isinstance(value, torch.Tensor): + if inputs: + return _InputTensor( + value.detach().to("cpu", copy=True), value.device, value.requires_grad + ) + return value.detach().clone().requires_grad_(value.requires_grad) + if is_dataclass(value) and not isinstance(value, type): + return replace( + value, + **{ + f.name: _snapshot(getattr(value, f.name), inputs=inputs) + for f in fields(value) + if f.init + }, + ) + if isinstance(value, tuple): + items = tuple(_snapshot(item, inputs=inputs) for item in value) + return type(value)(*items) if hasattr(value, "_fields") else items + if isinstance(value, list): + return [_snapshot(item, inputs=inputs) for item in value] + if isinstance(value, dict): + return {key: _snapshot(item, inputs=inputs) for key, item in value.items()} + return deepcopy(value) + + +def _restore_inputs(value: Any) -> Any: + if isinstance(value, _InputTensor): + return value.value.to(value.device, copy=True).requires_grad_( + value.requires_grad + ) + if is_dataclass(value) and not isinstance(value, type): + return replace( + value, + **{ + f.name: _restore_inputs(getattr(value, f.name)) + for f in fields(value) + if f.init + }, + ) + if isinstance(value, tuple): + items = tuple(_restore_inputs(item) for item in value) + return type(value)(*items) if hasattr(value, "_fields") else items + if isinstance(value, list): + return [_restore_inputs(item) for item in value] + if isinstance(value, dict): + return {key: _restore_inputs(item) for key, item in value.items()} + return value + + +def _tensors(value: Any) -> Iterable[torch.Tensor]: + if isinstance(value, torch.Tensor): + yield value + elif is_dataclass(value) and not isinstance(value, type): + for field in fields(value): + yield from _tensors(getattr(value, field.name)) + elif isinstance(value, dict): + for item in value.values(): + yield from _tensors(item) + elif isinstance(value, (tuple, list)): + for item in value: + yield from _tensors(item) + + +def _storage_sizes(tensors: Iterable[torch.Tensor]) -> tuple[int, int]: + sizes: dict[tuple[torch.device, int], int] = {} + for tensor in tensors: + storage = tensor.untyped_storage() + sizes[(tensor.device, storage.data_ptr())] = storage.nbytes() + return ( + sum(size for (device, _), size in sizes.items() if device.type != "cpu"), + sum(size for (device, _), size in sizes.items() if device.type == "cpu"), + ) + + +@dataclass +class _RNGState: + cpu: torch.Tensor + cuda: dict[int, torch.Tensor] + python: tuple[Any, ...] + tracker: Any = None + + @classmethod + def capture(cls, devices: Sequence[int], tracker: Any) -> _RNGState: + return cls( + torch.get_rng_state(), + {device: torch.cuda.get_rng_state(device) for device in devices}, + random.getstate(), + None if tracker is None else _snapshot(tracker.get_states()), + ) + + def restore(self, tracker: Any) -> None: + torch.set_rng_state(self.cpu) + for device, state in self.cuda.items(): + torch.cuda.set_rng_state(state, device) + random.setstate(self.python) + if tracker is not None: + tracker.set_states(_snapshot(self.tracker)) + + @contextmanager + def replay(self, tracker: Any): + ambient = self.capture(tuple(self.cuda), tracker) + try: + self.restore(tracker) + yield + finally: + ambient.restore(tracker) + + +@dataclass +class _TransferStats: + """Completed saved-storage copies, cumulative across released graphs.""" + + offload_bytes: int = 0 + offload_seconds: float = 0.0 + offload_count: int = 0 + offload_max_bytes: int = 0 + restore_bytes: int = 0 + restore_seconds: float = 0.0 + restore_count: int = 0 + restore_max_bytes: int = 0 + + @torch.compiler.disable + def copy(self, tensor: torch.Tensor, device: torch.device | str) -> torch.Tensor: + start = perf_counter() + if tensor.device.type == "cuda" and torch.device(device).type == "cpu": + # Pin the only host copy of each storage. Keep copies blocking so + # offload releases GPU ownership before returning, even on user + # streams, and unpack never exposes an unfinished restore. + result = torch.empty_like(tensor, device="cpu", pin_memory=True) + result.copy_(tensor) + else: + result = tensor.to(device, copy=True) + elapsed = perf_counter() - start + size = tensor.numel() * tensor.element_size() + if result.device.type == "cpu": + self.offload_bytes += size + self.offload_seconds += elapsed + self.offload_count += 1 + self.offload_max_bytes = max(self.offload_max_bytes, size) + else: + self.restore_bytes += size + self.restore_seconds += elapsed + self.restore_count += 1 + self.restore_max_bytes = max(self.restore_max_bytes, size) + return result + + +@dataclass +class _SavedTensor: + tensor: torch.Tensor + device: torch.device + managed: bool + source: StorageWeakRef + restored: dict[tuple[torch.device, StorageWeakRef], torch.Tensor] + transfer_stats: _TransferStats + + def view(self, storage, device) -> torch.Tensor: + tensor = self.tensor + result = torch.empty(0, dtype=tensor.dtype, device=device).set_( + storage, tensor.storage_offset(), tensor.size(), tensor.stride() + ) + if tensor.is_conj(): + result = result.conj() + if tensor.is_neg(): + result = torch._neg_view(result) + return result + + def offload(self, copies: dict[StorageWeakRef, torch.Tensor]) -> None: + if self.managed and self.tensor.device.type != "cpu": + tensor = self.tensor + storage = tensor.untyped_storage() + # Weak storage identity prevents allocator address reuse from + # confusing distinct activations without pinning CUDA storage. + key = self.source + if key not in copies: + raw = torch.empty(0, dtype=torch.uint8, device=tensor.device).set_( + storage, 0, (storage.nbytes(),), (1,) + ) + copies[key] = self.transfer_stats.copy(raw, "cpu") + self.tensor = self.view(copies[key].untyped_storage(), "cpu") + + def unpack(self) -> torch.Tensor: + if self.tensor.device == self.device: + # TE frees unpacked tensor data by rebinding Tensor.data. A retained + # graph needs its saved object's metadata intact for the next call. + if _get_current_graph_task_keep_graph(): + return self.tensor.detach() + return self.tensor + key = (self.device, self.source) + if key not in self.restored: + storage = self.tensor.untyped_storage() + raw = torch.empty(0, dtype=torch.uint8).set_( + storage, 0, (storage.nbytes(),), (1,) + ) + self.restored[key] = self.transfer_stats.copy(raw, self.device) + return self.view(self.restored[key].untyped_storage(), self.device) + + +@dataclass(frozen=True) +class GraphState: + retention: Retention + gpu_bytes: int + cpu_bytes: int + offload_bytes: int + replay_bytes: int + replayable: bool + replay_count: int + offloadable: bool + restore_workspace_bytes: int = 0 + checkpoint_versions: tuple[Any, ...] = () + non_offloadable_bytes: int | None = None + + +@dataclass +class _ForwardRecord: + execute: Callable[[Any], Sequence[torch.Tensor]] + inputs: Any + context_factory: Callable[[], AbstractContextManager[Any]] + validate_backward: Callable[[], None] | None + retention: Retention + checkpoint_versions: tuple[Any, ...] + options: Any + rng: _RNGState + rng_tracker: Any + keep_on_device: Callable[[torch.Tensor], bool] | None + transfer_stats: _TransferStats + outputs: tuple[torch.Tensor, ...] | None = None + saved: list[weakref.ReferenceType[_SavedTensor]] | None = None + resident: list[weakref.ReferenceType[torch.Tensor]] | None = None + metadata: tuple[tuple[torch.Size, torch.dtype, torch.device, bool], ...] = () + replay_count: int = 0 + autocast: tuple[tuple[str, bool, torch.dtype], ...] = () + corrections: Any = None + is_stale: Callable[[], bool] | None = None + current_context_factory: Callable[[], AbstractContextManager[Any]] | None = None + replay_with_current: bool = False + restored: dict[tuple[torch.device, StorageWeakRef], torch.Tensor] = field( + default_factory=dict + ) + execution_peak_bytes: int = 0 + + def run( + self, + *, + context_factory: Callable[[], AbstractContextManager[Any]] | None = None, + store: bool = True, + grad_enabled: bool = True, + ) -> tuple[torch.Tensor, ...]: + saved: list[weakref.ReferenceType[_SavedTensor]] = [] + resident: list[weakref.ReferenceType[torch.Tensor]] = [] + observed = False + copies: dict[StorageWeakRef, torch.Tensor] = {} + retention, keep_on_device, restored, transfer_stats = ( + self.retention, + self.keep_on_device, + self.restored, + self.transfer_stats, + ) + + def observe(tensors: Sequence[torch.Tensor]) -> None: + nonlocal observed + observed = True + resident.extend(weakref.ref(tensor) for tensor in tensors) + + def pack(tensor: torch.Tensor) -> _SavedTensor: + cell = _SavedTensor( + tensor.detach(), + tensor.device, + not (keep_on_device and keep_on_device(tensor)), + StorageWeakRef(tensor.untyped_storage()), + restored, + transfer_stats, + ) + if retention == "cpu": + cell.offload(copies) + saved.append(weakref.ref(cell)) + return cell + + with ( + torch.enable_grad(), + (context_factory or self.context_factory)(), + ExitStack() as stack, + ): + for device, enabled, dtype in self.autocast: + stack.enter_context( + torch.autocast(device, enabled=enabled, dtype=dtype) + ) + with ( + torch.set_grad_enabled(grad_enabled), + torch.autograd.graph.saved_tensors_hooks(pack, _SavedTensor.unpack), + observe_resident_tensors(observe), + ): + try: + outputs = tuple(self.execute(_restore_inputs(self.inputs))) + except BaseException: + # A failed forward has no usable graph. Release partial + # host allocations even if its traceback retains outputs. + for reference in saved: + if (cell := reference()) is not None: + cell.tensor = torch.empty(0) + raise + finally: + copies.clear() + if store: + # Outputs and external autograd contexts already own these storages. + # Moving their saved aliases would retain both CUDA and CPU copies, + # then allocate a duplicate CUDA restore beside the original storage. + output_storages = { + StorageWeakRef(value.untyped_storage()): value.untyped_storage() + for value in (*outputs, *(ref() for ref in resident)) + if value is not None + } + for reference in saved: + cell = reference() + if cell is not None and cell.managed and cell.source in output_storages: + cell.tensor = cell.view(output_storages[cell.source], cell.device) + cell.managed = False + self.outputs, self.saved = outputs, saved + self.resident = resident if observed else None + return outputs + + +class GraphCache: + """A rank-local cache. Its caller serializes model operations and policies.""" + + def __init__(self) -> None: + self._records: dict[ForwardHandle, _ForwardRecord] = {} + self.transfer_stats = _TransferStats() + + def run( + self, + execute: Callable[[Any], Sequence[torch.Tensor]], + inputs: Any, + *, + context_factory: Callable[[], AbstractContextManager[Any]] = nullcontext, + validate_backward: Callable[[], None] | None = None, + retention: Retention = "gpu", + checkpoint_versions: tuple[Any, ...] = (), + options: Any = None, + cuda_devices: Sequence[int] = (), + rng_tracker: Any = None, + keep_on_device: Callable[[torch.Tensor], bool] | None = None, + output_device: torch.device | str | None = None, + execution_peak_bytes: int = 0, + ) -> tuple[ForwardHandle, tuple[torch.Tensor, ...]]: + if retention not in ("gpu", "cpu", "replay"): + raise ValueError(f"Unknown graph retention {retention!r}") + if retention == "cpu" and not getattr(options, "allow_cpu_offload", True): + raise ValueError("CPU graph offload is disabled for this forward") + if retention == "replay" and not getattr(options, "allow_replay", True): + raise ValueError("Graph replay is disabled for this forward") + record = _ForwardRecord( + execute, + _snapshot(inputs, inputs=True), + context_factory, + validate_backward, + retention, + checkpoint_versions, + options, + _RNGState.capture(cuda_devices, rng_tracker), + rng_tracker, + keep_on_device, + self.transfer_stats, + ) + record.autocast = tuple( + ( + device, + torch.is_autocast_enabled(device), + torch.get_autocast_dtype(device), + ) + for device in ("cpu", "cuda") + ) + record.execution_peak_bytes = execution_peak_bytes + physical = record.run() + record.metadata = tuple( + (value.shape, value.dtype, value.device, value.requires_grad) + for value in physical + ) + # clone: a detached view may still pin a much larger model activation. + detached = tuple( + value.detach() + .to(device=output_device, copy=True) + .requires_grad_(value.requires_grad) + for value in physical + ) + handle = uuid4().hex + self._records[handle] = record + if retention == "replay": + self.evict(handle) + return handle, detached + + def handles(self) -> tuple[ForwardHandle, ...]: + return tuple(self._records) + + def set_corrections( + self, + handle: ForwardHandle, + context: Any, + *, + is_stale: Callable[[], bool], + current_context_factory: Callable[[], AbstractContextManager[Any]] + | None = None, + ) -> None: + record = self._records[handle] + record.corrections = context + record.is_stale = is_stale + record.current_context_factory = current_context_factory + + def state(self, handle: ForwardHandle) -> GraphState: + record = self._records[handle] + cells = [ + cell + for ref in record.saved or () + if (cell := ref()) is not None and cell.managed + ] + correction_tensors = ( + () if record.corrections is None else record.corrections.tensors + ) + resident = tuple( + value for ref in record.resident or () if (value := ref()) is not None + ) + gpu, cpu = _storage_sizes( + ( + *_tensors(record.inputs), + *_tensors(record.rng), + *correction_tensors, + *(cell.tensor for cell in cells), + *(record.outputs or ()), + *resident, + ) + ) + offload, _ = _storage_sizes(cell.tensor for cell in cells) + _, replay = _storage_sizes( + (*_tensors(record.inputs), *_tensors(record.rng), *correction_tensors) + ) + policy = getattr(record.options, "backward_state", "auto") + return GraphState( + record.retention, + gpu, + cpu, + offload, + replay, + getattr(record.options, "allow_replay", True) + and policy in ("auto", "replay"), + record.replay_count, + getattr(record.options, "allow_cpu_offload", True) + and policy in ("auto", "cpu"), + max(0, record.execution_peak_bytes - gpu), + record.checkpoint_versions, + None + if record.resident is None + else _storage_sizes((*(record.outputs or ()), *resident))[0], + ) + + def offload(self, handle: ForwardHandle) -> None: + record = self._records[handle] + if not getattr(record.options, "allow_cpu_offload", True): + raise RuntimeError("CPU graph offload is disabled for this forward") + if record.outputs is None: + return + copies: dict[StorageWeakRef, torch.Tensor] = {} + try: + for ref in record.saved or (): + if (cell := ref()) is not None: + cell.offload(copies) + finally: + copies.clear() + record.retention = "cpu" + + def evict( + self, handle: ForwardHandle, *, replay_with_current: bool | None = None + ) -> None: + record = self._records[handle] + if not getattr(record.options, "allow_replay", True): + raise RuntimeError("Graph replay is disabled for this forward") + if replay_with_current is not None: + if replay_with_current and record.current_context_factory is None: + raise ValueError("Current-weight replay requires a version context") + record.replay_with_current = replay_with_current + record.outputs = record.saved = record.resident = None + record.restored.clear() + record.retention = "replay" + + def release(self, handle: ForwardHandle) -> None: + if (record := self._records.pop(handle, None)) is not None: + # Saved-variable hooks can outlive their Python outputs. Break all + # ownership edges even when a caller retains a failure traceback. + record.outputs = record.saved = record.resident = None + record.restored.clear() + record.inputs = record.corrections = None + record.execute = lambda _: () + record.context_factory = nullcontext + record.validate_backward = record.current_context_factory = None + record.is_stale = record.keep_on_device = None + + def backward( + self, + handle: ForwardHandle, + gradients: Sequence[torch.Tensor | None], + *, + retain_graph: bool = False, + ) -> None: + self.backward_many(((handle, gradients),), retain_graph=retain_graph) + + def validate_many( + self, + packets: Sequence[tuple[ForwardHandle, Sequence[torch.Tensor | None]]], + ) -> None: + handles: set[str] = set() + for handle, gradients in packets: + if handle in handles: + raise ValueError("Duplicate forward handle in one backward") + handles.add(handle) + record = self._records[handle] + if record.validate_backward is not None: + record.validate_backward() + if len(gradients) != len(record.metadata): + raise ValueError("Cotangent count does not match forward outputs") + for gradient, (shape, dtype, _, requires_grad) in zip( + gradients, record.metadata, strict=True + ): + if gradient is not None and ( + not requires_grad + or gradient.shape != shape + or gradient.dtype != dtype + ): + raise ValueError( + "Cotangent shape/dtype or output requires_grad mismatch" + ) + if ( + record.corrections is not None + and record.is_stale is not None + and record.is_stale() + ): + if ( + record.corrections.requires_current(gradients) + and record.current_context_factory is None + ): + raise RuntimeError( + "Stale correction requires current model outputs but no current version context is available" + ) + + def backward_many( + self, + packets: Sequence[tuple[ForwardHandle, Sequence[torch.Tensor | None]]], + *, + retain_graph: bool = False, + coordinate: Callable[[Callable[[], Any]], Any] | None = None, + ) -> None: + self.validate_many(packets) + try: + self._backward_validated( + packets, + retain_graph=retain_graph, + coordinate=coordinate or (lambda function: function()), + ) + except BaseException: + # A replay/backward failure consumes the operation. The enclosing + # checkpoint transaction discards unpublished optimizer gradients. + for handle, _ in packets: + self.release(handle) + raise + + def _backward_validated( + self, + packets: Sequence[tuple[ForwardHandle, Sequence[torch.Tensor | None]]], + *, + retain_graph: bool, + coordinate: Callable[[Callable[[], Any]], Any], + ) -> None: + # Random wire handles differ across physical ranks. Creation order is + # shared by TP/CP peers and therefore fixes backward collective order. + ordinals = {handle: index for index, handle in enumerate(self._records)} + packets = sorted(packets, key=lambda packet: ordinals[packet[0]]) + prepared = {} + stale = { + handle: record.is_stale is not None and record.is_stale() + for handle, _ in packets + for record in (self._records[handle],) + } + # The coordinator uses this logical rank's TP/CP group, never all DP + # ranks: different DP owners may have different numbers of graphs. + for handle, gradients in packets: + record = self._records[handle] + prepared[handle] = coordinate( + lambda: self._prepare_correction(record, gradients, stale[handle]) + ) + for handle, gradients in packets: + record = self._records[handle] + pairs = coordinate( + lambda: self._prepare_backward( + record, gradients, stale[handle], prepared[handle] + ) + ) + try: + coordinate( + lambda: ( + torch.autograd.backward( + [output for output, _ in pairs], + [gradient for _, gradient in pairs], + retain_graph=retain_graph, + ) + if pairs + else None + ) + ) + finally: + record.restored.clear() + del pairs + if not retain_graph: + self.release(handle) + elif record.retention == "replay": + self.evict(handle) + + @staticmethod + def _prepare_correction(record, gradients, stale): + if not ( + stale + and record.corrections is not None + and record.corrections.requires_current(gradients) + and not (record.outputs is None and record.replay_with_current) + ): + return None + # Explicit always may add a no-grad forward. Stage every correction + # before physical backward; changed cotangents wait on CPU. + with record.rng.replay(record.rng_tracker): + current = record.run( + context_factory=record.current_context_factory, + store=False, + grad_enabled=False, + ) + corrected = record.corrections.correct(gradients, current) + return tuple( + value.to("cpu") if value is not None and value is not original else value + for value, original in zip(corrected, gradients, strict=True) + ) + + @staticmethod + def _prepare_backward(record, gradients, stale, prepared): + gradients = gradients if prepared is None else prepared + if not any(gradient is not None for gradient in gradients): + return [] + current_replay = stale and record.replay_with_current + if record.outputs is None: + with record.rng.replay(record.rng_tracker): + physical = record.run( + context_factory=record.current_context_factory + if current_replay + else None + ) + metadata = tuple( + (value.shape, value.dtype, value.device, value.requires_grad) + for value in physical + ) + if metadata != record.metadata: + record.outputs = record.saved = None + raise RuntimeError( + "Replayed output metadata differs from original forward" + ) + record.replay_count += 1 + if prepared is None and stale and record.corrections is not None: + if current_replay: + record.corrections.validate_replay(gradients, record.outputs) + gradients = record.corrections.correct( + gradients, record.outputs if current_replay else None + ) + assert record.outputs is not None + return [ + (output, gradient.to(output.device)) + for output, gradient in zip(record.outputs, gradients, strict=True) + if gradient is not None + ] diff --git a/src/art/trainer_rank/_heads.py b/src/art/trainer_rank/_heads.py new file mode 100644 index 000000000..c27ee7d33 --- /dev/null +++ b/src/art/trainer_rank/_heads.py @@ -0,0 +1,1363 @@ +"""Live checkpoint-owned modules and their serializable client state.""" + +from __future__ import annotations + +from collections.abc import Callable, Mapping +from contextvars import ContextVar +from copy import deepcopy +from dataclasses import dataclass, replace +from typing import TYPE_CHECKING, Any, Literal, SupportsIndex, cast +import weakref + +import torch + +if TYPE_CHECKING: + from ._impl import TrainerRank, _CustomObject, _CustomTensorTracker + +HeadKind = Literal["module", "parameter", "buffer"] +_parameter_transform: ContextVar[bool] = ContextVar( + "head_parameter_transform", default=False +) +_head_call: ContextVar[tuple[tuple[object, ...], dict[int, torch.Tensor]] | None] = ( + ContextVar("head_call", default=None) +) + + +def head_call_arguments(args: tuple[Any, ...], kwargs: dict[str, Any]) -> Any: + """Route preexisting live buffer aliases into their active module captures.""" + call = _head_call.get() + if call is None: + return None + from ._tensors import _map_tensors + + changed = False + + def replace(value: torch.Tensor) -> torch.Tensor: + nonlocal changed + if id(value) in call[1]: + changed = True + return call[1][id(value)] + return value + + result = _map_tensors(replace, (args, kwargs)) + return result if changed else None + + +TENSOR_METADATA_FUNCTIONS = frozenset( + { + "size", + "numel", + "nelement", + "dim", + "ndimension", + "stride", + "storage_offset", + "element_size", + "is_contiguous", + "is_floating_point", + "is_complex", + "is_signed", + "get_device", + } +) +_TENSOR_INSPECTION_PROPERTIES = frozenset( + { + "_backward_hooks", + "_base", + "_cdata", + "_grad", + "_grad_fn", + "_has_symbolic_sizes_strides", + "_post_accumulate_grad_hooks", + "_python_dispatch", + "_version", + "data", + "device", + "dtype", + "grad", + "grad_dtype", + "grad_fn", + "is_cpu", + "is_cuda", + "is_ipu", + "is_leaf", + "is_maia", + "is_meta", + "is_mkldnn", + "is_mps", + "is_mtia", + "is_nested", + "is_quantized", + "is_sparse", + "is_sparse_csr", + "is_vulkan", + "is_xla", + "is_xpu", + "itemsize", + "layout", + "name", + "names", + "nbytes", + "ndim", + "output_nr", + "requires_grad", + "retains_grad", + "shape", + "volatile", + } +) + + +def tensor_metadata_function(func: Callable[..., Any]) -> bool: + """Inspect handle metadata/state without capturing a weight-sized snapshot. + + Tensor-valued views (T, mT, H, mH, real, imag) and unknown descriptors must + use the ordinary computation path. Autograd/storage inspection retains its + existing handle semantics, including the client's separate .data guard. + """ + name = getattr(func, "__name__", "") + return name in TENSOR_METADATA_FUNCTIONS or ( + name == "__get__" + and getattr(getattr(func, "__self__", None), "__name__", None) + in _TENSOR_INSPECTION_PROPERTIES + ) + + +def mutates_tensor(func: Callable[..., Any], kwargs: Mapping[str, Any]) -> bool: + name = getattr(func, "__name__", "") + return ( + (name.endswith("_") and not name.endswith("__")) + or name + in { + "__setitem__", + "__set__", + "__iadd__", + "__isub__", + "__imul__", + "__itruediv__", + "__ifloordiv__", + "__imod__", + "__ipow__", + "__imatmul__", + "__iand__", + "__ior__", + "__ixor__", + "__ilshift__", + "__irshift__", + } + or kwargs.get("out") is not None + or kwargs.get("inplace") is True + ) + + +def tensor_mutation_targets( + func: Callable[..., Any], args: tuple[Any, ...], kwargs: Mapping[str, Any] +) -> set[int]: + from ._impl import _walk_objects + + if not mutates_tensor(func, kwargs): + return set() + values = kwargs["out"] if kwargs.get("out") is not None else args[:1] + return { + id(value) for value in _walk_objects(values) if isinstance(value, torch.Tensor) + } + + +def readonly_buffer_views(result: Any, snapshots: list[torch.Tensor]) -> Any: + from ._tensors import _map_tensors + + def wrap(value: torch.Tensor) -> torch.Tensor: + with torch._C.DisableTorchFunctionSubclass(): + if any( + getattr(torch._C, "_is_alias_of")(value, snapshot) + for snapshot in snapshots + ): + value.__class__ = _BufferSnapshotView + return value + + return _map_tensors(wrap, result) if snapshots else result + + +class _BufferSnapshotView(torch.Tensor): + """A readable snapshot view that cannot silently impersonate a live write.""" + + def __deepcopy__(self, memo: dict[int, Any]) -> torch.Tensor: + result = _plain(self) + memo[id(self)] = result + return result + + def __reduce_ex__(self, proto: SupportsIndex) -> Any: + return _plain(self).__reduce_ex__(proto) + + @classmethod + def __torch_function__(cls, func, types, args=(), kwargs=None): + from ._impl import _walk_objects + from ._tensors import _map_tensors + + kwargs = kwargs or {} + if getattr(func, "__name__", "") in { + "batch_norm", + "instance_norm", + "embedding", + "embedding_bag", + }: + raise RuntimeError( + "Stateful operations on buffer snapshot views are unsupported; use the live buffer or an explicit clone()" + ) + targets = tensor_mutation_targets(func, args, kwargs) + if any( + isinstance(value, cls) and id(value) in targets + for value in _walk_objects((args, kwargs)) + ): + raise RuntimeError( + "Views of live checkpoint buffers are read-only; use buffer[index] = value or buffer.copy_(), or clone() for a writable private copy" + ) + snapshots = [] + + def plain(value: torch.Tensor) -> torch.Tensor: + if isinstance(value, cls): + with torch._C.DisableTorchFunctionSubclass(): + value = value.as_subclass(torch.Tensor) + snapshots.append(value) + return value + + args, kwargs = _map_tensors(plain, (args, kwargs)) + return readonly_buffer_views(func(*args, **kwargs), snapshots) + + +def _plain(tensor: torch.Tensor) -> torch.Tensor: + with torch._C.DisableTorchFunctionSubclass(): + return tensor.as_subclass(torch.Tensor).detach().clone() + + +def _restore_module(module: torch.nn.Module) -> torch.nn.Module: + return module + + +def _preserve_module_aliases(module: torch.nn.Module, apply: Callable[[], Any]) -> None: + groups: dict[tuple[str, int], list[str]] = {} + for kind, values in ( + ("_parameters", module.named_parameters(remove_duplicate=False)), + ("_buffers", module.named_buffers(remove_duplicate=False)), + ): + for key, value in values: + groups.setdefault((kind, id(value)), []).append(key) + apply() + for (kind, _), keys in groups.items(): + prefix, _, key = keys[0].rpartition(".") + value = getattr(module.get_submodule(prefix), kind)[key] + for alias in keys[1:]: + prefix, _, key = alias.rpartition(".") + getattr(module.get_submodule(prefix), kind)[key] = value + + +def move_module(module: torch.nn.Module, device: torch.device) -> torch.nn.Module: + _preserve_module_aliases(module, lambda: module.to(device=device)) + return module + + +def head_staleness(trainer: TrainerRank) -> int: + from ._options import resolve_forward_options + + return resolve_forward_options( + getattr(trainer, "_forward_options", None) + ).max_gradient_staleness + + +class ModuleHandle(torch.nn.Module): + """A reusable module whose calls capture current checkpoint tensor versions. + + Each successful call publishes its buffer changes. A failed call leaves buffers + unchanged. Parameters and buffers retained by earlier calls are never modified. + The proxy delegates attributes, children, and named parameters to the source; + its type is ModuleHandle, so isinstance(handle, factory_class) is not preserved. + For an external activation-checkpoint closure, capture ``head.snapshot()`` + before calling ``torch.utils.checkpoint``. Snapshot buffers are private; + their mutations are not published to the checkpoint. Client and public logical + callback snapshots contain cotangent bridges: use ``use_reentrant=False``. + Internal physical snapshots support either mode under a gradient transaction. + """ + + def __init__( + self, + module: torch.nn.Module, + capture: Callable[[], tuple[dict[str, torch.Tensor], dict[str, torch.Tensor]]], + publish: Callable[[Mapping[str, torch.Tensor]], None], + ) -> None: + super().__init__() + object.__setattr__(self, "_source", module) + object.__setattr__(self, "_capture", capture) + object.__setattr__(self, "_publish", publish) + self._parameters = module._parameters + self._buffers = module._buffers + self._modules = module._modules + self._non_persistent_buffers_set = module._non_persistent_buffers_set + self.training = module.training + + def __getattr__(self, name: str) -> Any: + try: + return super().__getattr__(name) + except AttributeError: + return getattr(self._source, name) + + def __getitem__(self, key: Any) -> Any: + return self._source[key] + + def train(self, mode: bool = True) -> ModuleHandle: + self._source.train(mode) + self.training = mode + return self + + def _apply( + self, fn: Callable[[torch.Tensor], torch.Tensor], recurse: bool = True + ) -> ModuleHandle: + def convert(value: torch.Tensor) -> torch.Tensor: + result = fn(value) + if isinstance(value, _ClientBuffer): + return _ClientBuffer(result, value._head_owner) + if not isinstance(value, torch.nn.Parameter) and hasattr( + value, "_art_tracker" + ): + return type(value)(result, value._art_tracker) + return result + + token = _parameter_transform.set(True) + try: + _preserve_module_aliases( + self._source, lambda: self._source._apply(convert, recurse) + ) + finally: + _parameter_transform.reset(token) + self._parameters, self._buffers, self._modules = ( + self._source._parameters, + self._source._buffers, + self._source._modules, + ) + return self + + def requires_grad_(self, requires_grad: bool = True) -> ModuleHandle: + raise RuntimeError( + "Set checkpoint parameter trainability in the module factory" + ) + + def _call_impl(self, *args: Any, **kwargs: Any) -> Any: + if (call := _head_call.get()) is not None and any( + head is self for head in call[0] + ): + return super()._call_impl(*args, **kwargs) + return self._captured_call( + lambda: super(ModuleHandle, self)._call_impl(*args, **kwargs) + ) + + def forward(self, *args: Any, **kwargs: Any) -> Any: + if (call := _head_call.get()) is not None and any( + head is self for head in call[0] + ): + return self._source(*args, **kwargs) + return self._captured_call(lambda: self._source(*args, **kwargs)) + + def _captured_call(self, call: Callable[[], Any]) -> Any: + if torch._C._current_graph_task_id() >= 0: + raise RuntimeError( + "Live module called during backward recomputation; pass head.snapshot() to activation checkpointing (client and logical callback handles require use_reentrant=False)" + ) + parameters, buffers = self._capture() + parameter_versions = { + name: value._version for name, value in parameters.items() + } + # Hooks on the handle must see the same private tensors as forward. + # Internal checkpoint closures retain the captured source after bindings + # on the public handle are restored. Memo substitution avoids extra copies. + captured = self._captured_module(parameters, buffers) + original = self._source + active, aliases = _head_call.get() or ((), {}) + token = _head_call.set( + ( + (*active, self), + aliases + | {id(value): buffers[key] for key, value in original.named_buffers()}, + ) + ) + try: + object.__setattr__(self, "_source", captured) + self._parameters, self._buffers, self._modules = ( + captured._parameters, + captured._buffers, + captured._modules, + ) + result = call() + current_parameters = dict(captured.named_parameters()) + if current_parameters.keys() != parameters.keys() or any( + current_parameters[name] is not value + or value._version != parameter_versions[name] + for name, value in parameters.items() + ): + raise RuntimeError( + "Checkpoint module forward must not mutate parameters" + ) + updated_buffers = dict(captured.named_buffers()) + finally: + object.__setattr__(self, "_source", original) + self._parameters, self._buffers, self._modules = ( + original._parameters, + original._buffers, + original._modules, + ) + _head_call.reset(token) + self._publish(updated_buffers) + return result + + def _captured_module( + self, + parameters: Mapping[str, torch.Tensor], + buffers: Mapping[str, torch.Tensor], + ) -> torch.nn.Module: + memo = { + id(value): parameters[key] for key, value in self._source.named_parameters() + } | {id(value): buffers[key] for key, value in self._source.named_buffers()} + return deepcopy(self._source, memo) + + def snapshot(self) -> torch.nn.Module: + """Capture an external closure; cotangent bridges require use_reentrant=False.""" + if torch._C._current_graph_task_id() >= 0: + raise RuntimeError( + "Capture head.snapshot() before entering activation checkpointing" + ) + return self._captured_module(*self._capture()) + + def __deepcopy__(self, memo: dict[int, object]) -> torch.nn.Module: + return deepcopy(self._source, memo) + + def __reduce_ex__(self, protocol: SupportsIndex) -> tuple[Any, ...]: + return _restore_module, (self._source,) + + +class _NativeModuleState: + def __init__( + self, + trainer: TrainerRank, + tracker: _CustomTensorTracker, + module: torch.nn.Module, + ): + self.trainer = weakref.ref(trainer) + self.tracker = tracker + self.module = module + + def capture(self) -> tuple[dict[str, torch.Tensor], dict[str, torch.Tensor]]: + trainer = self.tracker.validate() + assert self.tracker.ref.name is not None + version = trainer._capture_checkpoint_version(self.tracker.ref.name) + from ._impl import _track_slot_graph_tensor + + parameters = {} + marker = torch.zeros((), dtype=torch.bool, device="cpu") + for key, parameter in self.module.named_parameters(): + with torch._C.DisableTorchFunctionSubclass(): + value = trainer._snapshot_parameter( + parameter, version, head_staleness(trainer) + ) + if value.requires_grad and torch.is_grad_enabled(): + value = _track_slot_graph_tensor(value, marker) + parameters[key] = value + if any(value.requires_grad for value in parameters.values()): + self.tracker.record(marker) + buffers = {key: _plain(value) for key, value in self.module.named_buffers()} + return parameters, buffers + + def publish(self, buffers: Mapping[str, torch.Tensor]) -> None: + self.tracker.validate() + staged = _stage_local_buffers(dict(self.module.named_buffers()), buffers) + with torch.no_grad(), torch._C.DisableTorchFunctionSubclass(): + for target, value in staged: + target.copy_(value) + if staged: + self.tracker.buffer_revision += 1 + + +def _stage_local_buffers( + targets: Mapping[str, torch.Tensor], values: Mapping[str, torch.Tensor] +) -> list[tuple[torch.Tensor, torch.Tensor]]: + if targets.keys() != values.keys(): + raise ValueError("Checkpoint module forward must preserve buffer names") + staged = [] + with torch.no_grad(), torch._C.DisableTorchFunctionSubclass(): + for key, target in targets.items(): + value = values[key] + if target.shape != value.shape or target.dtype != value.dtype: + raise ValueError( + "Checkpoint module forward must preserve buffer shape and dtype" + ) + value = value.to(device=target.device) + if not torch.equal(target, value): + prepared = _plain(target) + prepared.copy_(value) + staged.append((target, prepared)) + return staged + + +def native_module_handle( + trainer: TrainerRank, custom: _CustomObject, tracker: _CustomTensorTracker +) -> ModuleHandle: + state = _NativeModuleState(trainer, tracker, cast(torch.nn.Module, custom.value)) + return ModuleHandle(state.module, state.capture, state.publish) + + +@dataclass(frozen=True) +class HeadRegistration: + checkpoint: Any + name: str + kind: HeadKind + value: torch.nn.Module | torch.Tensor + + +@dataclass(frozen=True) +class HeadState: + version: Any + name: str + kind: HeadKind + parameters: dict[str, torch.Tensor] + buffers: dict[str, torch.Tensor] + buffer_revision: int + max_gradient_staleness: int = 2 + + +@dataclass(frozen=True) +class HeadBufferUpdate: + version: Any + name: str + buffer_revision: int + buffers: dict[str, torch.Tensor] + + +def _custom_tracker(custom: _CustomObject) -> _CustomTensorTracker: + if custom.kind == "module": + assert isinstance(custom.handle, ModuleHandle) + return custom.handle._capture.__self__.tracker + return cast(Any, custom.value)._art_tracker + + +def _custom_parameters(custom: _CustomObject) -> dict[str, torch.nn.Parameter]: + if custom.kind == "module": + return dict(cast(torch.nn.Module, custom.value).named_parameters()) + return ( + {"": cast(torch.nn.Parameter, custom.value)} + if custom.kind == "parameter" + else {} + ) + + +def _custom_buffers(custom: _CustomObject) -> dict[str, torch.Tensor]: + if custom.kind == "module": + return dict(cast(torch.nn.Module, custom.value).named_buffers()) + return {"": cast(torch.Tensor, custom.value)} if custom.kind == "buffer" else {} + + +def export_head(trainer: TrainerRank, checkpoint: str, name: str) -> HeadState: + custom = trainer._checkpoint_slots[checkpoint].custom[name] + parameters, buffers = _custom_parameters(custom), _custom_buffers(custom) + return HeadState( + trainer._capture_checkpoint_version(checkpoint), + name, + custom.kind, + { + key: _plain(value).cpu().requires_grad_(value.requires_grad) + for key, value in parameters.items() + }, + {key: _plain(value).cpu() for key, value in buffers.items()}, + _custom_tracker(custom).buffer_revision, + head_staleness(trainer), + ) + + +def execute_head_operation( + trainer: Any, + kind: str, + payload: Any, + *, + coordinate: Callable[[Callable[[], None]], None] | None = None, +) -> Any: + """Execute on every physical rank inside the owning trainer operation queue.""" + if hasattr(trainer, "_rank"): + return trainer._invoke("head", kind, payload) + if kind == "head_lookup": + checkpoint, name = payload + checkpoint = trainer._resolve_custom_checkpoint(checkpoint) + return checkpoint, export_head( + trainer, checkpoint, name + ) if name in trainer._checkpoint_slots[checkpoint].custom else None + if kind == "head_register": + registration: HeadRegistration = payload + checkpoint = trainer._resolve_custom_checkpoint(registration.checkpoint) + existing = trainer._checkpoint_slots[checkpoint].custom.get(registration.name) + if existing is not None: + _validate_registration(trainer, existing, registration) + trainer._custom_object( + registration.name, + registration.kind, + lambda: deepcopy(registration.value), + checkpoint=checkpoint, + ) + return export_head(trainer, checkpoint, registration.name) + if kind == "head_export": + return tuple( + export_head(trainer, checkpoint, name) + if checkpoint in trainer._checkpoint_slots + and name in trainer._checkpoint_slots[checkpoint].custom + else None + for checkpoint, name in payload + ) + if kind == "head_publish": + from . import _checkpoint + + staged, error = [], None + + def prepare() -> None: + nonlocal staged + staged = _stage_buffer_publications(trainer, payload) + + if coordinate is not None: + coordinate(prepare) + else: + try: + prepare() + except Exception as exc: + error = exc + _checkpoint.raise_distributed( + error, "publish custom buffers", trainer._checkpoint_group() + ) + with torch.no_grad(), torch._C.DisableTorchFunctionSubclass(): + for tracker, values in staged: + for target, value in values: + target.copy_(value) + tracker.buffer_revision += 1 + return None + raise ValueError(f"Unknown head operation: {kind!r}") + + +def _validate_registration( + trainer: TrainerRank, existing: _CustomObject, registration: HeadRegistration +) -> None: + from . import _checkpoint + from ._impl import _custom_signature, _custom_state, _CustomObject + + error = None + try: + value = registration.value + if existing.kind != registration.kind: + raise ValueError("Custom checkpoint object kind differs from registration") + if existing.kind == "module": + if not isinstance(value, torch.nn.Module): + raise TypeError("module() factory must return torch.nn.Module") + if trainer._checkpoint_slots[registration.checkpoint].snapshot: + value = deepcopy(value).requires_grad_(False) + candidate = _CustomObject("module", value, object()) + if _custom_signature( + registration.name, candidate, _custom_state(candidate) + ) != _custom_signature( + registration.name, existing, _custom_state(existing) + ): + raise ValueError( + "Custom module schema differs from its registered checkpoint state" + ) + elif ( + not isinstance(value, torch.Tensor) + or value.shape != cast(torch.Tensor, existing.value).shape + or value.dtype != cast(torch.Tensor, existing.value).dtype + ): + raise ValueError( + "Custom tensor shape or dtype differs from its registered checkpoint state" + ) + except Exception as exc: + error = exc + _checkpoint.raise_distributed( + error, "validate custom handle", trainer._checkpoint_group() + ) + + +def _stage_buffer_publications(trainer: TrainerRank, updates: Any) -> list[Any]: + staged = [] + for update in updates: + current = trainer._capture_checkpoint_version(update.version.checkpoint) + if current.generation != update.version.generation: + raise trainer._slot_state_error( + "Custom head checkpoint was replaced before buffer publication" + ) + custom = trainer._checkpoint_slots[update.version.checkpoint].custom[ + update.name + ] + tracker = _custom_tracker(custom) + if tracker.buffer_revision != update.buffer_revision: + raise RuntimeError( + f"Custom module {update.name!r} buffers changed before publication" + ) + targets = _custom_buffers(custom) + if update.buffers.keys() != targets.keys(): + raise ValueError("Custom module buffer keys changed before publication") + values = [] + for key, value in update.buffers.items(): + target = targets[key] + if value.shape != target.shape or value.dtype != target.dtype: + raise ValueError(f"Custom buffer {key!r} shape or dtype changed") + values.append((target, value.to(target.device).clone())) + staged.append((tracker, values)) + return staged + + +def head_gradient_targets( + trainer: TrainerRank, packet: Any, *, materialize: bool = True +) -> list[tuple[Any, int, torch.nn.Parameter, torch.Tensor]]: + """Translate one client head cotangent; caller commits all operation targets together.""" + import json + + from ._versions import CheckpointVersion + + if not packet.handle.startswith("head:"): + raise ValueError("Not a custom head cotangent") + metadata = json.loads(packet.handle[5:]) + version = CheckpointVersion( + metadata["checkpoint"], metadata["generation"], metadata["revision"] + ) + trainer._validate_checkpoint_version( + version, max_gradient_staleness=metadata["max_gradient_staleness"] + ) + custom = trainer._checkpoint_slots[version.checkpoint].custom[metadata["name"]] + parameters = _custom_parameters(custom) + if len(metadata["keys"]) != len(packet.gradients): + raise ValueError("Custom head cotangent count does not match parameters") + for key, gradient in zip(metadata["keys"], packet.gradients, strict=True): + parameter = parameters[key] + if gradient is not None and ( + gradient.shape != parameter.shape + or any( + tensor.layout != torch.strided + for tensor in (parameter, gradient, parameter.grad) + if tensor is not None + ) + ): + raise ValueError( + "Custom head cotangent shape/layout does not match parameter" + ) + return [ + ( + version, + metadata["max_gradient_staleness"], + parameters[key], + gradient.to(device=parameters[key].device, dtype=parameters[key].dtype) + if materialize + else gradient, + ) + for key, gradient in zip(metadata["keys"], packet.gradients, strict=True) + if gradient is not None + ] + + +class _ClientParameter(torch.nn.Parameter): + _head_owner: LiveHead + _head_key: str + + def __new__(cls, data: torch.Tensor, owner: LiveHead, key: str) -> _ClientParameter: + result = super().__new__( + cls, data, requires_grad=owner.state.parameters[key].requires_grad + ) + result._head_owner, result._head_key = owner, key + return result + + def __init__(self, data: torch.Tensor, owner: LiveHead, key: str) -> None: + pass + + def register_hook(self, hook: Any) -> Any: + """Run once on summed captured uses per backward; removal affects old graphs.""" + from ._parameter_hooks import register_parameter_hook + + self._head_owner._validate() + return register_parameter_hook(self, hook) + + def register_post_accumulate_grad_hook(self, hook: Any) -> Any: + from ._parameter_hooks import reject_post_accumulate_hook + + return reject_post_accumulate_hook() + + @property + def data(self) -> torch.Tensor: + raise RuntimeError( + "Client checkpoint parameters do not expose mutable .data; use detach() to read a snapshot" + ) + + @data.setter + def data(self, value: torch.Tensor) -> None: + if not _parameter_transform.get(): + raise RuntimeError( + "Client checkpoint parameters may only be changed by trainer.optim_step" + ) + with torch._C.DisableTorchFunctionSubclass(): + cast(Any, torch.Tensor.data).__set__(self, value) + + @classmethod + def __torch_function__( + cls, + func: Callable[..., Any], + types: tuple[type, ...], + args: tuple[Any, ...] = (), + kwargs: dict[str, Any] | None = None, + ) -> Any: + kwargs = kwargs or {} + name = getattr(func, "__name__", "") + if tensor_metadata_function(func) or name in { + "__format__", + "__hash__", + "__len__", + "__repr__", + "__str__", + }: + with torch._C.DisableTorchFunctionSubclass(): + return func(*args, **kwargs) + if torch._C._current_graph_task_id() >= 0: + raise RuntimeError( + "Live parameter used during backward recomputation; capture parameter.clone() before activation checkpointing and use use_reentrant=False for client or logical callback handles" + ) + if name == "__set__" and _parameter_transform.get(): + with torch._C.DisableTorchFunctionSubclass(): + return func(*args, **kwargs) + mutation_targets = tensor_mutation_targets(func, args, kwargs) + captures: dict[int, torch.Tensor] = {} + + def replace(value: Any) -> Any: + if isinstance(value, _ClientParameter): + if id(value) in mutation_targets: + raise RuntimeError( + "Client checkpoint parameters may only be changed by trainer.optim_step" + ) + if id(value) not in captures: + captures[id(value)] = value._head_owner.capture((value._head_key,))[ + value._head_key + ] + return captures[id(value)] + if isinstance(value, tuple): + items = tuple(replace(item) for item in value) + return type(value)(*items) if hasattr(value, "_fields") else items + if isinstance(value, list): + return [replace(item) for item in value] + if isinstance(value, dict): + return {key: replace(item) for key, item in value.items()} + return value + + return func(*replace(args), **replace(kwargs)) + + def __deepcopy__(self, memo: dict[int, Any]) -> torch.nn.Parameter: + with torch._C.DisableTorchFunctionSubclass(): + result = torch.nn.Parameter(_plain(self), requires_grad=self.requires_grad) + memo[id(self)] = result + return result + + def __reduce_ex__(self, proto: SupportsIndex) -> Any: + with torch._C.DisableTorchFunctionSubclass(): + return torch.nn.Parameter( + _plain(self), requires_grad=self.requires_grad + ).__reduce_ex__(proto) + + +class _ClientBuffer(torch.Tensor): + _head_owner: LiveHead + + @staticmethod + def __new__(cls, data: torch.Tensor, owner: LiveHead) -> _ClientBuffer: + result = torch.Tensor._make_subclass(cls, data.detach(), require_grad=False) + result._head_owner = owner + return result + + @property + def data(self) -> torch.Tensor: + raise RuntimeError("Use checkpoint buffer.copy_() to change its values") + + @data.setter + def data(self, value: torch.Tensor) -> None: + raise RuntimeError("Use checkpoint buffer.copy_() to change its values") + + @classmethod + def __torch_function__( + cls, + func: Callable[..., Any], + types: tuple[type, ...], + args: tuple[Any, ...] = (), + kwargs: dict[str, Any] | None = None, + ) -> Any: + kwargs = kwargs or {} + if (captured := head_call_arguments(args, kwargs)) is not None: + return func(*captured[0], **captured[1]) + name = getattr(func, "__name__", "") + from ._impl import _walk_objects + + mutation_targets = tensor_mutation_targets(func, args, kwargs) + mutating = any( + isinstance(value, _ClientBuffer) and id(value) in mutation_targets + for value in _walk_objects((args, kwargs)) + ) + if mutating and ( + kwargs.get("out") is not None + or name + not in { + "__setitem__", + "__iadd__", + "__isub__", + "__imul__", + "__itruediv__", + "__ifloordiv__", + "__imod__", + "__ipow__", + "__iand__", + "__ior__", + "__ixor__", + "__ilshift__", + "__irshift__", + "copy_", + "fill_", + "zero_", + "add_", + "sub_", + "mul_", + "div_", + "true_divide_", + "floor_divide_", + "remainder_", + "fmod_", + "pow_", + "lerp_", + "bitwise_and_", + "bitwise_or_", + "bitwise_xor_", + "bitwise_left_shift_", + "bitwise_right_shift_", + "masked_fill_", + "masked_scatter_", + "scatter_", + "scatter_add_", + "index_copy_", + "index_add_", + "index_fill_", + "index_put_", + "put_", + "clamp_", + "clamp_min_", + "clamp_max_", + } + ): + raise RuntimeError( + f"Unsupported checkpoint buffer mutation {name}; use buffer.copy_() with unchanged shape and dtype" + ) + copies, originals = {}, {} + if tensor_metadata_function(func): + for value in _walk_objects((args, kwargs)): + if isinstance(value, _ClientBuffer): + value._head_owner._validate() + with torch._C.DisableTorchFunctionSubclass(): + return func(*args, **kwargs) + + def replace(value: Any) -> Any: + if isinstance(value, _ClientBuffer): + value._head_owner._validate() + if id(value) not in copies: + with torch._C.DisableTorchFunctionSubclass(): + copies[id(value)] = ( + value.as_subclass(torch.Tensor) + if id(value) in mutation_targets + else _plain(value) + ) + originals[id(copies[id(value)])] = value + return copies[id(value)] + if isinstance(value, tuple): + items = tuple(replace(item) for item in value) + return type(value)(*items) if hasattr(value, "_fields") else items + if isinstance(value, list): + return [replace(item) for item in value] + if isinstance(value, dict): + return {key: replace(item) for key, item in value.items()} + return value + + result = func(*replace(args), **replace(kwargs)) + copied = [ + (original, copies[id(original)]) + for original in originals.values() + if id(original) not in mutation_targets + ] + staged = _stage_local_buffers( + {str(index): original for index, (original, _) in enumerate(copied)}, + {str(index): copy for index, (_, copy) in enumerate(copied)}, + ) + with torch.no_grad(), torch._C.DisableTorchFunctionSubclass(): + for original, copy in staged: + original.copy_(copy) + cast(_ClientBuffer, original)._head_owner.pending = True + if mutating: + for original in originals.values(): + if id(original) in mutation_targets: + original._head_owner.pending = True + result = originals.get(id(result), result) + return readonly_buffer_views(result, [copy for _, copy in copied]) + + def __deepcopy__(self, memo: dict[int, Any]) -> torch.Tensor: + result = _plain(self) + memo[id(self)] = result + return result + + def __reduce_ex__(self, proto: SupportsIndex) -> Any: + return _plain(self).__reduce_ex__(proto) + + +class LiveHead: + """Client/sandbox state shared by all references to one registered head.""" + + def __init__( + self, + state: HeadState, + value: torch.nn.Module | torch.Tensor, + collector: Any, + *, + max_gradient_staleness: int | None = None, + ): + self.state = state + self.collector = collector + self.max_gradient_staleness = ( + state.max_gradient_staleness + if max_gradient_staleness is None + else max_gradient_staleness + ) + self.pending = False + self.invalid = False + self.invalid_reason = "its checkpoint was replaced" + if state.kind == "module": + assert isinstance(value, torch.nn.Module) + value = deepcopy(value) + self.source = value + replaced = {} + for key, parameter in value.named_parameters(): + replaced[id(parameter)] = _ClientParameter( + state.parameters[key].to(parameter.device), self, key + ) + for child in value.modules(): + for key, parameter in child._parameters.items(): + if parameter is not None: + child._parameters[key] = replaced[id(parameter)] + buffers = { + id(buffer): _ClientBuffer(_plain(buffer), self) + for buffer in value.buffers() + } + for child in value.modules(): + for key, buffer in child._buffers.items(): + if buffer is not None: + child._buffers[key] = buffers[id(buffer)] + self.value = ModuleHandle(value, self.capture_module, self.publish) + elif state.kind == "parameter": + assert isinstance(value, torch.Tensor) + self.source = self.value = _ClientParameter( + state.parameters[""].to(value.device), self, "" + ) + else: + assert isinstance(value, torch.Tensor) + self.source = self.value = _ClientBuffer( + state.buffers[""].to(value.device).clone(), self + ) + self.refresh(state) + + def _validate(self) -> None: + if self.invalid: + raise RuntimeError( + f"Custom checkpoint object {self.state.name!r} is stale because {self.invalid_reason}" + ) + + def invalidate(self, reason: str) -> None: + self.invalid = True + self.invalid_reason = reason + + def parameters(self) -> dict[str, torch.Tensor]: + if self.state.kind == "module": + return dict(cast(torch.nn.Module, self.source).named_parameters()) + return ( + {"": cast(torch.Tensor, self.source)} + if self.state.kind == "parameter" + else {} + ) + + def buffers(self) -> dict[str, torch.Tensor]: + if self.state.kind == "module": + return dict(cast(torch.nn.Module, self.source).named_buffers()) + return ( + {"": cast(torch.Tensor, self.source)} if self.state.kind == "buffer" else {} + ) + + def capture(self, keys: tuple[str, ...] | None = None) -> dict[str, torch.Tensor]: + import json + from uuid import uuid4 + + from ._tensors import detach_tree + + self._validate() + state = self.state + keys = tuple(state.parameters) if keys is None else keys + version = state.version + handle = "head:" + json.dumps( + { + "checkpoint": version.checkpoint, + "generation": version.generation, + "revision": version.revision, + "name": state.name, + "keys": keys, + "max_gradient_staleness": self.max_gradient_staleness, + "capture": uuid4().hex, + }, + separators=(",", ":"), + ) + current = self.parameters() + parameters = { + key: state.parameters[key].to( + device=current[key].device, dtype=current[key].dtype + ) + for key in keys + } + if not torch.is_grad_enabled(): + return {key: value.detach().clone() for key, value in parameters.items()} + from ._parameter_hooks import parameter_hooks + + registries = self.collector._head_hooks + registries[handle] = tuple(parameter_hooks(current[key]) for key in keys) + try: + return self.collector.attach( + detach_tree(handle, parameters), + on_release=lambda: registries.pop(handle, None), + ) + except BaseException: + registries.pop(handle, None) + raise + + def capture_module(self) -> tuple[dict[str, torch.Tensor], dict[str, torch.Tensor]]: + parameters = self.capture() + return parameters, {key: _plain(value) for key, value in self.buffers().items()} + + def publish(self, buffers: Mapping[str, torch.Tensor]) -> None: + self._validate() + staged = _stage_local_buffers(self.buffers(), buffers) + with torch.no_grad(), torch._C.DisableTorchFunctionSubclass(): + for target, value in staged: + target.copy_(value) + self.pending = self.pending or bool(staged) + + def refresh(self, state: HeadState) -> None: + if state.version.generation != self.state.version.generation: + self.invalid = True + return + self._validate() + if state.version.revision < self.state.version.revision: + return + if ( + state.kind != self.state.kind + or state.parameters.keys() != self.state.parameters.keys() + or state.buffers.keys() != self.state.buffers.keys() + ): + raise ValueError("Custom checkpoint object schema changed") + targets = self.parameters() | self.buffers() + keep_buffers = ( + self.pending or state.buffer_revision < self.state.buffer_revision + ) + if keep_buffers: + state = replace( + state, + buffers=self.state.buffers, + buffer_revision=self.state.buffer_revision, + ) + with torch.no_grad(), torch._C.DisableTorchFunctionSubclass(): + for key, value in ( + state.parameters | ({} if keep_buffers else state.buffers) + ).items(): + targets[key].copy_(value.to(targets[key].device)) + self.state = state + + def take_publication(self) -> HeadBufferUpdate | None: + self._validate() + current = self.buffers() + if not self.pending and all( + torch.equal(_plain(value).cpu(), self.state.buffers[key]) + for key, value in current.items() + ): + return None + self.pending = False + result = HeadBufferUpdate( + self.state.version, + self.state.name, + self.state.buffer_revision, + { + key: _plain(value).to(device="cpu", dtype=self.state.buffers[key].dtype) + for key, value in current.items() + }, + ) + # Publication is ordered before subsequent operations by the owner. Keep + # local revision in sequence for calls made before that operation resolves. + self.state = replace( + self.state, + buffers=result.buffers, + buffer_revision=self.state.buffer_revision + 1, + ) + return result + + +def synchronize_head_buffers(trainer: TrainerRank, checkpoints: Any = None) -> None: + """Publish logical DP zero's persistent buffers at an ordered boundary. + + Parameter gradients still reduce across DP. Arbitrary counters and running + statistics are copied from the authority, never averaged. + """ + import torch.distributed as dist + + from . import _checkpoint + + if not (dist.is_available() and dist.is_initialized()): + return + group = trainer._checkpoint_group() + names = sorted(trainer._checkpoint_slots if checkpoints is None else checkpoints) + targets = {} + for checkpoint in names: + for name, custom in trainer._checkpoint_slots[checkpoint].custom.items(): + if custom.kind == "parameter": + continue + if custom.kind == "buffer": + buffers = _custom_buffers(custom) + else: + persistent = { + id(buffer) + for child in cast(torch.nn.Module, custom.value).modules() + for key, buffer in child._buffers.items() + if buffer is not None + and key not in child._non_persistent_buffers_set + } + buffers = { + key: value + for key, value in _custom_buffers(custom).items() + if id(value) in persistent + } + targets[(checkpoint, name)] = (_custom_tracker(custom), buffers) + revisions = _checkpoint._gather( + {key: tracker.buffer_revision for key, (tracker, _) in targets.items()}, group + ) + if any(peer.keys() != targets.keys() for peer in revisions): + raise trainer._slot_state_error( + "Custom buffer registrations differ across ranks" + ) + payload = ( + { + key: ( + tracker.buffer_revision, + {name: _plain(value).cpu() for name, value in buffers.items()}, + ) + for key, (tracker, buffers) in targets.items() + } + if dist.get_rank(group) == 0 + else None + ) + authoritative = _checkpoint._gather(payload, group)[0] + assert authoritative is not None + if authoritative.keys() != targets.keys(): + raise trainer._slot_state_error( + "Custom buffer registrations differ across ranks" + ) + staged, local_differences, error = {}, {}, None + try: + for key, (tracker, buffers) in targets.items(): + staged[key] = _stage_local_buffers(buffers, authoritative[key][1]) + local_differences[key] = tracker.buffer_revision != authoritative[key][ + 0 + ] or bool(staged[key]) + except Exception as exc: + error = exc + _checkpoint.raise_distributed(error, "validate synchronized buffers", group) + differences = _checkpoint._gather(local_differences, group) + with torch.no_grad(), torch._C.DisableTorchFunctionSubclass(): + for key, (tracker, _) in targets.items(): + for target, value in staged[key]: + target.copy_(value) + tracker.buffer_revision = max(peer[key] for peer in revisions) + int( + any(peer[key] for peer in differences) + ) + + +def _logical_heads(view: Any) -> dict[tuple[str, str], LiveHead]: + rank = view._rank + if not hasattr(rank, "_logical_head_handles"): + rank._logical_head_handles = {} + return rank._logical_head_handles + + +def logical_register_head( + view: Any, + kind: HeadKind, + name: str, + factory: Callable[[], Any], + *, + checkpoint: Any = ..., +) -> Any: + from ._options import resolve_forward_options + + if checkpoint is ...: + from ._impl import Unset + + checkpoint = Unset + + flush_logical_heads(view) + + checkpoint, state = view._invoke("head", "head_lookup", (checkpoint, name)) + registry = _logical_heads(view) + current = registry.get((checkpoint, name)) + if current is not None and not current.invalid and state is not None: + current.refresh(state) + if not current.invalid: + if state.kind != kind: + raise ValueError( + f"Checkpoint object {name!r} is already a {state.kind}" + ) + return current.value + if state is None: + value = factory() + state = view._invoke( + "head", "head_register", HeadRegistration(checkpoint, name, kind, value) + ) + else: + value = deepcopy(view._rank._checkpoint_slots[checkpoint].custom[name].value) + value = ( + move_module(value, view.device) + if isinstance(value, torch.nn.Module) + else value.to(view.device) + ) + maximum = resolve_forward_options( + getattr(view._rank, "_forward_options", None) + ).max_gradient_staleness + head = LiveHead( + state, value, view._executor.state.collector, max_gradient_staleness=maximum + ) + registry[(checkpoint, name)] = head + return head.value + + +def flush_logical_heads(view: Any) -> None: + updates = tuple( + update + for head in _logical_heads(view).values() + if not head.invalid and (update := head.take_publication()) is not None + ) + if updates: + try: + view._executor.invoke("head", "head_publish", updates) + except BaseException: + for update in updates: + _logical_heads(view)[ + (update.version.checkpoint, update.name) + ].invalidate("buffer publication failed; register the object again") + raise + + +def refresh_logical_heads(view: Any) -> None: + registry = _logical_heads(view) + keys = tuple(key for key, head in registry.items() if not head.invalid) + if keys: + states = view._executor.invoke("head", "head_export", keys) + for key, state in zip(keys, states, strict=True): + if state is None: + registry[key].invalid = True + else: + registry[key].refresh(state) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 4ee194dfb..7892565f4 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -15,9 +15,9 @@ from concurrent.futures import Future, ThreadPoolExecutor from contextlib import contextmanager from copy import deepcopy -from dataclasses import dataclass, replace +from dataclasses import dataclass, fields, is_dataclass, replace from dataclasses import field as dataclass_field -from functools import partial +from functools import lru_cache, partial import hashlib import logging import math @@ -54,6 +54,20 @@ _local_position_pairs, estimate_prefix_tree_packed_tokens, ) +from art.trainer_rank._memory_policy import ( + ForwardMemoryCost, + MemoryPlacement, + host_memory_budget, + local_rank_count, + placement_cost, +) +from art.trainer_rank._options import ( + ForwardOptions, + ResolvedForwardOptions, + Unset, + _Unset, + resolve_forward_options, +) from art.trainer_rank._planner_cost import ( COEFFICIENT_VERSION_FALLBACK, ModelGeometry, @@ -69,6 +83,11 @@ select_prefix_tree_layout, ) from art.trainer_rank._telemetry import phase as _telemetry_phase +from art.trainer_rank._versions import ( + CheckpointVersion, + CheckpointVersions, + VersionedGradient, +) if TYPE_CHECKING: from megatron.core.models.gpt.gpt_model import GPTModel @@ -78,7 +97,7 @@ ArtContextParallelState, ParallelTopology, ) - from art.megatron.lora import LoRASlotRef + from art.megatron.lora import LoRASlotRef, LoRAVersion from art.megatron.prefix_tree_state import PrefixTreeAttentionState from art.megatron.train import TrainingRuntime from art.trainer_rank._checkpoint import ( @@ -91,6 +110,8 @@ ) from art.trainer_rank._lora_export import _PreparedLoraExport + from ._heads import ModuleHandle + @dataclass(frozen=True) class AdamParams: @@ -170,11 +191,6 @@ class _AdapterConfig(TypedDict): hidden_size: NotRequired[int] -class _Unset: - pass - - -Unset = _Unset() type AdapterSelection = str | None | _Unset @@ -204,6 +220,7 @@ class ForwardInput(Generic[LogprobsT, TopKT, LogitsT, HiddenStatesT]): hidden_states: bool = False no_grad: bool | None = None checkpoint: AdapterSelection = Unset + options: ForwardOptions | None = None @overload def __new__( @@ -216,6 +233,7 @@ def __new__( hidden_states: Literal[False] = False, no_grad: bool | None = None, checkpoint: AdapterSelection = Unset, + options: ForwardOptions | None = None, ) -> "ForwardInput[None, None, None, None]": ... @overload @@ -229,6 +247,7 @@ def __new__( hidden_states: Literal[False] = False, no_grad: bool | None = None, checkpoint: AdapterSelection = Unset, + options: ForwardOptions | None = None, ) -> "ForwardInput[torch.Tensor, None, None, None]": ... @overload @@ -242,6 +261,7 @@ def __new__( hidden_states: Literal[False] = False, no_grad: bool | None = None, checkpoint: AdapterSelection = Unset, + options: ForwardOptions | None = None, ) -> "ForwardInput[None, TopK, None, None]": ... @overload @@ -255,6 +275,7 @@ def __new__( hidden_states: Literal[False] = False, no_grad: bool | None = None, checkpoint: AdapterSelection = Unset, + options: ForwardOptions | None = None, ) -> "ForwardInput[None, None, torch.Tensor, None]": ... @overload @@ -268,6 +289,7 @@ def __new__( hidden_states: Literal[True], no_grad: bool | None = None, checkpoint: AdapterSelection = Unset, + options: ForwardOptions | None = None, ) -> "ForwardInput[None, None, None, torch.Tensor]": ... @overload @@ -281,6 +303,7 @@ def __new__( hidden_states: Literal[False] = False, no_grad: bool | None = None, checkpoint: AdapterSelection = Unset, + options: ForwardOptions | None = None, ) -> "ForwardInput[torch.Tensor, TopK, None, None]": ... @overload @@ -294,6 +317,7 @@ def __new__( hidden_states: Literal[False] = False, no_grad: bool | None = None, checkpoint: AdapterSelection = Unset, + options: ForwardOptions | None = None, ) -> "ForwardInput[torch.Tensor, None, torch.Tensor, None]": ... @overload @@ -307,6 +331,7 @@ def __new__( hidden_states: Literal[True], no_grad: bool | None = None, checkpoint: AdapterSelection = Unset, + options: ForwardOptions | None = None, ) -> "ForwardInput[torch.Tensor, None, None, torch.Tensor]": ... @overload @@ -320,6 +345,7 @@ def __new__( hidden_states: Literal[False] = False, no_grad: bool | None = None, checkpoint: AdapterSelection = Unset, + options: ForwardOptions | None = None, ) -> "ForwardInput[None, TopK, torch.Tensor, None]": ... @overload @@ -333,6 +359,7 @@ def __new__( hidden_states: Literal[True], no_grad: bool | None = None, checkpoint: AdapterSelection = Unset, + options: ForwardOptions | None = None, ) -> "ForwardInput[None, TopK, None, torch.Tensor]": ... @overload @@ -346,6 +373,7 @@ def __new__( hidden_states: Literal[True], no_grad: bool | None = None, checkpoint: AdapterSelection = Unset, + options: ForwardOptions | None = None, ) -> "ForwardInput[None, None, torch.Tensor, torch.Tensor]": ... @overload @@ -359,6 +387,7 @@ def __new__( hidden_states: Literal[False] = False, no_grad: bool | None = None, checkpoint: AdapterSelection = Unset, + options: ForwardOptions | None = None, ) -> "ForwardInput[torch.Tensor, TopK, torch.Tensor, None]": ... @overload @@ -372,6 +401,7 @@ def __new__( hidden_states: Literal[True], no_grad: bool | None = None, checkpoint: AdapterSelection = Unset, + options: ForwardOptions | None = None, ) -> "ForwardInput[torch.Tensor, TopK, None, torch.Tensor]": ... @overload @@ -385,6 +415,7 @@ def __new__( hidden_states: Literal[True], no_grad: bool | None = None, checkpoint: AdapterSelection = Unset, + options: ForwardOptions | None = None, ) -> "ForwardInput[torch.Tensor, None, torch.Tensor, torch.Tensor]": ... @overload @@ -398,6 +429,7 @@ def __new__( hidden_states: Literal[True], no_grad: bool | None = None, checkpoint: AdapterSelection = Unset, + options: ForwardOptions | None = None, ) -> "ForwardInput[None, TopK, torch.Tensor, torch.Tensor]": ... @overload @@ -411,6 +443,7 @@ def __new__( hidden_states: Literal[True], no_grad: bool | None = None, checkpoint: AdapterSelection = Unset, + options: ForwardOptions | None = None, ) -> "ForwardInput[torch.Tensor, TopK, torch.Tensor, torch.Tensor]": ... @overload @@ -424,6 +457,7 @@ def __new__( hidden_states: bool = False, no_grad: bool | None = None, checkpoint: AdapterSelection = Unset, + options: ForwardOptions | None = None, ) -> "ForwardInput[torch.Tensor | None, TopK | None, torch.Tensor | None, torch.Tensor | None]": ... def __new__( @@ -436,6 +470,7 @@ def __new__( hidden_states: bool = False, no_grad: bool | None = None, checkpoint: AdapterSelection = Unset, + options: ForwardOptions | None = None, ) -> Self: return object.__new__(cls) @@ -449,6 +484,7 @@ def __init__( hidden_states: bool = False, no_grad: bool | None = None, checkpoint: AdapterSelection = Unset, + options: ForwardOptions | None = None, ) -> None: self.input_tokens = input_tokens self.target_tokens = target_tokens @@ -457,8 +493,12 @@ def __init__( self.hidden_states = hidden_states self.no_grad = no_grad self.checkpoint = checkpoint + self.options = options self.__post_init__() + def __getnewargs_ex__(self) -> tuple[tuple[()], dict[str, torch.Tensor]]: + return (), {"input_tokens": self.input_tokens} + def __post_init__(self) -> None: if self.top_k is not None and self.top_k < 1: raise ValueError("top_k must be >= 1") @@ -584,6 +624,10 @@ class _MemoryCheck: estimated_required_bytes: int available_bytes: int fits: bool + cpu_required_bytes: int = 0 + cpu_available_bytes: int = 0 + cpu_fits: bool = True + fallback_costs: dict[str, Any] | None = None @dataclass(frozen=True) @@ -612,23 +656,6 @@ class _CandidateMicroBatch(Generic[ForwardInputsT]): cold_start: bool -class _SlotGraphSentinel(torch.autograd.Function): - @staticmethod - def forward( - ctx: FunctionCtx, - tensor: torch.Tensor, - marker: torch.Tensor, - ) -> torch.Tensor: - ctx.save_for_backward(marker) - return tensor - - @staticmethod - def backward( - ctx: FunctionCtx, *grad_outputs: torch.Tensor - ) -> tuple[torch.Tensor, None]: - return grad_outputs[0], None - - class _GatherContextParallelRows(torch.autograd.Function): @staticmethod def forward( @@ -684,6 +711,17 @@ def finish() -> None: return grad_outputs[0], None +def _track_slot_graph_tensor( + tensor: torch.Tensor, marker: torch.Tensor +) -> torch.Tensor: + # This Function saves only a CPU, non-gradient control marker. Preserve its + # wrapper identity even under caller hooks, without intercepting activations. + with torch.autograd.graph.saved_tensors_hooks( + lambda value: value, lambda value: value + ): + return cast(torch.Tensor, _CustomSlotGraphSentinel.apply(tensor, marker)) + + @dataclass(eq=False) class _CustomTensorTracker: trainer: weakref.ReferenceType[TrainerRank] @@ -691,6 +729,7 @@ class _CustomTensorTracker: name: str generation: object active: bool = False + buffer_revision: int = 0 def validate(self) -> TrainerRank: trainer = self.trainer() @@ -778,6 +817,18 @@ def __init__( ) -> None: del data, tracker, requires_grad + def register_hook(self, hook: Any) -> Any: + """Run once on summed captured uses per backward; removal affects old graphs.""" + from ._parameter_hooks import register_parameter_hook + + self._art_tracker.validate() + return register_parameter_hook(self, hook) + + def register_post_accumulate_grad_hook(self, hook: Any) -> Any: + from ._parameter_hooks import reject_post_accumulate_hook + + return reject_post_accumulate_hook() + def __setattr__(self, name: str, value: object) -> None: if name == "grad": with torch._C.DisableTorchFunctionSubclass(): @@ -827,6 +878,7 @@ class _CustomObject: kind: Literal["module", "parameter", "buffer"] value: torch.nn.Module | torch.nn.Parameter | torch.Tensor generation: object + handle: torch.nn.Module | None = None @dataclass @@ -838,6 +890,7 @@ class _CheckpointSlot: custom: dict[str, _CustomObject] = dataclass_field(default_factory=dict) custom_payload: "PreparedCustomPayload | None" = None snapshot: bool = False + generation: int = 0 @dataclass(frozen=True) @@ -924,6 +977,7 @@ class _MemorySignature: request_mix: tuple[str, ...] grad_enabled: bool grad_modes: tuple[bool, ...] + memory_placement: tuple[tuple[str, str], ...] = () @dataclass(frozen=True) @@ -933,6 +987,7 @@ class _ForwardGroupPlan: request_indices: tuple[int, ...] items: tuple[_ForwardItem, ...] packed: PrefixTreePack + memory_placement: MemoryPlacement | None = None @dataclass(frozen=True) @@ -1029,7 +1084,7 @@ def ephemeral(self) -> int: _MEMORY_ERROR_SUGGESTION = ( "Use smaller top-level items, reduce output requests, or call " - "dp_rank_forward with already-DP-local smaller inputs." + "forward with already-DP-local smaller inputs." ) @@ -1047,7 +1102,13 @@ def _memory_error( f"logical_tokens={logical_tokens} " f"predicted_peak_gb={check.estimated_required_bytes / 1024**3:.3f} " f"usable_limit_gb={check.available_bytes / 1024**3:.3f}. " - f"{_MEMORY_ERROR_SUGGESTION}", + + ( + f"CPU retained bytes={check.cpu_required_bytes}, " + f"per-rank CPU headroom={check.cpu_available_bytes}. " + if not check.cpu_fits + else "" + ) + + f"{_MEMORY_ERROR_SUGGESTION}", predicted_peak_bytes=check.estimated_required_bytes, usable_limit_bytes=check.available_bytes, suggestion=_MEMORY_ERROR_SUGGESTION, @@ -1368,7 +1429,9 @@ def _moe_output_bytes_per_token( class TrainerRank: - def __init__(self, runtime: TrainingRuntime) -> None: + def __init__( + self, runtime: TrainingRuntime, *, options: ForwardOptions | None = None + ) -> None: pp_size = int(getattr(runtime.provider, "pipeline_model_parallel_size", 1) or 1) if pp_size > 1 or len(runtime.model) > 1: raise TrainerRankRuntimeSupportError( @@ -1383,6 +1446,8 @@ def __init__(self, runtime: TrainingRuntime) -> None: # TP calibrates itself online, and the fitted layout cost model prices # TP explicitly. The cold retained-activation floor also distinguishes # tensor/sequence-parallel storage from gathered LoRA inputs. + self._forward_options = options + resolve_forward_options(options) self.runtime: TrainingRuntime = runtime self.device: torch.device = next(runtime.model[0].parameters()).device self._param_dtype_size = _dtype_size(next(runtime.model[0].parameters()).dtype) @@ -1521,6 +1586,9 @@ def memory_field(name: str, default: Any = None) -> Any: self._hybridep_rows_high_water = 0 self._cache_recovery_state = _CacheRecoveryState() self._memory_profiles: dict[_MemorySignature, _MemoryProfile] = {} + self._graph_forward_times: OrderedDict[tuple[Any, ...], tuple[float, ...]] = ( + OrderedDict() + ) self._split_memory_floors: dict[bytes, int] = {} self._split_memory_floor_status = "not_observed" self._last_global_micro_batch_size: int | None = None @@ -1553,22 +1621,178 @@ def zero_grad(self) -> None: for slot in self._checkpoint_slots.values(): for param in slot.params: param.grad = None + self._version_state().clear() self._prune_slot_graphs() + def _version_state(self) -> CheckpointVersions: + state = getattr(self, "_checkpoint_versions", None) + if state is None: + state = self._checkpoint_versions = CheckpointVersions(self) + return state + + def _capture_checkpoint_version(self, name: str) -> CheckpointVersion: + return self._version_state().capture(name) + + def _validate_checkpoint_version( + self, version: CheckpointVersion, max_gradient_staleness: int = 2 + ) -> None: + self._version_state().validate(version, max_gradient_staleness) + + def _snapshot_parameter( + self, + parameter: torch.nn.Parameter, + version: CheckpointVersion, + max_gradient_staleness: int = 2, + ) -> torch.nn.Parameter: + return self._version_state().snapshot( + parameter, version, max_gradient_staleness + ) + + def _commit_versioned_gradients( + self, gradients: Sequence[VersionedGradient] + ) -> None: + self._version_state().accumulate(gradients) + + def _gradient_transaction( + self, *, before_commit: Callable[[Callable[[], None]], None] | None = None + ) -> Any: + return self._version_state().transaction(before_commit=before_commit) + + def _capture_lora_version( + self, + ref: LoRASlotRef | None, + max_gradient_staleness: int = 2, + *, + origin: CheckpointVersion | None = None, + ) -> LoRAVersion | None: + if ref is None or ref.name is None or not torch.is_grad_enabled(): + return None + from art.megatron.lora import LoRA, LoRASlot, LoRAVersion + + state = self._version_state() + weight_version = state.capture(ref.name) + version = weight_version if origin is None else origin + if version.checkpoint != ref.name: + raise ValueError("LoRA replay origin belongs to a different checkpoint") + state.validate(version, max_gradient_staleness) + key = (weight_version, version, max_gradient_staleness) + if cached := state.lora.get(key): + return cached + slots: dict[int, LoRASlot] = {} + for chunk in self.runtime.model: + for module in chunk.modules(): + if not isinstance(module, LoRA) or id(module) in slots: + continue + current = module._slot(ref) + if current is None: + continue + captured = slots[id(module)] = LoRASlot( + ref=ref, + a_t=current.A_T, + b_t=current.B_T, + alpha=current.alpha, + a_template=current.A_T, + b_template=current.B_T, + requires_grad=current.A_T.requires_grad, + ) + for snapshot, parameter in zip( + (captured.A_T, captured.B_T), + (current.A_T, current.B_T), + strict=True, + ): + snapshot.requires_grad_(parameter.requires_grad) + state.track(snapshot, parameter, version, max_gradient_staleness) + captured_version = LoRAVersion( + ref, + version, + slots, + lambda: state.validate(version, max_gradient_staleness), + weight_version, + ) + state.lora[key] = captured_version + return captured_version + + def _lora_version_capture_bytes( + self, + ref: LoRASlotRef | None, + max_gradient_staleness: int = 2, + *, + origin: CheckpointVersion | None = None, + ) -> int: + if ref is None or ref.name is None or not torch.is_grad_enabled(): + return 0 + state = self._version_state() + version = state.capture(ref.name) + key = (version, version if origin is None else origin, max_gradient_staleness) + if key in state.lora: + return 0 + return sum( + param.numel() * param.element_size() + for param in self._checkpoint_slots[ref.name].params + if not getattr(param, "_art_custom_checkpoint_param", False) + ) + + def _lora_gradient_staging_bytes(self, ref: LoRASlotRef | None) -> int: + """Reserve current, staged and replacement gradients across later backwards. + + Several forwards can be admitted before their first backward creates + ``.grad``. Existing gradients are already in the sampled memory baseline. + Registered custom heads share this transaction, including streamed remote + cotangents. Later registrations and arbitrary head activations are not + predicted by an earlier model forward. + """ + if ref is None or ref.name is None: + return 0 + return _gradient_staging_bytes(self._checkpoint_slots[ref.name].params) + + def _pending_backward_memory( + self, *, checkpoints: Iterable[str] = (), exclude_staging: Iterable[str] = () + ) -> tuple[int, int]: + """Return additional restore and gradient bytes beside sampled live storage.""" + cache = getattr(self, "_graph_cache", None) + states = () if cache is None else tuple(cache.state(h) for h in cache.handles()) + names = set(checkpoints) | { + version.checkpoint + for state in states + for version in getattr(state, "checkpoint_versions", ()) + } + excluded = set(exclude_staging) + staging = sum( + self._lora_gradient_staging_bytes(self._slot_ref(name)) + for name in names - excluded + if name in self._checkpoint_slots + ) + # Head losses can remain live without a model cache record. Preserve their + # registration reserve without charging unrelated unused LoRA targets. + staging += _gradient_staging_bytes( + parameter + for name, slot in self._checkpoint_slots.items() + if name not in names | excluded + for parameter in slot.params + if getattr(parameter, "_art_custom_checkpoint_param", False) + ) + return ( + max( + (getattr(state, "restore_workspace_bytes", 0) for state in states), + default=0, + ), + staging, + ) + def module( self, name: str, factory: Callable[[], ModuleT], *, checkpoint: AdapterSelection = Unset, - ) -> ModuleT: + ) -> ModuleHandle: """Return a checkpoint-owned module, registering it on first access. Registration is collective across TrainerRank processes. The returned module is bound to the resolved checkpoint and is not selected by later push/pop calls. """ value = self._custom_object(name, "module", factory, checkpoint=checkpoint) - return cast(ModuleT, value) + return cast("ModuleHandle", value) def parameter( self, @@ -1640,7 +1864,7 @@ def _custom_object( f"Checkpoint {checkpoint_name!r} already registers {name!r} " f"as a {existing.kind}, not a {kind}" ) - return existing.value + return existing.handle if existing.handle is not None else existing.value custom: _CustomObject | None = None try: value = factory() @@ -1648,7 +1872,9 @@ def _custom_object( if kind == "module": if not isinstance(value, torch.nn.Module): raise TypeError("module() factory must return torch.nn.Module") - value = value.to(device=self.device) + from ._heads import move_module + + value = move_module(deepcopy(value), self.device) if slot.snapshot: value.requires_grad_(False) elif kind == "parameter": @@ -1692,25 +1918,52 @@ def _custom_object( extended_optimizer = self._extend_dynamic_optimizer( checkpoint_name, named_params ) + self._admit_custom_gradient_storage(checkpoint_name, new_params) except BaseException as exc: error = exc - _checkpoint.raise_distributed( - error, f"stage custom checkpoint object {name!r}", group - ) - assert tracker is not None + try: + _checkpoint.raise_distributed( + error, f"stage custom checkpoint object {name!r}", group + ) + except BaseException: + # A retained registration traceback must not own rejected tensors. + value = custom = tracker = extended_optimizer = None + named_params = new_params = () + raise + assert tracker is not None and custom is not None if extended_optimizer is not None: slot.optimizer = extended_optimizer slot.custom[name] = custom slot.params += new_params tracker.active = True - return custom.value + return custom.handle if custom.handle is not None else custom.value + + def _admit_custom_gradient_storage( + self, checkpoint: str, parameters: Sequence[torch.nn.Parameter] + ) -> None: + """Price known new targets beside existing graph restore reservations.""" + try: + if not self._graph_memory_policy_enabled(): + return + workspace, staging = self._pending_backward_memory( + checkpoints=(checkpoint,) + ) + required = _gradient_staging_bytes(parameters) + staging + workspace + available = self._available_memory_bytes() + if required > available: + raise TrainerRankMemoryError( + f"Registering custom parameters needs {required} GPU bytes for " + f"gradient staging and existing graph restoration; available={available}" + ) + finally: + parameters = () 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) + name = checkpoint if name is None: raise TrainerRankSlotStateError( "Custom checkpoint objects require a loaded named checkpoint" @@ -2208,10 +2461,11 @@ def _validate_loaded_checkpoint_config( ) @overload - def forward_micro_batches( + def forward_batches( self, inputs: Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]], *, + options: ForwardOptions | None = None, checkpoint: AdapterSelection = Unset, no_grad: bool | None = None, yield_empty: bool = False, @@ -2223,12 +2477,13 @@ def forward_micro_batches( ]: ... @overload - def forward_micro_batches( + def forward_batches( self, inputs: Iterable[ Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]] ], *, + options: ForwardOptions | None = None, checkpoint: AdapterSelection = Unset, no_grad: bool | None = None, yield_empty: bool = False, @@ -2240,12 +2495,13 @@ def forward_micro_batches( ]: ... @overload - def forward_micro_batches( + def forward_batches( self, inputs: Iterable[ Iterable[Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]]] ], *, + options: ForwardOptions | None = None, checkpoint: AdapterSelection = Unset, no_grad: bool | None = None, yield_empty: bool = False, @@ -2257,7 +2513,7 @@ def forward_micro_batches( ]: ... @overload - def forward_micro_batches( + def forward_batches( self, inputs: Iterable[ Iterable[ @@ -2267,6 +2523,7 @@ def forward_micro_batches( ] ], *, + options: ForwardOptions | None = None, checkpoint: AdapterSelection = Unset, no_grad: bool | None = None, yield_empty: bool = False, @@ -2285,10 +2542,11 @@ def forward_micro_batches( ] ]: ... - def forward_micro_batches( + def forward_batches( self, inputs: Iterable[ForwardInputs], *, + options: ForwardOptions | None = None, checkpoint: AdapterSelection = Unset, no_grad: bool | None = None, yield_empty: bool = False, @@ -2296,13 +2554,27 @@ def forward_micro_batches( if not isinstance(yield_empty, bool): raise TypeError("yield_empty must be a bool") enabled = torch.is_grad_enabled() if no_grad is None else not no_grad - batches = self._forward_micro_batches( + inputs = cast( + Iterable[ForwardInputs], self._capture_forward_options(inputs, options) + ) + batches = self._forward_batches( inputs, checkpoint=checkpoint, yield_empty=yield_empty ) + return self._yield_forward_batches( + batches, enabled=enabled, yield_empty=yield_empty + ) + + def _yield_forward_batches( + self, + batches: Generator[MicroBatch[ForwardInputs, ForwardOutputs], None, None], + *, + enabled: bool, + yield_empty: bool, + ) -> Iterator[MicroBatch[ForwardInputs, ForwardOutputs]]: token = object() try: while True: - self._guard_forward_collective("forward_micro_batches") + self._guard_forward_collective("forward_batches") with torch.set_grad_enabled(enabled): try: batch = next(batches) @@ -2327,17 +2599,46 @@ def forward_micro_batches( finally: batches.close() + def _capture_forward_options( + self, inputs: ForwardInputs, options: ForwardOptions | None + ) -> ForwardInputs: + from ._graphs import _snapshot + + # Input enumeration already happens at submission; only execution is + # lazy. Own the submitted tensor storage before returning an iterator. + materialized = _snapshot(_materialize(inputs)) + constructor = getattr(self, "_forward_options", None) + from dataclasses import fields + + def capture(value: ForwardInputs) -> ForwardInputs: + if isinstance(value, ForwardInput): + if constructor is None and options is None: + return replace(value) + resolved = resolve_forward_options(constructor, options, value.options) + return replace( + value, + options=ForwardOptions( + **{ + field.name: getattr(resolved, field.name) + for field in fields(resolved) + } + ), + ) + return _rebuild_forward_tree(value, [capture(child) for child in value]) + + return capture(materialized) + def _guard_forward_collective(self, operation: str) -> None: for thread, start, stop in tuple(self._skipped_forward_waves.values()): if thread == threading.get_ident(): raise RuntimeError( - f"{operation} cannot run during forward_micro_batches wave " + f"{operation} cannot run during forward_batches wave " f"[{start}, {stop}): yield_empty=False skips some data-parallel " "ranks. Move collective calls after the iterator or use " "yield_empty=True on every rank." ) - def _forward_micro_batches( + def _forward_batches( self, inputs: Iterable[ForwardInputs], *, @@ -2368,7 +2669,7 @@ def _forward_micro_batches( self._run_flat_plan_with_memory_tracking( candidate.plan, check=candidate.check, - context="forward_micro_batches", + context="forward_batches", ) ) else: @@ -2376,7 +2677,7 @@ def _forward_micro_batches( self._execute_split_plan_with_memory_tracking( candidate.plan, check=candidate.check, - context="forward_micro_batches", + context="forward_batches", ) ) flat_outputs = iter(tracked_outputs) @@ -2435,21 +2736,33 @@ def _forward_micro_batches( start = stop @overload - def dp_rank_forward( + def forward( + self, + inputs: ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT], + *, + options: ForwardOptions | None = None, + checkpoint: AdapterSelection = Unset, + no_grad: bool | None = None, + ) -> ForwardOutput[LogprobsT, TopKT, LogitsT, HiddenStatesT]: ... + + @overload + def forward( self, inputs: Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]], *, + options: ForwardOptions | None = None, checkpoint: AdapterSelection = Unset, no_grad: bool | None = None, ) -> Sequence[ForwardOutput[LogprobsT, TopKT, LogitsT, HiddenStatesT]]: ... @overload - def dp_rank_forward( + def forward( self, inputs: Iterable[ Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]] ], *, + options: ForwardOptions | None = None, checkpoint: AdapterSelection = Unset, no_grad: bool | None = None, ) -> Sequence[ @@ -2457,12 +2770,13 @@ def dp_rank_forward( ]: ... @overload - def dp_rank_forward( + def forward( self, inputs: Iterable[ Iterable[Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]]] ], *, + options: ForwardOptions | None = None, checkpoint: AdapterSelection = Unset, no_grad: bool | None = None, ) -> Sequence[ @@ -2470,7 +2784,7 @@ def dp_rank_forward( ]: ... @overload - def dp_rank_forward( + def forward( self, inputs: Iterable[ Iterable[ @@ -2480,6 +2794,7 @@ def dp_rank_forward( ] ], *, + options: ForwardOptions | None = None, checkpoint: AdapterSelection = Unset, no_grad: bool | None = None, ) -> Sequence[ @@ -2488,24 +2803,25 @@ def dp_rank_forward( ] ]: ... - def dp_rank_forward( + def forward( self, inputs: ForwardInputs, *, + options: ForwardOptions | None = None, checkpoint: AdapterSelection = Unset, no_grad: bool | None = None, ) -> ForwardOutputs: - self._guard_forward_collective("dp_rank_forward") + self._guard_forward_collective("forward") enabled = torch.is_grad_enabled() if no_grad is None else not no_grad with torch.set_grad_enabled(enabled): self._reset_planning_telemetry() - materialized = _materialize(inputs) + materialized = self._capture_forward_options(inputs, options) requests = list(_flatten(materialized)) plan, check = self._plan_admissible_forward( - requests, checkpoint=checkpoint, context="dp_rank_forward" + requests, checkpoint=checkpoint, context="forward" ) tracked_outputs = self._execute_admitted_plan( - plan, check=check, context="dp_rank_forward" + plan, check=check, context="forward" ) return _unflatten(materialized, iter(tracked_outputs)) @@ -2616,7 +2932,10 @@ def _find_admissible_forward( plan = self._plan_flat_forward( requests, checkpoint=checkpoint, ensure_slots=False ) - check = self._memory_check(plan) + if self._graph_memory_policy_enabled(): + plan, check = self._admit_graph_memory(plan) + else: + check = self._memory_check(plan) if check.fits: return plan, check # Best effort before splitting: the memory-minimal (full sharing) @@ -2624,7 +2943,10 @@ def _find_admissible_forward( plan = self._plan_flat_forward( requests, checkpoint=checkpoint, memory_minimal=True, ensure_slots=False ) - check = self._memory_check(plan) + if self._graph_memory_policy_enabled(): + plan, check = self._admit_graph_memory(plan) + else: + check = self._memory_check(plan) if check.fits: return plan, check request_count = len(requests) @@ -2700,17 +3022,19 @@ def _admit_split_rung( more than one rung may need exact planning before one executes. """ - lower = [ - self._split_chunk_lower_cost( - [requests[index] for index in chunk], - [rows[index] for index in chunk], - checkpoint=checkpoint, - ) - for chunk in chunks - ] - check = self._split_rung_check(lower) - if not check.fits: - return None, check + managed = self._graph_memory_policy_enabled() + if not managed: + lower = [ + self._split_chunk_lower_cost( + [requests[index] for index in chunk], + [rows[index] for index in chunk], + checkpoint=checkpoint, + ) + for chunk in chunks + ] + check = self._split_rung_check(lower) + if not check.fits: + return None, check for memory_minimal in (False, True): plans = [ self._plan_flat_forward( @@ -2729,7 +3053,10 @@ def _admit_split_rung( request_indices=tuple(tuple(chunks[i]) for i in order), request_count=len(requests), ) - check = self._split_plan_memory_check(split, costs) + if managed: + split, check = self._admit_graph_memory(split) + else: + check = self._split_plan_memory_check(split, costs) if check.fits: return split, check return None, check @@ -2793,6 +3120,7 @@ def feed(value: Any) -> None: signature.request_mix, signature.grad_enabled, signature.grad_modes, + signature.memory_placement, p.packed_tokens, p.logical_tokens, p.inactive_logical_tokens, @@ -3080,12 +3408,16 @@ def _snapshot_planning_telemetry( "predicted_peak_bytes": check.estimated_required_bytes, "usable_limit_bytes": check.available_bytes, } + if check.fallback_costs is not None: + self._last_forward_telemetry_snapshot["fallback_costs"] = ( + check.fallback_costs + ) def last_forward_telemetry(self) -> dict[str, Any]: """Concise planner telemetry for the most recent planned forward. ``planning_ms`` is critical-path planning accumulated across the whole - public call (all waves of ``forward_micro_batches``, including the + public call (all waves of ``forward_batches``, including the synchronous cost of submitting speculative work); ``speculative_planning_ms`` is worker CPU time hidden under the caller's GPU work; ``selected_max_depth`` describes the most recently @@ -3100,13 +3432,13 @@ def last_forward_telemetry(self) -> dict[str, Any]: raise RuntimeError("no forward has completed planning yet") return dict(self._last_forward_telemetry_snapshot) - def dp_reduce( + def reduce( self, tensor: torch.Tensor, *, op: dist.ReduceOp.RedOpType = dist.ReduceOp.SUM, ) -> None: - self._guard_forward_collective("dp_reduce") + self._guard_forward_collective("reduce") from megatron.core import parallel_state as ps # Public outputs are CP-replicated; internal shard reductions still include CP. @@ -3116,6 +3448,36 @@ def dp_reduce( group=ps.get_data_parallel_group(with_context_parallel=False), ) + def backward( + self, + loss: torch.Tensor | Sequence[torch.Tensor], + gradient: torch.Tensor | Sequence[torch.Tensor | None] | None = None, + *, + retain_graph: bool = False, + ) -> None: + """Collect a complete local backward before committing model cotangents.""" + + from ._commands import _coordinate_call + + preflight = partial(_coordinate_call, group=self._forward_memory_group()) + with self._gradient_transaction(before_commit=preflight): + packets = [] + cache = self._forward_graph_cache() + + def collect() -> None: + packets.extend( + (packet.handle, packet.gradients) + for packet in self._forward_cotangent_collector().backward( + loss, gradient, retain_graph=retain_graph + ) + ) + cache.validate_many(packets) + + preflight(collect) + cache.backward_many( + packets, retain_graph=retain_graph, coordinate=preflight + ) + def optim_step( self, *, @@ -3203,11 +3565,15 @@ def optim_step( "optim", {"checkpoint_count": len(selected_checkpoints)}, ): - return self._dynamic_optim_step( + metrics = self._dynamic_optim_step( selected_checkpoints, params=params_by_checkpoint, scale_grads=scales_by_checkpoint, ) + from ._heads import synchronize_head_buffers + + synchronize_head_buffers(self, selected_checkpoints) + return metrics def _guard_optim_step_configuration( self, @@ -3429,7 +3795,7 @@ def _resolve_slot_ref( request.checkpoint if request.checkpoint is not Unset else checkpoint ) if selection is not Unset: - name = cast(str | None, selection) + name = selection if name is not None and name not in self._checkpoint_slots: raise TrainerRankSlotStateError( f"Forward selects unloaded checkpoint {name!r}" @@ -3529,6 +3895,16 @@ def _dynamic_optim_step( params: Mapping[str, AdamParams], scale_grads: Mapping[str, float], ) -> dict[str, float]: + from ._checkpoint import raise_distributed + + version_error: Exception | None = None + try: + self._version_state().validate_accumulated(checkpoint_names) + except Exception as exc: + version_error = exc + raise_distributed( + version_error, "validate gradient versions", self._checkpoint_group() + ) self.runtime.model_support_handler.zero_internal_padding_grads( self.runtime.model ) @@ -3567,6 +3943,7 @@ def _dynamic_optim_step( for param in self._checkpoint_slots[name].params: param.grad = None self._prune_slot_graphs(self._slot_ref(name)) + self._version_state().clear(checkpoint_names) return metrics previous = { name: ( @@ -3621,6 +3998,7 @@ def _dynamic_optim_step( model.grad = None self._prune_slot_graphs(self._slot_ref(name)) self._checkpoint_slots[name].revision += 1 + self._version_state().clear((name,)) return metrics def _dynamic_param_step_flags( @@ -3914,7 +4292,7 @@ def _select_next_micro_batch( lambda: self._search_next_micro_batch(items, start, checkpoint=checkpoint), lambda value: (value.plan, value.check), lambda value, check: replace(value, check=check), - context="forward_micro_batches", + context="forward_batches", sync_across_dp=True, ) @@ -4115,6 +4493,11 @@ def candidate(width: int) -> _CandidateMicroBatch[ForwardInputsT]: plan, sync_across_dp=True, sync_planning_errors=True ) ) + if self._graph_memory_policy_enabled(): + plan, check = self._admit_graph_memory(plan, sync_across_dp=True) + if not check.fits and width > min_width: + rejected_widths.add(width) + return candidate(max(min_width, width // 2)) cold_start = not self._all_ranks_have_memory_profile( packed_tokens=plan.packed_tokens, signature=plan.signature, @@ -4237,9 +4620,9 @@ def _validate_replicated_top_level_count( if len(set(configurations)) == 1: return raise ValueError( - "forward_micro_batches requires the same top-level input count and " + "forward_batches requires the same top-level input count and " "yield_empty setting on every " - "distributed rank. Pass already-DP-local inputs to dp_rank_forward instead. " + "distributed rank. Pass already-DP-local inputs to forward instead. " f"Observed (count, yield_empty) by rank: {configurations}." ) @@ -4707,7 +5090,7 @@ def _ensure_checkpoint_slots_for( checkpoint: AdapterSelection, ) -> None: self._ensure_checkpoint_slots( - cast(str, selection) + selection for request in requests if ( request.target_tokens is not None @@ -4735,7 +5118,9 @@ def _group_active_request_indices( ) -> tuple[tuple[tuple["LoRASlotRef | None", bool], tuple[int, ...]], ...]: if ensure_slots: self._ensure_checkpoint_slots_for(requests, checkpoint=checkpoint) - groups: dict[tuple[LoRASlotRef | None, bool], list[int]] = {} + groups: dict[ + tuple[LoRASlotRef | None, bool, ResolvedForwardOptions], list[int] + ] = {} for index, request in enumerate(requests): if ( request.target_tokens is not None @@ -4751,10 +5136,14 @@ def _group_active_request_indices( if request.no_grad is None else not request.no_grad ), + _resolved_request_policy(request.options), ), [], ).append(index) - return tuple((slot_ref, tuple(indices)) for slot_ref, indices in groups.items()) + return tuple( + ((slot_ref, grad), tuple(indices)) + for (slot_ref, grad, _options), indices in groups.items() + ) def _run_flat_plan_with_memory_tracking( self, @@ -4799,6 +5188,7 @@ def _run_flat_plan_with_memory_tracking( ) if seconds is not None and plan.packed_tokens > 0: try: + self._record_graph_forward_time(plan, seconds) self._record_recovery_work(context, seconds) except Exception: self._recovery_state().invalid = True @@ -4862,27 +5252,10 @@ def _execute_flat_plan(self, plan: _FlatForwardPlan) -> list[AnyForwardOutput]: ) try: for group_index, group in enumerate(plan.groups): - from art.megatron.lora import use_lora_slot - if hybridep is not None: self._set_hybridep_rows(hybridep[0][group_index]) with torch.set_grad_enabled(group.grad_enabled): - with use_lora_slot(group.slot_ref): - prepared = self._prepare_packed_forward(group.packed) - item_outputs = self._forward_packed(group.items, prepared) - item_outputs = [ - replace( - output, - checkpoint=( - None if group.slot_ref is None else group.slot_ref.name - ), - no_grad=not group.grad_enabled, - ) - for output in item_outputs - ] - item_outputs = self._track_slot_graph_outputs( - group.slot_ref, item_outputs - ) + item_outputs = self._execute_graph_group(group) for index, output in zip( group.request_indices, item_outputs, strict=True ): @@ -4892,6 +5265,195 @@ def _execute_flat_plan(self, plan: _FlatForwardPlan) -> list[AnyForwardOutput]: self._set_hybridep_rows(hybridep[1]) return outputs + def _forward_graph_cache(self): + from ._graphs import GraphCache + + if not hasattr(self, "_graph_cache"): + self._graph_cache = GraphCache() + return self._graph_cache + + def _forward_cotangent_collector(self): + from ._tensors import CotangentCollector + + if not hasattr(self, "_cotangent_collector"): + self._cotangent_collector = CotangentCollector() + return self._cotangent_collector + + def _execute_graph_group(self, group: _ForwardGroupPlan) -> list[AnyForwardOutput]: + from art.megatron.lora import use_lora_slot + + from ._corrections import capture_forward_corrections + from ._options import resolve_forward_options + from ._tensors import ( + TensorPacket, + flatten_tensors, + managed_tree, + unflatten_tensors, + ) + + options = resolve_forward_options( + getattr(self, "_forward_options", None), + input=group.items[0].request.options, + ) + placement = getattr(group, "memory_placement", None) + retention = ( + options.backward_state if placement is None else placement.backward_state + ) + retention = "gpu" if retention == "auto" else retention + output_device = ( + options.output_device if placement is None else placement.output_device + ) + output_device = "cpu" if output_device == "cpu" else None + ref = group.slot_ref + topology = self._topology() + spec = None + + def execute(captured: _ForwardGroupPlan) -> tuple[torch.Tensor, ...]: + nonlocal spec + if self._topology() != topology: + raise TrainerRankRuntimeSupportError( + "Forward replay requires its original parallel topology" + ) + hybrid = self._configure_hybridep((captured.packed,), topology=topology) + try: + if hybrid is not None: + self._set_hybridep_rows(hybrid[0][0]) + prepared = self._prepare_packed_forward(captured.packed) + outputs = self._forward_packed(captured.items, prepared) + outputs = [ + replace( + output, + checkpoint=None if ref is None else ref.name, + no_grad=not captured.grad_enabled, + ) + for output in outputs + ] + tensors, captured_spec = flatten_tensors(outputs) + if spec is not None and captured_spec != spec: + raise RuntimeError("Forward replay changed its output tree") + spec = captured_spec + return tensors + finally: + if hybrid is not None: + self._set_hybridep_rows(hybrid[1]) + + if not group.grad_enabled: + with torch.no_grad(), use_lora_slot(ref): + tensors = execute(group) + assert spec is not None + outputs = unflatten_tensors(spec, tensors) + return ( + outputs + if output_device is None + else managed_tree(outputs, device=output_device) + ) + + version = self._capture_lora_version(ref, options.max_gradient_staleness) + # Saved views of externally owned parameters must not duplicate whole + # frozen model weights or immutable LoRA captures into every CPU graph. + parameters = [ + tensor + for chunk in self.runtime.model + for tensor in ( + *chunk.parameters(), + *(buffer for _, buffer in chunk.named_buffers()), + ) + ] + if version is not None: + parameters.extend( + parameter + for slot in version.slots.values() + for parameter in slot.parameters() + ) + storages = { + (parameter.device, parameter.untyped_storage().data_ptr()) + for parameter in parameters + } + tracker = None + devices = () + if self.device.type == "cuda": + from megatron.core.tensor_parallel.random import get_cuda_rng_tracker + + tracker = get_cuda_rng_tracker() + devices = ( + self.device.index + if self.device.index is not None + else torch.cuda.current_device(), + ) + cache = self._forward_graph_cache() + handle, tensors = cache.run( + execute, + group, + context_factory=lambda: use_lora_slot(ref, version=version), + validate_backward=None if version is None else version.validate, + retention=retention, + checkpoint_versions=() if version is None else (version.version,), + options=options, + cuda_devices=devices, + rng_tracker=tracker, + keep_on_device=lambda tensor: ( + (tensor.device, tensor.untyped_storage().data_ptr()) in storages + ), + output_device=output_device, + execution_peak_bytes=getattr(placement, "execution_peak_bytes", 0), + ) + if topology.cp > 1 and retention != "replay": + residual = cache.state(handle).non_offloadable_bytes + if residual is not None: + profiles = getattr(self, "_graph_residency", None) + if profiles is None: + self._graph_residency = profiles = OrderedDict() + key = self._graph_residency_key(group) + profiles[key] = max(residual, profiles.pop(key, 0)) + if len(profiles) > 256: + profiles.popitem(last=False) + assert spec is not None + if version is not None: + + @contextmanager + def current_context(): + current = self._capture_lora_version( + ref, options.max_gradient_staleness, origin=version.version + ) + assert current is not None + previous_storages = storages.copy() + storages.update( + (parameter.device, parameter.untyped_storage().data_ptr()) + for slot in current.slots.values() + for parameter in slot.parameters() + ) + try: + with use_lora_slot(ref, version=current): + yield + finally: + storages.clear() + storages.update(previous_storages) + + cache.set_corrections( + handle, + capture_forward_corrections( + unflatten_tensors(spec, tensors), tensors, options + ), + is_stale=lambda: ( + self._capture_checkpoint_version( + version.version.checkpoint + ).revision + != version.weight_version.revision + ), + current_context_factory=current_context, + ) + packet = TensorPacket( + handle, spec, tensors, tuple(tensor.requires_grad for tensor in tensors) + ) + outputs = self._forward_cotangent_collector().attach( + packet, + managed=output_device is not None, + on_release=partial(cache.release, handle), + ) + # Track the caller graph outside saved-state hooks, which detach markers. + # Consumption also ends the lifetime of unused sibling outputs. + return self._track_slot_graph_outputs(ref, outputs) + def _track_slot_graph_outputs( self, ref: "LoRASlotRef | None", @@ -4909,8 +5471,8 @@ def track(tensor: torch.Tensor | None) -> torch.Tensor | None: 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)) + marker = torch.zeros((), dtype=torch.bool, device="cpu") + return _track_slot_graph_tensor(tensor, marker) tracked_outputs = [ ForwardOutput( @@ -4951,7 +5513,7 @@ def _forward_output_metadata( 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, selection) + name = selection enabled = ( torch.is_grad_enabled() if request.no_grad is None else not request.no_grad ) @@ -4966,7 +5528,7 @@ def _hybridep_graphs(self) -> list[weakref.ReferenceType[torch.Tensor]]: def _has_live_hybridep_graphs(self) -> bool: graphs = self._hybridep_graphs() - graphs[:] = [marker for marker in graphs if marker() is not None] + graphs[:] = [marker for marker in graphs if _graph_marker_is_live(marker)] return bool(graphs) def _slot_graphs( @@ -5121,6 +5683,386 @@ def _physical_tokens(self, packed_tokens: int) -> int: multiple = max(1, self._topology_key()[1]) return packed_tokens + (-packed_tokens % multiple) + def _graph_memory_policy_enabled(self) -> bool: + return self.device.type == "cuda" and hasattr(self, "_forward_graph_cache") + + def _available_cpu_memory_bytes(self) -> int: + world = ( + dist.get_world_size() + if dist.is_available() and dist.is_initialized() + else 1 + ) + return host_memory_budget( + local_world_size=local_rank_count(world_size=world) + ).available_bytes + + @staticmethod + def _graph_forward_time_key(plan: _FlatForwardPlan) -> tuple[Any, ...]: + return ( + replace(plan.signature, memory_placement=()), + plan.packed_tokens, + plan.logical_tokens, + tuple(group.packed.segments for group in plan.groups), + ) + + def _record_graph_forward_time( + self, plan: _FlatForwardPlan, seconds: float + ) -> None: + if ( + not math.isfinite(seconds) + or seconds <= 0 + or any( + group.memory_placement is not None + and group.memory_placement.backward_state != "gpu" + for group in plan.groups + ) + ): + return + profiles = getattr(self, "_graph_forward_times", None) + if profiles is None: + self._graph_forward_times = profiles = OrderedDict() + key = self._graph_forward_time_key(plan) + profiles[key] = (*profiles.pop(key, ())[-2:], seconds) + if len(profiles) > 256: + profiles.popitem(last=False) + + def _graph_residency_key(self, group: _ForwardGroupPlan) -> tuple[Any, ...]: + return ( + self._topology_key(), + group.packed.segments, + group.grad_enabled, + tuple( + ( + item.request.input_tokens.numel(), + None + if item.request.target_tokens is None + else item.request.target_tokens.numel(), + item.request.top_k, + item.request.logits, + item.request.hidden_states, + ) + for item in group.items + ), + ) + + def _graph_memory_units( + self, plan: _AnyForwardPlan + ) -> Iterator[ + tuple[int, tuple[int, ...], ForwardMemoryCost, ResolvedForwardOptions] + ]: + """Use existing aggregate profiles when all physical groups share policy.""" + flats = plan.subforwards if isinstance(plan, _SplitForwardPlan) else (plan,) + staged_slots = set() + for flat_index, flat in enumerate(flats): + policies = [ + _resolved_request_policy(g.items[0].request.options) + for g in flat.groups + ] + partitions = ( + [tuple(range(len(flat.groups)))] + if len(set(policies)) <= 1 + else [(i,) for i in range(len(flat.groups))] + ) + for indices in partitions: + if not indices: + continue + groups = tuple(flat.groups[i] for i in indices) + requests = [item.request for group in groups for item in group.items] + policy = policies[indices[0]] + priced = ( + flat + if len(indices) == len(flat.groups) + else replace( + flat, + groups=groups, + packed_tokens=sum( + self._physical_tokens(int(g.packed.tokens.numel())) + for g in groups + ), + logical_tokens=sum( + int(r.input_tokens.numel()) for r in requests + ), + inactive_logical_tokens=0, + output_bytes=self._estimate_group_request_output_bytes( + requests + ), + signature=self._memory_signature_from_requests( + requests, + slot_group_count=len(groups), + grad_modes=tuple(g.grad_enabled for g in groups), + ), + ) + ) + # CPU/replay observations cannot lower the GPU retention model. + priced = replace( + priced, signature=replace(priced.signature, memory_placement=()) + ) + cost = self._plan_cost(priced) + timings = getattr(self, "_graph_forward_times", {}).get( + self._graph_forward_time_key(priced), () + ) + output = priced.output_bytes + retained = max(output, cost.retained) + transient = max(0, cost.required - retained) + residual = 0 + if priced.signature.topology[2] > 1: + profiles = getattr(self, "_graph_residency", {}) + observed = [ + profiles.get(self._graph_residency_key(g)) for g in groups + ] + # CP owns raw stage graphs outside saved-variable hooks. + # Until this exact layout is observed, grant no release credit. + residual = ( + int(sum(observed) * _MEMORY_SAFETY_FACTOR) + if all(value is not None for value in observed) + else retained + ) + retained = max(retained, residual) + version_bytes = getattr(self, "_lora_version_capture_bytes", None) + persistent = ( + sum( + version_bytes(group.slot_ref, policy.max_gradient_staleness) + for group in groups + if group.grad_enabled + ) + if version_bytes is not None + else 0 + ) + staging = 0 + for group in groups: + if group.grad_enabled and group.slot_ref not in staged_slots: + staged_slots.add(group.slot_ref) + staging += self._lora_gradient_staging_bytes(group.slot_ref) + yield ( + flat_index, + indices, + ForwardMemoryCost( + peak_bytes=max(cost.required, retained + transient), + retained_bytes=retained, + cpu_resident_bytes=residual, + output_bytes=output, + replay_bytes=sum( + _snapshot_tensor_bytes(group) + + 64 * 1024 + + _correction_state_bytes(group, policy) + for group in groups + if group.grad_enabled + ), + backward_required=any(group.grad_enabled for group in groups), + persistent_bytes=persistent, + gradient_staging_bytes=staging, + replay_seconds=max(timings) if len(timings) >= 2 else None, + correction_workspace_bytes=cost.required + if any( + correction.policy == "always" + for correction in policy.stale_gradient_corrections + ) + else 0, + ), + policy, + ) + + def _graph_memory_candidates( + self, + units: Sequence[ + tuple[int, tuple[int, ...], ForwardMemoryCost, ResolvedForwardOptions] + ], + *, + sync_across_dp: bool, + ) -> Iterator[ + tuple[ + Literal["gpu", "cpu", "replay"], + Literal["model", "cpu"], + dict[str, Any] | None, + ] + ]: + # Avoid timing work and its collective entirely on the GPU headroom path. + yield "gpu", "model", None + yield "gpu", "cpu", None + eligible = [ + cost + for _, _, cost, options in units + if cost.backward_required + and options.backward_state == "auto" + and options.allow_cpu_offload + and options.allow_replay + ] + stats = getattr(getattr(self, "_graph_cache", None), "transfer_stats", None) + transfer_bytes = sum( + cost.retained_bytes - cost.output_bytes for cost in eligible + ) + trusted = bool(stats) and all( + cost.replay_seconds is not None for cost in eligible + ) + if stats is not None and trusted: + trusted = all( + math.isfinite(value) and value > 0 + for value in ( + stats.offload_bytes, + stats.offload_seconds, + stats.restore_bytes, + stats.restore_seconds, + ) + ) and ( + min(stats.offload_seconds, stats.restore_seconds) >= 0.001 + and transfer_bytes <= 2 * min(stats.offload_bytes, stats.restore_bytes) + ) + cpu_seconds = 0.0 + if stats is not None and trusted: + cpu_seconds = transfer_bytes * ( + stats.offload_seconds / stats.offload_bytes + + stats.restore_seconds / stats.restore_bytes + ) + replay_seconds = sum(cost.replay_seconds or 0.0 for cost in eligible) + cpu_seconds, replay_seconds, missing = self._recovery_reduce( + [cpu_seconds, replay_seconds, float(bool(eligible) and not trusted)], + op="MAX", + sync_across_dp=sync_across_dp, + ) + prefer_replay = ( + not missing and replay_seconds > 0 and replay_seconds * 1.1 < cpu_seconds + ) + evidence = { + "source": "insufficient_samples" + if missing + else "measured_forward_and_transfers", + "cpu_extra_seconds": cpu_seconds, + "replay_extra_seconds": replay_seconds, + "preferred": "replay" if prefer_replay else "cpu", + } + for state in ("replay", "cpu") if prefer_replay else ("cpu", "replay"): + yield state, "model", evidence + yield state, "cpu", evidence + + @overload + def _admit_graph_memory( + self, plan: _FlatForwardPlan, *, sync_across_dp: bool = False + ) -> tuple[_FlatForwardPlan, _MemoryCheck]: ... + + @overload + def _admit_graph_memory( + self, plan: _SplitForwardPlan, *, sync_across_dp: bool = False + ) -> tuple[_SplitForwardPlan, _MemoryCheck]: ... + + def _admit_graph_memory( + self, plan: _AnyForwardPlan, *, sync_across_dp: bool = False + ) -> tuple[_AnyForwardPlan, _MemoryCheck]: + """Try a bounded placement ladder without changing root/group structure.""" + units = list(self._graph_memory_units(plan)) + cpu_available = self._available_cpu_memory_bytes() + planned_checkpoints = { + group.slot_ref.name + for group in plan.groups + if group.grad_enabled + and group.slot_ref is not None + and group.slot_ref.name is not None + } + prior_workspace, prior_staging = self._pending_backward_memory( + exclude_staging=planned_checkpoints + ) + flats = plan.subforwards if isinstance(plan, _SplitForwardPlan) else (plan,) + for state, device, evidence in self._graph_memory_candidates( + units, sync_across_dp=sync_across_dp + ): + placements = [] + for _, _, cost, options in units: + selected_state = options.backward_state + if selected_state == "auto": + selected_state = state + if selected_state == "replay" and not options.allow_replay: + selected_state = "cpu" + if selected_state == "cpu" and not options.allow_cpu_offload: + selected_state = "gpu" + placements.append( + placement_cost( + (cost,), + backward_state=selected_state, + output_device=device + if options.output_device == "auto" + else options.output_device, + ) + ) + selected_groups = [list(flat.groups) for flat in flats] + for (flat_index, indices, _, _), placement in zip( + units, placements, strict=True + ): + for index in indices: + selected_groups[flat_index][index] = replace( + selected_groups[flat_index][index], memory_placement=placement + ) + selected_flats = [] + for flat, groups in zip(flats, selected_groups, strict=True): + modes = tuple( + ( + cast(MemoryPlacement, g.memory_placement).backward_state, + cast(MemoryPlacement, g.memory_placement).output_device, + ) + for g in groups + ) + selected_flats.append( + replace( + flat, + groups=tuple(groups), + signature=replace( + flat.signature, + memory_placement=modes + if any(mode != ("gpu", "model") for mode in modes) + else (), + ), + ) + ) + selected = ( + replace(plan, subforwards=tuple(selected_flats)) + if isinstance(plan, _SplitForwardPlan) + else selected_flats[0] + ) + required = ( + prior_staging + + sum(p.gpu_retained_bytes + p.gpu_backward_bytes for p in placements) + + max( + prior_workspace, + max( + ( + p.gpu_required_bytes + - p.gpu_retained_bytes + - p.gpu_backward_bytes + for p in placements + ), + default=0, + ), + ) + ) + if isinstance(selected, _SplitForwardPlan): + key = self._split_memory_key(selected) + required = max( + required, + 0 + if key is None + else int( + self._split_memory_floors.get(key, 0) * _MEMORY_SAFETY_FACTOR + ), + ) + check = self._memory_check_required(required, sync_across_dp=sync_across_dp) + cpu_required = sum(p.cpu_required_bytes for p in placements) + # One fixed reduction per candidate keeps policy choice identical on + # every physical participant, including locally empty DP ownership. + cpu_margin = self._recovery_reduce( + [float(cpu_available - cpu_required)], + op="MIN", + sync_across_dp=sync_across_dp, + )[0] + check = replace( + check, + fits=check.fits and cpu_margin >= 0, + cpu_required_bytes=cpu_required, + cpu_available_bytes=cpu_available, + cpu_fits=cpu_margin >= 0, + fallback_costs=evidence, + ) + if check.fits: + return selected, check + return selected, check + def _memory_check( self, forward: _FlatForwardPlan, @@ -5151,6 +6093,76 @@ def _admission_outcome(self, local: int) -> int: dist.all_reduce(value, op=dist.ReduceOp.MIN) return int(value.item()) + def _reclaim_graph_memory( + self, check: _MemoryCheck, *, sync_across_dp: bool + ) -> bool: + if not self._graph_memory_policy_enabled(): + return False + cache = getattr(self, "_graph_cache", None) + actions = [] + # DP partitions can own different numbers of graphs; only TP/CP peers + # coordinate individual records. WORLD sees one final success exchange. + with self._planning_status(sync_across_dp): + handles = () if cache is None else cache.handles() + counts = self._recovery_reduce( + [float(len(handles)), -float(len(handles))], + op="MAX", + sync_across_dp=False, + ) + if counts[0] != -counts[1]: + raise RuntimeError( + "Physical participants have different graph cache lengths" + ) + cpu_available = max( + 0, self._available_cpu_memory_bytes() - check.cpu_required_bytes + ) + for handle in handles: + assert cache is not None + state = cache.state(handle) + can_offload, can_replay, cpu_margin = self._recovery_reduce( + [ + float(state.offloadable and state.retention == "gpu"), + float(state.replayable and state.retention != "replay"), + float(cpu_available - state.offload_bytes), + ], + op="MIN", + sync_across_dp=False, + ) + if can_offload and cpu_margin >= 0 and check.cpu_fits: + actions.append((cache.offload, handle)) + cpu_available -= state.offload_bytes + elif can_replay: + actions.append((cache.evict, handle)) + # Finish all collective decisions before a transfer/allocation can + # fail, so peers still reach the final error exchange on failure. + error: BaseException | None = None + try: + for operation, handle in actions: + operation(handle) + if actions: + # Native allocator accounting only credits physical free bytes. + # Release reclaimed graph storage once, on this refusal path. + torch.cuda.empty_cache() + except BaseException as exc: + error = exc + try: + succeeded = self._recovery_reduce( + [float(error is None)], op="MIN", sync_across_dp=False + )[0] + except BaseException as exchange_error: + if error is None: + raise + raise self._memory_error_with_reduction_note(error, exchange_error) + if error is not None: + raise error + if not succeeded: + raise RuntimeError("Graph reclamation failed on another physical rank") + return bool( + self._recovery_reduce( + [float(bool(actions))], op="MAX", sync_across_dp=sync_across_dp + )[0] + ) + def _recover_admission( self, search: Callable[[], Any], @@ -5171,9 +6183,15 @@ def finish(value: Any) -> Any: return None plan, check = describe(value) if sync_across_dp: - check = self._memory_check_required( + refreshed = self._memory_check_required( check.estimated_required_bytes, sync_across_dp=True ) + check = replace( + check, + estimated_required_bytes=refreshed.estimated_required_bytes, + available_bytes=refreshed.available_bytes, + fits=refreshed.fits and check.cpu_fits, + ) if check.fits: return update(value, check) refused = _ForwardRefusal( @@ -5207,12 +6225,16 @@ def finish(value: Any) -> Any: self._snapshot_planning_telemetry(refused.plan, refused.check) raise refused.error(context) from original assert refused is not None - if self._try_cache_recovery( + reclaimed = self._reclaim_graph_memory( + refused.check, sync_across_dp=sync_across_dp + ) + recovered = self._try_cache_recovery( refused.check, sync_across_dp=sync_across_dp, owner=owner, started=started, - ): + ) + if reclaimed or recovered: value = search() result = finish(value) if result is not None: @@ -5272,7 +6294,7 @@ def _recovery_clock(self) -> float | None: return None def _record_recovery_work(self, context: str, seconds: float) -> None: - if context not in ("forward_micro_batches", "dp_rank_forward"): + if context not in ("forward_batches", "forward"): return state = self._recovery_state() try: @@ -6694,6 +7716,14 @@ 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 _gradient_staging_bytes(parameters: Iterable[torch.nn.Parameter]) -> int: + return sum( + param.numel() * param.element_size() * (3 - (param.grad is not None)) + for param in parameters + if param.requires_grad + ) + + def _custom_parameters(custom: _CustomObject) -> Iterator[torch.nn.Parameter]: if custom.kind == "module": yield from cast(torch.nn.Module, custom.value).parameters() @@ -6718,6 +7748,17 @@ def _tracked_tensor_function( args: tuple[object, ...], kwargs: dict[str, object], ) -> object: + from ._heads import ( + _stage_local_buffers, + head_call_arguments, + mutates_tensor, + readonly_buffer_views, + tensor_metadata_function, + tensor_mutation_targets, + ) + + if (captured := head_call_arguments(args, kwargs)) is not None: + return func(*captured[0], **captured[1]) del types tracked = tuple( value @@ -6727,9 +7768,9 @@ def _tracked_tensor_function( trackers = {id(value._art_tracker): value._art_tracker for value in tracked} for tracker in trackers.values(): tracker.validate() - if getattr(func, "__name__", "") in { + if tensor_metadata_function(func) or getattr(func, "__name__", "") in { "__format__", - "__get__", + "__set__", "__hash__", "__len__", "__repr__", @@ -6737,10 +7778,32 @@ def _tracked_tensor_function( }: with torch._C.DisableTorchFunctionSubclass(): return func(*args, **kwargs) + if torch._C._current_graph_task_id() >= 0 and any( + value._art_tracker.active for value in tracked + ): + raise RuntimeError( + "Live checkpoint tensor used during backward recomputation; capture " + "parameter.clone() or head.snapshot() before activation checkpointing" + ) markers: dict[int, torch.Tensor] = {} replacements: dict[int, torch.Tensor] = {} + mutating = mutates_tensor(func, kwargs) + mutation_targets = tensor_mutation_targets(func, args, kwargs) + from ._parameter_hooks import parameter_hook_active + + if parameter_hook_active.get() and any( + isinstance(value, _TrackedParameter) and id(value) in mutation_targets + for value in tracked + ): + raise RuntimeError("Parameter hooks must not mutate checkpoint parameters") + if getattr(func, "__name__", "") == "requires_grad_" and any( + value._art_tracker.active for value in tracked + ): + raise RuntimeError("Set checkpoint parameter trainability in its factory") + copied_buffers: list[tuple[torch.Tensor, torch.Tensor]] = [] + def replace(value: object) -> object: if isinstance(value, _TrackedParameter | _TrackedTensor): cached = replacements.get(id(value)) @@ -6750,18 +7813,31 @@ def replace(value: object) -> object: with torch._C.DisableTorchFunctionSubclass(): if ( tracker.active + and id(value) not in mutation_targets and isinstance(value, _TrackedParameter) and value.requires_grad and torch.is_grad_enabled() ): marker = markers.get(id(tracker)) if marker is None: - marker = torch.zeros((), dtype=torch.bool) + marker = torch.zeros((), dtype=torch.bool, device="cpu") markers[id(tracker)] = marker tracker.record(marker) - result = _CustomSlotGraphSentinel.apply( - value.as_subclass(torch.Tensor), marker + trainer = tracker.validate() + assert tracker.ref.name is not None + from ._heads import head_staleness + + snapshot = trainer._snapshot_parameter( + value, + trainer._capture_checkpoint_version(tracker.ref.name), + head_staleness(trainer), ) + result = _track_slot_graph_tensor(snapshot, marker) + elif tracker.active and id(value) not in mutation_targets: + # Detached/no-grad reads can still be saved by a later graph. + result = value.as_subclass(torch.Tensor).detach().clone() + if isinstance(value, _TrackedTensor): + copied_buffers.append((value, result)) else: result = value.as_subclass(torch.Tensor) replacements[id(value)] = result @@ -6779,13 +7855,35 @@ def replace(value: object) -> object: *cast(tuple[object, ...], replace(args)), **cast(dict[str, object], replace(kwargs)), ) + changed_buffers: set[_CustomTensorTracker] = set() + staged = _stage_local_buffers( + {str(index): target for index, (target, _) in enumerate(copied_buffers)}, + {str(index): value for index, (_, value) in enumerate(copied_buffers)}, + ) + with torch.no_grad(), torch._C.DisableTorchFunctionSubclass(): + for target, value in staged: + target.copy_(value) + changed_buffers.add(cast(_TrackedTensor, target)._art_tracker) + if mutating: + changed_buffers.update( + value._art_tracker + for value in tracked + if isinstance(value, _TrackedTensor) and id(value) in mutation_targets + ) + for tracker in changed_buffers: + tracker.buffer_revision += 1 + if mutating: + for value in tracked: + if result is replacements[id(value)]: + result = value + break if markers and not any( isinstance(value, torch.Tensor) and value.requires_grad for value in _walk_objects(result) ): for marker in markers.values(): marker.fill_(True) - return result + return readonly_buffer_views(result, [value for _, value in copied_buffers]) def _graph_marker_is_live( @@ -6848,7 +7946,10 @@ def _track_custom_object( value = _TrackedTensor(source.detach().clone(), tracker) buffers[id(source)] = value child._buffers[key] = value - return custom + from ._heads import native_module_handle + + trainer = tracker.validate() + return replace(custom, handle=native_module_handle(trainer, custom, tracker)) def _custom_layout( @@ -7247,10 +8348,57 @@ def _chunk_boundaries( def _select_positions(values: torch.Tensor, positions: torch.Tensor) -> torch.Tensor: if int(positions.numel()) == 0: - return values[:0] + return values[:0].clone() return values.index_select(0, positions.to(device=values.device)) +@lru_cache(maxsize=256) +def _resolved_request_policy(options: ForwardOptions | None) -> ResolvedForwardOptions: + return resolve_forward_options(input=options) + + +def _correction_state_bytes( + group: _ForwardGroupPlan, options: ResolvedForwardOptions +) -> int: + if not group.grad_enabled: + return 0 + topk = sum( + item.input_ids.numel() * (item.request.top_k or 0) for item in group.items + ) + logprobs = 4 * ( + topk + + ( + sum(item.labels.numel() for item in group.items if item.labels is not None) + if options.stale_gradient_corrections + else 0 + ) + ) + # Top-k probabilities and identities persist for replay validation even + # without corrections; explicit always also stages corrected cotangents. + return topk * 8 + logprobs * ( + 2 + if any(c.policy == "always" for c in options.stale_gradient_corrections) + else 1 + ) + + +def _snapshot_tensor_bytes(value: object) -> int: + """Graph replay snapshots copy each tensor occurrence in the captured plan.""" + if isinstance(value, torch.Tensor): + return value.numel() * value.element_size() + if is_dataclass(value) and not isinstance(value, type): + return sum( + _snapshot_tensor_bytes(getattr(value, f.name)) + for f in fields(value) + if f.init + ) + if isinstance(value, dict): + return sum(_snapshot_tensor_bytes(item) for item in value.values()) + if isinstance(value, (tuple, list)): + return sum(_snapshot_tensor_bytes(item) for item in value) + return 0 + + def _batch_seq_logits(logits: torch.Tensor, *, seq_len: int) -> torch.Tensor: if int(logits.ndim) != 3: raise RuntimeError( @@ -7268,7 +8416,19 @@ def _batch_seq_logits(logits: torch.Tensor, *, seq_len: int) -> torch.Tensor: def _materialize(inputs: ForwardInputs) -> ForwardInputs: if isinstance(inputs, ForwardInput): return inputs - return [_materialize(item) for item in _nested_forward_children(inputs)] + return _rebuild_forward_tree( + inputs, [_materialize(item) for item in _nested_forward_children(inputs)] + ) + + +def _rebuild_forward_tree(template: Any, children: list[Any]) -> Any: + if isinstance(template, tuple): + return ( + type(template)(*children) + if hasattr(template, "_fields") + else tuple(children) + ) + return children def _is_forward_input(inputs: ForwardInputs) -> TypeIs[AnyForwardInput]: @@ -7288,7 +8448,10 @@ def _unflatten( ) -> ForwardOutputs: if isinstance(template, ForwardInput): return next(outputs) - return [_unflatten(item, outputs) for item in _nested_forward_children(template)] + return _rebuild_forward_tree( + template, + [_unflatten(item, outputs) for item in _nested_forward_children(template)], + ) def _nested_forward_children(inputs: ForwardInputs) -> Iterator[ForwardInputs]: diff --git a/src/art/trainer_rank/_memory_policy.py b/src/art/trainer_rank/_memory_policy.py new file mode 100644 index 000000000..8af88108c --- /dev/null +++ b/src/art/trainer_rank/_memory_policy.py @@ -0,0 +1,352 @@ +"""Host memory limits and bounded graph/output placement admission. + +The calibrated GPU estimator remains authoritative for a physical forward. These +policies only change how much state coexists between physical forwards; they do +not assume offload or replay reduces the workspace needed to execute one. +""" + +from __future__ import annotations + +from collections.abc import Sequence +from dataclasses import dataclass +from functools import lru_cache +import os +from pathlib import Path +from typing import Literal + +Retention = Literal["gpu", "cpu", "replay"] +OutputDevice = Literal["model", "cpu"] +_HOST_RESERVE_FRACTION = 0.1 +_HOST_RESERVE_MAX_BYTES = 4 * 1024**3 + + +def choose_output_placements( + outputs: Sequence[tuple[int, Literal["auto", "model", "cpu"]]], + *, + gpu_available_bytes: int, +) -> tuple[OutputDevice, ...]: + """Place already gathered CPU outputs before making any model-device copy. + + The caller admits the CPU receive/serialization buffers before gathering. + Fresh GPU headroom excludes live storage and pending restore/staging reserves. + Explicit model placement is reserved before optional model placement. + """ + if any( + size < 0 or device not in ("auto", "model", "cpu") for size, device in outputs + ): + raise ValueError("Expected nonnegative output bytes and auto/model/cpu policy") + required = sum(size for size, device in outputs if device == "model") + if required > max(0, gpu_available_bytes): + raise MemoryError( + f"Gathered model-device outputs require {required} GPU bytes; " + f"only {max(0, gpu_available_bytes)} bytes are available" + ) + remaining = max(0, gpu_available_bytes) - required + placements: list[OutputDevice] = [] + for size, device in outputs: + if device == "auto": + device = "model" if size <= remaining else "cpu" + if device == "model": + remaining -= size + placements.append(device) + return tuple(placements) + + +@dataclass(frozen=True) +class MemoryScope: + """A shared capacity limit, with the ranks whose allocations it counts.""" + + name: str + limit_bytes: int + available_bytes: int + rank_count: int + + @property + def per_rank_available_bytes(self) -> int: + if self.rank_count < 1: + raise ValueError("Memory scope rank_count must be positive") + reserve = min( + _HOST_RESERVE_MAX_BYTES, + int(max(0, self.limit_bytes) * _HOST_RESERVE_FRACTION), + ) + return max(0, min(self.limit_bytes, self.available_bytes) - reserve) // ( + self.rank_count + ) + + +@dataclass(frozen=True) +class HostMemoryBudget: + scopes: tuple[MemoryScope, ...] + + @property + def available_bytes(self) -> int: + """Additional CPU allocation allowed per rank, excluding existing use.""" + return min((scope.per_rank_available_bytes for scope in self.scopes), default=0) + + +def _read_int(path: Path) -> int | None: + try: + return int(path.read_text().strip()) + except (OSError, ValueError): + return None + + +def _mount_path(value: str) -> str: + for code, char in ( + ("\\040", " "), + ("\\011", "\t"), + ("\\012", "\n"), + ("\\134", "\\"), + ): + value = value.replace(code, char) + return value + + +def _cgroup_paths(proc_root: Path) -> tuple[tuple[Path, Path, bool], ...]: + """Resolve membership through mount roots, including cgroup namespaces.""" + try: + return _resolve_cgroup_paths( + (proc_root / "self/cgroup").read_text(), + (proc_root / "self/mountinfo").read_text(), + ) + except OSError: + return () + + +@lru_cache(maxsize=8) +def _resolve_cgroup_paths( + membership_text: str, mounts: str +) -> tuple[tuple[Path, Path, bool], ...]: + # Cache only pure discovery. Membership/mount changes invalidate immediately; + # every available-memory, limit and usage counter is read on every query. + memberships = [line.split(":", 2) for line in membership_text.splitlines()] + result = [] + for line in mounts.splitlines(): + if " - cgroup" not in line: + continue + before, separator, after = line.partition(" - ") + fields, fs = before.split(), after.split() + if not separator or len(fields) < 5 or len(fs) < 3: + continue + v2 = fs[0] == "cgroup2" + if not v2 and not (fs[0] == "cgroup" and "memory" in fs[2].split(",")): + continue + mount_root, mount = Path(_mount_path(fields[3])), Path(_mount_path(fields[4])) + for membership in memberships: + if len(membership) != 3: + continue + _, controllers, name = membership + if not (controllers == "" if v2 else "memory" in controllers.split(",")): + continue + try: + relative = Path(name).relative_to(mount_root) + except ValueError: + # A namespaced membership can already be relative to its mount. + if name != "/": + continue + relative = Path() + current = mount / relative + if ".." not in current.parts: + result.append((current, mount, v2)) + return tuple(result) + + +def host_memory_budget( + *, local_world_size: int, proc_root: Path = Path("/proc") +) -> HostMemoryBudget: + """Read fresh host/cgroup headroom, conservatively shared by local ranks. + + Each scope subtracts existing allocations before division. A per-process + cgroup is also divided by the host rank count: conservative when ranks have + separate limits, safe when a pod or ancestor limit covers all local ranks. + Missing required counters grant no memory credit. + """ + if local_world_size < 1: + raise ValueError("local_world_size must be positive") + try: + memory = { + key: int(value.split()[0]) * 1024 + for key, value in ( + line.split(":", 1) + for line in (proc_root / "meminfo").read_text().splitlines() + if ":" in line + ) + } + total, available = memory["MemTotal"], memory["MemAvailable"] + except (OSError, ValueError, KeyError): + total, available = 0, 0 + scopes = [MemoryScope("host", total, available, local_world_size)] + visited = set() + for current, mount, v2 in _cgroup_paths(proc_root): + while True: + if current not in visited: + visited.add(current) + limit = _read_int( + current / ("memory.max" if v2 else "memory.limit_in_bytes") + ) + if limit is not None and limit < (1 << 60): + used = _read_int( + current / ("memory.current" if v2 else "memory.usage_in_bytes") + ) + scopes.append( + MemoryScope( + str(current), + limit, + 0 if used is None else max(0, limit - used), + local_world_size, + ) + ) + if current == mount: + break + current = current.parent + return HostMemoryBudget(tuple(scopes)) + + +def local_rank_count(*, world_size: int = 1) -> int: + """Launcher rank count; callers with explicit topology should pass it instead.""" + for name in ("LOCAL_WORLD_SIZE", "OMPI_COMM_WORLD_LOCAL_SIZE", "MPI_LOCALNRANKS"): + value = os.environ.get(name) + if value is not None: + count = int(value) + if count < 1: + raise ValueError(f"{name} must be positive") + return count + # WORLD_SIZE is conservative across hosts and safe without local topology. + return max(1, world_size, int(os.environ.get("WORLD_SIZE", "1"))) + + +@dataclass(frozen=True) +class ForwardMemoryCost: + peak_bytes: int + retained_bytes: int + output_bytes: int + replay_bytes: int = 0 + backward_required: bool = True + persistent_bytes: int = 0 + correction_workspace_bytes: int = 0 + gradient_staging_bytes: int = 0 + replay_seconds: float | None = None + cpu_resident_bytes: int = 0 + + def __post_init__(self) -> None: + if not 0 <= self.output_bytes <= self.retained_bytes <= self.peak_bytes: + raise ValueError("Expected 0 <= output <= retained <= peak bytes") + if ( + min( + self.replay_bytes, + self.persistent_bytes, + self.correction_workspace_bytes, + self.gradient_staging_bytes, + self.cpu_resident_bytes, + ) + < 0 + ): + raise ValueError( + "Replay, persistent, correction and staging bytes must be nonnegative" + ) + if self.cpu_resident_bytes > self.retained_bytes: + raise ValueError("CPU-mode GPU residency cannot exceed retained bytes") + + +@dataclass(frozen=True) +class MemoryPlacement: + backward_state: Retention + output_device: OutputDevice + gpu_required_bytes: int + cpu_required_bytes: int + gpu_retained_bytes: int + gpu_backward_bytes: int = 0 + execution_peak_bytes: int = 0 + + +def placement_cost( + costs: Sequence[ForwardMemoryCost], + *, + backward_state: Retention, + output_device: OutputDevice, +) -> MemoryPlacement: + """Bound one complete root's live state plus one child's restore workspace.""" + gpu_retained, cpu_retained, workspace, staging = 0, 0, 0, 0 + for cost in costs: + output_cpu = output_device == "cpu" + if not cost.backward_required: + gpu = 0 if output_cpu else cost.output_bytes + cpu = cost.output_bytes if output_cpu else 0 + transient = cost.peak_bytes - gpu + else: + # Graph outputs keep their original CUDA storage while a graph lives, + # even if caller-facing copies are on CPU. The detached caller copy + # has distinct storage to isolate caller mutation and oversized views. + physical_gpu = ( + cost.retained_bytes + if backward_state == "gpu" + else max(cost.output_bytes, cost.cpu_resident_bytes) + if backward_state == "cpu" + else 0 + ) + gpu = physical_gpu + (0 if output_cpu else cost.output_bytes) + cpu = cost.replay_bytes + (cost.output_bytes if output_cpu else 0) + if backward_state == "cpu": + cpu += cost.retained_bytes - cost.output_bytes + transient = cost.peak_bytes - physical_gpu + gpu_retained += gpu + cost.persistent_bytes + cpu_retained += cpu + workspace = max(workspace, transient, cost.correction_workspace_bytes) + staging += cost.gradient_staging_bytes + return MemoryPlacement( + backward_state, + output_device, + gpu_retained + workspace + staging, + cpu_retained, + gpu_retained, + staging, + max((cost.peak_bytes for cost in costs), default=0), + ) + + +def choose_memory_placement( + costs: Sequence[ForwardMemoryCost], + *, + gpu_available_bytes: int, + cpu_available_bytes: int, + backward_state: Literal["auto", "gpu", "cpu", "replay"] = "auto", + output_device: Literal["auto", "model", "cpu"] = "model", + allow_cpu_offload: bool = True, + allow_replay: bool = True, +) -> MemoryPlacement | None: + """Prefer retained GPU, then CPU saved state, then replay; O(children).""" + if backward_state not in ("auto", "gpu", "cpu", "replay"): + raise ValueError(f"Unknown backward_state: {backward_state}") + if output_device not in ("auto", "model", "cpu"): + raise ValueError(f"Unknown output_device: {output_device}") + if backward_state == "cpu" and not allow_cpu_offload: + raise ValueError("CPU backward state conflicts with allow_cpu_offload=False") + if backward_state == "replay" and not allow_replay: + raise ValueError("Replay backward state conflicts with allow_replay=False") + states: tuple[Retention, ...] = ( + tuple( + state + for state, enabled in ( + ("gpu", True), + ("cpu", allow_cpu_offload), + ("replay", allow_replay), + ) + if enabled + ) + if backward_state == "auto" + else (backward_state,) + ) + devices: tuple[OutputDevice, ...] = ( + ("model", "cpu") if output_device == "auto" else (output_device,) + ) + for state in states: + for device in devices: + placement = placement_cost( + costs, backward_state=state, output_device=device + ) + if ( + placement.gpu_required_bytes <= gpu_available_bytes + and placement.cpu_required_bytes <= cpu_available_bytes + ): + return placement + return None diff --git a/src/art/trainer_rank/_operations.py b/src/art/trainer_rank/_operations.py new file mode 100644 index 000000000..770747c0f --- /dev/null +++ b/src/art/trainer_rank/_operations.py @@ -0,0 +1,150 @@ +"""Identified operations shared by dedicated clients and native rank-zero views.""" + +from __future__ import annotations + +import asyncio +from copy import copy +from dataclasses import dataclass +import hashlib +import inspect +from typing import Any + +import cloudpickle + + +@dataclass(frozen=True) +class TrainerOperation: + id: str + kind: str + payload: bytes + + @classmethod + def capture(cls, id: str, kind: str, payload: Any) -> TrainerOperation: + """Freeze arguments before asynchronous submission can observe mutations.""" + return cls(id, kind, cloudpickle.dumps(payload)) + + +@dataclass +class _Outcome: + fingerprint: tuple[str, bytes] + completion: asyncio.Future[Any] | None + + +@dataclass(frozen=True) +class _Failure: + payload: bytes + + @classmethod + def capture(cls, error: BaseException) -> _Failure: + # Retaining the exception itself retains its traceback and activation + # frames. Serialization also gives each retry its own exception object. + try: + payload = cloudpickle.dumps(error) + if len(payload) > 65536 or not isinstance( + cloudpickle.loads(payload), BaseException + ): + raise ValueError("exception cannot be retained as a small outcome") + except BaseException: + try: + message = str(error)[:8192] + except BaseException: + message = "exception message unavailable" + payload = cloudpickle.dumps( + RuntimeError(f"{type(error).__qualname__}: {message}") + ) + return cls(payload) + + +class OperationResultReleasedError(RuntimeError): + """The operation completed, but its acknowledged result is no longer retained.""" + + +def _ledger(rank_zero: Any) -> dict[str, _Outcome]: + owner = rank_zero._rank + if not hasattr(owner, "_operation_outcomes"): + owner._operation_outcomes = {} + return owner._operation_outcomes + + +async def execute_operation(rank_zero: Any, operation: TrainerOperation) -> Any: + """Apply at most once, including concurrent retries and failed mutations. + + This ledger is scoped to the actor process lifetime. Worker loss fails the + session; callers must not transparently replay updates on a replacement. + """ + ledger = _ledger(rank_zero) + payload = cloudpickle.loads(operation.payload) + if operation.kind == "acknowledge": + for operation_id in payload: + if (outcome := ledger.get(operation_id)) is not None: + if outcome.completion is not None and not outcome.completion.done(): + raise RuntimeError( + "cannot acknowledge an unfinished trainer operation" + ) + outcome.completion = None + return None + fingerprint = (operation.kind, hashlib.sha256(operation.payload).digest()) + if (outcome := ledger.get(operation.id)) is not None: + if outcome.fingerprint != fingerprint: + raise ValueError( + "trainer operation identity was reused with different arguments" + ) + if outcome.completion is None: + raise OperationResultReleasedError(operation.id) + result = await asyncio.shield(outcome.completion) + if isinstance(result, _Failure): + raise cloudpickle.loads(result.payload) from None + return result + completion = asyncio.get_running_loop().create_future() + ledger[operation.id] = _Outcome(fingerprint, completion) + try: + if operation.kind in ("forward", "batches_next"): + # Only driver replies use CPU transport views. The physical command + # receives the original policy, and native callback views are intact. + transport = copy(rank_zero) + transport._transport_handles = [] + try: + tree = ( + transport.forward(**payload) + if operation.kind == "forward" + else transport.next_forward_batch(**payload) + ) + result = None if tree is None else transport.export_forward(tree) + except BaseException: + if transport._transport_handles: + transport._invoke("release", tuple(transport._transport_handles)) + raise + elif operation.kind == "backward": + result = rank_zero.backward_packets(**payload) + elif operation.kind == "optim_step": + result = rank_zero.optim_step(**payload) + elif operation.kind == "release": + result = rank_zero.release_forward(**payload) + elif operation.kind == "batches_open": + result = rank_zero.open_forward_batches(**payload) + elif operation.kind == "batches_close": + result = rank_zero.close_forward_batches(payload["handle"]) + pending = ledger.get(payload.get("pending_operation")) + if ( + pending is not None + and pending.completion is not None + and pending.completion.done() + ): + if ( + packet := pending.completion.result() + ) is not None and not isinstance(packet, _Failure): + rank_zero.release_forward((packet.handle,)) + elif operation.kind.startswith("head_"): + from ._heads import execute_head_operation + + result = execute_head_operation(rank_zero, operation.kind, payload) + else: + raise ValueError(f"unknown trainer operation: {operation.kind!r}") + if inspect.isawaitable(result): + result = await result + except BaseException as error: + completion.set_result(_Failure.capture(error)) + raise + else: + completion.set_result(result) + return result diff --git a/src/art/trainer_rank/_options.py b/src/art/trainer_rank/_options.py new file mode 100644 index 000000000..9535ae1d9 --- /dev/null +++ b/src/art/trainer_rank/_options.py @@ -0,0 +1,162 @@ +"""Immutable forward policy shared by native and remote trainers.""" + +from __future__ import annotations + +from collections.abc import Sequence +from dataclasses import dataclass, fields +from enum import Enum +import math +from typing import Literal + + +class _Unset(Enum): + VALUE = "Unset" + + def __repr__(self) -> str: + return "Unset" + + +Unset = _Unset.VALUE + + +@dataclass(frozen=True, kw_only=True) +class ImportanceSamplingGradientCorrection: + """Weight selected-token cotangents by clipped current/original probability. + + This is a sampling-distribution correction for score-function estimators, + not an exact correction for arbitrary losses or stale Jacobians. Clipping + introduces bias; token-local and top-k ratios do not correct a full sequence + distribution. Top-k correction requires current probabilities of the + original token IDs, without renormalizing over the selected tokens. + + ``when_available`` never adds a forward solely to obtain current logprobs. + ``always`` requires them, and fails if the runtime cannot supply them. + Only stale, active logprob cotangents are eligible. Hidden states, logits, + and custom-head gradients remain bounded and uncorrected under both policies. + """ + + clip_low: float = 0.0 + clip_high: float = 5.0 + policy: Literal["when_available", "always"] = "when_available" + + def __post_init__(self) -> None: + if not ( + math.isfinite(self.clip_low) + and math.isfinite(self.clip_high) + and 0 <= self.clip_low <= self.clip_high + ): + raise ValueError( + "correction clipping bounds must be finite and 0 <= low <= high" + ) + if self.policy not in ("when_available", "always"): + raise ValueError(f"unknown correction policy: {self.policy!r}") + + +@dataclass(frozen=True, kw_only=True) +class ForwardOptions: + """Per-field overrides; ``Unset`` inherits input > method > constructor. + + An empty correction collection disables inherited corrections. Collections + are snapshotted to tuples at construction. Existing checkpoint, grad mode, + and output selectors remain separate forward arguments. + + Staleness counts completed checkpoint updates since the original forward, + measured before the update consuming its gradient; replay does not reset it. + Zero requires current weights. ``backward_state`` forces a cache placement + unless set to ``auto``. ``allow_cpu_offload`` controls saved backward state; + output placement is independent. ``output_device="model"`` retains native + model-device outputs; ``auto`` permits planner placement, and ``cpu`` forces + CPU outputs. Remote drivers transport these outputs as CPU tensor proxies. + """ + + max_gradient_staleness: int | _Unset = Unset + stale_gradient_corrections: ( + Sequence[ImportanceSamplingGradientCorrection] | _Unset + ) = Unset + backward_state: Literal["auto", "gpu", "cpu", "replay"] | _Unset = Unset + allow_cpu_offload: bool | _Unset = Unset + allow_replay: bool | _Unset = Unset + output_device: Literal["auto", "model", "cpu"] | _Unset = Unset + + def __post_init__(self) -> None: + corrections = self.stale_gradient_corrections + if corrections is not Unset: + object.__setattr__(self, "stale_gradient_corrections", tuple(corrections)) + _validate_options(self) + + +@dataclass(frozen=True, kw_only=True) +class ResolvedForwardOptions: + """Concrete policy captured at submission, never read from mutable defaults.""" + + max_gradient_staleness: int = 2 + stale_gradient_corrections: tuple[ImportanceSamplingGradientCorrection, ...] = ( + ImportanceSamplingGradientCorrection(), + ) + backward_state: Literal["auto", "gpu", "cpu", "replay"] = "auto" + allow_cpu_offload: bool = True + allow_replay: bool = True + output_device: Literal["auto", "model", "cpu"] = "model" + + def __post_init__(self) -> None: + if any(getattr(self, field.name) is Unset for field in fields(self)): + raise ValueError("resolved options cannot contain Unset") + object.__setattr__( + self, "stale_gradient_corrections", tuple(self.stale_gradient_corrections) + ) + _validate_options(self) + if self.backward_state == "cpu" and not self.allow_cpu_offload: + raise ValueError("backward_state='cpu' requires allow_cpu_offload=True") + if self.backward_state == "replay" and not self.allow_replay: + raise ValueError("backward_state='replay' requires allow_replay=True") + + +def _validate_options(options: ForwardOptions | ResolvedForwardOptions) -> None: + age = options.max_gradient_staleness + if age is not Unset and (type(age) is not int or age < 0): + raise ValueError("max_gradient_staleness must be a nonnegative integer") + for name in ("allow_cpu_offload", "allow_replay"): + value = getattr(options, name) + if value is not Unset and type(value) is not bool: + raise ValueError(f"{name} must be a bool") + if options.backward_state is not Unset and options.backward_state not in ( + "auto", + "gpu", + "cpu", + "replay", + ): + raise ValueError(f"unknown backward_state: {options.backward_state!r}") + if options.output_device is not Unset and options.output_device not in ( + "auto", + "model", + "cpu", + ): + raise ValueError(f"unknown output_device: {options.output_device!r}") + corrections = options.stale_gradient_corrections + if corrections is not Unset: + if any( + not isinstance(correction, ImportanceSamplingGradientCorrection) + for correction in corrections + ): + raise TypeError("unsupported stale gradient correction") + if len(corrections) > 1: + raise ValueError("only one importance sampling correction may be specified") + + +def resolve_forward_options( + constructor: ForwardOptions | None = None, + method: ForwardOptions | None = None, + input: ForwardOptions | None = None, +) -> ResolvedForwardOptions: + """Resolve each field independently and validate the resulting policy.""" + values = {} + for options in (constructor, method, input): + if options is not None: + if not isinstance(options, ForwardOptions): + raise TypeError("options must be ForwardOptions or None") + values.update( + (field.name, value) + for field in fields(options) + if (value := getattr(options, field.name)) is not Unset + ) + return ResolvedForwardOptions(**values) diff --git a/src/art/trainer_rank/_parameter_hooks.py b/src/art/trainer_rank/_parameter_hooks.py new file mode 100644 index 000000000..124258314 --- /dev/null +++ b/src/art/trainer_rank/_parameter_hooks.py @@ -0,0 +1,141 @@ +"""Persistent live-parameter hooks applied to one backward's summed gradient.""" + +from __future__ import annotations + +from collections import OrderedDict +from contextvars import ContextVar +from dataclasses import dataclass, field +import json +from typing import Any +import weakref + +import torch +from torch.utils.hooks import RemovableHandle + +parameter_hook_active: ContextVar[bool] = ContextVar( + "parameter_hook_active", default=False +) + + +@dataclass(eq=False) +class ParameterHooks: + parameter: weakref.ReferenceType[torch.Tensor] + hooks: OrderedDict[int, Any] = field(default_factory=OrderedDict) + + def apply(self, gradient: torch.Tensor) -> torch.Tensor: + for hook in tuple(self.hooks.values()): + token = parameter_hook_active.set(True) + try: + result = hook(gradient) + finally: + parameter_hook_active.reset(token) + if result is None: + continue + if not isinstance(result, torch.Tensor) or ( + result.shape != gradient.shape + or result.dtype != gradient.dtype + or result.device != gradient.device + or result.layout != gradient.layout + ): + raise RuntimeError( + "Live parameter hook must return None or a gradient with unchanged shape, dtype, device and layout" + ) + gradient = result + return gradient + + +def parameter_hooks(parameter: torch.Tensor) -> ParameterHooks: + registry = getattr(parameter, "_art_parameter_hooks", None) + if registry is None: + registry = ParameterHooks(weakref.ref(parameter)) + setattr(parameter, "_art_parameter_hooks", registry) + return registry + + +def register_parameter_hook(parameter: torch.Tensor, hook: Any) -> RemovableHandle: + """Run once per explicit backward, before accumulating into existing .grad.""" + if not parameter.requires_grad: + raise RuntimeError( + "cannot register a hook on a tensor that doesn't require gradient" + ) + if not callable(hook): + raise TypeError("parameter hook must be callable") + registry = parameter_hooks(parameter) + handle = RemovableHandle(registry.hooks) + registry.hooks[handle.id] = hook + return handle + + +def reject_post_accumulate_hook(*args: Any, **kwargs: Any) -> Any: + raise RuntimeError( + "Live checkpoint parameters do not support register_post_accumulate_grad_hook; use register_hook before transactional gradient publication" + ) + + +def apply_parameter_hooks( + parameter: torch.Tensor, gradient: torch.Tensor +) -> torch.Tensor: + registry = getattr(parameter, "_art_parameter_hooks", None) + if registry is not None: + result = registry.apply(gradient) + if result is not gradient: + gradient.copy_(result) + return gradient + + +def apply_head_hooks( + packets: tuple[Any, ...], registries: dict[str, tuple[ParameterHooks, ...]] +) -> tuple[Any, ...]: + """Combine hooked client captures without weakening any origin's expiry.""" + from ._tensors import CotangentPacket + + groups: dict[ParameterHooks, list[tuple[int, int, Any]]] = {} + gradients = [list(packet.gradients) for packet in packets] + for index, packet in enumerate(packets): + for column, registry in enumerate(registries.get(packet.handle, ())): + if registry.hooks and gradients[index][column] is not None: + groups.setdefault(registry, []).append( + (index, column, json.loads(packet.handle[5:])) + ) + combined = [] + with torch.no_grad(): + for registry, entries in groups.items(): + oldest = min(entries, key=lambda entry: entry[2]["revision"]) + metadata = dict(oldest[2]) + gradient = gradients[oldest[0]][oldest[1]] + parameter = registry.parameter() + if parameter is not None: + gradient = gradient.to(device=parameter.device, dtype=parameter.dtype) + gradient = gradient.clone() + for index, column, _ in entries: + if (index, column) != oldest[:2]: + addition = gradients[index][column].to(gradient) + if ( + gradient.layout != torch.strided + and addition.layout == torch.strided + ): + gradient = addition + gradient + else: + gradient.add_(addition) + gradients[index][column] = None + gradient = registry.apply(gradient) + metadata["keys"] = [metadata["keys"][oldest[1]]] + metadata["capture"] = "hooks:" + metadata["capture"] + metadata["max_gradient_staleness"] = ( + min( + origin["revision"] + origin["max_gradient_staleness"] + for _, _, origin in entries + ) + - metadata["revision"] + ) + combined.append( + CotangentPacket( + "head:" + json.dumps(metadata, separators=(",", ":")), (gradient,) + ) + ) + # Keep original packets for validation even when their gradients were folded + # into a single packet. The combined expiry also governs later optim_step. + return tuple( + CotangentPacket(packet.handle, tuple(values)) + for packet, values in zip(packets, gradients, strict=True) + ) + tuple(combined) diff --git a/src/art/trainer_rank/_tensors.py b/src/art/trainer_rank/_tensors.py new file mode 100644 index 000000000..cffe6512e --- /dev/null +++ b/src/art/trainer_rank/_tensors.py @@ -0,0 +1,491 @@ +"""Detached output trees and transaction-scoped, first-order cotangents.""" + +from __future__ import annotations + +from collections import OrderedDict +from collections.abc import Callable, Sequence +from dataclasses import dataclass, fields, is_dataclass +from threading import Lock +from typing import Any +import weakref + +import torch + + +@dataclass(frozen=True) +class TensorTreeSpec: + kind: str + context: Any = None + children: tuple[TensorTreeSpec, ...] = () + + +def flatten_tensors(tree: Any) -> tuple[tuple[torch.Tensor, ...], TensorTreeSpec]: + """Flatten tensor leaves once per identity, preserving container structure.""" + from ._impl import Unset + + tensors: list[torch.Tensor] = [] + indices: dict[int, int] = {} + + def visit(value: Any) -> TensorTreeSpec: + if value is Unset: + return TensorTreeSpec("unset") + if isinstance(value, torch.Tensor): + if id(value) not in indices: + indices[id(value)] = len(tensors) + tensors.append(value) + return TensorTreeSpec("tensor", indices[id(value)]) + if is_dataclass(value) and not isinstance(value, type): + names = tuple(field.name for field in fields(value)) + children = tuple(visit(getattr(value, name)) for name in names) + if isinstance(value, dict): + return TensorTreeSpec( + "dataclass_dict", + ( + type(value), + names, + tuple(value), + getattr(value, "default_factory", None), + ), + children + tuple(visit(item) for item in value.values()), + ) + return TensorTreeSpec("dataclass", (type(value), names), children) + if isinstance(value, dict): + return TensorTreeSpec( + "dict", + (type(value), tuple(value), getattr(value, "default_factory", None)), + tuple(visit(item) for item in value.values()), + ) + if isinstance(value, (list, tuple)): + return TensorTreeSpec("sequence", type(value), tuple(map(visit, value))) + return TensorTreeSpec("constant", value) + + try: + spec = visit(tree) + return tuple(tensors), spec + finally: + # The recursive closure otherwise retains its tensor list until cyclic GC. + del visit + + +def unflatten_tensors(spec: TensorTreeSpec, tensors: Sequence[torch.Tensor]) -> Any: + if spec.kind == "unset": + from ._impl import Unset + + return Unset + if spec.kind == "tensor": + return tensors[spec.context] + if spec.kind == "constant": + return spec.context + values = [unflatten_tensors(child, tensors) for child in spec.children] + if spec.kind in {"dataclass", "dataclass_dict"}: + cls, names = spec.context[:2] + if spec.kind == "dataclass_dict": + # ModelOutput-style dataclasses have a builtin mapping allocation + # and may prohibit update(); restore their actual mapping separately. + base = OrderedDict if issubclass(cls, OrderedDict) else dict + result = base.__new__(cls) + base.__init__( + result, zip(spec.context[2], values[len(names) :], strict=True) + ) + else: + result = object.__new__(cls) + for name, value in zip(names, values[: len(names)], strict=True): + object.__setattr__(result, name, value) + return result + if spec.kind == "dict": + cls, keys, factory = spec.context + result = cls() if factory is None else cls(factory) + result.update(zip(keys, values, strict=True)) + return result + if spec.kind == "sequence": + cls = spec.context + return cls(*values) if hasattr(cls, "_fields") else cls(values) + raise ValueError(f"Unknown tensor tree node: {spec.kind!r}") + + +def _map_tensors(fn: Callable[[torch.Tensor], torch.Tensor], tree: Any) -> Any: + tensors, spec = flatten_tensors(tree) + return unflatten_tensors(spec, tuple(map(fn, tensors))) + + +def _plain(tensor: torch.Tensor) -> torch.Tensor: + if isinstance(tensor, ManagedTensor): + with torch.enable_grad(), torch._C.DisableTorchFunctionSubclass(): + return tensor.as_subclass(torch.Tensor) + return tensor + + +@dataclass(frozen=True) +class TensorPacket: + handle: str + spec: TensorTreeSpec + tensors: tuple[torch.Tensor, ...] + requires_grad: tuple[bool, ...] + + +@dataclass(frozen=True) +class CotangentPacket: + handle: str + gradients: tuple[torch.Tensor | None, ...] + + +def _validate_output_spec(spec: TensorTreeSpec) -> None: + def metadata(value: Any) -> None: + if type(value) is tuple: + for item in value: + metadata(item) + elif type(value) not in ( + type(None), + bool, + int, + float, + complex, + str, + bytes, + ) and not isinstance( + value, (torch.dtype, torch.device, torch.layout, torch.memory_format) + ): + raise TypeError( + f"Unsupported output tree metadata {type(value).__name__}; " + "use tensor leaves, dataclasses, dictionaries, lists, tuples and scalar metadata." + ) + + if spec.kind == "constant": + metadata(spec.context) + elif spec.kind == "dict": + metadata(spec.context[1]) + metadata(spec.context[2]) + elif spec.kind == "dataclass_dict": + metadata(spec.context[2]) + metadata(spec.context[3]) + for child in spec.children: + _validate_output_spec(child) + + +def detach_tree( + handle: str, tree: Any, *, device: torch.device | str | None = None +) -> TensorPacket: + """Snapshot supported output containers; reject opaque tensor-bearing objects.""" + tensors, spec = flatten_tensors(tree) + _validate_output_spec(spec) + return TensorPacket( + handle, + spec, + tuple( + _plain(tensor).detach().to(device=device, copy=True) for tensor in tensors + ), + tuple(tensor.requires_grad for tensor in tensors), + ) + + +class _OutputBridge(torch.autograd.Function): + @staticmethod + def forward(ctx, anchor, collector, handle, requires_grad, on_release, *values): + ctx.collector, ctx.handle = collector, handle + ctx.signature = tuple( + (value.shape, value.dtype, value.device, required) + for value, required in zip(values, requires_grad, strict=True) + ) + if on_release is not None: + weakref.finalize(ctx, on_release).atexit = False + ctx.set_materialize_grads(False) + # Accessed by backward so PyTorch enforces retain_graph on repeated use. + ctx.save_for_backward(anchor) + ctx.mark_non_differentiable( + *( + value + for value, required in zip(values, requires_grad, strict=True) + if not required + ) + ) + return tuple(values) + + @staticmethod + def backward(ctx, *gradients): + ctx.saved_tensors + ctx.collector._record(ctx.handle, ctx.signature, gradients) + return (None,) * (5 + len(gradients)) + + +class _CollectionRoot(torch.autograd.Function): + @staticmethod + def forward(ctx, collector, *outputs): + ctx.collector = collector + ctx.set_materialize_grads(False) + ctx.mark_non_differentiable( + *(value for value in outputs if not value.requires_grad) + ) + return tuple(outputs) + + @staticmethod + def backward(ctx, *gradients): + with ctx.collector._lock: + ctx.collector._graph_task_id = torch._C._current_graph_task_id() + return (None, *gradients) + + +class CotangentCollector: + """Collect every local cotangent before the caller commits any remote work. + + One instance belongs to one caller/owner, including its head snapshots. A + concurrent or nested remote backward on the same collector is rejected; pass + coupled losses together instead. Ordinary local recomputation is supported. Local + parameter gradients follow normal PyTorch accumulation semantics; remote + cotangents are discarded if any part of local backward raises. + """ + + def __init__(self) -> None: + self._lock = Lock() + self._pending: dict[str, tuple[torch.Tensor | None, ...]] | None = None + self._signatures: dict[str, tuple[Any, ...]] = {} + self._graph_task_id: int | None = None + self._head_hooks: dict[str, tuple[Any, ...]] = {} + + def attach( + self, + packet: TensorPacket, + *, + managed: bool = False, + on_release: Callable[[], None] | None = None, + ) -> Any: + """Attach outputs; optionally release their owner when the graph dies. + + The callback follows the autograd context, including dependent losses, + and must be nonblocking and safe after explicit graph consumption. For + packets without differentiable outputs, release happens immediately. + Grad mode applies normally. Differentiable outputs are read-only custom + Function views; clone before in-place changes, including under no_grad. + """ + _validate_output_spec(packet.spec) + if len(packet.tensors) != len(packet.requires_grad): + raise ValueError( + "Tensor packet values and requires_grad flags differ in length" + ) + values = tuple(_plain(tensor).detach() for tensor in packet.tensors) + if any(packet.requires_grad): + for value, required in zip(values, packet.requires_grad, strict=True): + if required and not (value.is_floating_point() or value.is_complex()): + raise ValueError( + "Only floating-point or complex outputs can require gradients" + ) + values = _OutputBridge.apply( + torch.empty(0, requires_grad=True), + self, + packet.handle, + packet.requires_grad, + on_release, + *values, + ) + elif on_release is not None: + on_release() + if managed: + values = tuple(map(managed_tensor, values)) + return unflatten_tensors(packet.spec, values) + + def _record( + self, + handle: str, + signature: tuple[Any, ...], + gradients: Sequence[torch.Tensor | None], + ) -> None: + copies = tuple( + None if grad is None else _plain(grad).detach().clone() + for grad in gradients + ) + with self._lock: + if ( + self._pending is None + or self._graph_task_id != torch._C._current_graph_task_id() + ): + raise RuntimeError( + "Use the owning trainer.backward(loss) to collect remote cotangents; " + "unscoped or nested remote backward is unsupported. Pass coupled losses together." + ) + if handle in self._signatures and self._signatures[handle] != signature: + raise ValueError( + f"Incompatible output signatures for repeated handle {handle!r}" + ) + self._signatures[handle] = signature + previous = self._pending.get(handle) + if previous is not None: + copies = tuple( + new + if old is None + else old + if new is None + else new + old + if old.layout != torch.strided and new.layout == torch.strided + else old + new + for old, new in zip(previous, copies, strict=True) + ) + self._pending[handle] = copies + + def backward( + self, + loss: torch.Tensor | Sequence[torch.Tensor], + gradient: torch.Tensor | Sequence[torch.Tensor | None] | None = None, + *, + retain_graph: bool = False, + ) -> tuple[CotangentPacket, ...]: + from ._heads import _ClientParameter + from ._impl import _TrackedParameter + + with self._lock: + if self._pending is not None: + raise RuntimeError( + "A backward collection is already active on this owner" + ) + self._pending = {} + try: + outputs = (loss,) if isinstance(loss, torch.Tensor) else tuple(loss) + # A single root records the engine task before any upstream hooks or + # bridges execute, including when local backward runs on CUDA workers. + with torch.enable_grad(): + # Direct live roots need the same version capture as arithmetic. + # Deduplicate aliases so a repeated root shares one snapshot. + outputs = _map_tensors( + lambda value: ( + value.clone() + if isinstance(value, _ClientParameter | _TrackedParameter) + else value + ), + outputs, + ) + roots = _CollectionRoot.apply(self, *map(_plain, outputs)) + head_hooks = self._head_hooks.copy() + torch.autograd.backward(roots, gradient, retain_graph=retain_graph) + with self._lock: + packets = tuple( + CotangentPacket(handle, grads) + for handle, grads in sorted(self._pending.items()) + ) + from ._parameter_hooks import apply_head_hooks + + return apply_head_hooks(packets, head_hooks) + finally: + with self._lock: + self._pending = None + self._signatures.clear() + self._graph_task_id = None + + +# Implicit copies are only safe for known pure operations. Stateful functional +# calls (for example BatchNorm and embedding(max_norm=...)) can mutate arguments +# without an in-place name or even a tensor version-counter change. +_MIXED_DEVICE_PURE_OPS = frozenset( + "__add__ __radd__ __sub__ __rsub__ __mul__ __rmul__ __truediv__ __rtruediv__ " + "__floordiv__ __rfloordiv__ __pow__ __rpow__ __mod__ __rmod__ __matmul__ __rmatmul__ " + "__eq__ __ne__ __lt__ __le__ __gt__ __ge__ __and__ __rand__ __or__ __ror__ __xor__ __rxor__ " + "__getitem__ add sub subtract mul multiply div divide true_divide floor_divide pow " + "remainder fmod maximum minimum fmax fmin eq ne lt le gt ge equal allclose isclose " + "logical_and logical_or logical_xor bitwise_and bitwise_or bitwise_xor " + "matmul mm bmm mv dot vdot inner outer addmm addbmm baddbmm addmv addr linear " + "einsum tensordot bilinear cat concat concatenate stack hstack vstack dstack " + "where lerp clamp clip masked_select gather take take_along_dim index_select " + "scatter scatter_add scatter_reduce index_add index_copy index_fill index_put " + "mse_loss l1_loss smooth_l1_loss huber_loss binary_cross_entropy " + "binary_cross_entropy_with_logits cross_entropy nll_loss kl_div poisson_nll_loss " + "cosine_similarity cosine_embedding_loss hinge_embedding_loss margin_ranking_loss " + "triplet_margin_loss pairwise_distance pdist cdist".split() +) + + +class ManagedTensor(torch.Tensor): + """Eager CPU/CUDA interop that copies operands through ordinary autograd. + + Computation follows managed placement; CPU wins when managed operands + disagree. Results remain managed. Mixed-device mutation and multiple CUDA + devices require an explicit move. Mixed-device support is limited to pure + arithmetic, linear algebra, indexing, tensor combination and common losses. + Stateful or unknown mixed-device operations require explicit placement. + Arbitrary subclasses and compiled graphs are outside this eager interface. + """ + + @classmethod + def __torch_function__(cls, func, types, args=(), kwargs=None): + # Let live parameter/buffer proxies capture their values before running + # the operation with subclass dispatch disabled. + if not all(issubclass(cls, other) for other in types): + return NotImplemented + name = getattr(func, "__name__", "") + identity_property = name == "__get__" and getattr( + getattr(func, "__self__", None), "__name__", "" + ) not in {"T", "mT", "H", "mH", "real", "imag"} + if ( + identity_property + or func + in ( + torch.autograd.grad, + torch.autograd.backward, + torch.Tensor.backward, + ) + or name + in { + "__set__", + "register_hook", + "register_post_accumulate_grad_hook", + "retain_grad", + "requires_grad_", + "detach_", + } + ): + # Autograd targets and metadata belong to the original tensor, not + # a fresh alias (which is absent from the user's existing graph). + with torch._C.DisableTorchFunctionSubclass(): + return func(*args, **(kwargs or {})) + kwargs = kwargs or {} + originals, _ = flatten_tensors((args, kwargs)) + identities = {id(tensor) for tensor in originals} + # Operations must receive the original TensorImpl/autograd identity. + # Unwrapping into aliases loses metadata mutations and existing hooks. + with torch._C.DisableTorchFunctionSubclass(): + devices = {tensor.device for tensor in originals} + if len(devices) > 1 and name not in {"to", "type_as"}: + accelerators = {device for device in devices if device.type != "cpu"} + if len(accelerators) != 1 or next(iter(accelerators)).type != "cuda": + raise RuntimeError( + "Managed tensors require an explicit move between accelerator devices" + ) + if name not in _MIXED_DEVICE_PURE_OPS or kwargs.get("out") is not None: + raise RuntimeError( + f"Managed mixed-device operation {name!r} across " + f"{', '.join(sorted(map(str, devices)))} may perform mutation " + "or is unsupported; use explicit tensor .to(device) placement." + ) + managed_devices = { + tensor.device + for tensor in originals + if isinstance(tensor, ManagedTensor) + } + device = ( + torch.device("cpu") + if torch.device("cpu") in managed_devices + else next(iter(managed_devices)) + ) + args, kwargs = _map_tensors( + lambda tensor: tensor.to(device), (args, kwargs) + ) + result = func(*args, **kwargs) + + def wrap(tensor: torch.Tensor) -> torch.Tensor: + # Keep no-op/in-place/out identities, including ordinary operands. + # New results already own the right autograd/view metadata; an + # as_subclass alias would turn even clone() into a view and lose + # its hooks when a later in-place operation rebases that view. + if id(tensor) not in identities: + tensor.__class__ = ManagedTensor + return tensor + + return _map_tensors(wrap, result) + + +def managed_tensor(tensor: torch.Tensor) -> torch.Tensor: + if isinstance(tensor, ManagedTensor): + return tensor + with torch.enable_grad(), torch._C.DisableTorchFunctionSubclass(): + return tensor.as_subclass(ManagedTensor) + + +def managed_tree(tree: Any, *, device: torch.device | str | None = None) -> Any: + """Place a nested output tree while preserving its original gradient paths.""" + return _map_tensors(lambda tensor: managed_tensor(tensor.to(device=device)), tree) diff --git a/src/art/trainer_rank/_versions.py b/src/art/trainer_rank/_versions.py new file mode 100644 index 000000000..4776e891b --- /dev/null +++ b/src/art/trainer_rank/_versions.py @@ -0,0 +1,336 @@ +"""Immutable forward identities and transactional routing to current parameters.""" + +from __future__ import annotations + +from collections.abc import Callable, Iterable, Iterator, Sequence +from contextlib import contextmanager +from dataclasses import dataclass, field +import threading +from typing import TYPE_CHECKING, Any, cast +import weakref + +import torch + +if TYPE_CHECKING: + from ._impl import TrainerRank + + +@dataclass(frozen=True) +class CheckpointVersion: + checkpoint: str + generation: int + revision: int + + +VersionedGradient = tuple[CheckpointVersion, int, torch.nn.Parameter, torch.Tensor] + + +@dataclass +class _GradientBatch: + gradients: dict[int, tuple[torch.nn.Parameter, torch.Tensor]] = field( + default_factory=dict + ) + origins: set[tuple[CheckpointVersion, int, int]] = field(default_factory=set) + failed: bool = False + + def add( + self, + version: CheckpointVersion, + maximum: int, + parameter: torch.nn.Parameter, + gradient: torch.Tensor, + ) -> None: + key = id(parameter) + if key in self.gradients: + current = self.gradients[key][1] + if current.layout != torch.strided and gradient.layout == torch.strided: + self.gradients[key] = parameter, gradient.detach() + current + else: + current.add_(gradient.detach()) + else: + self.gradients[key] = (parameter, gradient.detach().clone()) + self.origins.add((version, maximum, key)) + + def validations(self) -> Iterator[VersionedGradient]: + for version, maximum, key in self.origins: + yield version, maximum, *self.gradients[key] + + def clear(self) -> None: + self.gradients.clear() + self.origins.clear() + + +@dataclass +class _PreparedGradients: + parameters: list[tuple[torch.nn.Parameter, torch.Tensor, torch.Tensor | None]] + origins: dict[str, set[tuple[CheckpointVersion, int]]] + + def clear(self) -> None: + self.parameters.clear() + self.origins.clear() + + +class CheckpointVersions: + def __init__(self, trainer: TrainerRank) -> None: + self._trainer = weakref.ref(trainer) + self._transaction: _GradientBatch | None = None + self._lock = threading.RLock() + self._origins: dict[str, set[tuple[CheckpointVersion, int]]] = {} + self.generation = 0 + self.lora: weakref.WeakValueDictionary[ + tuple[CheckpointVersion, CheckpointVersion, int], Any + ] = weakref.WeakValueDictionary() + + def capture(self, name: str) -> CheckpointVersion: + trainer = self._trainer() + if trainer is None: + raise RuntimeError("TrainerRank no longer exists") + slot = trainer._checkpoint_slots[name] + return CheckpointVersion(name, slot.generation, slot.revision) + + def validate(self, version: CheckpointVersion, maximum: int = 2) -> None: + if isinstance(maximum, bool) or not isinstance(maximum, int) or maximum < 0: + raise ValueError("max_gradient_staleness must be a nonnegative integer") + trainer = self._trainer() + if trainer is None: + raise RuntimeError("TrainerRank no longer exists") + slot = trainer._checkpoint_slots.get(version.checkpoint) + if slot is None or slot.generation != version.generation: + raise trainer._slot_state_error( + f"Checkpoint {version.checkpoint!r} was replaced after forward" + ) + age = slot.revision - version.revision + if not 0 <= age <= maximum: + raise trainer._slot_state_error( + f"Checkpoint {version.checkpoint!r} gradient staleness {age} exceeds " + f"max_gradient_staleness={maximum} (forward revision " + f"{version.revision}, current revision {slot.revision})" + ) + + def snapshot( + self, + parameter: torch.nn.Parameter, + version: CheckpointVersion, + maximum: int = 2, + ) -> torch.nn.Parameter: + self.validate(version, maximum) + with torch._C.DisableTorchFunctionSubclass(): + result = torch.nn.Parameter( + parameter.detach().clone(), requires_grad=parameter.requires_grad + ) + self.track(result, parameter, version, maximum) + return result + + def track( + self, + snapshot: torch.nn.Parameter, + parameter: torch.nn.Parameter, + version: CheckpointVersion, + maximum: int = 2, + ) -> None: + if snapshot.requires_grad: + snapshot.register_hook( + lambda grad: self._stage(version, maximum, parameter, grad) + ) + snapshot.register_post_accumulate_grad_hook( + lambda parameter: setattr(parameter, "grad", None) + ) + + def _stage( + self, + version: CheckpointVersion, + maximum: int, + parameter: torch.nn.Parameter, + gradient: torch.Tensor, + ) -> None: + with self._lock: + batch = self._transaction + if batch is None: + raise RuntimeError( + "Backward through captured parameters requires TrainerRank.backward() " + "or an explicit _gradient_transaction()" + ) + try: + self.validate(version, maximum) + batch.add(version, maximum, parameter, gradient) + except BaseException: + batch.failed = True + batch.clear() + raise + + @contextmanager + def transaction( + self, *, before_commit: Callable[[Callable[[], None]], None] | None = None + ) -> Iterator[None]: + """Coordinate one exit phase on success or failure, then publish once.""" + owns_batch = self._transaction is None + batch = self._transaction = self._transaction or _GradientBatch() + error: BaseException | None = None + prepared: _PreparedGradients | None = None + + def validate() -> None: + nonlocal prepared + if error is not None: + raise error + if batch.failed: + raise RuntimeError("A nested gradient transaction failed") + if owns_batch: + prepared = self._prepare_batch(batch) + else: + self.validate_gradients(batch.validations()) + + try: + try: + yield + except BaseException as exc: + error = exc + try: + if before_commit is None: + validate() + else: + before_commit(validate) + if error is not None: + raise error + if owns_batch: + if prepared is None: + raise RuntimeError( + "before_commit must invoke its validation callback" + ) + self._publish(prepared) + except BaseException: + batch.failed = True + if error is not None: + raise error + raise + finally: + if prepared is not None: + prepared.clear() + if owns_batch: + batch.clear() + self._transaction = None + + def validate_gradients(self, gradients: Iterable[VersionedGradient]) -> None: + with torch._C.DisableTorchFunctionSubclass(): + trainer = self._trainer() + if trainer is None: + raise RuntimeError("TrainerRank no longer exists") + targets: dict[str, set[int]] = {} + parameter = gradient = None + try: + for version, maximum, parameter, gradient in gradients: + self.validate(version, maximum) + if version.checkpoint not in targets: + targets[version.checkpoint] = { + id(current) + for current in trainer._checkpoint_slots[ + version.checkpoint + ].params + } + if id(parameter) not in targets[version.checkpoint]: + raise trainer._slot_state_error( + f"Checkpoint {version.checkpoint!r} gradient target was replaced" + ) + if ( + gradient.shape != parameter.shape + or gradient.device != parameter.device + or gradient.dtype != parameter.dtype + ): + raise ValueError( + "Versioned gradient shape/device/dtype differs from parameter" + ) + if any( + tensor.layout != torch.strided + for tensor in (parameter, gradient, parameter.grad) + if tensor is not None + ): + raise ValueError( + "Versioned gradients require strided tensor layouts" + ) + finally: + parameter = gradient = None + + def accumulate(self, gradients: Sequence[VersionedGradient]) -> None: + with self._lock: + self.validate_gradients(gradients) + if self._transaction is None: + self.commit(gradients) + else: + for entry in gradients: + self._transaction.add(*entry) + + def commit(self, gradients: Sequence[VersionedGradient]) -> None: + batch = _GradientBatch() + try: + self.validate_gradients(gradients) + for entry in gradients: + batch.add(*entry) + self._commit_batch(batch) + finally: + batch.clear() + + def _prepare_batch(self, batch: _GradientBatch) -> _PreparedGradients: + from ._parameter_hooks import apply_parameter_hooks + + prepared = _PreparedGradients([], {}) + parameter = gradient = previous = combined = None + try: + with self._lock, torch.no_grad(): + self.validate_gradients(batch.validations()) + # Prepare every allocation before coordinated exit/publication. + for parameter, gradient in batch.gradients.values(): + gradient = apply_parameter_hooks(parameter, gradient) + with torch._C.DisableTorchFunctionSubclass(): + previous = parameter.grad + combined = gradient if previous is None else previous + gradient + prepared.parameters.append((parameter, combined, previous)) + prepared.origins = { + name: origins.copy() for name, origins in self._origins.items() + } + for version, maximum, _ in batch.origins: + prepared.origins.setdefault(version.checkpoint, set()).add( + (version, maximum) + ) + return prepared + except BaseException: + prepared.clear() + raise + finally: + # A retained exception traceback must not own unpublished tensors. + parameter = gradient = previous = combined = None + + def _publish(self, prepared: _PreparedGradients) -> None: + with self._lock, torch.no_grad(), torch._C.DisableTorchFunctionSubclass(): + parameter = gradient = previous = None + try: + for parameter, gradient, _ in prepared.parameters: + parameter.grad = gradient + except BaseException: + for parameter, gradient, previous in prepared.parameters: + cast(Any, torch.Tensor.grad).__set__(parameter, previous) + raise + else: + self._origins, prepared.origins = prepared.origins, {} + finally: + parameter = gradient = previous = None + + def _commit_batch(self, batch: _GradientBatch) -> None: + prepared = None + try: + prepared = self._prepare_batch(batch) + self._publish(prepared) + finally: + batch.clear() + if prepared is not None: + prepared.clear() + + def validate_accumulated(self, names: Sequence[str]) -> None: + for name in names: + for version, maximum in self._origins.get(name, ()): + self.validate(version, maximum) + + def clear(self, names: Sequence[str] | None = None) -> None: + if names is None: + self._origins.clear() + else: + for name in names: + self._origins.pop(name, None) diff --git a/tests/acceptance/trainer_rank_planner/test_public_contract.py b/tests/acceptance/trainer_rank_planner/test_public_contract.py index 80780de81..37daad4f0 100644 --- a/tests/acceptance/trainer_rank_planner/test_public_contract.py +++ b/tests/acceptance/trainer_rank_planner/test_public_contract.py @@ -4,10 +4,9 @@ contract (research thread behavior spec, frozen 2026-08-31): - ``TrainerRank`` exposes no prefix-sharing depth, microbatch width, - head-chunk, or memory-safety policy knob. Its constructor accepts only the - training runtime. -- ``forward_micro_batches`` and ``dp_rank_forward`` accept only - ``inputs``, ``checkpoint``, and ``no_grad``, plus ``yield_empty`` on the iterator. + head-chunk, or memory-safety policy knob. Its constructor accepts the training runtime and immutable forward policy. +- ``forward_batches`` and ``forward`` accept only + ``inputs``, ``options``, ``checkpoint``, and ``no_grad``, plus ``yield_empty`` on the iterator. - ``TrainerRankMemoryError`` reports only a predicted peak, the usable limit, and an actionable reduction suggestion. It carries no infeasibility proof. @@ -60,11 +59,10 @@ def _parameters(callable_: Any) -> dict[str, inspect.Parameter]: } -def test_constructor_accepts_only_the_training_runtime() -> None: +def test_constructor_accepts_runtime_and_forward_options() -> None: parameters = _parameters(trainer_rank.TrainerRank.__init__) - assert list(parameters) == ["runtime"], ( - "TrainerRank must accept exactly one constructor argument (the training" - f" runtime); found {sorted(parameters)}" + assert list(parameters) == ["runtime", "options"], ( + f"Unexpected TrainerRank constructor parameters: {sorted(parameters)}" ) @@ -78,12 +76,12 @@ def test_constructor_rejects_policy_knob(knob: str) -> None: ), "TrainerRank.__init__ must not accept **kwargs (knobs could pass silently)" -@pytest.mark.parametrize("method_name", ("forward_micro_batches", "dp_rank_forward")) +@pytest.mark.parametrize("method_name", ("forward_batches", "forward")) def test_forward_method_signatures_are_knob_free(method_name: str) -> None: method = getattr(trainer_rank.TrainerRank, method_name) parameters = _parameters(method) - allowed = {"inputs", "checkpoint", "no_grad"} - if method_name == "forward_micro_batches": + allowed = {"inputs", "options", "checkpoint", "no_grad"} + if method_name == "forward_batches": allowed.add("yield_empty") assert parameters["yield_empty"].default is False assert set(parameters) <= allowed, ( @@ -117,7 +115,7 @@ def test_memory_error_reports_actionable_fields_without_proof() -> None: def test_no_public_test_anchor_hook() -> None: """Forced layout anchors are test-only; they must not be public API.""" - for method_name in ("__init__", "forward_micro_batches", "dp_rank_forward"): + for method_name in ("__init__", "forward_batches", "forward"): parameters = _parameters(getattr(trainer_rank.TrainerRank, method_name)) leaked = [ name diff --git a/tests/integration/megatron/cp_attn/test_cpu_offload_residency.py b/tests/integration/megatron/cp_attn/test_cpu_offload_residency.py new file mode 100644 index 000000000..ed7b1910f --- /dev/null +++ b/tests/integration/megatron/cp_attn/test_cpu_offload_residency.py @@ -0,0 +1,253 @@ +"""Actual CP2 graph residency and constrained complete-root placement.""" + +from dataclasses import asdict, replace +from datetime import timedelta +import gc +import json +import os +from typing import cast +from unittest.mock import patch + +import pytest +import torch +import torch.distributed as dist +import torch.multiprocessing as mp + +pytest.importorskip("megatron.core") + +from art.megatron.context_parallel import executor # noqa: E402 +from art.megatron.context_parallel.runtime import ( + prepare_megatron_context_parallel_state, +) # noqa: E402 +from art.megatron.context_parallel.types import ( # noqa: E402 + ContextParallelConfig, + ParallelTopology, +) +from art.megatron.flex_attn import compiled # noqa: E402 +from art.megatron.runtime.compile_cache import configure_reusable_backward # noqa: E402 +from art.preprocessing.pack import PackedTensors # noqa: E402 +from art.trainer_rank._graphs import GraphCache # noqa: E402 +from art.trainer_rank._memory_policy import ( # noqa: E402 + ForwardMemoryCost, + choose_memory_placement, + placement_cost, +) + + +@pytest.mark.skipif( + os.environ.get("ART_GRAPH_GPU_TEST") != "1" or torch.cuda.device_count() < 2, + reason="requires two reserved GPUs", +) +def test_cp_cpu_residency_constrains_complete_root(tmp_path): + mp.spawn( + _worker, args=(f"file://{tmp_path / 'cp-residency'}",), nprocs=2, join=True + ) + + +def _worker(rank, rendezvous): + torch.set_num_threads(2) + torch.cuda.set_device(rank) + device = torch.device("cuda", rank) + configure_reusable_backward() + dist.init_process_group( + "nccl", + init_method=rendezvous, + rank=rank, + world_size=2, + timeout=timedelta(seconds=120), + device_id=device, + ) + try: + with ( + patch.object(compiled, "_FORCED_FLEX_BACKEND", "TRITON"), + patch.object( + compiled, + "sparse_compiled_flex_attention", + compiled.triton_sparse_compiled_flex_attention, + ), + ): + _check(rank, device) + finally: + dist.destroy_process_group() + + +def _check(rank, device): + length, heads, dim = 512, 2, 64 + micro = cast( + PackedTensors, + { + "tokens": torch.arange(length)[None], + "group_ids": torch.ones((1, length), dtype=torch.long), + "parent_ids": torch.ones((1, length), dtype=torch.long), + "input_pos": torch.arange(length)[None], + }, + ) + state, plan, _, _ = prepare_megatron_context_parallel_state( + micro=micro, + topology=ParallelTopology(cp=2), + config=ContextParallelConfig( + planner_chunk_size=128, planner_owned_token_ms=1.0 + ), + cp_group=dist.group.WORLD, + cp_rank=rank, + target_device=device, + ) + indices = torch.tensor( + [ + index + for start, end, _local_start in plan.token_layout_index.ownership_ranges_by_rank[ + rank + ] + for index in range(start, end) + ], + device=device, + ) + assert indices.numel() > 0 + assert indices.numel() == sum(plan.local_valid_lengths) + executor.prepare_context_parallel_execution_state(state=state, device=device) + torch.manual_seed(841) + full = tuple(torch.randn(length, 1, heads, dim).to(device) for _ in range(3)) + local = tuple(value.index_select(0, indices) for value in full) + weight = torch.nn.Parameter(torch.ones((), device=device)) + reference_weight = torch.ones((), device=device, requires_grad=True) + q, k, v = (value[:, 0].transpose(0, 1) * reference_weight for value in full) + scores = (q @ k.transpose(-1, -2)) * dim**-0.5 + mask = torch.ones(length, length, dtype=torch.bool, device=device).tril() + reference = (scores.masked_fill(~mask, -torch.inf).softmax(-1) @ v).sum() + (expected,) = torch.autograd.grad(reference, reference_weight) + expected = expected.detach() + del q, k, v, scores, mask, reference, reference_weight, full + cache = GraphCache() + + def execute(inputs): + q, k, v = (value * weight for value in inputs) + return ( + executor.run_context_parallel( + query=q, + key=k, + value=v, + state=state, + scale=dim**-0.5, + enable_gqa=False, + compile_enabled=True, + ).sum(), + ) + + def run(retention, count=1): + weight.grad = None + gc.collect() + torch.cuda.synchronize() + baseline = torch.cuda.memory_allocated() + torch.cuda.reset_peak_memory_stats() + records = [ + cache.run( + execute, + local, + retention=retention, + output_device="cpu", + cuda_devices=[rank], + keep_on_device=lambda tensor: ( + tensor.untyped_storage().data_ptr() + == weight.untyped_storage().data_ptr() + ), + ) + for _ in range(count) + ] + torch.cuda.synchronize() + retained = torch.cuda.memory_allocated() - baseline + states = [cache.state(handle) for handle, _ in records] + for handle, outputs in records: + cache.backward(handle, tuple(torch.ones_like(value) for value in outputs)) + torch.cuda.synchronize() + peak = torch.cuda.max_memory_allocated() - baseline + assert weight.grad is not None + observed = weight.grad.detach().clone() + dist.all_reduce(observed) + torch.testing.assert_close(observed, expected * count, atol=3e-4, rtol=3e-4) + assert not cache.handles() + return dict(retained=retained, peak=peak, states=states) + + run("gpu") # Compile and initialize communication before measurements. + run("cpu") + gpu, cpu = run("gpu"), run("cpu") + print( + "CP_RESIDENCY_PROBE=" + + json.dumps( + dict( + rank=rank, + gpu_retained=gpu["retained"], + cpu_retained=cpu["retained"], + reported_cpu_state=asdict(cpu["states"][0]), + ) + ), + flush=True, + ) + resident = getattr(cpu["states"][0], "non_offloadable_bytes", None) + assert resident is not None and resident > 0 + assert cpu["retained"] <= cpu["states"][0].gpu_bytes + 64 * 1024 + assert cpu["retained"] < gpu["retained"] + peak = int(max(gpu["peak"], cpu["peak"]) * 1.2) + 64 * 1024 + cost = ForwardMemoryCost( + peak_bytes=peak, + retained_bytes=max(int(gpu["states"][0].gpu_bytes * 1.2), int(resident * 1.2)), + output_bytes=4, + cpu_resident_bytes=int(resident * 1.2), + ) + count = 8 + corrected = placement_cost( + [cost] * count, backward_state="cpu", output_device="cpu" + ) + old = placement_cost( + [replace(cost, cpu_resident_bytes=0)] * count, + backward_state="cpu", + output_device="cpu", + ) + cap = old.gpu_required_bytes + resident + assert old.gpu_required_bytes <= cap < corrected.gpu_required_bytes + assert ( + choose_memory_placement( + [cost] * count, + gpu_available_bytes=cap, + cpu_available_bytes=1 << 40, + backward_state="cpu", + output_device="cpu", + ) + is None + ) + replay_plan = choose_memory_placement( + [cost] * count, + gpu_available_bytes=cap, + cpu_available_bytes=1 << 40, + backward_state="replay", + output_device="cpu", + ) + assert replay_plan is not None + replay = run("replay", count) + assert replay["peak"] <= cap + partial_plan = choose_memory_placement( + [cost] * count, + gpu_available_bytes=corrected.gpu_required_bytes, + cpu_available_bytes=1 << 40, + output_device="cpu", + ) + assert partial_plan is not None and partial_plan.backward_state == "cpu" + partial = run("cpu", count) + assert cap < partial["peak"] <= corrected.gpu_required_bytes + print( + "CP_RESIDENCY=" + + json.dumps( + dict( + rank=rank, + gpu_retained=gpu["retained"], + cpu_retained=cpu["retained"], + reported_cpu_state=asdict(cpu["states"][0]), + old_required=old.gpu_required_bytes, + corrected_required=corrected.gpu_required_bytes, + cap=cap, + replay_peak=replay["peak"], + partial_cpu_peak=partial["peak"], + children=count, + ) + ), + flush=True, + ) diff --git a/tests/integration/megatron/cp_attn/test_retained_backward.py b/tests/integration/megatron/cp_attn/test_retained_backward.py new file mode 100644 index 000000000..0ece0cc27 --- /dev/null +++ b/tests/integration/megatron/cp_attn/test_retained_backward.py @@ -0,0 +1,206 @@ +"""Actual CP collectives and compiled attention against a dense manual oracle.""" + +from datetime import timedelta +import os +from pathlib import Path +from typing import cast +from unittest.mock import patch +import weakref + +import pytest +import torch +import torch.distributed as dist +import torch.multiprocessing as mp + +pytest.importorskip("megatron.core") + +from art.megatron.context_parallel import executor # noqa: E402 +from art.megatron.context_parallel.runtime import ( # noqa: E402 + prepare_megatron_context_parallel_state, +) +from art.megatron.context_parallel.types import ( # noqa: E402 + ContextParallelConfig, + ParallelTopology, +) +from art.megatron.flex_attn import compiled # noqa: E402 +from art.megatron.runtime.compile_cache import configure_reusable_backward # noqa: E402 +from art.preprocessing.pack import PackedTensors # noqa: E402 +from art.trainer_rank._graphs import GraphCache # noqa: E402 + + +def test_cp_retained_failure_releases_original_records(monkeypatch): + contexts, saved = [], [] + + def recorded(*, query, key, value, **kwargs): + output = query * key + value + saved.append(weakref.ref(output)) + return output, output.detach(), output.detach(), [{"stage_out": output}] + + def fail(*, replay_records, **kwargs): + # Stage cleanup already consumed a per-backward dictionary when a + # later operation fails. The untouched originals must also be released. + replay_records[0].clear() + raise RuntimeError("injected CP backward failure") + + monkeypatch.setattr(executor, "_run_context_parallel_forward_recorded", recorded) + monkeypatch.setattr(executor, "_run_context_parallel_backward", fail) + weight = torch.nn.Parameter(torch.tensor(2.0)) + + def execute(_): + output = executor.ArtContextParallelFn.apply( + weight, weight, weight, None, None, 1.0, False, True, None, () + ) + contexts.append(output.grad_fn) + return (output,) + + cache = GraphCache() + handle, (output,) = cache.run(execute, ()) + with pytest.raises(RuntimeError, match="injected CP backward failure"): + cache.backward(handle, (torch.ones_like(output),), retain_graph=True) + assert cache.handles() == () + assert getattr(contexts[0], "replay_records") is None + assert all(reference() is None for reference in saved) + assert weight.grad is None + + +@pytest.mark.skipif( + os.environ.get("ART_GRAPH_GPU_TEST") != "1" or torch.cuda.device_count() < 2, + reason="requires two reserved GPUs", +) +@pytest.mark.parametrize("backend,dim", [("TRITON", 64), ("FLASH", 64), ("FLASH", 128)]) +def test_cp_retained_backward_matches_dense_attention( + tmp_path: Path, backend: str, dim: int +): + mp.spawn( + _worker, + args=(f"file://{tmp_path / 'rendezvous'}", backend, dim), + nprocs=2, + join=True, + ) + + +def _worker(rank: int, init_method: str, backend: str, dim: int) -> None: + # Exercise group ranks independently from CUDA device numbering. + device = torch.device("cuda", 1 - rank) + torch.cuda.set_device(device) + configure_reusable_backward() + dist.init_process_group( + "nccl", + init_method=init_method, + rank=rank, + world_size=2, + timeout=timedelta(seconds=90), + device_id=device, + ) + try: + if backend == "TRITON": + with ( + patch.object(compiled, "_FORCED_FLEX_BACKEND", "TRITON"), + patch.object( + compiled, + "sparse_compiled_flex_attention", + compiled.triton_sparse_compiled_flex_attention, + ), + ): + _check_repeated_backward(rank, device, backend, dim) + else: + _check_repeated_backward(rank, device, backend, dim) + finally: + dist.destroy_process_group() + + +def _check_repeated_backward( + rank: int, device: torch.device, backend: str, dim: int +) -> None: + length, heads = 512, 2 + micro = cast( + PackedTensors, + { + "tokens": torch.arange(length)[None], + "group_ids": torch.ones((1, length), dtype=torch.long), + "parent_ids": torch.ones((1, length), dtype=torch.long), + "input_pos": torch.arange(length)[None], + }, + ) + state, plan, _, _ = prepare_megatron_context_parallel_state( + micro=micro, + topology=ParallelTopology(cp=2), + config=ContextParallelConfig( + planner_chunk_size=128, planner_owned_token_ms=1.0 + ), + cp_group=dist.group.WORLD, + cp_rank=rank, + target_device=device, + ) + indices = torch.tensor( + [ + index + for start, end, _ in plan.token_layout_index.ownership_ranges_by_rank[rank] + for index in range(start, end) + ], + device=device, + ) + assert indices.numel() > 0 + assert indices.numel() == sum(plan.local_valid_lengths) + executor.prepare_context_parallel_execution_state(state=state, device=device) + torch.manual_seed(841) + dtype = torch.bfloat16 if backend == "FLASH" else torch.float32 + full = tuple( + torch.randn(length, 1, heads, dim, dtype=dtype).to(device) for _ in range(3) + ) + local = tuple(value.index_select(0, indices).requires_grad_() for value in full) + refs = tuple(value.float().requires_grad_() for value in full) + q, k, v = (value[:, 0].transpose(0, 1) for value in refs) + scores = (q @ k.transpose(-1, -2)) * dim**-0.5 + mask = torch.ones(length, length, dtype=torch.bool, device=device).tril() + reference = (scores.masked_fill(~mask, -torch.inf).softmax(-1) @ v).transpose(0, 1)[ + :, None + ] + with patch.object( + executor, "_forward_stage_records", wraps=executor._forward_stage_records + ) as forwards: + output = executor.run_context_parallel( + query=local[0], + key=local[1], + value=local[2], + state=state, + scale=dim**-0.5, + enable_gqa=False, + compile_enabled=True, + ) + context = output.grad_fn + assert context is not None + records = getattr(context, "replay_records") + saved_refs = [weakref.ref(record["stage_out"]) for record in records] + del records + atol, rtol = (0.012, 0.025) if backend == "FLASH" else (3e-5, 3e-4) + torch.testing.assert_close( + output.float(), reference.index_select(0, indices), atol=atol, rtol=rtol + ) + for step, retain in enumerate((True, True, False)): + torch.manual_seed(920 + step) + cotangent = torch.randn(length, 1, heads, dim, dtype=dtype).to(device) + expected = torch.autograd.grad( + reference, refs, cotangent.float(), retain_graph=True + ) + actual = torch.autograd.grad( + output, local, cotangent.index_select(0, indices), retain_graph=retain + ) + torch.cuda.synchronize(device) + for observed, wanted in zip(actual, expected, strict=True): + torch.testing.assert_close( + observed.float(), + wanted.index_select(0, indices), + atol=atol, + rtol=rtol, + ) + assert forwards.call_count == 1 + if retain: + assert getattr(context, "replay_records") + else: + assert getattr(context, "replay_records") is None + assert all(ref() is None for ref in saved_refs) + print( + f"rank={rank} policy={backend} dim={dim} backward={step + 1} retain={retain} passed", + flush=True, + ) diff --git a/tests/integration/megatron/lora/test_dynamic_lora_slots.py b/tests/integration/megatron/lora/test_dynamic_lora_slots.py index cd9d47f60..0782d35e0 100644 --- a/tests/integration/megatron/lora/test_dynamic_lora_slots.py +++ b/tests/integration/megatron/lora/test_dynamic_lora_slots.py @@ -140,7 +140,7 @@ def test_trainer_rank_custom_objects_train_and_become_stale_on_cuda() -> None: output = head.score(torch.randn(3, 4, device=device))["value"] * gain + running with pytest.raises(TrainerRankSlotStateError, match="live backward graph"): trainer._guard_slot_can_load(trainer._slot_ref("A")) - output.sum().backward() + trainer.backward(output.sum()) before = tuple(param.detach().clone() for param in head.parameters()) + ( gain.detach().clone(), ) @@ -270,7 +270,8 @@ def _custom_parameter_reduction_worker( checkpoint="A", ) torch.testing.assert_close(parameter, torch.tensor(1.0, device=device)) - (parameter * float(rank + 1)).backward() + trainer.backward(parameter * float(rank + 1)) + assert parameter.grad is not None (reduced,) = trainer._reduce_dynamic_grads((parameter,), scale_grads=1.0) expected = {"dp": 3.0, "tp": 1.5, "cp": 1.5, "tp_cp": 2.5}[topology] torch.testing.assert_close(reduced, torch.tensor(expected, device=device)) diff --git a/tests/integration/megatron/lora/test_lora_versions.py b/tests/integration/megatron/lora/test_lora_versions.py new file mode 100644 index 000000000..07478d7b5 --- /dev/null +++ b/tests/integration/megatron/lora/test_lora_versions.py @@ -0,0 +1,206 @@ +from __future__ import annotations + +import gc +import weakref + +import pytest + +torch = pytest.importorskip("torch") +pytest.importorskip("megatron.core") + +from art.megatron.lora import LoRA, LoRASlotRef, use_lora_slot # noqa: E402 +from art.trainer_rank import TrainerRankSlotStateError # noqa: E402 +from art.trainer_rank._checkpoint import ( # noqa: E402 + discard_snapshot_checkpoint, + snapshot_checkpoint, +) + +from .test_dynamic_lora_slots import ( # noqa: E402 + _adapter, + _install_checkpoint, + _single_rank_model_parallel, + _trainer_for, +) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +@pytest.mark.parametrize("train_a", (False, True)) +def test_native_capture_preserves_independent_parameter_trainability( + train_a: bool, +) -> None: + with _single_rank_model_parallel(): + device = torch.device("cuda") + lora = LoRA("dense", 4, 5, 2, 32, torch.float32, device) + trainer = _trainer_for(lora, device) + _install_checkpoint(trainer, "A", _adapter("dense", rank=2, seed=1)) + ref = LoRASlotRef("checkpoint", "A") + current = lora._slot(ref) + assert current is not None + current.A_T.requires_grad_(train_a) + current.B_T.requires_grad_(not train_a) + x = torch.ones(2, 4, device=device) + with use_lora_slot(ref): + lora(x).sum().backward() + expected = [ + None if parameter.grad is None else parameter.grad.clone() + for parameter in (current.A_T, current.B_T) + ] + trainer.zero_grad() + capture = trainer._capture_lora_version(ref) + assert capture is not None + captured = capture.slots[id(lora)] + assert (captured.A_T.requires_grad, captured.B_T.requires_grad) == ( + train_a, + not train_a, + ) + with use_lora_slot(ref, version=capture): + output = lora(x) + with trainer._gradient_transaction(): + output.sum().backward() + for parameter, gradient in zip( + (current.A_T, current.B_T), expected, strict=True + ): + if gradient is None: + assert parameter.grad is None + else: + torch.testing.assert_close(parameter.grad, gradient) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +def test_native_version_storage_reuse_accounting_and_checkpoint_lifetime() -> None: + with _single_rank_model_parallel(): + device = torch.device("cuda") + lora = LoRA("dense", 4, 5, 2, 32, torch.float32, device) + trainer = _trainer_for(lora, device) + _install_checkpoint(trainer, "A", _adapter("dense", rank=2, seed=1)) + ref = LoRASlotRef("checkpoint", "A") + expected_bytes = (4 * 2 + 2 * 5) * 4 + custom = torch.nn.Parameter(torch.ones(100, device=device)) + setattr(custom, "_art_custom_checkpoint_param", True) + trainer._checkpoint_slots["A"].params += (custom,) + assert trainer._lora_version_capture_bytes(ref) == expected_bytes + custom_bytes = custom.numel() * custom.element_size() + assert trainer._lora_gradient_staging_bytes(ref) == 3 * ( + expected_bytes + custom_bytes + ) + capture = trainer._capture_lora_version(ref) + assert capture is not None and capture.nbytes == expected_bytes + assert trainer._capture_lora_version(ref) is capture + assert trainer._lora_version_capture_bytes(ref) == 0 + assert all( + old is not current + for old in capture.slots[id(lora)].parameters() + for current in lora.parameters() + ) + from torch.utils.checkpoint import checkpoint + + with use_lora_slot(ref, version=capture): + output = checkpoint( + lora, torch.ones(2, 4, device=device), use_reentrant=False + ) + reference = weakref.ref(capture) + del capture + gc.collect() + assert reference() is not None + trainer._checkpoint_slots["A"].revision += 1 + assert trainer._lora_version_capture_bytes(ref) == expected_bytes + with trainer._gradient_transaction(): + output.sum().backward() + assert ( + trainer._lora_gradient_staging_bytes(ref) + == 2 * expected_bytes + 3 * custom_bytes + ) + del output + gc.collect() + assert reference() is None + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +def test_native_capture_rejects_staleness_and_checkpoint_replacement() -> None: + with _single_rank_model_parallel(): + device = torch.device("cuda") + lora = LoRA("dense", 4, 5, 2, 32, torch.float32, device) + trainer = _trainer_for(lora, device) + _install_checkpoint(trainer, "A", _adapter("dense", rank=2, seed=1)) + ref = LoRASlotRef("checkpoint", "A") + capture = trainer._capture_lora_version(ref, max_gradient_staleness=0) + with use_lora_slot(ref, version=capture): + output = lora(torch.ones(2, 4, device=device)) + trainer._checkpoint_slots["A"].revision += 1 + with pytest.raises(TrainerRankSlotStateError, match="staleness 1"): + with trainer._gradient_transaction(): + output.sum().backward(retain_graph=True) + assert all(p.grad is None for p in trainer._checkpoint_slots["A"].params) + trainer._checkpoint_slots["A"].generation += 1 + trainer._checkpoint_slots["A"].revision = 0 + with pytest.raises(TrainerRankSlotStateError, match="replaced"): + with trainer._gradient_transaction(): + output.sum().backward() + assert all(p.grad is None for p in trainer._checkpoint_slots["A"].params) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +def test_newer_weight_replay_keeps_original_gradient_age() -> None: + with _single_rank_model_parallel(): + device = torch.device("cuda") + lora = LoRA("dense", 4, 5, 2, 32, torch.float32, device) + trainer = _trainer_for(lora, device) + _install_checkpoint(trainer, "A", _adapter("dense", rank=2, seed=1)) + ref = LoRASlotRef("checkpoint", "A") + origin = trainer._capture_checkpoint_version("A") + trainer._checkpoint_slots["A"].revision = origin.revision + 2 + with torch.no_grad(): + for parameter in trainer._checkpoint_slots["A"].params: + parameter.add_(1) + capture = trainer._capture_lora_version(ref, origin=origin) + assert capture is not None + assert capture.version == origin + assert capture.weight_version.revision == origin.revision + 2 + for old, current in zip( + capture.slots[id(lora)].parameters(), + trainer._checkpoint_slots["A"].params, + strict=True, + ): + torch.testing.assert_close(old, current) + with use_lora_slot(ref, version=capture): + output = lora(torch.ones(2, 4, device=device)) + trainer._checkpoint_slots["A"].revision = origin.revision + 3 + with pytest.raises(TrainerRankSlotStateError, match="staleness 3"): + with trainer._gradient_transaction(): + output.sum().backward() + assert all(p.grad is None for p in trainer._checkpoint_slots["A"].params) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +def test_discarded_snapshot_name_cannot_reuse_old_capture() -> None: + with _single_rank_model_parallel(): + device = torch.device("cuda") + lora = LoRA("dense", 4, 5, 2, 32, torch.float32, device) + trainer = _trainer_for(lora, device) + _install_checkpoint(trainer, "A", _adapter("dense", rank=2, seed=1)) + trainer._checkpoint_slots["A"].config = { + "base_model_name_or_path": "test/model", + "r": 2, + "lora_alpha": 32.0, + "target_modules": ["dense"], + } + snapshot_checkpoint(trainer, "A", "saved") + ref = LoRASlotRef("checkpoint", "saved") + old = trainer._capture_lora_version(ref) + assert old is not None + with use_lora_slot(ref, version=old): + old_output = lora(torch.ones(2, 4, device=device)) + discard_snapshot_checkpoint(trainer, "saved") + with torch.no_grad(): + for parameter in trainer._checkpoint_slots["A"].params: + parameter.add_(2) + snapshot_checkpoint(trainer, "A", "saved") + assert trainer._lora_version_capture_bytes(ref) == old.nbytes + new = trainer._capture_lora_version(ref) + assert new is not None and new is not old + assert new.version.generation > old.version.generation + with pytest.raises(TrainerRankSlotStateError, match="replaced"): + old.validate() + with use_lora_slot(ref, version=new): + new_output = lora(torch.ones(2, 4, device=device)) + assert not torch.equal(old_output, new_output) diff --git a/tests/integration/megatron/lora/test_trainer_v1_graph_cache.py b/tests/integration/megatron/lora/test_trainer_v1_graph_cache.py new file mode 100644 index 000000000..92d2975cb --- /dev/null +++ b/tests/integration/megatron/lora/test_trainer_v1_graph_cache.py @@ -0,0 +1,219 @@ +"""Native LoRA/group-executor cache oracle; distributed/full-model gates are separate.""" + +import gc +import os + +import pytest + +torch = pytest.importorskip("torch") +pytest.importorskip("megatron.core") + +from art.megatron.lora import LoRA, LoRASlotRef, use_lora_slot # noqa: E402 +from art.megatron.prefix_tree_packing import prefix_tree_pack # noqa: E402 +from art.trainer_rank import ( # noqa: E402 + AdamParams, + ForwardInput, + ForwardOptions, + ForwardOutput, +) +from art.trainer_rank._impl import _ForwardGroupPlan # noqa: E402 + +from .test_dynamic_lora_slots import ( # noqa: E402 + _adapter, + _install_checkpoint, + _single_rank_model_parallel, + _trainer_for, +) + +pytestmark = pytest.mark.skipif( + os.environ.get("ART_GRAPH_GPU_TEST") != "1", reason="requires a reserved GPU" +) + + +@pytest.mark.parametrize("retention", ["gpu", "cpu", "replay"]) +@pytest.mark.parametrize("output_device", ["model", "cpu"]) +def test_group_cache_routes_old_gradients_after_optimizer_update( + retention, output_device, monkeypatch +): + with _single_rank_model_parallel(): + torch.manual_seed(42) + device = torch.device("cuda") + lora = LoRA("dense", 4, 5, 2, 32, torch.float32, device) + trainer = _trainer_for(lora, device) + # Megatron DDP uses `buffers` for a list of gradient-buffer objects. + wrapper = torch.nn.Module() + wrapper.add_module("module", lora) + setattr(wrapper, "buffers", []) + trainer.runtime.model = [wrapper] + _install_checkpoint(trainer, "A", _adapter("dense", rank=2, seed=91)) + initial_revision = trainer._checkpoint_slots["A"].revision + ref = LoRASlotRef("checkpoint", "A") + originals = tuple(trainer._checkpoint_slots["A"].params) + references = tuple( + value.detach().clone().requires_grad_() for value in originals + ) + monkeypatch.setattr(trainer, "_topology", lambda: (1, 1, 1, 1)) + monkeypatch.setattr( + trainer, + "_prepare_packed_forward", + lambda packed: packed.tokens.to(device).float().reshape(-1, 4) / 13, + ) + + def physical(items, inputs): + from torch.utils.checkpoint import checkpoint + + values = checkpoint( + lambda x: lora(torch.nn.functional.dropout(x, 0.25)), + inputs, + use_reentrant=False, + ) + return [ForwardOutput(None, None, None, values)] + + monkeypatch.setattr(trainer, "_forward_packed", physical) + outputs, expected_outputs = [], [] + for offset in (0, 3): + tokens = torch.arange(12) + offset + request = ForwardInput( + input_tokens=tokens, + hidden_states=True, + options=ForwardOptions( + backward_state=retention, output_device=output_device + ), + ) + group = _ForwardGroupPlan( + ref, + True, + (0,), + (trainer._forward_item(request),), + prefix_tree_pack([tokens], max_depth=0), + ) + rng = torch.cuda.get_rng_state() + output = trainer._execute_graph_group(group)[0].hidden_states + assert output is not None + after = torch.cuda.get_rng_state() + torch.cuda.set_rng_state(rng) + x = tokens.to(device).float().reshape(-1, 4) / 13 + expected = ( + (torch.nn.functional.dropout(x, 0.25) @ references[0]) @ references[1] + ) * 16 + torch.cuda.set_rng_state(after) + torch.testing.assert_close( + output.to(device), expected, atol=3e-5, rtol=3e-5 + ) + outputs.append(output) + expected_outputs.append(expected) + tokens.fill_(99) + loss = (outputs[0] * outputs[1].tanh()).mean() + expected_loss = (expected_outputs[0] * expected_outputs[1].tanh()).mean() + expected_gradients = torch.autograd.grad(expected_loss, references) + with use_lora_slot(ref, version=trainer._capture_lora_version(ref)): + update = lora(torch.ones(3, 4, device=device)).square().mean() + with trainer._gradient_transaction(): + update.backward() + trainer.optim_step( + params=AdamParams(learning_rate=0.02, grad_clip_norm=0), checkpoints=["A"] + ) + assert trainer._checkpoint_slots["A"].revision == initial_revision + 1 + rng = torch.cuda.get_rng_state() + with trainer._gradient_transaction(): + packets = trainer._forward_cotangent_collector().backward(loss) + trainer._forward_graph_cache().backward_many( + [(packet.handle, packet.gradients) for packet in packets] + ) + for actual, expected in zip(originals, expected_gradients, strict=True): + torch.testing.assert_close(actual.grad, expected, atol=2e-4, rtol=3e-5) + assert torch.equal(rng, torch.cuda.get_rng_state()) + assert trainer._forward_graph_cache().handles() == () + + # Abandoned caller graphs also release cache records, including versions. + unused = trainer._execute_graph_group(group)[0].hidden_states + assert trainer._forward_graph_cache().handles() + del unused + gc.collect() + assert trainer._forward_graph_cache().handles() == () + + +@pytest.mark.parametrize("mode", ["always", "current_replay"]) +def test_native_stale_logprob_correction_keeps_original_gradient_age(mode, monkeypatch): + from art.trainer_rank import ( + ImportanceSamplingGradientCorrection, + TrainerRankSlotStateError, + ) + + with _single_rank_model_parallel(): + torch.manual_seed(19) + device = torch.device("cuda") + lora = LoRA("dense", 4, 5, 2, 32, torch.float32, device) + trainer = _trainer_for(lora, device) + _install_checkpoint(trainer, "A", _adapter("dense", rank=2, seed=91)) + ref = LoRASlotRef("checkpoint", "A") + origin = trainer._capture_checkpoint_version("A") + parameters = trainer._checkpoint_slots["A"].params + historical = [value.detach().clone().requires_grad_() for value in parameters] + tokens = torch.arange(12) + x = tokens.to(device).float().reshape(-1, 4) / 13 + monkeypatch.setattr(trainer, "_topology", lambda: (1, 1, 1, 1)) + monkeypatch.setattr( + trainer, + "_prepare_packed_forward", + lambda packed: packed.tokens.to(device).float().reshape(-1, 4) / 13, + ) + executions = [] + + def physical(items, inputs): + executions.append(torch.is_grad_enabled()) + return [ForwardOutput(lora(inputs).log_softmax(-1)[:, 0], None, None, None)] + + monkeypatch.setattr(trainer, "_forward_packed", physical) + request = ForwardInput( + input_tokens=tokens, + target_tokens=torch.zeros_like(tokens), + options=ForwardOptions( + stale_gradient_corrections=( + ImportanceSamplingGradientCorrection( + policy="always" if mode == "always" else "when_available" + ), + ) + ), + ) + group = _ForwardGroupPlan( + ref, + True, + (0,), + (trainer._forward_item(request),), + prefix_tree_pack([tokens], max_depth=0), + ) + output = trainer._execute_graph_group(group)[0].target_logprobs + assert output is not None + old_logprobs = (((x @ historical[0]) @ historical[1]) * 16).log_softmax(-1)[ + :, 0 + ] + with use_lora_slot(ref, version=trainer._capture_lora_version(ref)): + update = lora(torch.ones_like(x)).square().mean() + with trainer._gradient_transaction(): + update.backward() + trainer.optim_step( + params=AdamParams(learning_rate=0.001, grad_clip_norm=0), checkpoints=["A"] + ) + current = [value.detach().clone().requires_grad_() for value in parameters] + new_logprobs = (((x @ current[0]) @ current[1]) * 16).log_softmax(-1)[:, 0] + weights = (new_logprobs.detach() - old_logprobs.detach()).exp().clamp(0, 5) + expected = torch.autograd.grad( + ((old_logprobs if mode == "always" else new_logprobs) * weights).sum(), + historical if mode == "always" else current, + ) + cache = trainer._forward_graph_cache() + if mode == "current_replay": + cache.evict(cache.handles()[0], replay_with_current=True) + with trainer._gradient_transaction(): + packets = trainer._forward_cotangent_collector().backward(output.sum()) + cache.backward_many( + [(packet.handle, packet.gradients) for packet in packets] + ) + for parameter, gradient in zip(parameters, expected, strict=True): + torch.testing.assert_close(parameter.grad, gradient, atol=2e-4, rtol=5e-5) + assert executions == [True, mode == "current_replay"] + assert trainer._version_state()._origins["A"] == {(origin, 2)} + trainer._checkpoint_slots["A"].revision += 2 + with pytest.raises(TrainerRankSlotStateError, match="staleness 3"): + trainer._version_state().validate_accumulated(["A"]) diff --git a/tests/integration/megatron/lora/test_trainer_v1_versions.py b/tests/integration/megatron/lora/test_trainer_v1_versions.py new file mode 100644 index 000000000..d80d82664 --- /dev/null +++ b/tests/integration/megatron/lora/test_trainer_v1_versions.py @@ -0,0 +1,171 @@ +"""Independent CUDA matrix/Adam oracle for historical native LoRA graphs. + +This deliberately exercises the production LoRA and TrainerRank optimizer but +uses explicit float64 matrix products and Adam equations for expected values. +Full-model and distributed acceptance are additional gates, not implied here. +""" + +from __future__ import annotations + +import json + +import pytest + +torch = pytest.importorskip("torch") +pytest.importorskip("megatron.core") + +from art.megatron.lora import LoRA, LoRASlotRef, use_lora_slot # noqa: E402 +from art.trainer_rank import AdamParams # noqa: E402 + +from .test_dynamic_lora_slots import ( # noqa: E402 + _adapter, + _install_checkpoint, + _single_rank_model_parallel, + _trainer_for, +) + + +def _coupled_loss(first, second): + return (first * second.tanh()).mean() + 0.13 * first.square().mean() + + +def _adam(parameters, gradients, moments, step, params): + """Explicit AdamW, independent of the runtime's optimizer implementation.""" + result = [] + for parameter, gradient, (first, second) in zip( + parameters, gradients, moments, strict=True + ): + first.mul_(params.beta1).add_(gradient, alpha=1 - params.beta1) + second.mul_(params.beta2).addcmul_(gradient, gradient, value=1 - params.beta2) + numerator = first / (1 - params.beta1**step) + denominator = (second / (1 - params.beta2**step)).sqrt() + 1e-8 + result.append( + parameter * (1 - params.learning_rate * params.weight_decay) + - params.learning_rate * numerator / denominator + ) + return result + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +@pytest.mark.parametrize("recompute", ["none", "torch", "reentrant", "megatron"]) +def test_native_old_lora_graph_matches_matrix_and_adam_oracle(recompute, artifact_dir): + with _single_rank_model_parallel(): + torch.manual_seed(1709) + device = torch.device("cuda") + lora = LoRA("dense", 4, 5, 2, 32, torch.float32, device) + trainer = _trainer_for(lora, device) + _install_checkpoint(trainer, "A", _adapter("dense", rank=2, seed=91)) + ref = LoRASlotRef("checkpoint", "A") + originals = tuple(trainer._checkpoint_slots["A"].params) + reference = [p.detach().double().requires_grad_() for p in originals] + scale = 16.0 + params = AdamParams( + learning_rate=0.05, + beta1=0.2, + beta2=0.5, + weight_decay=0.01, + grad_clip_norm=0.0, + ) + moments = [(torch.zeros_like(p), torch.zeros_like(p)) for p in reference] + inputs = [ + ( + torch.arange(12, device=device).reshape(3, 4) / 13 + offset + ).requires_grad_() + for offset in (0.1, -0.3) + ] + reference_inputs = [x.detach().double().requires_grad_() for x in inputs] + + def native(x): + return lora(x * torch.nn.functional.dropout(torch.ones_like(x), 0.25)) + + def explicit(x): + mask = torch.nn.functional.dropout( + torch.ones_like(x, dtype=torch.float32), 0.25 + ) + return ((x * mask) @ reference[0]) @ reference[1] * scale + + def checkpoint(x): + if recompute == "none": + return native(x) + if recompute == "megatron": + from megatron.core.tensor_parallel.random import checkpoint + + return checkpoint(native, False, x) + from torch.utils.checkpoint import checkpoint + + return checkpoint(native, x, use_reentrant=recompute == "reentrant") + + capture = trainer._capture_lora_version(ref, max_gradient_staleness=2) + rng = torch.cuda.get_rng_state() + with use_lora_slot(ref, version=capture): + outputs = [checkpoint(x) for x in inputs] + unused = checkpoint(inputs[0] + 0.7) + after_forward = torch.cuda.get_rng_state() + torch.cuda.set_rng_state(rng) + reference_outputs = [explicit(x) for x in reference_inputs] + torch.cuda.set_rng_state(after_forward) + for actual, expected in zip(outputs, reference_outputs, strict=True): + torch.testing.assert_close(actual.double(), expected, atol=3e-5, rtol=2e-5) + old_loss = _coupled_loss(*outputs) + old_reference_loss = _coupled_loss(*reference_outputs) + + # A separate, completed step mutates current optimizer weights while + # both coupled old forwards and an unused output remain alive. + update_input = torch.linspace(-0.3, 0.8, 20, device=device).reshape(5, 4) + with use_lora_slot( + ref, version=trainer._capture_lora_version(ref, max_gradient_staleness=2) + ): + update_loss = lora(update_input).square().mean() / scale**2 + update_reference = ( + ((update_input.double() @ reference[0]) @ reference[1]).square().mean() + ) + update_grads = torch.autograd.grad(update_reference, reference) + with trainer._gradient_transaction(): + update_loss.backward() + trainer.optim_step(params=params, checkpoints=["A"]) + expected_current = _adam(reference, update_grads, moments, 1, params) + current = tuple(trainer._checkpoint_slots["A"].params) + for actual, expected in zip(current, expected_current, strict=True): + torch.testing.assert_close(actual.double(), expected, atol=2e-6, rtol=2e-5) + assert any(not torch.equal(a, b) for a, b in zip(current, reference)) + + expected_gradients = torch.autograd.grad( + old_reference_loss, [*reference, *reference_inputs] + ) + torch.rand(17, device=device) + before_backward = torch.cuda.get_rng_state() + with trainer._gradient_transaction(): + old_loss.backward() + assert torch.equal(before_backward, torch.cuda.get_rng_state()) + assert unused.grad_fn is not None + errors = [] + for actual, expected in zip( + [p.grad for p in current] + [x.grad for x in inputs], + expected_gradients, + strict=True, + ): + assert actual is not None + torch.testing.assert_close(actual.double(), expected, atol=2e-4, rtol=3e-5) + errors.append(float((actual.double() - expected).abs().max())) + trainer.optim_step(params=params, checkpoints=["A"]) + expected_final = _adam( + expected_current, expected_gradients[:2], moments, 2, params + ) + for actual, expected in zip( + trainer._checkpoint_slots["A"].params, expected_final, strict=True + ): + torch.testing.assert_close(actual.double(), expected, atol=2e-6, rtol=2e-5) + (artifact_dir / "oracle.json").write_text( + json.dumps( + { + "recompute": recompute, + "device": torch.cuda.get_device_name(), + "torch": torch.__version__, + "max_abs_gradient_errors": errors, + "original_version_age": 1, + "optimizer_steps": 2, + }, + indent=2, + ) + + "\n" + ) diff --git a/tests/integration/megatron/model_support/test_compile_flags.py b/tests/integration/megatron/model_support/test_compile_flags.py index 29353e225..55d09d3af 100644 --- a/tests/integration/megatron/model_support/test_compile_flags.py +++ b/tests/integration/megatron/model_support/test_compile_flags.py @@ -4,6 +4,7 @@ import pytest import torch from torch._dynamo.testing import CompileCounter +from torch._functorch import config as functorch_config from art.megatron.flex_attn.compiled import _needs_blackwell_wide_head_tile from art.megatron.model_support.handlers.gemma4 import ( @@ -35,11 +36,23 @@ def test_dynamic_projection_parameters_reuse_compiled_graph() -> None: torch._dynamo.reset() counter = CompileCounter() try: - with torch._dynamo.config.patch( - force_parameter_static_shapes=True, recompile_limit=32 + with ( + torch._dynamo.config.patch( + force_parameter_static_shapes=True, recompile_limit=32 + ), + functorch_config.patch(donated_buffer=True), + cast(Any, torch.compiler.config).patch(cache_key_tag="existing-tag"), ): _configure_dynamo() assert not torch._dynamo.config.force_parameter_static_shapes + assert not functorch_config.donated_buffer + assert torch.compiler.config.cache_key_tag == ( + "existing-tag|art-retained-backward-v1" + ) + _configure_dynamo() + assert torch.compiler.config.cache_key_tag == ( + "existing-tag|art-retained-backward-v1" + ) compiled = [ torch.compile(_DynamicProjection(width), backend=counter) for width in (8, 4, 16, 32, 12, 20, 24, 28, 36, 40) @@ -78,9 +91,17 @@ def test_disabled_training_compile_does_not_change_dynamo_policy( lambda: pytest.fail("disabled compilation must not mutate Dynamo config"), ) - assert not compile_module.configure_training_compile( - model=[], provider=object(), provider_bundle=cast(Any, bundle) - ) + with ( + functorch_config.patch(donated_buffer=True), + cast(Any, torch.compiler.config).patch(cache_key_tag="existing-tag"), + ): + assert not compile_module.configure_training_compile( + model=[], provider=object(), provider_bundle=cast(Any, bundle) + ) + assert not functorch_config.donated_buffer + assert torch.compiler.config.cache_key_tag == ( + "existing-tag|art-retained-backward-v1" + ) def test_wide_head_tile_workaround_is_blackwell_only(monkeypatch) -> None: diff --git a/tests/unit/test_trainer_batch_input_capture.py b/tests/unit/test_trainer_batch_input_capture.py new file mode 100644 index 000000000..e1353eb47 --- /dev/null +++ b/tests/unit/test_trainer_batch_input_capture.py @@ -0,0 +1,168 @@ +"""Batch iterators own submitted inputs while executing one wave at a time.""" + +from __future__ import annotations + +import asyncio +from typing import Any, cast + +import pytest +from test_trainer_rank_commands import _Rank +import torch + +from art.trainer_rank import ( + ForwardInput, + ForwardOptions, + MicroBatch, + MicroBatchStats, + TrainerRank, + Unset, + run_rank_callback, +) + + +class _CapturingRank(_Rank): + _capture_forward_options = TrainerRank._capture_forward_options + + +@pytest.mark.parametrize("surface", ["native", "rank", "zero", "persistent"]) +@pytest.mark.parametrize("policy", [None, ForwardOptions(max_gradient_staleness=1)]) +def test_batches_snapshot_tokens_targets_and_structure_before_first_pull( + surface, policy +): + calls = [] + rank: Any + if surface == "native": + rank = object.__new__(TrainerRank) + rank._skipped_forward_waves = {} + + def batches(inputs, **kwargs): + for index, item in enumerate(inputs): + calls.append(index) + yield MicroBatch( + [item], + [], + [index], + MicroBatchStats(index, index + 1, 1, 1, 0, 0, 0, 0, 0, False), + ) + + rank._forward_batches = batches + else: + rank = cast(Any, _CapturingRank()) + forward = rank.forward + + def record(inputs, **kwargs): + if isinstance(inputs, ForwardInput): + calls.append(inputs.input_tokens.tolist()) + return forward(inputs, **kwargs) + + rank.forward = record + rank._forward_options = policy + tokens = torch.tensor([1, 2, 3], dtype=torch.int32) + targets = torch.arange(12, dtype=torch.int64).reshape(3, 4)[:, ::2] + request = ForwardInput(input_tokens=tokens, target_tokens=targets, checkpoint=None) + later = ForwardInput(input_tokens=torch.tensor([4, 5]), checkpoint=Unset) + roots = [(request,), [later]] + expected_tokens, expected_targets = tokens.clone(), targets.clone() + + def mutate_sources(): + tokens.fill_(9) + targets.fill_(-100) + request.input_tokens = torch.tensor([11]) + request.target_tokens = None + later.input_tokens.fill_(8) + roots.clear() + + def check_first(batch): + assert batch.indices == [0] + (captured,) = batch.inputs[0] + assert captured is not request + assert captured.checkpoint is None + torch.testing.assert_close(captured.input_tokens, expected_tokens) + torch.testing.assert_close(captured.target_tokens, expected_targets) + assert captured.input_tokens.untyped_storage() is not tokens.untyped_storage() + assert captured.target_tokens.untyped_storage() is not targets.untyped_storage() + + def check_second(batch): + assert batch.indices == [1] + (captured,) = batch.inputs[0] + assert captured.checkpoint is Unset + torch.testing.assert_close(captured.input_tokens, torch.tensor([4, 5])) + + if surface == "native": + iterator = rank.forward_batches(roots, yield_empty=True) + assert calls == [] + mutate_sources() + check_first(next(iterator)) + later.input_tokens.fill_(99) + check_second(next(iterator)) + assert list(iterator) == [] + return + + if surface == "persistent": + + def run(callback): + return asyncio.run(run_rank_callback(rank, callback, mode="zero")).value + + handle = run(lambda view: view.open_forward_batches(roots)) + assert calls == [] + mutate_sources() + check_first(run(lambda view: view.next_forward_batch(handle))) + later.input_tokens.fill_(99) + check_second(run(lambda view: view.next_forward_batch(handle))) + assert run(lambda view: view.next_forward_batch(handle)) is None + else: + + def callback(view): + iterator = view.forward_batches(roots) + assert calls == [] + mutate_sources() + check_first(next(iterator)) + later.input_tokens.fill_(99) + check_second(next(iterator)) + assert list(iterator) == [] + + asyncio.run(run_rank_callback(rank, callback, mode=surface)) + assert calls == [[1, 2, 3], [4, 5]] + + +@pytest.mark.parametrize("logical", [False, True]) +def test_generator_structure_is_consumed_once_at_submission_without_executing_a_wave( + logical, +): + enumerated, executed = [], [] + rank: Any = _CapturingRank() if logical else object.__new__(TrainerRank) + rank._forward_options = None + rank._skipped_forward_waves = {} + + def batches(inputs, **kwargs): + executed.append(True) + yield MicroBatch( + inputs, [], [0], MicroBatchStats(0, 1, 1, 1, 0, 0, 0, 0, 0, False) + ) + + if logical: + rank.forward_batches = batches + else: + rank._forward_batches = batches + + def inputs(): + enumerated.append("outer") + + def nested(): + enumerated.append("inner") + yield ForwardInput(input_tokens=torch.tensor([1, 2, 3])) + + yield nested() + + def check(view): + iterator = view.forward_batches(inputs()) + assert enumerated == ["outer", "inner"] + assert executed == [] + iterator.close() + assert enumerated == ["outer", "inner"] + assert executed == [] + + if logical: + asyncio.run(run_rank_callback(rank, check, mode="zero")) + else: + check(rank) diff --git a/tests/unit/test_trainer_command_transport.py b/tests/unit/test_trainer_command_transport.py new file mode 100644 index 000000000..56ed7f1be --- /dev/null +++ b/tests/unit/test_trainer_command_transport.py @@ -0,0 +1,238 @@ +"""Commands cross ranks through CPU storage, independent of CUDA ordinals.""" + +from __future__ import annotations + +from dataclasses import dataclass +from datetime import timedelta +import gc +import sys +from types import ModuleType, SimpleNamespace +from typing import Any, Literal, cast +import weakref + +import pytest +import torch +import torch.distributed as dist +import torch.multiprocessing as mp + +from art.trainer_rank import ForwardInput, ForwardOutput, TrainerRank +from art.trainer_rank._commands import _Command, _encode_command, _Executor +from art.trainer_rank._heads import HeadRegistration +from art.trainer_rank._impl import _CheckpointSlot +from art.trainer_rank._tensors import managed_tensor + + +def _payload(device: str) -> Any: + # Local types and closures exercise CloudPickler inside Torch's persistent + # storage protocol, including tensors that a container walk cannot find. + @dataclass + class Payload: + base: torch.Tensor + view: torch.Tensor + module: Any + parameter: torch.nn.Parameter + subclass: torch.Tensor + managed: torch.Tensor + captured: Any + + class TaggedTensor(torch.Tensor): + pass + + base = torch.arange(6, device=device, dtype=torch.float64, requires_grad=True) + captured = base[1::2] + + class Head(torch.nn.Module): + offset: torch.Tensor + + def __init__(self) -> None: + super().__init__() + self.weight = torch.nn.Parameter(torch.ones(3, device=device)) + self.tied = self.weight + self.register_buffer("offset", base.detach()[::2]) + self.register_buffer("alias", self.offset) + + def forward(self, value: torch.Tensor) -> torch.Tensor: + return value * self.weight + self.offset + + module = Head() + + def hook(_module: Any, _args: Any, output: torch.Tensor) -> torch.Tensor: + # Module.to moves registered state, as in ordinary eager PyTorch. Hooks + # that capture constants must explicitly follow the output placement. + return output + captured.to(output) + + module.register_forward_hook(hook) + return Payload( + base, + captured, + module, + module.weight, + base.detach().as_subclass(TaggedTensor), + managed_tensor(base.detach()), + lambda: captured, + ) + + +def _check_payload(value: Any) -> None: + assert value.base.device.type == value.view.device.type == "cpu" + assert value.base.requires_grad and value.view.requires_grad + assert value.base.dtype == torch.float64 + assert value.view.shape == (3,) and value.view.stride() == (2,) + assert value.view.storage_offset() == 1 + assert value.view.untyped_storage() is value.base.untyped_storage() + assert value.captured() is value.view + assert value.module.weight is value.module.tied is value.parameter + assert value.parameter.device.type == "cpu" and value.parameter.requires_grad + assert value.module.offset is value.module.alias + assert not value.module.offset.requires_grad + assert value.module.offset.untyped_storage() is value.base.untyped_storage() + assert type(value.subclass).__name__ == "TaggedTensor" + assert value.subclass.device.type == value.managed.device.type == "cpu" + assert not value.subclass.requires_grad and not value.managed.requires_grad + assert value.subclass.untyped_storage() is value.base.untyped_storage() + assert value.managed.untyped_storage() is value.base.untyped_storage() + torch.testing.assert_close( + value.module(torch.ones(3)), torch.tensor([2.0, 6.0, 10.0], dtype=torch.float64) + ) + + +def test_command_codec_preserves_nested_types_aliases_and_closure_storage() -> None: + from test_trainer_rank_commands import _Rank + + executor = _Executor(cast(Any, _Rank()), "zero") + source = _payload("cpu") + result = executor._decode(_encode_command(_Command(1, "test", (source,), {}, True))) + assert result.operation == "test" and result.grad_enabled + _check_payload(result.args[0]) + assert result.args[0].base.untyped_storage() is not source.base.untyped_storage() + + +_decoded_refs: list[weakref.ReferenceType[torch.Tensor]] = [] + + +def _remember_restore(tensor: torch.Tensor) -> None: + _decoded_refs.append(weakref.ref(tensor)) + + +class _Remember: + def __init__(self, tensor: torch.Tensor) -> None: + self.tensor = tensor + + def __reduce__(self): + return _remember_restore, (self.tensor,) + + +def _transport_worker(physical: int, rendezvous: str, cuda: bool) -> None: + from test_trainer_rank_commands import _BadRestore + from test_trainer_rank_custom_tensors import _config, _runtime + + # Deliberately reverse physical rank and CUDA ordinal, as real actors can. + device = torch.device(f"cuda:{1 - physical}" if cuda else "cpu") + if cuda: + torch.cuda.set_device(device) + dist.init_process_group( + "gloo", + init_method=f"file://{rendezvous}", + rank=physical, + world_size=2, + timeout=timedelta(seconds=30), + ) + try: + ps = SimpleNamespace( + get_tensor_model_parallel_rank=lambda: physical, + get_context_parallel_rank=lambda: 0, + ) + core, megatron = ModuleType("megatron.core"), ModuleType("megatron") + setattr(core, "parallel_state", ps) + setattr(megatron, "core", core) + sys.modules.update({"megatron": megatron, "megatron.core": core}) + runtime = _runtime(torch.nn.Linear(1, 1).to(device)) + runtime.rank, runtime.world_size = physical, 2 + rank: Any = TrainerRank(runtime) + rank._dp_rank_and_size = lambda: (0, 1) + rank._checkpoint_slots["student"] = _CheckpointSlot(config=cast(Any, _config())) + executor = _Executor(rank, "zero") + source = _payload(str(device)) if physical == 0 else None + foreign_before = torch.cuda.memory_allocated(physical) if cuda else 0 + + def broadcast(sequence: int, operation: str, *args: Any) -> _Command: + command = _Command(sequence, operation, args, {}, True) + return executor._broadcast(command if physical == 0 else None) + + result = broadcast(1, "test", source) + assert result.operation == "test", result.args + _check_payload(result.args[0]) + if cuda: + assert torch.cuda.memory_allocated(physical) == foreign_before + + # Real registration handlers put CPU-decoded modules and Parameters on + # each native rank's own model device, preserving tied state and hooks. + kinds: tuple[Literal["module", "parameter"], ...] = ("module", "parameter") + for index, kind in enumerate(kinds, start=2): + value = None if source is None else getattr(source, kind) + command = broadcast( + index, + "head", + "head_register", + HeadRegistration("student", kind, kind, cast(Any, value)), + ) + assert command.operation == "head", command.args + executor._dispatch(command) + registered = rank._checkpoint_slots["student"].custom[kind].value + if kind == "module": + assert registered.weight.device == registered.offset.device == device + assert registered.weight is registered.tied + assert registered.offset is registered.alias + torch.testing.assert_close( + registered(torch.ones(3, device=device)), + torch.tensor([2.0, 6.0, 10.0], device=device, dtype=torch.float64), + ) + else: + assert registered.device == device and registered.requires_grad + + # Exercise the public ForwardInput command shape and an actual forward + # dispatch; the lightweight kernel follows native per-rank placement. + def forward(inputs: ForwardInput) -> ForwardOutput: + assert inputs.input_tokens.device.type == "cpu" + values = inputs.input_tokens.to(device=device, dtype=torch.float32) + return ForwardOutput(None, None, None, values * runtime.model[0].weight) + + rank.forward = forward + inputs = ForwardInput(input_tokens=torch.tensor([1, 2, 3], device=device)) + command = broadcast(4, "forward", inputs) + executor._dispatch(command) + assert executor.state.graphs + assert all( + t.device == device for ts in executor.state.graphs.values() for t in ts + ) + + # A peer-local restore exception must be coordinated before dispatch, + # release partially decoded tensors, and leave the next command usable. + before = set(executor.state.graphs) + bad = broadcast( + 5, "forward", _Remember(torch.ones(5, device=device)), _BadRestore() + ) + assert bad.operation == "error" + assert "peer deserialization failure" in bad.args[0] + with pytest.raises(RuntimeError, match="deserialization failed"): + executor._execute(bad) + gc.collect() + assert _decoded_refs and all(ref() is None for ref in _decoded_refs) + assert set(executor.state.graphs) == before + assert ( + broadcast(6, "test", torch.ones(2, device=device)).args[0].device.type + == "cpu" + ) + if cuda: + assert torch.cuda.memory_allocated(physical) == foreign_before + finally: + dist.destroy_process_group() + + +@pytest.mark.parametrize("cuda", [False, True], ids=["cpu", "reversed-cuda-indices"]) +def test_commands_decode_on_cpu_and_native_handlers_place_locally( + tmp_path, cuda: bool +) -> None: + if cuda and torch.cuda.device_count() < 2: + pytest.skip("requires two CUDA devices") + mp.spawn(_transport_worker, args=(str(tmp_path / "init"), cuda), nprocs=2) diff --git a/tests/unit/test_trainer_driver_transport.py b/tests/unit/test_trainer_driver_transport.py new file mode 100644 index 000000000..2b9938af1 --- /dev/null +++ b/tests/unit/test_trainer_driver_transport.py @@ -0,0 +1,255 @@ +"""Driver CPU transport stays separate from native output placement.""" + +from __future__ import annotations + +import asyncio +from dataclasses import replace +import gc +from typing import Any +import weakref + +import pytest +from test_trainer_rank_commands import _input, _Rank +import torch + +from art.trainer_rank import ForwardInput, ForwardOptions, ForwardOutput +from art.trainer_rank._commands import _Executor, _view +from art.trainer_rank._operations import TrainerOperation, execute_operation +from art.trainer_rank._options import resolve_forward_options +from art.trainer_rank._tensors import CotangentCollector, flatten_tensors + + +class _TransportRank(_Rank): + def __init__(self, device="cpu"): + super().__init__() + self.device = torch.device(device) + self.weight = torch.nn.Parameter(self.weight.detach().to(device)) + self.device = self.weight.device + self.policies = [] + self.worker_reject = False + + def forward(self, tree, **kwargs): + if not isinstance(tree, ForwardInput): + return [self.forward(child, **kwargs) for child in tree] + policy = resolve_forward_options( + method=kwargs.get("options"), input=tree.options + ) + self.policies.append(policy) + if self.worker_reject: + raise MemoryError("worker admission rejected") + with torch.set_grad_enabled(not kwargs.get("no_grad", False)): + value = ( + tree.input_tokens.to(self.device).float() * self.weight.clone().square() + ) + if policy.output_device == "cpu": + value = value.cpu() + return ForwardOutput(None, None, None, value) + + +def _operation(view, kind, payload, identity=None): + return execute_operation( + view, TrainerOperation.capture(identity or str(id(payload)), kind, payload) + ) + + +@pytest.mark.parametrize("device", ["cpu", "cuda"]) +@pytest.mark.parametrize("batches", [False, True]) +def test_driver_cpu_exports_preserve_old_gradients_and_worker_policy(device, batches): + if device == "cuda" and not torch.cuda.is_available(): + pytest.skip("requires CUDA") + + async def run(): + rank: Any = _TransportRank(device) + view = _view(_Executor(rank, "zero")) + client = CotangentCollector() + request = replace( + _input(3), + input_tokens=torch.tensor([3], device=device), + options=ForwardOptions(output_device="model"), + ) + # Physical worker succeeds, but no aggregate GPU copy can be admitted. + rank._available_memory_bytes = lambda: 0 + if device == "cuda": + with pytest.raises(MemoryError, match="Gathered model-device outputs"): + view.forward(request) + seen = [] + attach = view._attach + + def check_cpu(packet): + assert all(packet.cpu) + assert all(tensor.device.type == "cpu" for tensor in packet.packet.tensors) + seen.append(packet.packet.handle) + before = torch.cuda.memory_allocated() if device == "cuda" else 0 + result = attach(packet) + if device == "cuda": + assert torch.cuda.memory_allocated() == before + return result + + # A shallow operation view preserves this bound observer; its delegate + # inspects the actual placement before collector.attach can copy anything. + setattr(view, "_attach", check_cpu) + + async def forward(identity): + if batches: + handle = await _operation( + view, "batches_open", {"inputs": [request]}, identity + ":open" + ) + packet = await _operation( + view, "batches_next", {"handle": handle}, identity + ) + await _operation( + view, "batches_close", {"handle": handle}, identity + ":close" + ) + return client.attach(packet, managed=True).outputs[0].hidden_states + packet = await _operation(view, "forward", {"inputs": request}, identity) + return client.attach(packet, managed=True).hidden_states + + old = await forward("old") + with torch.no_grad(): + rank.weight.add_(1) + fresh = await forward("fresh") + assert view._transport_handles is None + assert old.device.type == fresh.device.type == "cpu" + assert len(seen) == 2 + assert all(policy.output_device == "model" for policy in rank.policies) + state = rank._rank_command_state + assert len(state.exports) == 2 + assert all( + tensor.device.type == "cpu" + for tensors in state.exports.values() + for tensor in tensors + ) + assert all( + tensor.device == rank.device + for tensors in state.graphs.values() + for tensor in tensors + ) + head = torch.nn.Parameter(torch.tensor(5.0, device=device)) + loss = (old * fresh * head).sum() + for retained in (True, False): + packets = client.backward(loss, retain_graph=retained) + await _operation( + view, + "backward", + {"packets": packets, "retain_graph": retained}, + "backward:" + str(retained), + ) + assert len(state.exports) == (2 if retained else 0) + assert rank.weight.grad is not None and head.grad is not None + # 3*w_old^2 * 3*w_new^2 * head, differentiated into the same parameter. + torch.testing.assert_close( + rank.weight.grad, torch.tensor(5400.0, device=device) + ) + torch.testing.assert_close(head.grad, torch.tensor(648.0, device=device)) + assert not state.graphs + setattr(view, "_attach", attach) + rank._available_memory_bytes = lambda: 1 << 60 + native = view.forward(request) + assert native.hidden_states.device == rank.device + view.backward(native.hidden_states.sum()) + assert not state.graphs + + asyncio.run(run()) + + +@pytest.mark.parametrize("policy", ["model", "cpu", "auto"]) +def test_transport_preserves_worker_admission_and_native_view(policy): + async def run(): + rank: Any = _TransportRank() + view = _view(_Executor(rank, "zero")) + rank.worker_reject = True + request = replace(_input(3), options=ForwardOptions(output_device=policy)) + with pytest.raises(MemoryError, match="worker admission rejected"): + await _operation(view, "forward", {"inputs": request}) + assert rank.policies[0].output_device == policy + assert view._transport_handles is None + assert not rank._rank_command_state.exports + assert not rank._rank_command_state.graphs + + asyncio.run(run()) + + +@pytest.mark.parametrize("failure_kind", ["budget", "allocation"]) +def test_export_failure_releases_only_failed_operation_and_counts_live_exports( + monkeypatch, failure_kind +): + async def run(): + rank: Any = _TransportRank() + view = _view(_Executor(rank, "zero")) + state = rank._rank_command_state + # Additional headroom excludes the CPU payloads held by earlier exports, + # mirroring fresh host/cgroup accounting used by the real rank. + total = 64 + observed = [] + + def available(): + used = sum( + t.untyped_storage().nbytes() + for values in state.exports.values() + for t in values + ) + observed.append(used) + return total - used + + rank._available_cpu_memory_bytes = available + request = _input(3) + good = await _operation(view, "forward", {"inputs": request}, "good") + assert ( + sum(t.numel() * t.element_size() for t in state.exports[good.handle]) == 4 + ) + original_graphs = set(state.graphs) + # The worker snapshot fits. A budget drop at export must fail before + # its clone and clean only the newly created physical graphs. + calls = 0 + + def constrained(): + nonlocal calls + calls += 1 + return available() if calls == 1 else 0 + + if failure_kind == "budget": + rank._available_cpu_memory_bytes = constrained + else: + from art.trainer_rank import _tensors + + detach = _tensors.detach_tree + + def reject_clone(handle, *args, **kwargs): + if handle.startswith("client:"): + raise MemoryError("export snapshot allocation failed") + return detach(handle, *args, **kwargs) + + monkeypatch.setattr(_tensors, "detach_tree", reject_clone) + with pytest.raises(MemoryError, match="export snapshot") as failure: + await _operation(view, "forward", {"inputs": request}, "failure") + assert failure.value.__traceback__ is not None + assert set(state.exports) == {good.handle} + assert set(state.graphs) == original_graphs + assert 4 in observed + rank._available_cpu_memory_bytes = available + await _operation(view, "release", {"handles": [good.handle]}, "release") + assert not state.exports + assert not state.graphs + assert available() == total + + asyncio.run(run()) + + +def test_transport_release_drops_cpu_payload_without_collecting_cycles(): + async def run(): + rank: Any = _TransportRank() + view = _view(_Executor(rank, "zero")) + packet = await _operation(view, "forward", {"inputs": _input(3)}, "forward") + state = rank._rank_command_state + references = [weakref.ref(t) for t in state.exports[packet.handle]] + await _operation(view, "release", {"handles": [packet.handle]}, "release") + assert all(ref() is None for ref in references) + assert not state.exports and not state.graphs + + enabled = gc.isenabled() + gc.disable() + try: + asyncio.run(run()) + finally: + if enabled: + gc.enable() diff --git a/tests/unit/test_trainer_live_parameter_roots.py b/tests/unit/test_trainer_live_parameter_roots.py new file mode 100644 index 000000000..4fb0f7ade --- /dev/null +++ b/tests/unit/test_trainer_live_parameter_roots.py @@ -0,0 +1,137 @@ +"""Direct live parameter roots use the same snapshots as parameter arithmetic.""" + +from __future__ import annotations + +from dataclasses import replace +from typing import Any + +import pytest +from test_trainer_rank_custom_tensors import _trainer +import torch + +from art.trainer_rank._heads import LiveHead, export_head, head_gradient_targets +from art.trainer_rank._tensors import CotangentCollector + + +class _FailBackward(torch.autograd.Function): + @staticmethod + def forward(ctx, value): + return value + + @staticmethod + def backward(ctx, *gradients): + raise RuntimeError("local backward failed after a direct parameter root") + + +@pytest.mark.parametrize("existing", [False, True]) +def test_native_direct_root_does_not_publish_when_another_root_fails(existing): + trainer, rank = _trainer("student") + parameter = rank.parameter( + "weight", lambda: torch.tensor(2.0), checkpoint="student" + ) + if existing: + parameter.grad = torch.tensor(7.0) + original = parameter.grad + bad = _FailBackward.apply(torch.tensor(1.0, requires_grad=True)) + with pytest.raises(RuntimeError, match="local backward failed"): + trainer.backward([bad, parameter]) + assert parameter.grad is original + if existing: + torch.testing.assert_close(parameter.grad, torch.tensor(7.0)) + # Failed local collection is reusable and a later direct root commits once. + trainer.backward(parameter) + torch.testing.assert_close(parameter.grad, torch.tensor(8.0 if existing else 1.0)) + + +@pytest.mark.parametrize("dtype", [torch.float32, torch.float64, torch.complex64]) +@pytest.mark.parametrize("surface", ["native", "client"]) +def test_direct_roots_preserve_aliases_explicit_gradients_and_dtype(dtype, surface): + trainer, rank = _trainer("student") + parameter = rank.parameter( + "weight", lambda: torch.tensor([2.0, 3.0], dtype=dtype), checkpoint="student" + ) + collector = CotangentCollector() + live = LiveHead( + export_head(trainer, "student", "weight"), parameter.detach(), collector + ) + root: Any = parameter if surface == "native" else live.value + gradients = ( + torch.tensor([1.0, 2.0], dtype=dtype), + torch.tensor([3.0, 5.0], dtype=dtype), + ) + with torch.no_grad(): + if surface == "native": + trainer.backward((root, root), gradients, retain_graph=True) + else: + packets = collector.backward((root, root), gradients, retain_graph=True) + assert len(packets) == 1 + assert len(packets[0].gradients) == 1 + torch.testing.assert_close(packets[0].gradients[0], sum(gradients)) + trainer._commit_versioned_gradients( + head_gradient_targets(trainer, packets[0]) + ) + assert parameter.grad is not None and parameter.grad.dtype == dtype + torch.testing.assert_close(parameter.grad, sum(gradients)) + # Every direct use captures the current version independently of retain_graph. + if surface == "native": + trainer.backward(root, gradients[0]) + else: + (packet,) = collector.backward(root, gradients[0]) + trainer._commit_versioned_gradients(head_gradient_targets(trainer, packet)) + assert root.grad is None + torch.testing.assert_close(parameter.grad, sum(gradients) + gradients[0]) + + +def test_direct_live_root_and_old_arithmetic_keep_their_own_versions(): + trainer, rank = _trainer("student") + parameter = rank.parameter( + "weight", lambda: torch.tensor(2.0), checkpoint="student" + ) + collector = CotangentCollector() + live = LiveHead( + export_head(trainer, "student", "weight"), parameter.detach(), collector + ) + root: Any = live.value + old = root.square() + parameter.data.fill_(5) + trainer._checkpoint_slots["student"].revision += 1 + live.refresh(export_head(trainer, "student", "weight")) + packets = collector.backward((old, root)) + assert len(packets) == 2 + targets = [ + target + for packet in packets + for target in head_gradient_targets(trainer, packet) + ] + assert {target[0].revision for target in targets} == {0, 1} + trainer._commit_versioned_gradients(targets) + torch.testing.assert_close(parameter.grad, torch.tensor(5.0)) + + +def test_client_direct_root_failure_discards_cotangents_and_revalidates_handle(): + trainer, rank = _trainer("student") + parameter = rank.parameter( + "weight", lambda: torch.tensor(2.0), checkpoint="student" + ) + collector = CotangentCollector() + live = LiveHead( + export_head(trainer, "student", "weight"), parameter.detach(), collector + ) + root: Any = live.value + bad = _FailBackward.apply(torch.tensor(1.0, requires_grad=True)) + with pytest.raises(RuntimeError, match="local backward failed"): + collector.backward([bad, root]) + assert root.grad is None and parameter.grad is None + (packet,) = collector.backward(root) + torch.testing.assert_close(packet.gradients[0], torch.tensor(1.0)) + live.refresh(replace(live.state, version=replace(live.state.version, generation=1))) + live.invalidate("checkpoint replaced") + with pytest.raises(RuntimeError, match="stale"): + collector.backward(root) + + +def test_ordinary_local_parameter_root_keeps_pytorch_accumulation(): + parameter = torch.nn.Parameter(torch.tensor(3.0, dtype=torch.float64)) + collector = CotangentCollector() + assert collector.backward(parameter) == () + torch.testing.assert_close(parameter.grad, torch.tensor(1.0, dtype=torch.float64)) diff --git a/tests/unit/test_trainer_operations.py b/tests/unit/test_trainer_operations.py new file mode 100644 index 000000000..781928e99 --- /dev/null +++ b/tests/unit/test_trainer_operations.py @@ -0,0 +1,325 @@ +import asyncio +import gc +from types import SimpleNamespace +from typing import Any, cast +import weakref + +import pytest +import torch + +from art.trainer_rank._operations import ( + OperationResultReleasedError, + TrainerOperation, + execute_operation, +) + + +def test_update_identity_replays_outcome_without_applying_again(): + async def run(): + calls = [] + rank = SimpleNamespace( + _rank=SimpleNamespace(), + optim_step=lambda **kwargs: calls.append(kwargs) or {"step": len(calls)}, + ) + operation = TrainerOperation.capture( + "step-1", "optim_step", {"params": {"lr": 1}} + ) + assert await execute_operation(rank, operation) == {"step": 1} + assert await execute_operation(rank, operation) == {"step": 1} + assert len(calls) == 1 + with pytest.raises(ValueError, match="different arguments"): + await execute_operation( + rank, + TrainerOperation.capture("step-1", "optim_step", {"params": {"lr": 2}}), + ) + await execute_operation( + rank, TrainerOperation.capture("ack", "acknowledge", ["step-1"]) + ) + with pytest.raises(OperationResultReleasedError): + await execute_operation(rank, operation) + assert len(calls) == 1 + + asyncio.run(run()) + + +def test_failed_gradient_identity_preserves_original_error(): + from art.trainer_rank import TrainerRankSlotStateError + + async def run(): + stale = TrainerRankSlotStateError("original forward is stale") + calls = [] + + def backward_packets(**kwargs): + calls.append(kwargs) + raise stale + + rank = SimpleNamespace( + _rank=SimpleNamespace(), backward_packets=backward_packets + ) + operation = TrainerOperation.capture("backward-1", "backward", {"packets": ()}) + for _ in range(2): + with pytest.raises(TrainerRankSlotStateError) as error: + await execute_operation(rank, operation) + assert type(error.value) is type(stale) + assert str(error.value) == str(stale) + assert len(calls) == 1 + + asyncio.run(run()) + + +def test_concurrent_retry_waiter_cancellation_does_not_cancel_update(): + async def run(): + entered, release = asyncio.Event(), asyncio.Event() + calls = 0 + + async def optim_step(): + nonlocal calls + calls += 1 + entered.set() + await release.wait() + return calls + + rank = SimpleNamespace(_rank=SimpleNamespace(), optim_step=optim_step) + operation = TrainerOperation.capture("step-1", "optim_step", {}) + original = asyncio.create_task(execute_operation(rank, operation)) + await entered.wait() + retry = asyncio.create_task(execute_operation(rank, operation)) + await asyncio.sleep(0) + retry.cancel() + with pytest.raises(asyncio.CancelledError): + await retry + release.set() + assert await original == 1 + assert await execute_operation(rank, operation) == 1 + + asyncio.run(run()) + + +def test_operation_captures_tensor_arguments_at_submission(): + async def run(): + source = torch.tensor([2.0]) + operation = TrainerOperation.capture("forward-1", "forward", {"inputs": source}) + source.add_(10) + rank = SimpleNamespace( + _rank=SimpleNamespace(), + forward=lambda inputs: inputs * 3, + export_forward=lambda output: output, + ) + assert torch.equal( + await execute_operation(rank, operation), torch.tensor([6.0]) + ) + + asyncio.run(run()) + + +def test_batch_pulls_and_close_are_identified_without_advancing_twice(): + async def run(): + events = [] + rank = SimpleNamespace( + _rank=SimpleNamespace(), + open_forward_batches=lambda **kwargs: events.append("open") or "iterator", + next_forward_batch=lambda **kwargs: events.append("next") or "batch", + export_forward=lambda batch: SimpleNamespace(handle="packet", batch=batch), + close_forward_batches=lambda handle: events.append(("close", handle)), + release_forward=lambda handles: events.append(("release", tuple(handles))), + ) + opening = TrainerOperation.capture("open", "batches_open", {"inputs": []}) + assert await execute_operation(rank, opening) == "iterator" + assert await execute_operation(rank, opening) == "iterator" + next_wave = TrainerOperation.capture( + "next", "batches_next", {"handle": "iterator"} + ) + first = await execute_operation(rank, next_wave) + assert await execute_operation(rank, next_wave) is first + close = TrainerOperation.capture( + "close", + "batches_close", + {"handle": "iterator", "pending_operation": "next"}, + ) + await execute_operation(rank, close) + await execute_operation(rank, close) + assert events == [ + "open", + "next", + ("close", "iterator"), + ("release", ("packet",)), + ] + + asyncio.run(run()) + + +@pytest.mark.parametrize("unserializable", [False, True]) +def test_failed_outcome_releases_traceback_activations_and_replays_error( + unserializable, +): + from art.trainer_rank import TrainerRankMemoryError + + class UnserializableError(RuntimeError): + def __reduce__(self): + raise TypeError("cannot serialize this error") + + async def run(): + references = [] + + def forward(): + activation = torch.ones(64, requires_grad=True).square() + references.append(weakref.ref(activation)) + if unserializable: + raise UnserializableError("failed forward") + raise TrainerRankMemoryError( + "failed forward", predicted_peak_bytes=123, usable_limit_bytes=100 + ) + + rank = SimpleNamespace( + _rank=SimpleNamespace(), + forward=forward, + export_forward=lambda output: output, + ) + operation = TrainerOperation.capture("failed", "forward", {}) + for _ in range(3): + try: + await execute_operation(rank, operation) + except RuntimeError as error: + assert "failed forward" in str(error) + if not unserializable: + assert isinstance(error, TrainerRankMemoryError) + assert error.predicted_peak_bytes == 123 + assert error.usable_limit_bytes == 100 + else: + pytest.fail("failed operation unexpectedly succeeded") + gc.collect() + assert references[0]() is None + assert len(references) == 1 + assert rank._rank._operation_outcomes["failed"].completion.exception() is None + + asyncio.run(run()) + + +def test_concurrent_failed_retry_has_independent_error_without_retained_traceback(): + async def run(): + entered, release = asyncio.Event(), asyncio.Event() + references = [] + + async def optim_step(): + activation = torch.ones(64, requires_grad=True).square() + references.append(weakref.ref(activation)) + entered.set() + await release.wait() + raise ValueError("update rejected") + + rank = SimpleNamespace(_rank=SimpleNamespace(), optim_step=optim_step) + operation = TrainerOperation.capture("failed", "optim_step", {}) + original = asyncio.create_task(execute_operation(rank, operation)) + await entered.wait() + retry = asyncio.create_task(execute_operation(rank, operation)) + await asyncio.sleep(0) + release.set() + results = await asyncio.gather(original, retry, return_exceptions=True) + assert all(isinstance(error, ValueError) for error in results) + assert results[0] is not results[1] + assert len(references) == 1 + del original, retry, results + await asyncio.sleep(0) + gc.collect() + assert references[0]() is None + + asyncio.run(run()) + + +@pytest.mark.parametrize("retain_graph", [False, True]) +@pytest.mark.parametrize("failure_stage", ["collect", "remote"]) +def test_failed_exported_backward_replays_once_and_releases_only_consumed_graphs( + retain_graph, + failure_stage, +): + from art.trainer_rank import TrainerRankZero + from art.trainer_rank._tensors import CotangentCollector, detach_tree + + async def run(): + state = SimpleNamespace(collector=CotangentCollector(), sequence=0, exports={}) + owner = SimpleNamespace() + view = TrainerRankZero( + cast(Any, SimpleNamespace(rank=owner, state=state, dp_rank=0)) + ) + references, attempts = [], [] + + def export(handle): + graph = state.collector.attach( + detach_tree(handle, torch.tensor(2.0, requires_grad=True)) + ) + references.append(weakref.ref(graph)) + return view.export_forward(graph) + + packet = export("used") + unrelated = export("unrelated") + client = CotangentCollector() + loss = client.attach(packet).square() + packets = client.backward(loss, retain_graph=retain_graph) + + def fail(*args, **kwargs): + attempts.append("attempt") + raise ValueError("remote gradient rejected") + + hook = None + if failure_stage == "collect": + hook = references[0]().grad_fn.register_hook(fail) + setattr(view, "_submit_backward", lambda *args, **kwargs: None) + else: + setattr(view, "_submit_backward", fail) + operation = TrainerOperation.capture( + "backward", "backward", {"packets": packets, "retain_graph": retain_graph} + ) + for _ in range(2): + try: + await execute_operation(view, operation) + except ValueError as error: + assert str(error) == "remote gradient rejected" + else: + pytest.fail("failed backward unexpectedly succeeded") + gc.collect() + assert attempts == ["attempt"] + assert (packet.handle in state.exports) is retain_graph + assert (references[0]() is not None) is retain_graph + assert unrelated.handle in state.exports and references[1]() is not None + if hook is not None: + hook.remove() + if retain_graph: + setattr(view, "_submit_backward", lambda *args, **kwargs: None) + await execute_operation( + view, + TrainerOperation.capture( + "retry-intentionally", "backward", {"packets": packets} + ), + ) + gc.collect() + assert packet.handle not in state.exports and references[0]() is None + + asyncio.run(run()) + + +def test_malformed_nonretained_backward_preserves_unrelated_exports(): + from art.trainer_rank import TrainerRankZero + from art.trainer_rank._tensors import CotangentCollector, CotangentPacket + + async def run(): + for handle, gradients in (("known", ()), ("missing", (torch.ones(1),))): + state = SimpleNamespace( + collector=CotangentCollector(), + exports={"known": (torch.ones(1),), "unrelated": (torch.ones(1),)}, + ) + view = TrainerRankZero( + cast(Any, SimpleNamespace(rank=SimpleNamespace(), state=state)) + ) + with pytest.raises((ValueError, KeyError)): + await execute_operation( + view, + TrainerOperation.capture( + "invalid", + "backward", + {"packets": (CotangentPacket(handle, gradients),)}, + ), + ) + assert "unrelated" in state.exports + assert ("known" in state.exports) is (handle == "missing") + + asyncio.run(run()) diff --git a/tests/unit/test_trainer_rank_active_memory.py b/tests/unit/test_trainer_rank_active_memory.py index 25c1e8c89..8f9c6a490 100644 --- a/tests/unit/test_trainer_rank_active_memory.py +++ b/tests/unit/test_trainer_rank_active_memory.py @@ -111,7 +111,7 @@ def run(plan, **kwargs): ], None monkeypatch.setattr(rank, "_run_flat_plan_with_memory_tracking", run) - batches = list(rank.forward_micro_batches([_requests(inactive_length=8001)])) + batches = list(rank.forward_batches([_requests(inactive_length=8001)])) assert len(batches) == 1 batch = batches[0] assert batch.indices == (0,) @@ -247,7 +247,7 @@ def unexpected_execution(*args, **kwargs): with pytest.raises( TrainerRankMemoryError, match="single request cannot be split" ): - rank.dp_rank_forward([request(length)]) + rank.forward([request(length)]) assert rank.last_forward_telemetry()["predicted_peak_bytes"] >= 10_000 diff --git a/tests/unit/test_trainer_rank_cache_recovery.py b/tests/unit/test_trainer_rank_cache_recovery.py index 3b6b08f8c..7ed41d6c0 100644 --- a/tests/unit/test_trainer_rank_cache_recovery.py +++ b/tests/unit/test_trainer_rank_cache_recovery.py @@ -126,7 +126,7 @@ def search(): search, lambda v: v, lambda v, c: (v[0], c), - context="forward_micro_batches" if sync else "dp_rank_forward", + context="forward_batches" if sync else "forward", sync_across_dp=sync, ) except BaseException as e: @@ -155,7 +155,9 @@ def make(self): self.addCleanup(patcher.stop) q = object.__new__(_impl.TrainerRank) q.device = types.SimpleNamespace(type="cuda") + q._graph_memory_policy_enabled = lambda: False q._update_peak_memory_profile = lambda *a: None + q._record_graph_forward_time = lambda *a: None q._execute_flat_plan = lambda p: [object() for _ in range(p.request_count)] q._telemetry_signature = lambda p: {} q._telemetry_plan_signature = lambda p: {} @@ -263,7 +265,7 @@ def test_quota_stops_with_persistent_first_debt(self): def test_completed_forward_earns_next_trial(self): q, c, k, n = self.make() run(q, [fail(n), success(n)]) - q._record_recovery_work("dp_rank_forward", 2.0) + q._record_recovery_work("forward", 2.0) c.free = 40 v, e, count = run(q, [fail(n), success(n)]) self.assertIsNone(e) @@ -345,7 +347,7 @@ def test_foreign_owner_is_retained(self): def test_cap_refusal_in_later_trial_does_not_use_work(self): q, c, k, n = self.make() run(q, [fail(n), success(n)]) - q._record_recovery_work("dp_rank_forward", 100) + q._record_recovery_work("forward", 100) c.free = 40 os.environ["CONTROL_ART_HOOK"] = "1" os.environ["CONTROL_ART_LIMIT"] = "20" @@ -376,8 +378,8 @@ def sample(d): def test_overflowing_work_disables_recovery(self): q, c, k, n = self.make() - q._record_recovery_work("dp_rank_forward", 1e308) - q._record_recovery_work("dp_rank_forward", 1e308) + q._record_recovery_work("forward", 1e308) + q._record_recovery_work("forward", 1e308) self.assertTrue(q._recovery_state().invalid) v, e, count = run(q, [fail(n), success(n)]) self.assertNotIn("release", c.events) @@ -409,26 +411,29 @@ def timer(): p = plan() p.request_count = 1 outputs, baseline = q._run_flat_plan_with_memory_tracking( - p, check=n["_MemoryCheck"](80, 170, True), context="dp_rank_forward" + p, check=n["_MemoryCheck"](80, 170, True), context="forward" ) self.assertEqual(len(outputs), 1) self.assertTrue(q._recovery_state().invalid) self.assertEqual(q._recovery_state().work, 0) def test_forward_recording_failure_preserves_success(self): - q, c, k, n = self.make() - p = plan() - p.request_count = 1 - - def record(*a): - raise ValueError("recording only") - - q._record_recovery_work = record - outputs, baseline = q._run_flat_plan_with_memory_tracking( - p, check=n["_MemoryCheck"](80, 170, True), context="dp_rank_forward" - ) - self.assertEqual(len(outputs), 1) - self.assertTrue(q._recovery_state().invalid) + for recorder in ("_record_graph_forward_time", "_record_recovery_work"): + with self.subTest(recorder=recorder): + q, c, k, n = self.make() + p = plan() + p.request_count = 1 + + def record(*a): + raise ValueError("recording only") + + setattr(q, recorder, record) + outputs, baseline = q._run_flat_plan_with_memory_tracking( + p, check=n["_MemoryCheck"](80, 170, True), context="forward" + ) + self.assertEqual(len(outputs), 1) + self.assertTrue(q._recovery_state().invalid) + self.assertEqual(q._recovery_state().work, 0) def test_execution_oom_keeps_admission_and_cause(self): q, c, k, n = self.make() @@ -442,9 +447,7 @@ def execute(p): q._execute_flat_plan = execute try: - q._run_flat_plan_with_memory_tracking( - p, check=check, context="dp_rank_forward" - ) + q._run_flat_plan_with_memory_tracking(p, check=check, context="forward") except Refusal as e: self.assertIs(e.__cause__, original) self.assertEqual(e.usable_limit_bytes, check.available_bytes) @@ -480,7 +483,7 @@ def execute(p): outputs, baseline, peak = q._execute_split_plan_with_memory_tracking( split, check=n["_MemoryCheck"](80, 170, True), - context="dp_rank_forward", + context="forward", ) except Partial: self.assertEqual(fail_at, 2) @@ -503,14 +506,14 @@ def test_split_mapping_error_rolls_back(self): state.work = 1.0 with self.assertRaises(ValueError): q._execute_split_plan_with_memory_tracking( - split, check=n["_MemoryCheck"](80, 170, True), context="dp_rank_forward" + split, check=n["_MemoryCheck"](80, 170, True), context="forward" ) self.assertEqual(state.work, 1.0) def test_cross_entrypoint_progress_with_work_and_no_new_first_trial(self): q, c, k, n = self.make() run(q, [fail(n), success(n)], sync=True) - q._record_recovery_work("dp_rank_forward", 2.0) + q._record_recovery_work("forward", 2.0) c.free = 40 v, e, count = run(q, [fail(n), success(n)], sync=False) self.assertIsNone(e) @@ -526,9 +529,7 @@ def forbidden(*a, **kw): raise AssertionError("unnecessary added admission collective") q._memory_check_required = forbidden - result = q._plan_admissible_forward( - [], checkpoint=None, context="dp_rank_forward" - ) + result = q._plan_admissible_forward([], checkpoint=None, context="forward") self.assertIs(result[1], value[1]) self.assertNotIn("release", c.events) @@ -547,7 +548,7 @@ def search(): search, lambda v: v, lambda v, c: (v[0], c), - context="forward_micro_batches", + context="forward_batches", sync_across_dp=True, ) self.assertTrue(result[1].fits) diff --git a/tests/unit/test_trainer_rank_calibration_harness.py b/tests/unit/test_trainer_rank_calibration_harness.py index 36d2f65a9..cf94022e1 100644 --- a/tests/unit/test_trainer_rank_calibration_harness.py +++ b/tests/unit/test_trainer_rank_calibration_harness.py @@ -248,19 +248,28 @@ def test_gdn_planner_variants_bracket_the_chain_decision() -> None: assert variant in driver._PLANNER_VARIANTS -def test_contract_accepts_the_yield_empty_flag_only_when_it_is_off_by_default() -> None: - """PR #864 added ``yield_empty`` to the public forwards as a keyword-only flag - that defaults to False; the contract phase tolerates exactly that.""" +def test_contract_accepts_inherited_options_and_disabled_yield_empty() -> None: + """Forward options inherit by default; empty batches remain opt-in.""" import inspect import art.trainer_rank as trainer_rank - for method_name in ("forward_micro_batches", "dp_rank_forward"): + driver.phase_contract() + for method_name in ("forward_batches", "forward"): parameters = driver._public_parameters( getattr(trainer_rank.TrainerRank, method_name) ) - assert set(parameters) <= {"inputs", "checkpoint", "no_grad", "yield_empty"} + assert set(parameters) <= { + "inputs", + "checkpoint", + "no_grad", + "yield_empty", + "options", + } + options = parameters["options"] + assert options.kind is inspect.Parameter.KEYWORD_ONLY + assert options.default is None flag = parameters.get("yield_empty") if flag is not None: assert flag.kind is inspect.Parameter.KEYWORD_ONLY diff --git a/tests/unit/test_trainer_rank_callback_cleanup.py b/tests/unit/test_trainer_rank_callback_cleanup.py new file mode 100644 index 000000000..44a607663 --- /dev/null +++ b/tests/unit/test_trainer_rank_callback_cleanup.py @@ -0,0 +1,179 @@ +"""Independent DP2 x TP2 coverage of callback-boundary graph ownership.""" + +from __future__ import annotations + +import asyncio +from datetime import timedelta +import gc +import sys +from types import ModuleType, SimpleNamespace +from typing import Any + +import pytest +from test_trainer_rank_commands import _input, _loss_tree, _Rank +import torch +import torch.distributed as dist +import torch.multiprocessing as mp + +from art.trainer_rank import run_rank_callback, run_rank_callback_stream + + +def _ownership(rank: Any) -> list[Any]: + gc.collect() + state = rank._rank_command_state + rows: list[Any] = [None] * dist.get_world_size() + dist.all_gather_object(rows, (tuple(state.graphs), tuple(state.released))) + return rows + + +def _empty(rank: Any, boundary: str) -> None: + rows = _ownership(rank) + assert all(not graphs and not released for graphs, released in rows), ( + boundary, + rows, + ) + + +async def _cases(rank: Any, physical: int) -> None: + inputs = [_input(2), _input(3)] + for mode, other in (("zero", "rank"), ("rank", "zero")): + leader = physical == 0 if mode == "zero" else physical % 2 == 0 + + async def call(callback, selected=mode): + return (await run_rank_callback(rank, callback, mode=selected)).value + + # The proxy dies after its creating command session has already STOPped. + returned = await call(lambda view: view.forward(inputs)) + assert any(graphs for graphs, _ in _ownership(rank)) + returned = None + gc.collect() + await call(lambda view: view.zero_grad(), other) + _empty(rank, f"{mode} -> {other} post-STOP release") + + # No later command is required to reclaim a callback-local result. + await call(lambda view: _loss_tree(view.forward(inputs)).item()) + _empty(rank, f"{mode} unused local output") + + # The common release boundary must not revoke genuinely live proxies. + returned = await call(lambda view: view.forward(inputs)) + await call(lambda view: view.zero_grad(), other) + assert any(graphs for graphs, _ in _ownership(rank)) + await call(lambda view: view.backward(_loss_tree(returned))) + expected = (2 if rank.dp == 0 else 3) if mode == "zero" else 5 + torch.testing.assert_close(rank.weight.grad, torch.tensor(float(expected))) + returned = None + gc.collect() + await call(lambda _view: None, other) + _empty(rank, f"{mode} retained output consumed in original mode") + + for ending in ("close", "error", "cancel", "live"): + + async def generate(view): + yield view.forward(inputs) + if ending == "error": + raise RuntimeError("cleanup stream failure") + if ending == "cancel": + task = asyncio.current_task() + assert task is not None + asyncio.get_running_loop().call_soon(task.cancel) + await asyncio.sleep(10) + + stream = run_rank_callback_stream(rank, generate, mode=mode) + yielded = await anext(stream) + if leader: + assert _loss_tree(yielded.value).item() == 10 + else: + assert yielded.logical_rank is None + kept = yielded.value if ending == "live" else None + del yielded + gc.collect() + if leader and ending == "error": + with pytest.raises(RuntimeError, match="cleanup stream failure"): + await anext(stream) + elif leader and ending == "cancel": + pending = asyncio.create_task(anext(stream)) + with pytest.raises(asyncio.CancelledError): + await pending + await stream.aclose() + if ending == "live": + # Closing the producing generator cannot revoke an output that + # its caller still holds, even after the other view runs. + await call(lambda view: view.zero_grad(), other) + assert any(graphs for graphs, _ in _ownership(rank)) + await call(lambda view: view.backward(_loss_tree(kept))) + torch.testing.assert_close( + rank.weight.grad, torch.tensor(float(expected)) + ) + kept = None + gc.collect() + await call(lambda _view: None, other) + if ending in ("error", "cancel"): + # Failure reaches the controller before all peers join cleanup; + # the next callback must finish that cleanup before admission. + await call(lambda _view: None, other) + _empty(rank, f"{mode} generator {ending}") + await call(lambda view: view.zero_grad(), other) + _empty(rank, f"{mode} generator {ending} followed by {other}") + + # Followers cancelled during a session must still join cleanup, and the + # next callback in the other mode must see an intact communicator. + entered = asyncio.Event() + zero_grad = rank.zero_grad + + def entered_zero_grad(): + zero_grad() + entered.set() + + rank.zero_grad = entered_zero_grad + + async def suspended(view): + unused = view.forward(inputs) + view.zero_grad() + del unused + gc.collect() + await asyncio.sleep(0.05) + + pending = asyncio.create_task(run_rank_callback(rank, suspended, mode=mode)) + await entered.wait() + if not leader: + pending.cancel() + (result,) = await asyncio.gather(pending, return_exceptions=True) + rank.zero_grad = zero_grad + if leader: + assert not isinstance(result, BaseException), result + else: + assert isinstance(result, asyncio.CancelledError), result + await call(lambda view: view.zero_grad(), other) + _empty(rank, f"{mode} cancellation followed by {other}") + + +def _worker(physical: int, rendezvous: str) -> None: + torch.set_num_threads(1) + dist.init_process_group( + "gloo", + init_method=f"file://{rendezvous}", + rank=physical, + world_size=4, + timeout=timedelta(seconds=20), + ) + try: + groups = [dist.new_group([0, 1]), dist.new_group([2, 3])] + dp, tp = divmod(physical, 2) + ps = SimpleNamespace( + get_tensor_model_parallel_rank=lambda: tp, + get_context_parallel_rank=lambda: 0, + get_data_parallel_rank=lambda: dp, + get_data_parallel_world_size=lambda: 2, + get_tensor_and_context_parallel_group=lambda **kwargs: groups[dp], + ) + megatron, core = ModuleType("megatron"), ModuleType("megatron.core") + setattr(core, "parallel_state", ps) + setattr(megatron, "core", core) + sys.modules.update({"megatron": megatron, "megatron.core": core}) + asyncio.run(_cases(_Rank(dp, 2), physical)) + finally: + dist.destroy_process_group() + + +def test_gloo_dp2_tp2_callback_cleanup_across_modes(tmp_path): + mp.spawn(_worker, args=(str(tmp_path / "cleanup-init"),), nprocs=4, join=True) diff --git a/tests/unit/test_trainer_rank_callback_failures.py b/tests/unit/test_trainer_rank_callback_failures.py new file mode 100644 index 000000000..578224489 --- /dev/null +++ b/tests/unit/test_trainer_rank_callback_failures.py @@ -0,0 +1,234 @@ +"""Controller cancellation must remain possible while global cleanup is pending.""" + +from __future__ import annotations + +import asyncio +from datetime import timedelta +import gc +from multiprocessing.connection import Connection +import sys +from types import ModuleType, SimpleNamespace +from typing import Any + +import pytest +from test_trainer_rank_commands import _input, _loss_tree, _Rank +import torch +import torch.distributed as dist +import torch.multiprocessing as mp + +from art.trainer_rank import run_rank_callback, run_rank_callback_stream + + +def _messages(connection: Connection) -> asyncio.Queue[Any]: + queue: asyncio.Queue[Any] = asyncio.Queue() + loop = asyncio.get_running_loop() + + def receive(): + try: + queue.put_nowait(connection.recv()) + except EOFError: + loop.remove_reader(connection.fileno()) + + loop.add_reader(connection.fileno(), receive) + return queue + + +async def _serve(rank: Any, connection: Connection, kind: str) -> None: + commands = _messages(connection) + release = asyncio.Event() + + def forward(view): + value = _loss_tree(view.forward([_input(rank.dp + 2)])).item() + gc.collect() + return value + + async def callback(view): + forward(view) + connection.send(("entered", rank.dp)) + await release.wait() + raise RuntimeError("DP0 user failure") + + async def generate(view): + try: + yield forward(view) + connection.send(("entered", rank.dp)) + await release.wait() + raise RuntimeError("DP0 user failure") + finally: + if kind == "stream_close" and rank.dp == 0: + raise RuntimeError("DP0 generator close failure") + + async def execute(): + try: + if kind == "ordinary": + await run_rank_callback(rank, callback, mode="rank") + else: + stream = run_rank_callback_stream(rank, generate, mode="rank") + try: + if kind == "stream_close": + await anext(stream) + connection.send(("entered", rank.dp)) + await release.wait() + await stream.aclose() + else: + async for _ in stream: + pass + finally: + await stream.aclose() + except asyncio.CancelledError: + connection.send(("cancelled", rank.dp)) + except Exception as error: + connection.send(("error", str(error))) + else: + connection.send(("unexpected_success", rank.dp)) + + pending = asyncio.create_task(execute()) + try: + while True: + command = await commands.get() + if command == "raise": + release.set() + elif command == "cancel": + pending.cancel() + elif command == "followup": + await pending + # Entry must join the previous cleanup before any new session. + result = await run_rank_callback( + rank, lambda view: (view.zero_grad(), 17)[1], mode="zero" + ) + state = rank._rank_command_state + assert not state.graphs and not state.released + connection.send(("followup", result.value)) + elif command == "stop": + return + else: + raise AssertionError(command) + finally: + pending.cancel() + await asyncio.gather(pending, return_exceptions=True) + asyncio.get_running_loop().remove_reader(connection.fileno()) + + +def _worker(physical: int, rendezvous: str, connection: Connection, kind: str) -> None: + torch.set_num_threads(1) + dist.init_process_group( + "gloo", + init_method=f"file://{rendezvous}", + rank=physical, + world_size=2, + timeout=timedelta(seconds=12), + ) + try: + groups = [dist.new_group([0]), dist.new_group([1])] + ps = SimpleNamespace( + get_tensor_model_parallel_rank=lambda: 0, + get_context_parallel_rank=lambda: 0, + get_data_parallel_rank=lambda: physical, + get_data_parallel_world_size=lambda: 2, + get_tensor_and_context_parallel_group=lambda **kwargs: groups[physical], + ) + megatron, core = ModuleType("megatron"), ModuleType("megatron.core") + setattr(core, "parallel_state", ps) + setattr(megatron, "core", core) + sys.modules.update({"megatron": megatron, "megatron.core": core}) + asyncio.run(_serve(_Rank(physical, 2), connection, kind)) + finally: + connection.close() + dist.destroy_process_group() + + +async def _gather(executions): + # Caladan _execution.gather uses FIRST_EXCEPTION, then a grace period, + # then cancellation/drain. A failure hidden behind cleanup defeats step 1. + done, pending = await asyncio.wait(executions, return_when=asyncio.FIRST_EXCEPTION) + error = next(task.exception() for task in done if task.exception() is not None) + if pending: + _, pending = await asyncio.wait(pending, timeout=0.05) + for task in pending: + task.cancel() + await asyncio.gather(*pending, return_exceptions=True) + assert error is not None + raise error + + +async def _controller(connections: list[Connection]) -> None: + messages = [_messages(connection) for connection in connections] + executions = [] + cancelled = [] + + async def remote(index: int): + try: + result = await messages[index].get() + except asyncio.CancelledError: + connections[index].send("cancel") + result = await asyncio.shield(messages[index].get()) + assert result == ("cancelled", index), result + cancelled.append(index) + raise + assert result[0] == "error", result + raise RuntimeError(result[1]) + + try: + entered = await asyncio.wait_for( + asyncio.gather(*(queue.get() for queue in messages)), timeout=20 + ) + assert entered == [("entered", 0), ("entered", 1)] + executions = [asyncio.create_task(remote(index)) for index in range(2)] + connections[0].send("raise") + # This deadline measures failure visibility, not worker startup. Only + # the controller may cancel DP1; its user callback never completes. + controller = asyncio.create_task(_gather(executions)) + done, _ = await asyncio.wait([controller], timeout=2) + try: + assert done, "DP0 failure hidden behind blocked DP1 callback cleanup" + with pytest.raises(RuntimeError, match="DP0 .*failure"): + await controller + assert cancelled == [1] + for connection in connections: + connection.send("followup") + followup = await asyncio.wait_for( + asyncio.gather(*(queue.get() for queue in messages)), timeout=10 + ) + assert followup == [("followup", 17), ("followup", None)] + finally: + for task in executions: + task.cancel() + await asyncio.wait_for( + asyncio.gather(*executions, return_exceptions=True), timeout=15 + ) + controller.cancel() + await asyncio.gather(controller, return_exceptions=True) + finally: + for connection in connections: + connection.send("stop") + asyncio.get_running_loop().remove_reader(connection.fileno()) + + +@pytest.mark.parametrize("kind", ["ordinary", "stream", "stream_close"]) +def test_gloo_dp2_first_exception_cancels_peer_and_reuses_actor(tmp_path, kind): + context = mp.get_context("spawn") + pairs = [context.Pipe() for _ in range(2)] + processes = [ + context.Process( + target=_worker, + args=(rank, str(tmp_path / "init"), pairs[rank][1], kind), + ) + for rank in range(2) + ] + try: + for process in processes: + process.start() + for _, child in pairs: + child.close() + asyncio.run(_controller([parent for parent, _ in pairs])) + for process in processes: + process.join(timeout=15) + assert process.exitcode == 0 + finally: + for process in processes: + if process.is_alive(): + process.terminate() + process.join(timeout=5) + for parent, child in pairs: + parent.close() + child.close() diff --git a/tests/unit/test_trainer_rank_callback_lifecycle.py b/tests/unit/test_trainer_rank_callback_lifecycle.py new file mode 100644 index 000000000..4d2f51b2e --- /dev/null +++ b/tests/unit/test_trainer_rank_callback_lifecycle.py @@ -0,0 +1,182 @@ +"""Checkpoint scopes and delayed cleanup stay within their callback session.""" + +import asyncio +from datetime import timedelta +import gc +import sys +from types import ModuleType, SimpleNamespace +from typing import Any + +import pytest +from test_trainer_rank_commands import _input, _Rank +import torch.distributed as dist +import torch.multiprocessing as mp + +from art.trainer_rank import run_rank_callback + + +class _CheckpointRank(_Rank): + def __init__(self): + super().__init__() + self._slot_stack = [] + + @staticmethod + def _checkpoint_source(checkpoint): + return checkpoint, checkpoint + + @staticmethod + def _slot_ref(path): + return path + + def _push_checkpoint_sync(self, path, directory): + self._slot_stack.append(path) + + def pop_checkpoint(self): + self._slot_stack.pop() + + +@pytest.mark.parametrize("mode", ["rank", "zero"]) +def test_callback_checkpoint_context_restores_nested_and_exceptional_stack(mode): + rank: Any = _CheckpointRank() + + def callback(view): + with view.push_checkpoint("outer") as outer: + assert rank._slot_stack == ["outer"] + with pytest.raises(ValueError, match="body failed"): + with view.push_checkpoint("inner"): + assert rank._slot_stack == ["outer", "inner"] + raise ValueError("body failed") + assert rank._slot_stack == ["outer"] + assert not rank._slot_stack + with pytest.raises(RuntimeError, match="entered twice"): + outer.__enter__() + with pytest.raises(BaseExceptionGroup) as failure: + with view.push_checkpoint("outer"): + view._push_checkpoint_sync("changed", None) + raise ValueError("body failed") + assert [type(error) for error in failure.value.exceptions] == [ + ValueError, + RuntimeError, + ] + assert rank._slot_stack == ["outer", "changed"] + view.pop_checkpoint() + view.pop_checkpoint() + + asyncio.run(run_rank_callback(rank, callback, mode=mode)) + assert not rank._slot_stack + + +@pytest.mark.parametrize("mode", ["rank", "zero"]) +def test_retained_callback_iterator_closes_before_stop_and_cannot_reenter(mode): + rank: Any = _CheckpointRank() + views, held = [], [] + + def callback(view): + views.append(view) + batches = view.forward_batches([_input(1), _input(2)]) + held.append(batches) + next(batches) + raise ValueError("keep traceback alive") + + with pytest.raises(ValueError) as failure: + asyncio.run(run_rank_callback(rank, callback, mode=mode)) + assert failure.value.__traceback__ is not None + assert rank.closed == 1 + sequence = rank._rank_command_state.sequence + held.pop().close() + assert rank._rank_command_state.sequence == sequence + + def following(view): + with pytest.raises(RuntimeError, match="session has stopped"): + views[0].zero_grad() + view.zero_grad() + + asyncio.run(run_rank_callback(rank, following, mode=mode)) + + +def _lifecycle_worker(physical, rendezvous): + dist.init_process_group( + "gloo", + init_method=f"file://{rendezvous}", + rank=physical, + world_size=2, + timeout=timedelta(seconds=15), + ) + try: + megatron, core = ModuleType("megatron"), ModuleType("megatron.core") + setattr( + core, + "parallel_state", + SimpleNamespace( + get_tensor_model_parallel_rank=lambda: physical, + get_context_parallel_rank=lambda: 0, + get_tensor_and_context_parallel_group=lambda: dist.group.WORLD, + ), + ) + setattr(megatron, "core", core) + sys.modules.update({"megatron": megatron, "megatron.core": core}) + for mode in ("zero", "rank"): + rank: Any = _CheckpointRank() + held = [] + + def fail(view): + held.append(view) + with view.push_checkpoint("outer"): + with view.push_checkpoint("inner"): + batches = view.forward_batches([_input(1), _input(2)]) + next(batches) + raise ValueError("retain callback traceback") + + failure = None + try: + asyncio.run(run_rank_callback(rank, fail, mode=mode)) + except ValueError as error: + failure = error + assert (failure is not None) == (physical == 0) + assert not rank._slot_stack + assert rank.closed == 1 + dist.barrier() + if failure is not None: + # The traceback is the only owner of the suspended user iterator. + # Releasing it after STOP must not broadcast to absent receivers. + failure.__traceback__ = None + failure = None + gc.collect() + dist.barrier() + + def following(view): + with pytest.raises(RuntimeError, match="session has stopped"): + held[0].zero_grad() + with view.push_checkpoint("next"): + view.zero_grad() + + asyncio.run(run_rank_callback(rank, following, mode=mode)) + assert not rank._slot_stack + dist.barrier() + + zero_grad = rank.zero_grad + + def change_stack(): + if physical == 1: + rank._slot_stack.append("changed") + + rank.zero_grad = change_stack + + def mismatched(view): + with pytest.raises(RuntimeError, match="stack changed"): + with view.push_checkpoint("outer"): + view.zero_grad() + + asyncio.run(run_rank_callback(rank, mismatched, mode=mode)) + # Validate every peer before any peer mutates its stack. + assert rank._slot_stack == (["outer", "changed"] if physical else ["outer"]) + rank._slot_stack.clear() + rank.zero_grad = zero_grad + asyncio.run(run_rank_callback(rank, following, mode=mode)) + dist.barrier() + finally: + dist.destroy_process_group() + + +def test_delayed_callback_cleanup_and_checkpoint_scopes_leave_gloo_reusable(tmp_path): + mp.spawn(_lifecycle_worker, args=(str(tmp_path / "lifecycle"),), nprocs=2) diff --git a/tests/unit/test_trainer_rank_commands.py b/tests/unit/test_trainer_rank_commands.py new file mode 100644 index 000000000..9cb010cc5 --- /dev/null +++ b/tests/unit/test_trainer_rank_commands.py @@ -0,0 +1,809 @@ +"""Command participation tests; native model numerics are a separate GPU gate.""" + +from __future__ import annotations + +import asyncio +from datetime import timedelta +import gc +import sys +import threading +from types import ModuleType, SimpleNamespace +from typing import Any + +import pytest +import torch +import torch.distributed as dist +import torch.multiprocessing as mp + +from art.trainer_rank import ( + ForwardInput, + ForwardOutput, + MicroBatch, + MicroBatchStats, + TrainerRank, + TrainerRankZero, + run_rank_callback, + run_rank_callback_stream, +) + + +class _Rank: + device = torch.device("cpu") + hidden_size = 1 + + def __init__(self, dp: int = 0, size: int = 1) -> None: + self.dp, self.size = dp, size + self.weight = torch.nn.Parameter(torch.tensor(2.0)) + self.closed = 0 + self.steps = 0 + + def _dp_rank_and_size(self): + return self.dp, self.size + + def forward(self, tree, **kwargs): + if isinstance(tree, ForwardInput): + enabled = ( + torch.is_grad_enabled() + if kwargs.get("no_grad") is None + else not kwargs["no_grad"] + ) + with torch.set_grad_enabled(enabled): + value = tree.input_tokens.float() * self.weight + return ForwardOutput(None, None, None, value) + from art.trainer_rank._impl import _rebuild_forward_tree + + return _rebuild_forward_tree( + tree, [self.forward(child, **kwargs) for child in tree] + ) + + def forward_batches(self, inputs, **kwargs): + try: + # One global root per wave guarantees empty DP partitions. + for index, item in enumerate(inputs): + owned = index % self.size == self.dp + yield MicroBatch( + [item] if owned else [], + [self.forward(item, **kwargs)] if owned else [], + [index] if owned else [], + MicroBatchStats( + index, index + 1, 1, int(owned), 0, 0, 0, 0, 0, False + ), + ) + finally: + self.closed += 1 + + def zero_grad(self): + self.weight.grad = None + + def optim_step(self, **kwargs): + self.steps += 1 + return {"steps": self.steps} + + +def _input(value): + return ForwardInput(input_tokens=torch.tensor([value])) + + +def _suspended_abort_worker(physical, rendezvous): + from test_trainer_rank_custom_tensors import _trainer + + dist.init_process_group( + "gloo", + init_method=f"file://{rendezvous}", + rank=physical, + world_size=2, + timeout=timedelta(seconds=8), + ) + try: + native, _ = _trainer("student") + native._checkpoint_process_group = dist.new_group( + backend="gloo", timeout=timedelta(seconds=8) + ) + native._checkpoint_finalize_process_group = dist.new_group( + backend="gloo", timeout=timedelta(seconds=8) + ) + rank: Any = _Rank() + actor_thread = threading.get_ident() + + async def run(mode): + entered, release = asyncio.Event(), asyncio.Event() + + def zero_grad(): + assert threading.get_ident() == actor_thread + entered.set() + + rank.zero_grad = zero_grad + + async def callback(view): + view.zero_grad() + await release.wait() + view.zero_grad() + + async def generator(view): + view.zero_grad() + try: + yield "suspended" + finally: + view.zero_grad() + + async def consume(): + if mode != "stream": + return await run_rank_callback(rank, callback, mode="zero") + stream = run_rank_callback_stream(rank, generator, mode="zero") + result = await anext(stream) + if physical == 0: + assert result.value == "suspended" + await release.wait() + await stream.aclose() + + pending = asyncio.create_task(consume()) + + if mode == "shutdown": + await entered.wait() + # asyncio.run cancels every pending Task, not just the callback. + # The control receive must remain joinable during that drain. + return + + async def abort(): + await entered.wait() + await asyncio.sleep(0.05) + if mode == "cancel" and physical == 1: + pending.cancel() + await asyncio.sleep(0) + # The same synchronous physical abort dispatched by Caladan, + # on the actor loop while the logical callback is suspended. + native.abort_checkpoint_save("unprepared-save") + assert not pending.done() + release.set() + + await abort() + result = (await asyncio.gather(pending, return_exceptions=True))[0] + if mode == "cancel" and physical == 1: + assert isinstance(result, asyncio.CancelledError) + elif isinstance(result, BaseException): + raise result + + for mode in ("async", "stream", "cancel", "shutdown"): + asyncio.run(run(mode)) + finally: + dist.destroy_process_group() + + +def test_suspended_callbacks_leave_physical_checkpoint_abort_responsive(tmp_path): + mp.spawn(_suspended_abort_worker, args=(str(tmp_path / "abort-init"),), nprocs=2) + + +def _loss_tree(tree): + if isinstance(tree, ForwardOutput): + return tree.hidden_states.sum() + return sum(_loss_tree(item) for item in tree) + + +def test_native_facade_owns_backward_and_hides_zero_reduce(): + rank: Any = _Rank() + + def callback(view): + assert isinstance(view, TrainerRankZero) + assert not hasattr(view, "reduce") + outputs = view.forward([[_input(2), [_input(3)]]]) + extra = view.forward(_input(5)) + unused = view.forward(_input(100)) + view.backward((_loss_tree(outputs) + _loss_tree(extra)).square()) + assert rank.weight.grad.item() == 400 + assert unused.hidden_states.item() == 200 + assert view.optim_step() == {"steps": 1} + return outputs[0][1][0].hidden_states.item() + + result = asyncio.run(run_rank_callback(rank, callback, mode="zero")) + assert result.logical_rank == 0 + assert result.value == 6 + assert rank.closed == 3 + + +def test_client_packet_registry_survives_callbacks(): + rank: Any = _Rank() + from art.trainer_rank._tensors import CotangentCollector + + packet = asyncio.run( + run_rank_callback( + rank, + lambda view: view.export_forward(view.forward([_input(3)])), + mode="zero", + ) + ).value + collector = CotangentCollector() + output = collector.attach(packet) + gradients = collector.backward(output[0].hidden_states.square().sum()) + asyncio.run( + run_rank_callback( + rank, lambda view: view.backward_packets(gradients), mode="zero" + ) + ) + assert rank.weight.grad.item() == 36 + + +def test_rank_facade_is_trainer_rank_and_dispatches_inherited_methods(): + rank: Any = _Rank() + + def callback(view): + assert isinstance(view, TrainerRank) + view.zero_grad() + return view.optim_step() + + assert asyncio.run(run_rank_callback(rank, callback)).value == {"steps": 1} + + +def test_stream_forwards_sends_and_closes(): + rank: Any = _Rank() + closed = [] + + def callback(view): + try: + sent = yield view.forward(_input(3)).hidden_states.item() + yield sent * 2 + finally: + closed.append(True) + + async def run(): + stream = run_rank_callback_stream(rank, callback, mode="zero") + assert (await anext(stream)).value == 6 + assert (await stream.asend(9)).value == 18 + await stream.aclose() + + asyncio.run(run()) + assert closed == [True] + + +class _Unserializable: + def __reduce_ex__(self, protocol): + raise TypeError("intentional serialization failure") + + +def _rank_local_restore_failure(): + if dist.get_rank() == 1: + raise ValueError("intentional peer deserialization failure") + return None + + +class _BadRestore: + def __reduce_ex__(self, protocol): + return _rank_local_restore_failure, () + + +class _FailPhysicalBackward(torch.autograd.Function): + @staticmethod + def forward(ctx, value): + return value + + @staticmethod + def backward(ctx, gradient): + if dist.get_rank() == 1: + raise RuntimeError("intentional physical backward failure") + return gradient + + +def _distributed_worker(physical, rendezvous, output): + dist.init_process_group( + "gloo", + init_method=f"file://{rendezvous}", + rank=physical, + world_size=4, + timeout=timedelta(seconds=30), + ) + try: + groups = [dist.new_group([0, 1]), dist.new_group([2, 3])] + dp_groups = [dist.new_group([0, 2]), dist.new_group([1, 3])] + dp, tp = divmod(physical, 2) + ps = SimpleNamespace( + get_tensor_model_parallel_rank=lambda: tp, + get_context_parallel_rank=lambda: 0, + get_data_parallel_rank=lambda: dp, + get_data_parallel_world_size=lambda: 2, + get_tensor_and_context_parallel_group=lambda **kwargs: groups[dp], + ) + megatron = ModuleType("megatron") + core = ModuleType("megatron.core") + setattr(core, "parallel_state", ps) + setattr(megatron, "core", core) + sys.modules.update({"megatron": megatron, "megatron.core": core}) + rank: Any = _Rank(dp, 2) + counts = [0, 0] + + def zero(view): + counts[0] += 1 + result = view.forward([[_input(2), [_input(3)]], _input(7)]) + second = view.forward(_input(11)) + view.backward((_loss_tree(result) + _loss_tree(second)).square()) + # Fail one physical rank's packet preflight before any backward. + from art.trainer_rank._tensors import CotangentPacket + + with pytest.raises((RuntimeError, ValueError)): + view._invoke( + "backward", + (CotangentPacket("zero:missing:dp:1", (None,)),), + retain_graph=False, + ) + with pytest.raises(RuntimeError, match="serialization failed"): + view._invoke("zero_grad", _Unserializable()) + with pytest.raises(RuntimeError, match="deserialization failed"): + view._invoke("zero_grad", _BadRestore()) + iterator = view.forward_batches([_input(1), _input(2)]) + next(iterator) + iterator.close() + return [_loss_tree(root).item() for root in result] + + result = asyncio.run(run_rank_callback(rank, zero, mode="zero")) + # Global scalar loss is (2 * (2+3+7+11)) ** 2. + expected = 2 * 46 * (16 if dp == 0 else 7) + torch.testing.assert_close(rank.weight.grad, torch.tensor(float(expected))) + assert rank.closed == 3 + + def logical(view): + counts[1] += 1 + view.zero_grad() + roots = [] if dp == 1 else [[_input(5)]] + out = view.forward(roots) + if out: + view.backward(_loss_tree(out)) + view.optim_step() + return dp + + per_dp = asyncio.run(run_rank_callback(rank, logical)) + + # Persistent streams release the command scope between waves. Every + # physical participant must retain its iterator while serving other jobs. + def run(callback): + return asyncio.run(run_rank_callback(rank, callback, mode="zero")).value + + handle = run(lambda view: view.open_forward_batches([_input(3), _input(5)])) + assert rank._rank_command_state.iterators + run(lambda view: view.zero_grad()) + first = run(lambda view: view.next_forward_batch(handle)) + run(lambda view: view.backward(_loss_tree(first.outputs))) + run(lambda view: view.optim_step()) + second = run(lambda view: view.next_forward_batch(handle)) + run(lambda view: view.backward(_loss_tree(second.outputs))) + assert run(lambda view: view.next_forward_batch(handle)) is None + run(lambda view: view.close_forward_batches(handle)) + assert not rank._rank_command_state.iterators + assert rank.weight.grad.item() == (3 if dp == 0 else 5) + + from art.trainer_rank import _commands + + decoded_outputs = [] + loads = _commands.cloudpickle.loads + + def track_loads(payload): + value = loads(payload) + if ( + isinstance(value, tuple) + and len(value) == 2 + and isinstance(value[1], _commands._OutputPacket) + ): + decoded_outputs.append(value[1].packet.tensors) + return value + + setattr(_commands.cloudpickle, "loads", track_loads) + try: + large = run( + lambda view: view.forward( + [ + [ + ForwardInput( + input_tokens=torch.ones(131072, dtype=torch.long) + ) + ], + [ + ForwardInput( + input_tokens=torch.ones(131072, dtype=torch.long) + ) + ], + ], + no_grad=True, + ) + ) + finally: + setattr(_commands.cloudpickle, "loads", loads) + assert bool(decoded_outputs) is (physical == 0) + if physical == 0: + assert ( + sum( + t.numel() * t.element_size() + for tensors in decoded_outputs + for t in tensors + ) + == 2 * 131072 * 4 + ) + assert large[1][0].hidden_states.shape == (131072,) + rank._available_cpu_memory_bytes = lambda: 0 if physical == 0 else 1 << 60 + + def host_refusal(view): + with pytest.raises(MemoryError, match="CPU bytes"): + view.forward([_input(3)], no_grad=True) + view.zero_grad() + + run(host_refusal) + del rank._available_cpu_memory_bytes + + retained_before_failure = set(rank._rank_command_state.graphs) + + def fail_result_decode(payload): + value = loads(payload) + if ( + physical == 0 + and isinstance(value, tuple) + and len(value) == 2 + and isinstance(value[1], _commands._OutputPacket) + ): + raise ValueError("intentional result decode failure") + return value + + def decode_refusal(view): + with pytest.raises(ValueError, match="result decode failure"): + view.forward([_input(11)]) + + setattr(_commands.cloudpickle, "loads", fail_result_decode) + try: + run(decode_refusal) + finally: + setattr(_commands.cloudpickle, "loads", loads) + assert set(rank._rank_command_state.graphs) <= retained_before_failure + assert not rank._rank_command_state.iterators + run(lambda view: view.backward(_loss_tree(view.forward([_input(11)])))) + if dp == 0: + assert rank.weight.grad.item() == 11 + + # Snapshot fits, but serialization's transient buffers do not. Refuse + # before pickle allocates those buffers on any participant. + dumps, serialized = _commands.cloudpickle.dumps, [] + + def track_dumps(value): + if isinstance(value, tuple) and len(value) == 2: + serialized.append(True) + return dumps(value) + + rank._available_cpu_memory_bytes = lambda: 1024 if physical == 0 else 1 << 60 + setattr(_commands.cloudpickle, "dumps", track_dumps) + try: + run(host_refusal) + finally: + setattr(_commands.cloudpickle, "dumps", dumps) + del rank._available_cpu_memory_bytes + assert not serialized + + forward = rank.forward + retained_before_failure = set(rank._rank_command_state.graphs) + + def failing_forward(tree, **kwargs): + from dataclasses import replace + + result = forward(tree, **kwargs) + if isinstance(result, ForwardOutput): + result = replace( + result, + hidden_states=_FailPhysicalBackward.apply(result.hidden_states), + ) + return result + + def backward_refusal(view): + with pytest.raises(RuntimeError, match="physical backward failure"): + view.backward(_loss_tree(view.forward([_input(7)]))) + + rank.forward = failing_forward + try: + run(backward_refusal) + finally: + rank.forward = forward + assert set(rank._rank_command_state.graphs) <= retained_before_failure + run(lambda view: view.zero_grad()) + del sys.modules["megatron"], sys.modules["megatron.core"] + import megatron.core as real_core + from test_trainer_rank_custom_tensors import _trainer + + setattr(real_core, "parallel_state", ps) + native, _ = _trainer("student") + factories = [] + + def head_callback(view): + class LocalHead(torch.nn.Module): + def __init__(self): + super().__init__() + factories.append(True) + self.weight = torch.nn.Parameter(torch.tensor(2.0)) + + def forward(self, value): + return self.weight.square() * value + + head = view.module("head", LocalHead, checkpoint="student") + view.backward(head(torch.tensor(3.0))) + + asyncio.run(run_rank_callback(native, head_callback, mode="zero")) + weight = native._checkpoint_slots["student"].custom["head"].value.weight + gradient = ( + torch.zeros_like(weight) if weight.grad is None else weight.grad.clone() + ) + torch.testing.assert_close(gradient, torch.tensor(12.0 if dp == 0 else 0.0)) + dist.all_reduce(gradient, group=dp_groups[tp]) + torch.testing.assert_close(gradient, torch.tensor(12.0)) + assert len(factories) == int(physical == 0) + # Only DP0 has a logical handle after Zero. Refresh must stay inside + # its TP group when the next callback is independently per-DP. + asyncio.run(run_rank_callback(native, lambda view: view.zero_grad())) + + gathered = [None] * 4 + dist.all_gather_object(gathered, (counts, result, per_dp, rank.steps)) + if physical == 0: + torch.save(gathered, output) + except BaseException: + import traceback + + traceback.print_exc() + raise + finally: + dist.destroy_process_group() + + +def test_gloo_dp2_tp2_participation_and_gradients(tmp_path): + output = tmp_path / "result.pt" + mp.spawn( + _distributed_worker, + args=(str(tmp_path / "init"), str(output)), + nprocs=4, + join=True, + ) + rows = torch.load(output, weights_only=False) + assert [row[0] for row in rows] == [[1, 1], [0, 0], [0, 1], [0, 0]] + assert [row[1].logical_rank for row in rows] == [0, None, None, None] + assert rows[0][1].value == [10, 14] + assert [row[2].logical_rank for row in rows] == [0, None, 1, None] + assert [row[3] for row in rows] == [2, 2, 2, 2] + + +def test_released_client_graph_releases_physical_bridge(): + rank: Any = _Rank() + packet = asyncio.run( + run_rank_callback( + rank, lambda view: view.export_forward(view.forward(_input(3))), mode="zero" + ) + ).value + assert rank._rank_command_state.graphs + asyncio.run( + run_rank_callback( + rank, lambda view: view.release_forward([packet.handle]), mode="zero" + ) + ) + assert not rank._rank_command_state.exports + gc.collect() + asyncio.run( + run_rank_callback(rank, lambda view: view.release_forward([]), mode="zero") + ) + assert not rank._rank_command_state.graphs + + +def test_forward_batches_captures_policy_before_iteration(): + from art.trainer_rank import ForwardOptions + + rank = object.__new__(TrainerRank) + rank._forward_options = ForwardOptions(max_gradient_staleness=1, allow_replay=False) + rank._skipped_forward_waves = {} + request = _input(3) + request.options = ForwardOptions(max_gradient_staleness=0) + + def batches(inputs, **kwargs): + yield MicroBatch( + inputs, [], [0], MicroBatchStats(0, 1, 1, 1, 0, 0, 0, 0, 0, False) + ) + + rank._forward_batches = batches + iterator = rank.forward_batches( + [request], options=ForwardOptions(allow_cpu_offload=False), yield_empty=True + ) + request.options = ForwardOptions(max_gradient_staleness=4) + rank._forward_options = ForwardOptions(max_gradient_staleness=7) + captured = next(iterator).inputs[0].options + assert captured.max_gradient_staleness == 0 + assert captured.allow_cpu_offload is False + assert captured.allow_replay is False + iterator.close() + + +@pytest.mark.parametrize("mode", ["rank", "zero"]) +def test_tuple_root_and_nested_tuple_shape(mode): + rank: Any = _Rank() + inputs = ([_input(2), (_input(3),)],) + result = asyncio.run( + run_rank_callback(rank, lambda view: view.forward(inputs), mode=mode) + ).value + assert isinstance(result, tuple) + assert isinstance(result[0], list) + assert isinstance(result[0][1], tuple) + assert result[0][1][0].hidden_states.item() == 6 + + +def test_logical_head_factory_runs_once_and_head_only_client_backward(): + from test_trainer_rank_custom_tensors import _trainer + + from art.trainer_rank._heads import LiveHead + from art.trainer_rank._tensors import CotangentCollector + + trainer, _ = _trainer("student") + calls = [] + + def factory(): + calls.append(True) + return torch.tensor(2.0) + + def register(view): + parameter = view.parameter("gain", factory, checkpoint="student") + assert view.parameter("gain", factory, checkpoint="student") is parameter + return view._invoke("head", "head_export", (("student", "gain"),))[0] + + state = asyncio.run(run_rank_callback(trainer, register, mode="zero")).value + assert calls == [True] + collector = CotangentCollector() + client = LiveHead(state, torch.tensor(2.0), collector) + packets = collector.backward(client.value.square() * 3) + asyncio.run( + run_rank_callback( + trainer, lambda view: view.backward_packets(packets), mode="zero" + ) + ) + parameter = trainer._checkpoint_slots["student"].custom["gain"].value + torch.testing.assert_close(parameter.grad, torch.tensor(12.0)) + + +def test_logical_native_head_backward_commits_after_local_autograd(): + from test_trainer_rank_custom_tensors import _trainer + + trainer, _ = _trainer("student") + + class LocalHead(torch.nn.Module): + def __init__(self): + super().__init__() + self.weight = torch.nn.Parameter(torch.tensor(2.0)) + self.register_buffer("count", torch.tensor(0)) + + def forward(self, value): + self.count.add_(1) + return value * self.weight.square() + + def callback(view): + head = view.module("head", LocalHead, checkpoint="student") + view.backward(head(torch.tensor(3.0))) + + asyncio.run(run_rank_callback(trainer, callback, mode="zero")) + native = trainer._checkpoint_slots["student"].custom["head"].value + torch.testing.assert_close(native.weight.grad, torch.tensor(12.0)) + assert native.count.item() == 1 + + +@pytest.mark.parametrize("enabled", [True, False]) +def test_logical_iterator_captures_ambient_grad_mode(enabled): + rank: Any = _Rank() + + def callback(view): + with torch.set_grad_enabled(enabled): + iterator = view.forward_batches([_input(3)]) + with torch.set_grad_enabled(not enabled): + result = next(iterator).outputs[0].hidden_states + assert result.requires_grad is enabled + iterator.close() + + asyncio.run(run_rank_callback(rank, callback, mode="zero")) + + +def test_persistent_iterator_binds_policy_and_checkpoint_and_pulls_one_wave(): + from art.trainer_rank import ForwardOptions + from art.trainer_rank._tensors import CotangentCollector + + class Rank(_Rank): + _capture_forward_options = TrainerRank._capture_forward_options + + def __init__(self): + super().__init__() + self._forward_options = ForwardOptions(allow_replay=False) + self._default_slot_ref = SimpleNamespace(name="original") + self.seen = [] + + def forward(self, tree, **kwargs): + if isinstance(tree, ForwardInput): + self.seen.append((tree.options, kwargs["checkpoint"])) + return super().forward(tree, **kwargs) + + def optim_step(self, **kwargs): + with torch.no_grad(): + self.weight.add_(1) + self.zero_grad() + + rank: Any = Rank() + + def run(callback): + return asyncio.run(run_rank_callback(rank, callback, mode="zero")).value + + first = _input(3) + first.options = ForwardOptions(max_gradient_staleness=0) + handle = run( + lambda view: view.open_forward_batches( + [first, _input(5), _input(7)], + options=ForwardOptions(allow_cpu_offload=False), + ) + ) + assert rank.seen == [] + first.options = ForwardOptions(max_gradient_staleness=4) + rank._forward_options = ForwardOptions(allow_replay=True) + rank._default_slot_ref = SimpleNamespace(name="later") + collector = CotangentCollector() + with torch.no_grad(): + packet = run(lambda view: view.export_forward(view.next_forward_batch(handle))) + batch = collector.attach(packet) + assert batch.indices == [0] + assert len(rank.seen) == 1 + policy, checkpoint = rank.seen[0] + assert checkpoint == "original" + assert policy.max_gradient_staleness == 0 + assert policy.allow_cpu_offload is False + assert policy.allow_replay is False + cotangents = collector.backward(_loss_tree(batch.outputs).square()) + run(lambda view: view.backward_packets(cotangents)) + assert rank.weight.grad.item() == 36 + run(lambda view: view.optim_step()) + packet = run(lambda view: view.export_forward(view.next_forward_batch(handle))) + batch = collector.attach(packet) + assert batch.indices == [1] + assert _loss_tree(batch.outputs).item() == 15 + cotangents = collector.backward(_loss_tree(batch.outputs)) + run(lambda view: view.backward_packets(cotangents)) + assert rank.weight.grad.item() == 5 + run(lambda view: view.close_forward_batches(handle)) + run(lambda view: view.close_forward_batches(handle)) + assert run(lambda view: view.next_forward_batch(handle)) is None + assert len(rank.seen) == 2 + assert rank.closed == 1 + assert not rank._rank_command_state.iterators + assert not rank._rank_command_state.batch_inputs + + +def test_persistent_iterator_captures_no_grad_before_next_callback(): + rank: Any = _Rank() + + def run(callback): + return asyncio.run(run_rank_callback(rank, callback, mode="zero")).value + + with torch.no_grad(): + handle = run(lambda view: view.open_forward_batches([_input(3)])) + batch = run(lambda view: view.next_forward_batch(handle)) + assert not batch.outputs[0].hidden_states.requires_grad + assert run(lambda view: view.next_forward_batch(handle)) is None + assert rank.closed == 1 + + +def test_nested_aggregate_outputs_admit_before_any_model_copy(): + from dataclasses import replace + + from art.trainer_rank import ForwardOptions + from art.trainer_rank._commands import _Executor, _OutputPacket, _view + from art.trainer_rank._tensors import detach_tree + + rank: Any = _Rank() + rank._available_memory_bytes = lambda: 1024 * 1024 + view = _view(_Executor(rank, "zero")) + request = ForwardInput( + input_tokens=torch.tensor([1]), options=ForwardOptions(output_device="auto") + ) + outputs = [ + _OutputPacket( + detach_tree( + f"zero:{index}:dp:{index}", + [[ForwardOutput(None, None, None, torch.ones(262144))]], + ), + (False,), + False, + ) + for index in range(2) + ] + planned = view._place_outputs([(output, [[request]]) for output in outputs]) + assert [output.cpu for output in planned] == [(False,), (True,)] + assert planned[1].managed + request = replace(request, options=ForwardOptions(output_device="model")) + with pytest.raises(MemoryError): + view._place_outputs([(output, [[request]]) for output in outputs]) diff --git a/tests/unit/test_trainer_rank_corrections.py b/tests/unit/test_trainer_rank_corrections.py new file mode 100644 index 000000000..6a6f85c6d --- /dev/null +++ b/tests/unit/test_trainer_rank_corrections.py @@ -0,0 +1,230 @@ +from __future__ import annotations + +import math + +import pytest +import torch + +from art.trainer_rank import ( + ForwardOutput, + ImportanceSamplingGradientCorrection, + ResolvedForwardOptions, + TopK, +) +from art.trainer_rank._corrections import capture_forward_corrections + + +def _fixture(policy="when_available", *, logits=False): + target = torch.tensor([-0.5, -1.0], requires_grad=True) + top_k = torch.tensor([[-0.5, -1.0]], requires_grad=True) + tokens = torch.tensor([[2, 0]]) + hidden = torch.zeros(2, 3, requires_grad=True) + logits_tensor = torch.zeros(1, 4, requires_grad=True) if logits else None + output = ForwardOutput(target, TopK(top_k, tokens), logits_tensor, hidden) + # Deliberately different from dataclass order: caller defines flat indices. + tensors = (hidden, target, tokens, top_k) + ( + () if logits_tensor is None else (logits_tensor,) + ) + context = capture_forward_corrections( + {"nested": [[output]]}, + tensors, + ResolvedForwardOptions( + stale_gradient_corrections=( + ImportanceSamplingGradientCorrection(policy=policy), + ) + ), + ) + return context, tensors + + +def test_capture_owns_only_original_logprobs_and_ids_on_cpu() -> None: + context, tensors = _fixture() + assert len(context.tensors) == 3 + assert all( + tensor.device.type == "cpu" and not tensor.requires_grad + for tensor in context.tensors + ) + original = tuple(tensor.clone() for tensor in context.tensors) + with torch.no_grad(): + tensors[1].fill_(-10) + tensors[2].fill_(3) + tensors[3].fill_(-10) + for actual, expected in zip(context.tensors, original, strict=True): + torch.testing.assert_close(actual, expected) + + +def test_context_applies_only_logprob_gradients_and_preserves_unused_outputs() -> None: + context, tensors = _fixture() + gradients = (torch.ones_like(tensors[0]), torch.ones_like(tensors[1]), None, None) + current = (tensors[0], tensors[1] - 1, tensors[2], tensors[3]) + result = context.correct(gradients, current) + assert result[0] is gradients[0] + assert result[2:] == (None, None) + torch.testing.assert_close(result[1], torch.full_like(tensors[1], math.exp(-1.0))) + torch.testing.assert_close(gradients[1], torch.ones_like(tensors[1])) + assert context.requires_current(gradients) is False + assert context.correct(gradients)[1] is gradients[1] + + +def test_always_requires_data_only_for_active_eligible_outputs() -> None: + context, tensors = _fixture("always") + assert not context.requires_current((torch.ones_like(tensors[0]), None, None, None)) + hidden_grad = torch.ones_like(tensors[0]) + assert context.correct((hidden_grad, None, None, None))[0] is hidden_grad + assert context.requires_current((None, torch.ones_like(tensors[1]), None, None)) + with pytest.raises(RuntimeError, match="requires current logprobs"): + context.correct((None, torch.ones_like(tensors[1]), None, None)) + # Merely requesting an unused unsupported output does not require correction. + assert not context.requires_current((None, None, None, None)) + + +def test_reordered_top_k_is_aligned_by_original_token_ids() -> None: + context, tensors = _fixture() + current = (tensors[0], tensors[1], tensors[2].flip(-1), (tensors[3] - 1).flip(-1)) + result = context.correct((None, None, None, torch.ones_like(tensors[3])), current) + torch.testing.assert_close(result[3], torch.exp(torch.full_like(tensors[3], -1))) + + +@pytest.mark.parametrize("policy", ["when_available", "always"]) +def test_changed_top_k_membership_is_unavailable_without_original_id_logits( + policy, +) -> None: + context, tensors = _fixture(policy) + current = (tensors[0], tensors[1], torch.tensor([[1, 3]]), tensors[3]) + gradients = (None, None, None, torch.ones_like(tensors[3])) + if policy == "always": + with pytest.raises(RuntimeError, match="requires current logprobs"): + context.correct(gradients, current) + else: + assert context.correct(gradients, current)[3] is gradients[3] + + +def test_available_full_logits_correct_changed_top_k_without_another_forward() -> None: + context, tensors = _fixture("always", logits=True) + current_logits = torch.tensor([[1.0, 0.0, -0.5, 0.5]]) + current = ( + tensors[0], + tensors[1], + torch.tensor([[1, 3]]), + tensors[3], + current_logits, + ) + result = context.correct( + (None, None, None, torch.ones_like(tensors[3]), None), current + ) + expected = ( + current_logits.log_softmax(-1).gather(-1, tensors[2]) - tensors[3] + ).exp() + torch.testing.assert_close(result[3], expected) + + +def test_context_checks_packet_layout_and_does_not_partially_modify_gradients() -> None: + context, tensors = _fixture("always") + target_grad, top_grad = torch.ones_like(tensors[1]), torch.ones_like(tensors[3]) + with pytest.raises(ValueError, match="output count"): + context.correct((target_grad,)) + with pytest.raises(RuntimeError, match="requires current logprobs"): + context.correct( + (None, target_grad, None, top_grad), + (tensors[0], tensors[1] - 1, torch.tensor([[1, 3]]), tensors[3]), + ) + torch.testing.assert_close(target_grad, torch.ones_like(target_grad)) + torch.testing.assert_close(top_grad, torch.ones_like(top_grad)) + + +@pytest.mark.parametrize("with_logits", [False, True]) +def test_top_k_ties_do_not_assume_stable_membership(with_logits: bool) -> None: + context, tensors = _fixture("always", logits=with_logits) + tied_logits = torch.zeros(1, 4) + tied_logprobs = torch.full((1, 2), -math.log(4)) + current = (tensors[0], tensors[1], torch.tensor([[1, 3]]), tied_logprobs) + gradients = (None, None, None, torch.ones_like(tensors[3])) + if with_logits: + result = context.correct(gradients + (None,), current + (tied_logits,)) + torch.testing.assert_close(result[3], (tied_logprobs - tensors[3]).exp()) + else: + with pytest.raises(RuntimeError, match="requires current logprobs"): + context.correct(gradients, current) + + +def test_zero_cotangent_masks_do_not_require_defined_sampling_support() -> None: + from art.trainer_rank._corrections import correct_logprob_cotangent + + original = torch.tensor([-0.5, -float("inf"), float("nan")]) + current = torch.tensor([-1.5, -float("inf"), float("nan")]) + gradient = torch.tensor([1.0, 0.0, 0.0]) + corrected = correct_logprob_cotangent( + gradient, + original_logprobs=original, + current_logprobs=current, + correction=ImportanceSamplingGradientCorrection(policy="always"), + ) + torch.testing.assert_close(corrected, torch.tensor([math.exp(-1), 0, 0])) + torch.testing.assert_close(gradient, torch.tensor([1.0, 0.0, 0.0])) + + +def test_zero_cotangents_do_not_require_current_data_under_always() -> None: + context, tensors = _fixture("always") + zero = torch.zeros_like(tensors[1]) + gradients = (None, zero, None, None) + assert not context.requires_current(gradients) + assert context.correct(gradients)[1] is zero + + +def test_only_active_top_k_ids_require_current_membership() -> None: + context, tensors = _fixture("always") + current = (tensors[0], tensors[1], torch.tensor([[2, 3]]), tensors[3] - 1) + result = context.correct((None, None, None, torch.tensor([[1.0, 0.0]])), current) + torch.testing.assert_close(result[3], torch.tensor([[math.exp(-1), 0.0]])) + + +def test_bad_cotangent_shape_is_rejected_before_requesting_replay() -> None: + context, _ = _fixture("always") + with pytest.raises(ValueError, match="shape must match"): + context.requires_current((None, torch.ones(1), None, None)) + + +@pytest.mark.parametrize("corrections", [(), (ImportanceSamplingGradientCorrection(),)]) +def test_current_replay_rejects_changed_active_top_k_events_even_without_correction( + corrections, +) -> None: + values = torch.tensor([[-0.5, -1.0]], requires_grad=True) + tokens = torch.tensor([[2, 0]]) + output = ForwardOutput(None, TopK(values, tokens), None, None) + context = capture_forward_corrections( + output, + (values, tokens), + ResolvedForwardOptions(stale_gradient_corrections=corrections), + ) + gradients = (torch.ones_like(values), None) + context.validate_replay(gradients, (values - 1, tokens)) + for changed in (tokens.flip(-1), torch.tensor([[2, 3]])): + with pytest.raises(RuntimeError, match="changed active top-k token identities"): + context.validate_replay(gradients, (values - 1, changed)) + torch.testing.assert_close(gradients[0], torch.ones_like(values)) + + +def test_current_replay_ignores_inactive_top_k_event_changes() -> None: + context, tensors = _fixture("always") + current = (tensors[0], tensors[1], torch.tensor([[2, 3]]), tensors[3] - 1) + context.validate_replay((None, None, None, torch.tensor([[1.0, 0.0]])), current) + context.validate_replay((None, None, None, torch.zeros_like(tensors[3])), current) + context.validate_replay((None, None, None, None), current) + + +def test_current_ratio_evaluation_can_reorder_but_physical_replay_cannot() -> None: + context, tensors = _fixture("always", logits=True) + gradients = (None, None, None, torch.ones_like(tensors[3]), None) + current = ( + tensors[0], + tensors[1], + tensors[2].flip(-1), + (tensors[3] - 1).flip(-1), + tensors[4], + ) + torch.testing.assert_close( + context.correct(gradients, current)[3], + torch.exp(torch.full_like(tensors[3], -1)), + ) + with pytest.raises(RuntimeError, match="changed active top-k token identities"): + context.validate_replay(gradients, current) diff --git a/tests/unit/test_trainer_rank_custom_tensors.py b/tests/unit/test_trainer_rank_custom_tensors.py index d676e8f6f..b46811dc7 100644 --- a/tests/unit/test_trainer_rank_custom_tensors.py +++ b/tests/unit/test_trainer_rank_custom_tensors.py @@ -21,6 +21,7 @@ from art.trainer_rank import ( AdamParams, MaterializedCheckpoint, + ModuleHandle, TrainerRank, TrainerRankSlotStateError, Unset, @@ -44,7 +45,7 @@ def module( factory: Callable[[], ModuleT], *, checkpoint: str | object = Unset, - ) -> ModuleT: ... + ) -> ModuleHandle: ... def parameter( self, @@ -252,7 +253,8 @@ def _distributed_custom_grad_flags_worker( used = api.parameter("used", lambda: torch.tensor(1.0), checkpoint="student") api.parameter("unused", lambda: torch.tensor(2.0), checkpoint="student") if rank == 0: - (used * 3).backward() + with trainer._gradient_transaction(): + (used * 3).backward() assert trainer._dynamic_param_step_flags( trainer._checkpoint_slots["student"].params ) == (True, False) @@ -279,8 +281,8 @@ def value_head() -> _ValueHead: class_head = rank.module("class_head", CountingHead, checkpoint="student") lambda_head = rank.module("lambda_head", value_head, checkpoint="student") - assert isinstance(class_head, CountingHead) - assert isinstance(lambda_head, _ValueHead) + assert isinstance(class_head, torch.nn.Module) + assert isinstance(lambda_head, torch.nn.Module) assert rank.module("class_head", CountingHead, checkpoint="student") is class_head assert rank.module("lambda_head", value_head, checkpoint="student") is lambda_head assert class_calls == 1 @@ -400,7 +402,8 @@ def test_custom_module_outputs_participate_in_checkpoint_graph_guards() -> None: with pytest.raises(TrainerRankSlotStateError, match="live backward graph"): trainer._guard_slot_can_load(ref) - output.sum().backward() + with trainer._gradient_transaction(): + output.sum().backward() trainer.zero_grad() trainer._guard_slot_can_load(ref) @@ -423,7 +426,8 @@ def test_custom_module_tracks_direct_parameter_use_and_custom_outputs( assert isinstance(output, _HeadOutput) with pytest.raises(TrainerRankSlotStateError, match="live backward graph"): trainer._guard_slot_can_load(trainer._slot_ref("student")) - output.value.sum().backward() + with trainer._gradient_transaction(): + output.value.sum().backward() trainer.zero_grad() trainer._guard_slot_can_load(trainer._slot_ref("student")) @@ -439,7 +443,8 @@ def test_custom_module_tracks_direct_weight_operations() -> None: with pytest.raises(TrainerRankSlotStateError, match="live backward graph"): trainer._guard_slot_can_load(trainer._slot_ref("student")) - loss.backward() + with trainer._gradient_transaction(): + loss.backward() trainer.zero_grad() trainer._guard_slot_can_load(trainer._slot_ref("student")) @@ -454,11 +459,13 @@ def test_custom_parameter_graph_guard_tracks_retain_and_abandonment() -> None: ref = trainer._slot_ref("student") loss = parameter.square().sum() - loss.backward(retain_graph=True) + with trainer._gradient_transaction(): + loss.backward(retain_graph=True) trainer.zero_grad() with pytest.raises(TrainerRankSlotStateError, match="live backward graph"): trainer._guard_slot_can_load(ref) - loss.backward() + with trainer._gradient_transaction(): + loss.backward() trainer.zero_grad() trainer._guard_slot_can_load(ref) @@ -480,7 +487,8 @@ def test_custom_module_graph_tracking_allows_in_place_layers() -> None: ) output = head(torch.ones(1, 3, requires_grad=True)) - output.sum().backward() + with _trainer_rank._gradient_transaction(): + output.sum().backward() assert head[0].weight.grad is not None @@ -496,7 +504,8 @@ def test_custom_parameter_outputs_participate_in_checkpoint_graph_guards() -> No with pytest.raises(TrainerRankSlotStateError, match="live backward graph"): trainer._guard_slot_can_load(ref) - output.backward() + with trainer._gradient_transaction(): + output.backward() trainer.zero_grad() trainer._guard_slot_can_load(ref) @@ -507,7 +516,8 @@ def test_checkpoint_load_rejects_custom_grads_and_stale_objects() -> None: parameter = rank.parameter( "temperature", lambda: torch.tensor(2.0), checkpoint="student" ) - (head(torch.ones(1)) + parameter).sum().backward() + with trainer._gradient_transaction(): + (head(torch.ones(1)) + parameter).sum().backward() ref = trainer._slot_ref("student") with pytest.raises(TrainerRankSlotStateError, match="accumulated gradients"): @@ -651,7 +661,8 @@ def test_selected_optimizer_step_updates_only_its_custom_checkpoint( for param in params ), ) - (a * 3).backward() + with trainer._gradient_transaction(): + (a * 3).backward() before_b = b.detach().clone() trainer.optim_step( params=AdamParams(learning_rate=1e-2, weight_decay=0.0), @@ -677,7 +688,8 @@ def test_optimizer_skips_unused_custom_parameters( for param in params ), ) - (used * 2).backward() + with trainer._gradient_transaction(): + (used * 2).backward() trainer.optim_step( params=AdamParams(learning_rate=0.1, weight_decay=0.5), @@ -959,7 +971,8 @@ def test_custom_tensor_names_cannot_corrupt_lora_optimizer_metadata( for param in params ), ) - head.q_proj.lora_A.weight.sum().backward() + with trainer._gradient_transaction(): + head.q_proj.lora_A.weight.sum().backward() trainer.optim_step( params=AdamParams(learning_rate=1e-3, weight_decay=0.0), checkpoints=["student"], @@ -1271,7 +1284,7 @@ def _empty_real_lora_trainer() -> tuple[TrainerRank, _CustomTensorAPI]: def _register_custom_tensors( rank: _CustomTensorAPI, -) -> tuple[_ValueHead, torch.nn.Parameter, torch.Tensor]: +) -> tuple[ModuleHandle, torch.nn.Parameter, torch.Tensor]: head = rank.module("value_head", lambda: _ValueHead(3), checkpoint="student") temperature = rank.parameter( "temperature", lambda: torch.tensor(0.5), checkpoint="student" @@ -1284,7 +1297,7 @@ def _register_custom_tensors( def _step_custom_tensors( trainer: TrainerRank, - head: _ValueHead, + head: ModuleHandle, temperature: torch.nn.Parameter, monkeypatch: pytest.MonkeyPatch, *, @@ -1301,7 +1314,8 @@ def _step_custom_tensors( ), ) hidden = torch.tensor([[0.25, -0.5, 1.0]]) - (head(hidden).sum() + temperature * scale).backward() + with trainer._gradient_transaction(): + (head(hidden).sum() + temperature * scale).backward() trainer.optim_step( params=AdamParams( learning_rate=1e-3, diff --git a/tests/unit/test_trainer_rank_graph_order.py b/tests/unit/test_trainer_rank_graph_order.py new file mode 100644 index 000000000..866654135 --- /dev/null +++ b/tests/unit/test_trainer_rank_graph_order.py @@ -0,0 +1,116 @@ +"""Physical peers must enter cached/replayed backward collectives identically.""" + +from datetime import timedelta +from types import SimpleNamespace + +import pytest +import torch +import torch.distributed as dist +import torch.multiprocessing as mp + +from art.trainer_rank import TrainerRank, _graphs +from art.trainer_rank._impl import _CheckpointSlot + + +class _Collective(torch.autograd.Function): + @staticmethod + def forward(ctx, value, tag): + ctx.tag = tag + return value.clone() + + @staticmethod + def backward(ctx, *gradients): + tag = torch.tensor([ctx.tag]) + tags = [torch.empty_like(tag) for _ in range(dist.get_world_size())] + dist.all_gather(tags, tag) + assert all(item.item() == ctx.tag for item in tags), ( + "different physical graph order" + ) + gradient = gradients[0].clone() + dist.all_reduce(gradient) + return gradient, None + + +def _worker(rank, rendezvous, fail_replay): + dist.init_process_group( + "gloo", + init_method=rendezvous, + rank=rank, + world_size=2, + timeout=timedelta(seconds=30), + ) + try: + names = iter(("z", "a") if rank == 0 else ("a", "z")) + setattr(_graphs, "uuid4", lambda: SimpleNamespace(hex=next(names))) + cache = _graphs.GraphCache() + parameters, packets = [], [] + trainer = TrainerRank.__new__(TrainerRank) + trainer._checkpoint_slots = {"student": _CheckpointSlot(params=())} + trainer.runtime = SimpleNamespace(model=[], optimizer=None) + for tag in range(2): + parameter = torch.nn.Parameter(torch.tensor(2.0)) + parameters.append(parameter) + trainer._checkpoint_slots["student"].params = tuple(parameters) + snapshot = trainer._snapshot_parameter( + parameter, trainer._capture_checkpoint_version("student") + ) + calls = [0] + + def execute(x, snapshot=snapshot, tag=tag, calls=calls): + calls[0] += 1 + output = _Collective.apply(snapshot * x, tag) + if fail_replay and rank == 0 and tag == 1 and calls[0] > 1: + output = output.expand(2) + return (output,) + + handle, _ = cache.run( + execute, + torch.tensor(3.0), + retention="replay", + ) + packets.append((handle, (torch.tensor(1.0),))) + + def coordinate(function): + result, error = None, None + try: + result = function() + except Exception as exc: + error = str(exc) + errors = [None, None] + dist.all_gather_object(errors, error) + if any(errors): + raise RuntimeError(str(errors)) + return result + + for parameter in parameters: + parameter.grad = torch.tensor(7.0) + + def backward(): + with trainer._gradient_transaction(before_commit=coordinate): + cache.backward_many(sorted(packets), coordinate=coordinate) + + if fail_replay: + with pytest.raises(RuntimeError, match="metadata differs"): + backward() + assert not trainer._version_state()._origins + else: + backward() + for parameter in parameters: + torch.testing.assert_close( + parameter.grad, torch.tensor(7.0 if fail_replay else 13.0) + ) + assert not cache.handles() + completed = torch.tensor(1) + dist.all_reduce(completed) + assert completed.item() == 2 + finally: + dist.destroy_process_group() + + +@pytest.mark.parametrize("fail_replay", [False, True]) +def test_backward_collectives_follow_creation_order_despite_different_handles( + tmp_path, fail_replay +): + mp.spawn( + _worker, args=(f"file://{tmp_path / 'order'}", fail_replay), nprocs=2, join=True + ) diff --git a/tests/unit/test_trainer_rank_graphs.py b/tests/unit/test_trainer_rank_graphs.py new file mode 100644 index 000000000..44bf1e721 --- /dev/null +++ b/tests/unit/test_trainer_rank_graphs.py @@ -0,0 +1,525 @@ +from __future__ import annotations + +import gc +from types import SimpleNamespace + +import pytest +import torch +from torch.utils.checkpoint import checkpoint + +from art.trainer_rank._graphs import GraphCache + + +def test_quantized_te_retained_backward_rejects_before_destructive_call(): + from art.megatron.compile_workarounds import _preserve_te_backward_metadata + + calls = [] + + class QuantizedSavedState(torch.autograd.Function): + @staticmethod + def forward(ctx, weight): + ctx.tensor_objects = [object()] + return weight.square() + + @staticmethod + def backward(ctx, *gradients): + calls.append(True) + return gradients[0] + + _preserve_te_backward_metadata(QuantizedSavedState) + parameter = torch.nn.Parameter(torch.tensor(2.0)) + cache = GraphCache() + handle, (output,) = cache.run(lambda _: (QuantizedSavedState.apply(parameter),), ()) + with pytest.raises(RuntimeError, match="quantized saved tensors"): + cache.backward(handle, (torch.ones_like(output),), retain_graph=True) + assert calls == [] + assert parameter.grad is None + assert cache.handles() == () + + +def test_retained_unpack_preserves_saved_view_when_consumer_clears_data(): + seen = [] + + class ClearsSavedTensor(torch.autograd.Function): + @staticmethod + def forward(ctx, weight, x): + ctx.save_for_backward(x.t()[1:, ::2]) + return weight * x.t()[1:, ::2].sum() + + @staticmethod + def backward(ctx, *cotangents): + (value,) = ctx.saved_tensors + seen.append((value.stride(), value.storage_offset(), value.data_ptr())) + result = cotangents[0] * value.sum() + value.data = torch.empty(0) + return result, None + + weight = torch.nn.Parameter(torch.tensor(2.0)) + inputs = torch.arange(24.0).reshape(4, 6) + cache = GraphCache() + handle, (output,) = cache.run( + lambda x: (ClearsSavedTensor.apply(weight, x),), inputs + ) + for retain in (True, True, False): + weight.grad = None + cache.backward(handle, (torch.ones_like(output),), retain_graph=retain) + torch.testing.assert_close(weight.grad, inputs.t()[1:, ::2].sum()) + assert seen[0] == seen[1] == seen[2] + assert seen[0][:2] == (inputs.t()[1:, ::2].stride(), 1) + assert cache.handles() == () + + +@pytest.mark.parametrize("retention", ["gpu", "cpu", "replay"]) +@pytest.mark.parametrize("recompute", [False, True]) +def test_original_inputs_weights_and_dropout_survive_update(retention, recompute): + torch.manual_seed(7) + inputs = torch.randn(13, 5, dtype=torch.float64) + weight = torch.nn.Parameter(torch.randn(5, 3, dtype=torch.float64)) + historical = torch.nn.Parameter(weight.detach().clone()) + reference = torch.nn.Parameter(weight.detach().clone()) + executions = [] + + def model(x, w): + return torch.nn.functional.dropout(x @ w.square(), 0.3, training=True).sin() + + def execute(captured): + executions.append(True) + y = ( + checkpoint(lambda x: model(x, historical), captured, use_reentrant=False) + if recompute + else model(captured, historical) + ) + return y, y.square(), torch.ones(2, dtype=torch.long) + + rng = torch.get_rng_state() + cache = GraphCache() + handle, (y, unused, tokens) = cache.run(execute, inputs, retention=retention) + torch.set_rng_state(rng) + expected = model(inputs.clone(), reference) + torch.testing.assert_close(y, expected) + # The caller receives leaves without a path back into the model graph. + assert y.is_leaf and y.grad_fn is None + assert unused.is_leaf and not tokens.requires_grad + inputs.fill_(100) + with torch.no_grad(): + weight.add_(9) + ambient = torch.get_rng_state() + (expected.cos().sum()).backward() + cache.backward(handle, (-y.sin(), None, None)) + torch.testing.assert_close(historical.grad, reference.grad) + assert torch.equal(torch.get_rng_state(), ambient) + assert len(executions) == (2 if retention == "replay" else 1) + assert cache.handles() == () + + +def test_evict_frees_saved_state_while_caller_output_remains(): + cache = GraphCache() + weight = torch.nn.Parameter(torch.tensor(2.0)) + + def execute(x): + activation = x.sin() + return (activation * weight,) + + handle, (output,) = cache.run(execute, torch.arange(10.0)) + references = tuple(cache._records[handle].saved or ()) + assert any(reference() is not None for reference in references) + cache.evict(handle) + gc.collect() + assert all(reference() is None for reference in references) + assert output.shape == (10,) + cache.backward(handle, (torch.ones_like(output),)) + torch.testing.assert_close(weight.grad, torch.arange(10.0).sin().sum()) + + +def test_preflight_all_records_before_any_replay_or_gradient_mutation(): + cache = GraphCache() + parameter = torch.nn.Parameter(torch.tensor(2.0)) + executions = [] + + def execute(x): + executions.append(True) + return (x * parameter,) + + def stale(): + raise RuntimeError("stale original version") + + first, _ = cache.run(execute, torch.tensor(3.0), retention="replay") + second, _ = cache.run(execute, torch.tensor(4.0), validate_backward=stale) + with pytest.raises(RuntimeError, match="stale original"): + cache.backward_many( + ((first, (torch.tensor(1.0),)), (second, (torch.tensor(1.0),))) + ) + assert len(executions) == 2 + assert parameter.grad is None + assert len(cache.handles()) == 2 + + +def test_retain_graph_replay_preserves_origin_and_does_not_keep_replayed_graph(): + cache = GraphCache() + parameter = torch.nn.Parameter(torch.tensor(2.0)) + validated = [] + handle, _ = cache.run( + lambda x: (x * parameter,), + torch.tensor(3.0), + retention="replay", + validate_backward=lambda: validated.append(0), + ) + cache.backward(handle, (torch.tensor(1.0),), retain_graph=True) + assert cache.state(handle).retention == "replay" + assert cache.state(handle).replay_count == 1 + cache.backward(handle, (torch.tensor(2.0),)) + assert parameter.grad is not None and parameter.grad.item() == 9 + assert validated == [0, 0] + + +def test_unused_graph_does_not_replay(): + cache = GraphCache() + executions = [] + parameter = torch.nn.Parameter(torch.tensor(2.0)) + + def execute(x): + executions.append(True) + return (x * parameter,) + + handle, _ = cache.run(execute, torch.tensor(3.0), retention="replay") + cache.backward(handle, (None,)) + assert executions == [True] + assert parameter.grad is None + + +@pytest.mark.parametrize( + "retention,options", + [ + ("cpu", SimpleNamespace(allow_cpu_offload=False)), + ("replay", SimpleNamespace(allow_replay=False)), + ], +) +def test_disabled_policy_rejects_before_execute(retention, options): + cache = GraphCache() + with pytest.raises(ValueError, match="disabled"): + cache.run( + lambda x: pytest.fail("should not execute"), + None, + retention=retention, + options=options, + ) + + +def test_bad_cotangent_rejects_before_replay(): + cache = GraphCache() + parameter = torch.nn.Parameter(torch.tensor(2.0)) + handle, _ = cache.run( + lambda x: (x * parameter,), torch.tensor(3.0), retention="replay" + ) + with pytest.raises(ValueError, match="mismatch"): + cache.backward(handle, (torch.ones(2),)) + assert cache.state(handle).replay_count == 0 + assert parameter.grad is None + + +def test_replay_backward_releases_each_child_before_replaying_next(): + cache = GraphCache() + parameter = torch.nn.Parameter(torch.tensor(2.0)) + handles = [] + replaying = False + + def execute(x): + if replaying and x.item() == 4: + assert handles[0] not in cache.handles() + return (x * parameter,) + + for value in (3.0, 4.0): + handle, _ = cache.run(execute, torch.tensor(value), retention="replay") + handles.append(handle) + replaying = True + cache.backward_many([(handle, (torch.tensor(1.0),)) for handle in handles]) + assert parameter.grad is not None and parameter.grad.item() == 7 + + +def test_tracker_rng_and_ambient_state_are_restored(): + tracker = SimpleNamespace(states={"stream": torch.tensor([17])}) + tracker.get_states = lambda: tracker.states + tracker.set_states = lambda value: setattr(tracker, "states", value) + seen = [] + parameter = torch.nn.Parameter(torch.tensor(2.0)) + + def execute(x): + seen.append(tracker.states["stream"].item()) + tracker.states["stream"].add_(1) + return (x * parameter,) + + cache = GraphCache() + handle, _ = cache.run( + execute, torch.tensor(3.0), retention="replay", rng_tracker=tracker + ) + tracker.states["stream"].fill_(99) + cache.backward(handle, (torch.tensor(1.0),)) + assert seen == [17, 17] + assert tracker.states["stream"].item() == 99 + + +def test_backward_uses_creation_order_not_wire_handle_order(): + cache = GraphCache() + seen, packets = [], [] + for index in range(3): + parameter = torch.nn.Parameter(torch.tensor(2.0)) + parameter.register_hook(lambda gradient, index=index: seen.append(index)) + handle, _ = cache.run( + lambda x, parameter=parameter: (parameter * x,), torch.tensor(3.0) + ) + packets.append((handle, (torch.tensor(1.0),))) + cache.backward_many(list(reversed(packets))) + assert seen == [0, 1, 2] + + +def _corrected_cache(*, retention="gpu", policy="when_available", stale=True): + from contextlib import contextmanager + + from art.trainer_rank import ForwardOutput + from art.trainer_rank._corrections import capture_forward_corrections + from art.trainer_rank._options import ( + ImportanceSamplingGradientCorrection, + ResolvedForwardOptions, + ) + + cache = GraphCache() + original = torch.nn.Parameter(torch.tensor(-1.0)) + current = torch.nn.Parameter(torch.tensor(-0.5)) + selected = [original] + executions = [] + + @contextmanager + def use(parameter): + previous, selected[0] = selected[0], parameter + try: + yield + finally: + selected[0] = previous + + def execute(x): + executions.append((selected[0] is current, torch.is_grad_enabled())) + return (selected[0].square() * x,) + + handle, outputs = cache.run(execute, torch.tensor(-1.0), retention=retention) + context = capture_forward_corrections( + ForwardOutput(outputs[0], None, None, None), + outputs, + ResolvedForwardOptions( + stale_gradient_corrections=( + ImportanceSamplingGradientCorrection(policy=policy), + ) + ), + ) + cache.set_corrections( + handle, + context, + is_stale=lambda: stale, + current_context_factory=lambda: use(current), + ) + return cache, handle, original, current, executions, context + + +@pytest.mark.parametrize("retention", ["gpu", "replay"]) +def test_opportunistic_exact_replay_never_adds_correction_forward(retention): + cache, handle, original, current, executions, _ = _corrected_cache( + retention=retention + ) + cache.backward(handle, (torch.tensor(1.0),)) + assert len(executions) == (1 if retention == "gpu" else 2) + assert original.grad is not None and original.grad.item() == 2 + assert current.grad is None + + +def test_always_correction_uses_no_grad_current_evaluation_and_old_jacobian(): + cache, handle, original, current, executions, _ = _corrected_cache(policy="always") + cache.backward(handle, (torch.tensor(1.0),)) + torch.testing.assert_close( + original.grad, torch.tensor(2.0) * torch.tensor(0.75).exp() + ) + assert current.grad is None + assert executions == [(False, True), (True, False)] + + +def test_newer_replay_opportunistically_corrects_current_jacobian(): + cache, handle, original, current, executions, _ = _corrected_cache() + cache.evict(handle, replay_with_current=True) + cache.backward(handle, (torch.tensor(1.0),)) + assert original.grad is None + torch.testing.assert_close(current.grad, torch.tensor(0.75).exp()) + assert executions == [(False, True), (True, True)] + + +@pytest.mark.parametrize("corrections", [False, True]) +def test_current_replay_rejects_changed_selected_token_events(corrections): + from contextlib import nullcontext + + from art.trainer_rank import ForwardOutput, TopK + from art.trainer_rank._corrections import capture_forward_corrections + from art.trainer_rank._options import ( + ImportanceSamplingGradientCorrection, + ResolvedForwardOptions, + ) + + cache = GraphCache() + parameter = torch.nn.Parameter(torch.tensor([1.0, 2.0])) + tokens = [torch.tensor([0, 1])] + + def execute(_): + return parameter.log_softmax(-1)[tokens[0]], tokens[0] + + handle, outputs = cache.run(execute, None) + context = capture_forward_corrections( + ForwardOutput(None, TopK(outputs[0], outputs[1]), None, None), + outputs, + ResolvedForwardOptions( + stale_gradient_corrections=(ImportanceSamplingGradientCorrection(),) + if corrections + else (), + ), + ) + cache.set_corrections( + handle, context, is_stale=lambda: True, current_context_factory=nullcontext + ) + cache.evict(handle, replay_with_current=True) + tokens[0] = torch.tensor([1, 0]) + with pytest.raises(RuntimeError, match="token identit"): + cache.backward(handle, (torch.ones(2), None)) + assert parameter.grad is None + assert not cache.handles() + + +@pytest.mark.parametrize("checkpointing", [False, True]) +def test_abandoned_release_frees_physical_record_without_autograd(checkpointing): + import gc + import weakref + + from torch.utils.checkpoint import checkpoint + + cache = GraphCache() + parameter = torch.nn.Parameter(torch.randn(4, 4)) + + def execute(x): + def compute(value): + return (parameter @ value).sin() + + return ( + checkpoint(compute, x, use_reentrant=False) + if checkpointing + else compute(x), + ) + + handle, outputs = cache.run(execute, torch.ones(4, 4)) + record = weakref.ref(cache._records[handle]) + original = cache._records[handle].outputs + assert original is not None + physical = weakref.ref(original[0]) + del original + cache.release(handle) + gc.collect() + assert record() is None and physical() is None + assert outputs[0].requires_grad + + +def test_backward_releases_unused_differentiable_output_branches(): + import gc + import weakref + + cache = GraphCache() + parameter = torch.nn.Parameter(torch.randn(4, 4)) + handle, outputs = cache.run( + lambda x: ((parameter @ x).sin().sum(), (parameter @ x).cos()), + torch.ones(4, 4), + ) + record = weakref.ref(cache._records[handle]) + cache.backward(handle, (torch.ones_like(outputs[0]), None)) + gc.collect() + assert record() is None + assert outputs[1].requires_grad + + +def test_restore_workspace_estimate_survives_eviction(): + cache = GraphCache() + parameter = torch.nn.Parameter(torch.tensor(2.0)) + handle, _ = cache.run( + lambda x: (parameter * x,), + torch.tensor(3.0), + execution_peak_bytes=1024, + checkpoint_versions=("original",), + ) + cache.evict(handle) + assert cache.state(handle).restore_workspace_bytes == 1024 + assert cache.state(handle).checkpoint_versions == ("original",) + cache.release(handle) + + +def test_all_correction_availability_preflights_before_any_replay(): + cache, first, original, _, executions, context = _corrected_cache( + retention="replay", policy="always" + ) + second, _ = cache.run( + lambda x: (original * x,), torch.tensor(-1.0), retention="replay" + ) + cache.set_corrections(second, context, is_stale=lambda: True) + with pytest.raises(RuntimeError, match="no current version context"): + cache.backward_many( + [(first, (torch.tensor(1.0),)), (second, (torch.tensor(1.0),))] + ) + assert executions == [(False, True)] + assert original.grad is None + + +def test_fresh_always_does_not_require_current_forward(): + cache, handle, original, current, executions, context = _corrected_cache( + policy="always", stale=False + ) + cache.set_corrections(handle, context, is_stale=lambda: False) + cache.backward(handle, (torch.tensor(1.0),)) + assert original.grad is not None and original.grad.item() == 2 + assert current.grad is None and len(executions) == 1 + + +def test_replay_restores_original_autocast_context(): + cache = GraphCache() + parameter = torch.nn.Parameter(torch.ones(3, 3)) + with torch.autocast("cpu", dtype=torch.bfloat16): + handle, (output,) = cache.run( + lambda x: (x @ parameter,), torch.ones(2, 3), retention="replay" + ) + assert output.dtype == torch.bfloat16 + cache.backward(handle, (torch.ones_like(output),)) + torch.testing.assert_close(parameter.grad, torch.full_like(parameter, 2.0)) + + +def test_replay_failure_discards_transaction_and_releases_participating_records(): + from art.trainer_rank import TrainerRank + from art.trainer_rank._impl import _CheckpointSlot + + trainer = TrainerRank.__new__(TrainerRank) + parameter = torch.nn.Parameter(torch.tensor(2.0)) + trainer._checkpoint_slots = {"student": _CheckpointSlot(params=(parameter,))} + trainer.runtime = SimpleNamespace(model=[], optimizer=None) + snapshot = trainer._snapshot_parameter( + parameter, trainer._capture_checkpoint_version("student") + ) + cache = GraphCache() + executions = [0] + + def failing(x): + executions[0] += 1 + result = snapshot * x + return (result if executions[0] == 1 else result.expand(2),) + + first, _ = cache.run( + lambda x: (snapshot * x,), torch.tensor(3.0), retention="replay" + ) + second, _ = cache.run(failing, torch.tensor(4.0), retention="replay") + parameter.grad = torch.tensor(7.0) + with pytest.raises(RuntimeError, match="metadata differs"): + with trainer._gradient_transaction(): + cache.backward_many( + [(first, (torch.tensor(1.0),)), (second, (torch.tensor(1.0),))] + ) + assert parameter.grad.item() == 7 + assert snapshot.grad is None + assert not trainer._version_state()._origins + assert cache.handles() == () diff --git a/tests/unit/test_trainer_rank_graphs_cuda.py b/tests/unit/test_trainer_rank_graphs_cuda.py new file mode 100644 index 000000000..811a7014b --- /dev/null +++ b/tests/unit/test_trainer_rank_graphs_cuda.py @@ -0,0 +1,596 @@ +"""Opt-in cache memory oracle; run on the validation lane's reserved GPU.""" + +import gc +import os +from types import SimpleNamespace +from typing import Any, cast +import weakref + +import pytest +import torch +from torch.multiprocessing.reductions import StorageWeakRef + +from art.trainer_rank._graphs import GraphCache + +pytestmark = pytest.mark.skipif( + os.environ.get("ART_GRAPH_GPU_TEST") != "1", reason="requires a reserved GPU" +) + + +@pytest.fixture +def cyclic_gc_disabled(): + enabled = gc.isenabled() + gc.disable() + try: + yield + finally: + if enabled: + gc.enable() + + +@pytest.mark.usefixtures("cyclic_gc_disabled") +@pytest.mark.parametrize("finish", ["backward", "release", "evict"]) +@pytest.mark.parametrize("retention", ["gpu", "cpu", "replay"]) +@pytest.mark.parametrize("multiple_hooks", [False, True]) +def test_selective_checkpoint_recomputes_each_retained_backward_and_releases( + finish, retention, multiple_hooks +): + from megatron.core.tensor_parallel.random import CheckpointWithoutOutput + + from art.megatron.compile_workarounds import install_reusable_checkpoint_backward + + install_reusable_checkpoint_backward() + inputs = torch.linspace(0.2, 0.8, 8, device="cuda") + weight = torch.nn.Parameter(torch.linspace(0.1, 0.7, 8, device="cuda")) + calls = [] + owners = [] + physical_outputs = [] + + def block(x): + calls.append(True) + return x.sin() + + def execute(x): + checkpoint = CheckpointWithoutOutput(fp8=None) + # Native TransformerLayer keeps this controller on a live module. + owners.append(checkpoint) + hidden = checkpoint.checkpoint(block, x * weight) + physical_outputs.append(weakref.ref(hidden)) + output = hidden.square() + extra = hidden * 2 if multiple_hooks else None + checkpoint.discard_output_and_register_recompute(output) + if extra is not None: + checkpoint.discard_output_and_register_recompute(extra) + output = output + extra + return (output,) + + cache = GraphCache() + handle, (output,) = cache.run( + execute, + inputs, + retention=retention, + options=SimpleNamespace( + allow_replay=retention == "replay" or finish == "evict" + ), + ) + z = inputs * weight.detach() + expected = 2 * z.sin() * z.cos() * inputs + if multiple_hooks: + expected += 2 * z.cos() * inputs + consume = finish == "backward" + for retain in (True, True, False) if consume else (True,): + weight.grad = None + cache.backward(handle, (torch.ones_like(output),), retain_graph=retain) + torch.testing.assert_close(weight.grad, expected) + backwards = 3 if consume else 1 + assert len(calls) == 1 + backwards * (2 if retention == "replay" else 1) + if finish == "evict": + cache.evict(handle) + assert all(ref() is None for ref in physical_outputs) + cache.release(handle) + assert all(ref() is None for ref in physical_outputs) + assert all( + getattr(owner, field) is None + for owner in owners + for field in ("run_function", "rng_states", "outputs", "ctx") + ) + assert cache.handles() == () + + +@pytest.mark.parametrize("no_grad", [False, True]) +def test_selective_checkpoint_failure_and_no_grad_clear_module_owner(no_grad): + from megatron.core.tensor_parallel.random import CheckpointWithoutOutput + + from art.megatron.compile_workarounds import install_reusable_checkpoint_backward + + install_reusable_checkpoint_backward() + checkpoint = CheckpointWithoutOutput(fp8=None) + weight = torch.nn.Parameter(torch.ones(8, device="cuda")) + calls = [] + + def block(x): + calls.append(True) + if len(calls) > 1: + raise RuntimeError("injected selective recompute failure") + return x.sin() + + def execute(x): + with torch.set_grad_enabled(not no_grad): + hidden = checkpoint.checkpoint(block, x * weight) + output = hidden.square() + checkpoint.discard_output_and_register_recompute(output) + if no_grad: + checkpoint.discard_output_and_register_recompute(output) + return (output,) + + cache = GraphCache() + handle, (output,) = cache.run(execute, torch.ones_like(weight)) + owner = getattr(checkpoint, "_art_recompute_owner") + if no_grad: + assert owner() is None + cache.release(handle) + else: + with pytest.raises(RuntimeError, match="injected selective recompute failure"): + cache.backward(handle, (torch.ones_like(output),), retain_graph=True) + if owner() is not None: + assert owner().ctx is None + assert owner().run_function is None + assert checkpoint.ctx is None + assert checkpoint.outputs is None + assert checkpoint.run_function is None + assert checkpoint.rng_states is None + assert cache.handles() == () + + +@pytest.mark.parametrize("retention", ["gpu", "cpu", "replay"]) +def test_gpu_saved_state_offload_eviction_and_replay(retention): + torch.manual_seed(7) + cache = GraphCache() + weight = torch.nn.Parameter(torch.randn(2048, device="cuda")) + original = torch.randn(2048, 2048, device="cuda") + reference = torch.nn.Parameter(weight.detach().clone()) + expected = (original.sin() * reference).sum(1) + expected.sum().backward() + storage_key = weight.untyped_storage().data_ptr() + + def execute(x): + return ((x.sin() * weight).sum(1),) + + handle, (output,) = cache.run( + execute, + original, + retention=retention, + cuda_devices=[torch.cuda.current_device()], + execution_peak_bytes=3 * original.numel() * original.element_size(), + keep_on_device=lambda tensor: ( + tensor.untyped_storage().data_ptr() == storage_key + ), + ) + torch.testing.assert_close(output, expected) + if retention == "gpu": + torch.cuda.synchronize() + before = torch.cuda.memory_allocated() + state = cache.state(handle) + assert state.offload_bytes >= original.numel() * original.element_size() + cache.offload(handle) + gc.collect() + torch.cuda.synchronize() + after = torch.cuda.memory_allocated() + assert before - after >= original.numel() * original.element_size() + assert cache.state(handle).gpu_bytes <= output.numel() * output.element_size() + cache.evict(handle) + assert cache.state(handle).gpu_bytes == 0 + assert cache.state(handle).restore_workspace_bytes == ( + 3 * original.numel() * original.element_size() + ) + original.fill_(100) + ambient = torch.cuda.get_rng_state() + cache.backward(handle, (torch.ones_like(output),)) + torch.testing.assert_close(weight.grad, reference.grad) + assert torch.equal(torch.cuda.get_rng_state(), ambient) + + +@pytest.mark.parametrize( + ("retention", "cache_state"), + [ + ("gpu", "cold"), + ("gpu", "warm"), + ("gpu", "artifact"), + ("cpu", "artifact"), + ("replay", "artifact"), + ], +) +def test_compiled_retained_backward_preserves_original_gradient_after_update( + retention, cache_state, tmp_path, monkeypatch +): + from torch._dynamo.utils import counters + from torch._functorch import config as functorch_config + + from art.megatron.training.compile import _configure_dynamo + from art.trainer_rank import TrainerRank + from art.trainer_rank._impl import _CheckpointSlot + + # Isolate disk artifacts, including an incompatible donating predecessor. + monkeypatch.setenv("TORCHINDUCTOR_CACHE_DIR", str(tmp_path / "inductor")) + monkeypatch.setenv("TRITON_CACHE_DIR", str(tmp_path / "triton")) + torch.compiler.reset() + torch.manual_seed(71) + inputs = torch.randn(48, 32, device="cuda", dtype=torch.float64) + weight = torch.nn.Parameter(torch.randn(32, 16, device="cuda", dtype=torch.float64)) + + def model(x, w): + return (x @ w).sin().square() + + def warm(compiled): + compiled(inputs, weight).sum().backward() + weight.grad = None + + with ( + functorch_config.patch(donated_buffer=True, enable_autograd_cache=True), + torch._dynamo.config.patch(force_parameter_static_shapes=False), + cast(Any, torch.compiler.config).patch( + cache_key_tag=f"graph-oracle:{tmp_path}" + ), + ): + try: + warm(torch.compile(model, fullgraph=True)) + donated = torch.compiler.save_cache_artifacts() + assert donated is not None + torch.compiler.reset() + assert torch.compiler.load_cache_artifacts(donated[0]) is not None + _configure_dynamo() + compiled = torch.compile(model, fullgraph=True) + # A prior non-retaining backward must not compile a donating kernel + # that rejects the later retained backward of the same graph shape. + if cache_state != "cold": + misses = counters["aot_autograd"]["autograd_cache_miss"] + warm(compiled) + assert counters["aot_autograd"]["autograd_cache_miss"] > misses + if cache_state == "artifact": + reusable = torch.compiler.save_cache_artifacts() + assert reusable is not None + torch.compiler.reset() + assert torch.compiler.load_cache_artifacts(reusable[0]) is not None + compiled = torch.compile(model, fullgraph=True) + hits = counters["aot_autograd"]["autograd_cache_hit"] + warm(compiled) + assert counters["aot_autograd"]["autograd_cache_hit"] > hits + + trainer = TrainerRank.__new__(TrainerRank) + trainer.runtime = SimpleNamespace(model=[], optimizer=None) + trainer._checkpoint_slots = {"student": _CheckpointSlot(params=(weight,))} + version = trainer._capture_checkpoint_version("student") + original = trainer._snapshot_parameter(weight, version) + z = inputs @ original.detach() + first = torch.randn_like(z) + second = torch.randn_like(z) + expected = [inputs.T @ (2 * z.sin() * z.cos() * g) for g in (first, second)] + cache = GraphCache() + handle, (output,) = cache.run( + lambda x: (compiled(x, original),), + inputs, + retention=retention, + options=SimpleNamespace(allow_replay=retention == "replay"), + cuda_devices=[torch.cuda.current_device()], + ) + torch.testing.assert_close(output, z.sin().square()) + with trainer._gradient_transaction(): + cache.backward(handle, (first,), retain_graph=True) + torch.testing.assert_close(weight.grad, expected[0]) + assert cache.handles() == (handle,) + weight.grad = None + with torch.no_grad(): + weight.add_(0.25) + trainer._checkpoint_slots["student"].revision += 1 + inputs.fill_(999) + with trainer._gradient_transaction(): + cache.backward(handle, (second,)) + torch.testing.assert_close(weight.grad, expected[1]) + assert original.grad is None + assert cache.handles() == () + finally: + torch.compiler.reset() + + +@pytest.mark.parametrize("compiled", [False, True]) +@pytest.mark.parametrize("retention", ["gpu", "cpu"]) +@pytest.mark.parametrize("operator", ["rmsnorm", "linear"]) +def test_te_three_backwards_match_manual_gradient(compiled, retention, operator): + from transformer_engine.pytorch import Linear + from transformer_engine.pytorch.ops import RMSNorm + + from art.megatron.compile_workarounds import install_te_reusable_backward + from art.megatron.runtime.compile_cache import configure_reusable_backward + + configure_reusable_backward() + install_te_reusable_backward() + torch.compiler.reset() + torch.manual_seed(109) + inputs = torch.randn(48, 32, device="cuda") + weight = torch.nn.Parameter(torch.randn(32, device="cuda")) + operation = ( + RMSNorm(32, device="cuda", dtype=torch.float32) + if operator == "rmsnorm" + else Linear(32, 32, bias=False, device="cuda", params_dtype=torch.float32) + ) + base_weight = cast(torch.Tensor, operation.weight) + base_weight.requires_grad_(False) + if operator == "linear": + # Exact dyadic inputs give the same oracle with TE's TF32 GEMM and + # PyTorch's FP32 GEMM, without relaxing the retained-gradient check. + inputs.copy_(torch.randint(-4, 5, inputs.shape, device="cuda") / 4) + with torch.no_grad(): + weight.copy_(torch.randint(-4, 5, weight.shape, device="cuda") / 4) + base_weight.copy_( + torch.randint(-4, 5, base_weight.shape, device="cuda") / 4 + ) + + def model(x): + return operation(x * weight) + + physical = torch.compile(model) if compiled else model + executions = [] + + def execute(x): + executions.append(True) + return (physical(x),) + + cache = GraphCache() + handle, (output,) = cache.run( + execute, + inputs, + retention=retention, + options=SimpleNamespace(allow_replay=False), + ) + z = inputs * weight.detach() + inverse = (z.square().mean(-1, keepdim=True) + 1e-5).rsqrt() + torch.testing.assert_close( + output, z * inverse if operator == "rmsnorm" else z @ base_weight.T + ) + for retain in (True, True, False): + cotangent = torch.randn_like(output) + if operator == "linear": + cotangent.copy_(torch.randint(-4, 5, output.shape, device="cuda") / 4) + if operator == "rmsnorm": + derivative = cotangent * inverse - z * inverse.pow(3) * ( + cotangent * z + ).mean(-1, keepdim=True) + else: + derivative = cotangent @ base_weight + weight.grad = None + cache.backward(handle, (cotangent,), retain_graph=retain) + torch.testing.assert_close(weight.grad, (inputs * derivative).sum(0)) + assert cache.handles() == ((handle,) if retain else ()) + assert len(executions) == 1 + torch.compiler.reset() + + +def test_te_retained_backward_failure_clears_unpacked_context(monkeypatch): + from transformer_engine.pytorch.ops import RMSNorm + from transformer_engine.pytorch.ops.fuser import _OperationFuserAutogradFunction + + from art.megatron.compile_workarounds import install_te_reusable_backward + + install_te_reusable_backward() + installed = _OperationFuserAutogradFunction.backward + install_te_reusable_backward() + assert _OperationFuserAutogradFunction.backward is installed + norm = RMSNorm(32, device="cuda", dtype=torch.float32) + norm.weight.requires_grad_(False) + weight = torch.nn.Parameter(torch.ones(32, device="cuda")) + contexts = [] + + def fail(ctx, gradient): + assert ctx.saved_tensors is not None + contexts.append(ctx) + raise RuntimeError("injected TE backward failure") + + monkeypatch.setattr(norm, "op_backward", fail) + cache = GraphCache() + handle, (output,) = cache.run( + lambda x: (norm(x * weight),), torch.ones(48, 32, device="cuda") + ) + with pytest.raises(RuntimeError, match="injected TE backward failure"): + cache.backward(handle, (torch.ones_like(output),), retain_graph=True) + assert cache.handles() == () + assert len(contexts) == 1 + assert contexts[0].saved_tensors is None + assert contexts[0]._saved_tensors_range is not None + + +def test_cpu_saved_views_share_storage_and_exclude_weight_views(): + cache = GraphCache() + weight = torch.nn.Parameter(torch.randn(8, device="cuda")) + inputs = torch.randn(4, 8, device="cuda") + storage_key = weight.untyped_storage().data_ptr() + + def execute(x): + # Both multiplications save aliased input storage for weight gradients. + return ((x * weight).sum(), (x[1:] * weight).sum()) + + handle, outputs = cache.run( + execute, + inputs, + retention="cpu", + cuda_devices=[torch.cuda.current_device()], + keep_on_device=lambda tensor: ( + tensor.untyped_storage().data_ptr() == storage_key + ), + ) + cells = [ + cell + for ref in cache._records[handle].saved or () + if (cell := ref()) is not None + ] + saved_inputs = [cell.tensor for cell in cells if cell.managed] + assert len(saved_inputs) == 2 + assert all(value.is_pinned() for value in saved_inputs) + assert ( + saved_inputs[0].untyped_storage().data_ptr() + == saved_inputs[1].untyped_storage().data_ptr() + ) + assert saved_inputs[1].storage_offset() == 8 + size = inputs.numel() * inputs.element_size() + assert cache.transfer_stats.offload_bytes == size + assert cache.transfer_stats.offload_count == 1 + assert cache.transfer_stats.offload_max_bytes == size + assert cache.transfer_stats.offload_seconds > 0 + assert cache.transfer_stats.restore_count == 0 + cache.backward(handle, tuple(torch.ones_like(value) for value in outputs)) + assert cache.handles() == () + assert cache.transfer_stats.restore_bytes == size + assert cache.transfer_stats.restore_count == 1 + assert cache.transfer_stats.restore_max_bytes == size + assert cache.transfer_stats.restore_seconds > 0 + torch.testing.assert_close(weight.grad, inputs.sum(0) + inputs[1:].sum(0)) + + +def test_saved_alias_restore_uses_one_storage_and_bounded_peak(): + class Aliases(torch.autograd.Function): + @staticmethod + def forward(ctx, weight, value): + ctx.save_for_backward(*(value[:, offset:] for offset in range(24))) + return weight * value.sum() + + @staticmethod + def backward(ctx, *gradients): + saved = ctx.saved_tensors + assert len({value.untyped_storage().data_ptr() for value in saved}) == 1 + return gradients[0] * saved[0].sum(), None + + weight = torch.nn.Parameter(torch.tensor(2.0, device="cuda")) + inputs = torch.ones(2048, 2048, device="cuda") + cache = GraphCache() + handle, (output,) = cache.run( + lambda x: (Aliases.apply(weight, x),), inputs, retention="cpu" + ) + size = inputs.numel() * inputs.element_size() + for retain in (True, False): + torch.cuda.synchronize() + baseline = torch.cuda.memory_allocated() + torch.cuda.reset_peak_memory_stats() + cache.backward(handle, (torch.ones_like(output),), retain_graph=retain) + torch.cuda.synchronize() + increase = torch.cuda.max_memory_allocated() - baseline + assert increase < size * 2, (increase, size) + if retain: + assert cache._records[handle].restored == {} + assert ( + cache.state(handle).gpu_bytes <= output.numel() * output.element_size() + ) + torch.testing.assert_close( + weight.grad, torch.tensor(2 * inputs.numel(), device="cuda", dtype=weight.dtype) + ) + assert cache.transfer_stats.offload_count == 1 + assert cache.transfer_stats.offload_bytes == size + assert cache.transfer_stats.restore_count == 2 + assert cache.transfer_stats.restore_bytes == size * 2 + + +def test_saved_physical_output_alias_is_not_offloaded_or_duplicated(): + cache = GraphCache() + weight = torch.nn.Parameter(torch.ones(2048, device="cuda")) + handle, (output,) = cache.run( + lambda x: ((weight * x).exp(),), + torch.ones(2048, device="cuda"), + retention="cpu", + ) + cells = [ + cell + for ref in cache._records[handle].saved or () + if (cell := ref()) is not None + ] + physical_outputs = cache._records[handle].outputs + assert physical_outputs is not None + physical = physical_outputs[0] + assert any( + not cell.managed + and cell.tensor.untyped_storage().data_ptr() + == physical.untyped_storage().data_ptr() + for cell in cells + ) + cache.backward(handle, (torch.ones_like(output),)) + torch.testing.assert_close(weight.grad, torch.full_like(weight, torch.e)) + + +@pytest.mark.parametrize("retention", ["gpu", "cpu"]) +def test_pinned_offload_completes_on_user_stream_and_preserves_strided_views(retention): + cache = GraphCache() + source = torch.arange(8192.0, device="cuda").reshape(64, 128) + weight = torch.nn.Parameter(torch.ones(64, device="cuda")) + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(stream): + handle, (output,) = cache.run( + lambda x: ((x.t()[3::7] * weight).sum(),), + source, + retention=retention, + ) + if retention == "gpu": + assert cache.transfer_stats.offload_count == 0 + cache.offload(handle) + saved = [ + cell.tensor + for ref in cache._records[handle].saved or () + if (cell := ref()) is not None and cell.managed + ] + assert len(saved) == 1 and saved[0].is_pinned() + # A blocking D2H copy is immediately readable on CPU, without waiting + # on the producer stream separately. Noncontiguous view metadata stays. + expected = torch.arange(8192.0).reshape(64, 128).t()[3::7] + assert saved[0].stride() == expected.stride() + torch.testing.assert_close(saved[0], expected) + torch.cuda.current_stream().wait_stream(stream) + cache.backward(handle, (torch.ones_like(output),)) + assert weight.grad is not None + torch.testing.assert_close(weight.grad.cpu(), expected.sum(0)) + assert cache.transfer_stats.offload_count == 1 + assert cache.transfer_stats.restore_count == 1 + + +@pytest.mark.parametrize("late", [False, True]) +def test_failed_pinned_allocation_releases_partial_saved_state(monkeypatch, late): + allocate = torch.empty_like + storages = [] + + def fail_second(tensor, **kwargs): + if kwargs.get("pin_memory"): + if storages: + raise RuntimeError("injected pinned allocation failure") + result = allocate(tensor, **kwargs) + storages.append(StorageWeakRef(result.untyped_storage())) + return result + return allocate(tensor, **kwargs) + + monkeypatch.setattr(torch, "empty_like", fail_second) + cache = GraphCache() + weight = torch.nn.Parameter(torch.ones(32, device="cuda")) + + def execute(x): + partial = x * weight + return (partial * x.square(),) + + was_enabled = gc.isenabled() + gc.disable() + try: + handle = None + if late: + handle, _ = cache.run(execute, torch.ones_like(weight)) + with pytest.raises(RuntimeError, match="injected pinned") as failure: + if handle is None: + cache.run(execute, torch.ones_like(weight), retention="cpu") + else: + cache.offload(handle) + assert failure.value.__traceback__ is not None + if handle is not None: + # A late failure leaves a valid, partially offloaded graph. Its + # storage remains owned until the caller explicitly releases it. + assert not storages[0].expired() + cache.release(handle) + assert len(storages) == 1 and storages[0].expired() + assert cache.handles() == () + assert cache.transfer_stats.offload_count == 1 + assert cache.transfer_stats.restore_count == 0 + finally: + if was_enabled: + gc.enable() diff --git a/tests/unit/test_trainer_rank_head_memory.py b/tests/unit/test_trainer_rank_head_memory.py new file mode 100644 index 000000000..ec8eeac72 --- /dev/null +++ b/tests/unit/test_trainer_rank_head_memory.py @@ -0,0 +1,198 @@ +"""Registered head admission and the real native/remote gradient transaction.""" + +from dataclasses import replace +import weakref + +import pytest +from test_trainer_rank_memory_admission import _requests, rank # noqa: F401 +import torch + +from art.trainer_rank import ForwardOptions, TrainerRankMemoryError, _impl +from art.trainer_rank._commands import _Executor +from art.trainer_rank._heads import LiveHead, export_head +from art.trainer_rank._tensors import CotangentCollector + + +class _Head(torch.nn.Module): + def __init__(self): + super().__init__() + self.weight = torch.nn.Parameter(torch.ones(4, 4)) + self.tied = self.weight + self.frozen = torch.nn.Parameter(torch.ones(16), requires_grad=False) + self.register_buffer("buffer", torch.ones(16)) + + def forward(self, inputs): + return inputs @ self.weight + + +def _register(rank, kind="module"): + rank._checkpoint_slots["student"] = _impl._CheckpointSlot() + if kind == "module": + return rank.module("head", _Head, checkpoint="student") + return rank.parameter("head", lambda: torch.ones(4, 4), checkpoint="student") + + +def _plan(rank): + requests = _requests(ForwardOptions(backward_state="replay", output_device="cpu"))[ + :1 + ] + plan = rank._plan_flat_forward(requests) + return replace( + plan, groups=(replace(plan.groups[0], slot_ref=rank._slot_ref("student")),) + ) + + +@pytest.mark.parametrize("kind", ["module", "parameter"]) +@pytest.mark.parametrize("existing_gradient", [False, True]) +def test_registered_head_forward_admission_and_native_backward( + rank, monkeypatch, kind, existing_gradient +): + head = _register(rank, kind) + parameter = rank._checkpoint_slots["student"].params[0] + if existing_gradient: + parameter.grad = torch.zeros_like(parameter) + expected = 64 * (2 if existing_gradient else 3) + assert rank._lora_gradient_staging_bytes(rank._slot_ref("student")) == expected + assert rank._lora_version_capture_bytes(rank._slot_ref("student")) == 0 + monkeypatch.setattr(rank, "_available_memory_bytes", lambda: 150) + _, check = rank._admit_graph_memory(_plan(rank)) + assert not check.fits and check.estimated_required_bytes == 100 + expected + monkeypatch.setattr(rank, "_available_memory_bytes", lambda: 100 + expected) + assert all(rank._admit_graph_memory(_plan(rank))[1].fits for _ in range(2)) + cache = rank._forward_graph_cache() + losses = [] + for factor in (2, 3): + handle, tensors = cache.run( + lambda inputs: (inputs.sin(),), + torch.ones(4, requires_grad=True), + retention="replay", + ) + from art.trainer_rank._tensors import detach_tree + + value = rank._forward_cotangent_collector().attach( + detach_tree(handle, tensors) + )[0] + result = head(value) if kind == "module" else value @ head + losses.append(result.sum() * factor) + for loss in losses: + rank.backward(loss) + torch.testing.assert_close( + parameter.grad, + torch.full_like(parameter, 5 * torch.sin(torch.tensor(1.0)).item()), + ) + assert not cache.handles() + + +def test_late_registration_admits_known_staging_and_preserves_old_graph( + rank, monkeypatch +): + rank._checkpoint_slots["student"] = _impl._CheckpointSlot() + cache = rank._forward_graph_cache() + handle, outputs = cache.run( + lambda value: (value.sin(),), + torch.ones(4, requires_grad=True), + retention="replay", + execution_peak_bytes=100, + checkpoint_versions=(rank._capture_checkpoint_version("student"),), + ) + monkeypatch.setattr(rank, "_available_memory_bytes", lambda: 250) + references = [] + tag = rank._tag_custom_parameters + + def watch(parameters): + references.extend(weakref.ref(parameter) for parameter in parameters) + tag(parameters) + + monkeypatch.setattr(rank, "_tag_custom_parameters", watch) + with pytest.raises(TrainerRankMemoryError, match="292 GPU bytes") as failure: + rank.module("head", _Head, checkpoint="student") + assert failure.value is not None + assert all(reference() is None for reference in references) + assert rank._checkpoint_slots["student"].params == () + assert not rank._checkpoint_slots["student"].custom + assert cache.handles() == (handle,) + monkeypatch.setattr(rank, "_available_memory_bytes", lambda: 292) + head = rank.module("head", _Head, checkpoint="student") + from art.trainer_rank._tensors import detach_tree + + value = rank._forward_cotangent_collector().attach(detach_tree(handle, outputs))[0] + rank.backward(head(value).sum()) + torch.testing.assert_close( + head.weight.grad, + torch.full_like(head.weight, torch.sin(torch.tensor(1.0)).item()), + ) + + +def test_remote_repeated_head_targets_stream_and_preserve_source_packets( + rank, monkeypatch +): + _register(rank, "parameter") + parameter = rank._checkpoint_slots["student"].params[0] + collector = CotangentCollector() + live = LiveHead(export_head(rank, "student", "head"), torch.ones(4, 4), collector) + assert isinstance(live.value, torch.Tensor) + packets = collector.backward( + torch.stack([(live.value * factor).sum() for factor in (2, 3, 4)]).sum() + ) + sources = tuple( + g.clone() for packet in packets if (g := packet.gradients[0]) is not None + ) + assert len(sources) == len(packets) + commit = rank._commit_versioned_gradients + sizes = [] + + def record(gradients): + sizes.append(len(gradients)) + commit(gradients) + + monkeypatch.setattr(rank, "_commit_versioned_gradients", record) + _Executor(rank, "zero")._backward(packets, retain_graph=False) + assert sizes == [1, 1, 1] + torch.testing.assert_close(parameter.grad, torch.full_like(parameter, 9)) + for packet, source in zip(packets, sources, strict=True): + torch.testing.assert_close(packet.gradients[0], source) + + +def test_frozen_and_buffer_only_registration_needs_no_gradient_reserve( + rank, monkeypatch +): + rank._checkpoint_slots["student"] = _impl._CheckpointSlot() + monkeypatch.setattr(rank, "_available_memory_bytes", lambda: 0) + rank.buffer("buffer", lambda: torch.ones(16), checkpoint="student") + rank.module("frozen", lambda: _Head().requires_grad_(False), checkpoint="student") + assert rank._checkpoint_slots["student"].params == () + + +@pytest.mark.parametrize("kind", ["buffer", "frozen"]) +def test_nontrainable_registration_preserves_pending_graph_workspace( + rank, monkeypatch, kind +): + rank._checkpoint_slots["student"] = _impl._CheckpointSlot() + cache = rank._forward_graph_cache() + handle, outputs = cache.run( + lambda value: (value.sin(),), + torch.ones(4, requires_grad=True), + retention="replay", + execution_peak_bytes=100, + checkpoint_versions=(rank._capture_checkpoint_version("student"),), + ) + + def register(): + if kind == "buffer": + return rank.buffer("head", lambda: torch.ones(16), checkpoint="student") + return rank.module( + "head", lambda: _Head().requires_grad_(False), checkpoint="student" + ) + + monkeypatch.setattr(rank, "_available_memory_bytes", lambda: 99) + with pytest.raises(TrainerRankMemoryError, match="100 GPU bytes"): + register() + assert not rank._checkpoint_slots["student"].custom + monkeypatch.setattr(rank, "_available_memory_bytes", lambda: 100) + register() + assert rank._checkpoint_slots["student"].params == () + from art.trainer_rank._tensors import detach_tree + + value = rank._forward_cotangent_collector().attach(detach_tree(handle, outputs))[0] + rank.backward(value.sum()) + assert not cache.handles() diff --git a/tests/unit/test_trainer_rank_head_memory_cuda.py b/tests/unit/test_trainer_rank_head_memory_cuda.py new file mode 100644 index 000000000..9d2bed82d --- /dev/null +++ b/tests/unit/test_trainer_rank_head_memory_cuda.py @@ -0,0 +1,129 @@ +"""Allocator assertions; requires a validation-owned GPU reservation.""" + +import json +import os + +import pytest +from test_trainer_rank_custom_tensors import _trainer +import torch +from torch.multiprocessing.reductions import StorageWeakRef + +from art.trainer_rank import TrainerRankMemoryError +from art.trainer_rank._commands import _Executor +from art.trainer_rank._heads import LiveHead, export_head +from art.trainer_rank._tensors import CotangentCollector + +pytestmark = pytest.mark.skipif( + os.environ.get("ART_GRAPH_GPU_TEST") != "1", reason="reserved GPU required" +) + + +@pytest.mark.parametrize("existing_gradient", [False, True]) +def test_repeated_remote_head_cotangents_fit_original_gradient_reserve( + existing_gradient, +): + torch.set_num_threads(2) + trainer, api = _trainer("student") + trainer.device = torch.device("cuda") + parameter = api.parameter( + "head", lambda: torch.ones(4 * 1024**2), checkpoint="student" + ) + if existing_gradient: + parameter.grad = torch.ones_like(parameter) + collector = CotangentCollector() + live = LiveHead( + export_head(trainer, "student", "head"), torch.ones(parameter.shape), collector + ) + assert isinstance(live.value, torch.Tensor) + packets = collector.backward( + torch.stack([(live.value * factor).sum() for factor in range(1, 9)]).sum() + ) + assert all( + gradient.device.type == "cpu" + for packet in packets + for gradient in packet.gradients + if gradient is not None + ) + reserve = trainer._lora_gradient_staging_bytes(trainer._slot_ref("student")) + torch.cuda.synchronize() + baseline = torch.cuda.memory_allocated() + torch.cuda.reset_peak_memory_stats() + executor = _Executor(trainer, "zero") + for _ in range(2): + executor._backward(packets, retain_graph=False) + torch.cuda.synchronize() + peak = torch.cuda.max_memory_allocated() - baseline + assert peak <= reserve + torch.testing.assert_close( + parameter.grad, torch.full_like(parameter, 73 if existing_gradient else 72) + ) + print( + "REMOTE_HEAD_RESERVATION=" + + json.dumps( + dict( + existing_gradient=existing_gradient, + peak=peak, + reserve=reserve, + repeated_captures=8, + ) + ) + ) + + +@pytest.mark.parametrize("kind", ["parameter", "buffer", "frozen"]) +def test_rejected_late_registration_releases_gpu_storage_with_live_traceback( + monkeypatch, kind +): + trainer, api = _trainer("student") + trainer.device = torch.device("cuda") + cache = trainer._forward_graph_cache() + handle, _ = cache.run( + lambda value: (value.sin(),), + torch.ones(1024, device="cuda", requires_grad=True), + retention="replay", + execution_peak_bytes=1024**2, + checkpoint_versions=(trainer._capture_checkpoint_version("student"),), + ) + monkeypatch.setattr(trainer, "_available_memory_bytes", lambda: 1024**2 - 1) + storages = [] + initialize = trainer._initialize_custom_object + + def watch(checkpoint, name, custom): + initialize(checkpoint, name, custom) + tensors = ( + custom.value.parameters() + if isinstance(custom.value, torch.nn.Module) + else (custom.value,) + ) + storages.extend(StorageWeakRef(tensor.untyped_storage()) for tensor in tensors) + + monkeypatch.setattr(trainer, "_initialize_custom_object", watch) + torch.cuda.synchronize() + baseline = torch.cuda.memory_allocated() + with pytest.raises(TrainerRankMemoryError) as failure: + if kind == "parameter": + api.parameter("late", lambda: torch.ones(4 * 1024**2), checkpoint="student") + elif kind == "buffer": + trainer.buffer( + "late", lambda: torch.ones(4 * 1024**2), checkpoint="student" + ) + else: + api.module( + "late", + lambda: torch.nn.Linear(1024, 4096, bias=False).requires_grad_(False), + checkpoint="student", + ) + assert failure.value is not None + torch.cuda.synchronize() + assert all(storage.expired() for storage in storages) + assert torch.cuda.memory_allocated() <= baseline + assert not trainer._checkpoint_slots["student"].custom + cache.release(handle) + print( + "LATE_HEAD_RELEASE=" + + json.dumps( + dict( + kind=kind, retained_extra_bytes=torch.cuda.memory_allocated() - baseline + ) + ) + ) diff --git a/tests/unit/test_trainer_rank_head_recompute.py b/tests/unit/test_trainer_rank_head_recompute.py index 1a2c8d9a8..45d7b1083 100644 --- a/tests/unit/test_trainer_rank_head_recompute.py +++ b/tests/unit/test_trainer_rank_head_recompute.py @@ -326,10 +326,10 @@ def run(local): dist.all_reduce(expected_loss) expected_loss /= cp_size trainer = actual[0] - trainer.dp_reduce(actual[2]) + trainer.reduce(actual[2]) torch.testing.assert_close(actual[2], expected_loss) count = torch.tensor(sum(len(row) for row in tokens), device=device) - trainer.dp_reduce(count) + trainer.reduce(count) assert count.item() == dp_size * sum(len(row) for row in tokens) if mode in ("frozen", "no_grad"): return diff --git a/tests/unit/test_trainer_rank_live_heads.py b/tests/unit/test_trainer_rank_live_heads.py new file mode 100644 index 000000000..7ff5a3354 --- /dev/null +++ b/tests/unit/test_trainer_rank_live_heads.py @@ -0,0 +1,1448 @@ +from __future__ import annotations + +import asyncio +from copy import deepcopy +from dataclasses import replace + +import pytest +from test_trainer_rank_custom_tensors import _trainer +import torch +from torch.utils.checkpoint import checkpoint + +from art.trainer_rank import AdamParams, ModuleHandle, run_rank_callback +from art.trainer_rank._heads import ( + HeadBufferUpdate, + HeadRegistration, + LiveHead, + execute_head_operation, + export_head, + head_gradient_targets, +) +from art.trainer_rank._tensors import CotangentCollector + + +class TiedHead(torch.nn.Module): + offset: torch.Tensor + other_offset: torch.Tensor + + def __init__(self, checkpointed: bool | None = None): + super().__init__() + self.left = torch.nn.Parameter(torch.tensor(2.0)) + self.right = self.left + self.register_buffer("offset", torch.tensor(1.0)) + self.register_buffer("other_offset", self.offset) + self.checkpointed = checkpointed + + def forward(self, value): + def compute(x): + assert self.left is self.right + assert self.offset is self.other_offset + return x * self.left.square() + self.right + self.offset + + return ( + compute(value) + if self.checkpointed is None + else checkpoint(compute, value, use_reentrant=self.checkpointed) + ) + + +def _module(live: LiveHead) -> ModuleHandle: + assert isinstance(live.value, ModuleHandle) + return live.value + + +def _tensor(live: LiveHead) -> torch.Tensor: + assert isinstance(live.value, torch.Tensor) + return live.value + + +def _step(trainer, monkeypatch): + monkeypatch.setattr( + trainer, + "_reduce_dynamic_grads", + lambda params, **kwargs: tuple( + torch.zeros_like(p, dtype=torch.float32) + if p.grad is None + else p.grad.float() + for p in params + ), + ) + trainer.optim_step( + params=AdamParams(learning_rate=0.1, weight_decay=0.0), checkpoints=["student"] + ) + + +@pytest.mark.parametrize("checkpointed", (None, False, True)) +def test_native_old_head_graph_uses_original_tied_weights_after_step( + monkeypatch, checkpointed +): + trainer, rank = _trainer("student") + head = rank.module("head", lambda: TiedHead(checkpointed), checkpoint="student") + old_input = torch.tensor(3.0, requires_grad=True) + old_loss = head(old_input) + with trainer._gradient_transaction(): + head(torch.tensor(1.0, requires_grad=True)).backward() + _step(trainer, monkeypatch) + assert head.left.item() != 2 + latest = head(torch.tensor(3.0)) + torch.testing.assert_close( + latest, 3 * head.left.detach().square() + head.left.detach() + 1 + ) + with trainer._gradient_transaction(): + old_loss.backward() + torch.testing.assert_close(head.left.grad, torch.tensor(13.0)) + torch.testing.assert_close(old_input.grad, torch.tensor(4.0)) + assert head.left is head.right + + +def test_batchnorm_buffers_publish_once_and_old_graph_keeps_original_buffers(): + trainer, rank = _trainer("student") + head = rank.module( + "bn", lambda: torch.nn.BatchNorm1d(2, dtype=torch.float64), checkpoint="student" + ) + initial = deepcopy(head).eval() + head.eval() + x = torch.tensor([[1.0, 3.0], [2.0, 5.0]], dtype=torch.float64, requires_grad=True) + old = head(x).square().sum() + expected_input = x.detach().clone().requires_grad_() + expected = initial(expected_input).square().sum() + head.train() + head(torch.tensor([[4.0, 2.0], [8.0, 10.0]], dtype=torch.float64)) + assert head.num_batches_tracked.item() == 1 + assert not torch.equal(head.running_mean, initial.running_mean) + with trainer._gradient_transaction(): + old.backward() + with trainer._gradient_transaction(): + expected.backward() + torch.testing.assert_close(x.grad, expected_input.grad) + torch.testing.assert_close(head.weight.grad, initial.weight.grad) + + +def test_failed_module_call_does_not_publish_buffers(): + class Failing(torch.nn.BatchNorm1d): + def forward(self, input): + super().forward(input) + raise RuntimeError("failed after buffer mutation") + + trainer, rank = _trainer("student") + head = rank.module("bn", lambda: Failing(2), checkpoint="student") + with pytest.raises(RuntimeError, match="failed after"): + head(torch.ones(4, 2)) + assert head.num_batches_tracked.item() == 0 + torch.testing.assert_close(head.running_mean, torch.zeros(2)) + + +def test_client_tied_head_old_backward_after_native_refresh(monkeypatch): + trainer, rank = _trainer("student") + native = rank.module("head", TiedHead, checkpoint="student") + collector = CotangentCollector() + live = LiveHead(export_head(trainer, "student", "head"), TiedHead(), collector) + old_input = torch.tensor(3.0, requires_grad=True) + old = _module(live)(old_input) + with trainer._gradient_transaction(): + native(torch.tensor(1.0)).backward() + _step(trainer, monkeypatch) + live.refresh(export_head(trainer, "student", "head")) + assert _module(live).left is _module(live).right + assert _module(live)(torch.tensor(3.0)).item() != old.item() + packets = collector.backward(old) + assert len(packets) == 1 + trainer._commit_versioned_gradients(head_gradient_targets(trainer, packets[0])) + torch.testing.assert_close(native.left.grad, torch.tensor(13.0)) + torch.testing.assert_close(old_input.grad, torch.tensor(4.0)) + + +def test_live_parameter_reuses_handle_and_earlier_capture_retains_version(): + trainer, rank = _trainer("student") + parameter = rank.parameter("gain", lambda: torch.tensor(2.0), checkpoint="student") + collector = CotangentCollector() + live = LiveHead( + export_head(trainer, "student", "gain"), torch.tensor(2.0), collector + ) + handle = _tensor(live) + old = handle.square() * 3 + parameter.data.fill_(4) + trainer._checkpoint_slots["student"].revision += 1 + live.refresh(export_head(trainer, "student", "gain")) + assert _tensor(live) is handle + assert (handle * 2).item() == 8 + packets = collector.backward(old) + trainer._commit_versioned_gradients(head_gradient_targets(trainer, packets[0])) + torch.testing.assert_close(parameter.grad, torch.tensor(12.0)) + + +def test_client_buffer_publication_conflict_is_atomic(): + trainer, rank = _trainer("student") + native = rank.module("bn", lambda: torch.nn.BatchNorm1d(2), checkpoint="student") + live = LiveHead( + export_head(trainer, "student", "bn"), + torch.nn.BatchNorm1d(2), + CotangentCollector(), + ) + _module(live)(torch.ones(4, 2)) + update = live.take_publication() + assert update is not None + execute_head_operation(trainer, "head_publish", (update,)) + live.refresh(export_head(trainer, "student", "bn")) + assert native.num_batches_tracked.item() == 1 + native(torch.zeros(4, 2)) + _module(live)(torch.ones(4, 2) * 5) + stale = live.take_publication() + assert stale is not None + before = native.running_mean.detach().clone() + with pytest.raises(RuntimeError, match="buffers changed"): + execute_head_operation(trainer, "head_publish", (stale,)) + torch.testing.assert_close(native.running_mean, before) + + +def test_replacement_invalidates_live_handle_and_old_gradient(): + trainer, rank = _trainer("student") + rank.parameter("gain", lambda: torch.tensor(2.0), checkpoint="student") + collector = CotangentCollector() + state = export_head(trainer, "student", "gain") + live = LiveHead(state, torch.tensor(2.0), collector) + old = _tensor(live).square() + trainer._checkpoint_slots["student"].generation += 1 + live.refresh(export_head(trainer, "student", "gain")) + with pytest.raises(RuntimeError, match="stale"): + _tensor(live) * 2 + with pytest.raises(RuntimeError, match="replaced"): + head_gradient_targets(trainer, collector.backward(old)[0]) + + +def test_registration_materializes_factory_state_and_persistent_buffers(): + trainer, _ = _trainer("student") + source = TiedHead() + state = execute_head_operation( + trainer, "head_register", HeadRegistration("student", "head", "module", source) + ) + assert list(state.parameters) == ["left"] + assert list(state.buffers) == ["offset"] + registered = trainer._checkpoint_slots["student"].custom["head"].value + assert isinstance(registered, torch.nn.Module) + assert registered.left is not source.left + torch.testing.assert_close(state.parameters["left"], source.left) + + +@pytest.mark.parametrize("reentrant", (False, True)) +def test_external_checkpoint_requires_explicit_snapshot(reentrant): + trainer, rank = _trainer("student") + head = rank.module("head", TiedHead, checkpoint="student") + x = torch.tensor(3.0, requires_grad=True) + old = checkpoint(head, x, use_reentrant=reentrant) + head.left.data.fill_(4) + with pytest.raises(RuntimeError, match="head.snapshot"): + with trainer._gradient_transaction(): + old.backward() + assert head.left.grad is None + + +@pytest.mark.parametrize("reentrant", (False, True)) +def test_explicit_snapshot_supports_external_checkpoint(reentrant): + trainer, rank = _trainer("student") + head = rank.module("head", TiedHead, checkpoint="student") + captured = head.snapshot() + x = torch.tensor(3.0, requires_grad=True) + old = checkpoint(captured, x, use_reentrant=reentrant) + head.left.data.fill_(4) + with trainer._gradient_transaction(): + old.backward() + torch.testing.assert_close(head.left.grad, torch.tensor(13.0)) + torch.testing.assert_close(x.grad, torch.tensor(4.0)) + + +def test_forward_hooks_see_captured_parameters_and_ties(): + trainer, rank = _trainer("student") + observed = [] + + def factory(): + source = TiedHead() + source.register_forward_pre_hook( + lambda module, inputs: observed.append(module.left.item()) + ) + return source + + head = rank.module("head", factory, checkpoint="student") + head(torch.tensor(1.0)) + head.left.data.fill_(4) + head(torch.tensor(1.0)) + assert observed == [2, 4] + + +def test_constructor_staleness_applies_to_heads_before_mutating_gradients(): + from art.trainer_rank._options import ForwardOptions + + trainer, rank = _trainer("student") + setattr(trainer, "_forward_options", ForwardOptions(max_gradient_staleness=0)) + head = rank.module("head", TiedHead, checkpoint="student") + old = head(torch.tensor(3.0)) + trainer._checkpoint_slots["student"].revision += 1 + with pytest.raises(RuntimeError, match="staleness"): + with trainer._gradient_transaction(): + old.backward() + assert head.left.grad is None + + +def _buffer_authority_worker(process_rank, init_method): + from datetime import timedelta + + import torch.distributed as dist + + from art.trainer_rank._heads import synchronize_head_buffers + + dist.init_process_group( + "gloo", + rank=process_rank, + world_size=2, + init_method=init_method, + timeout=timedelta(seconds=30), + ) + try: + trainer, rank = _trainer("student") + head = rank.module("bn", lambda: torch.nn.BatchNorm1d(2), checkpoint="student") + for _ in range(process_rank + 1): + head(torch.full((4, 2), float(process_rank + 1))) + synchronize_head_buffers(trainer) + torch.testing.assert_close(head.running_mean, torch.full((2,), 0.1)) + assert head.num_batches_tracked.item() == 1 + finally: + dist.destroy_process_group() + + +def test_distributed_persistent_buffers_use_dp_zero_authority(tmp_path): + torch.multiprocessing.spawn( + _buffer_authority_worker, + args=(f"file://{tmp_path / 'heads'}",), + nprocs=2, + join=True, + ) + + +def test_remote_reregistration_rejects_changed_ties(): + trainer, rank = _trainer("student") + rank.module("head", TiedHead, checkpoint="student") + untied = TiedHead() + untied.right = torch.nn.Parameter(torch.tensor(2.0)) + with pytest.raises(ValueError, match="schema differs"): + execute_head_operation( + trainer, + "head_register", + HeadRegistration("student", "head", "module", untied), + ) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +@pytest.mark.parametrize("client", (False, True)) +def test_cuda_checkpoint_head_across_completed_optimizer_update(monkeypatch, client): + from art.trainer_rank._tensors import managed_tensor + + trainer, rank = _trainer("student") + trainer.device = torch.device("cuda", 0) + native = rank.module("head", lambda: TiedHead(False), checkpoint="student") + collector = CotangentCollector() + live = ( + LiveHead(export_head(trainer, "student", "head"), TiedHead(False), collector) + if client + else None + ) + head = native if live is None else _module(live) + original_input = torch.tensor(3.0, device="cuda", requires_grad=True) + old = head(managed_tensor(original_input) if client else original_input) + with trainer._gradient_transaction(): + native(torch.tensor(1.0, device="cuda", requires_grad=True)).backward() + _step(trainer, monkeypatch) + if live is not None: + live.refresh(export_head(trainer, "student", "head")) + packets = collector.backward(old) + trainer._commit_versioned_gradients(head_gradient_targets(trainer, packets[0])) + else: + with trainer._gradient_transaction(): + old.backward() + torch.testing.assert_close(native.left.grad, torch.tensor(13.0, device="cuda")) + torch.testing.assert_close(original_input.grad, torch.tensor(4.0, device="cuda")) + + +def test_client_buffer_reads_keep_old_graph_and_replacement_invalidates_handle(): + trainer, rank = _trainer("student") + native = rank.buffer("scale", lambda: torch.tensor(2.0), checkpoint="student") + collector = CotangentCollector() + live = LiveHead( + export_head(trainer, "student", "scale"), torch.tensor(2.0), collector + ) + x = torch.tensor(3.0, requires_grad=True) + old = x * _tensor(live) + native.fill_(4) + live.refresh(export_head(trainer, "student", "scale")) + collector.backward(old) + torch.testing.assert_close(x.grad, torch.tensor(2.0)) + _tensor(live).add_(1) + update = live.take_publication() + assert update is not None + assert update.buffers[""].item() == 5 + trainer._checkpoint_slots["student"].generation += 1 + live.refresh(export_head(trainer, "student", "scale")) + with pytest.raises(RuntimeError, match="stale"): + _tensor(live) + 1 + + +def test_client_module_explicit_dtype_move_retains_ties_and_live_parameters(): + trainer, rank = _trainer("student") + native = rank.module("head", TiedHead, checkpoint="student") + live = LiveHead( + export_head(trainer, "student", "head"), TiedHead(), CotangentCollector() + ) + head = _module(live).to("cpu").to(dtype=torch.float64) + assert head.left.dtype == torch.float64 + assert head.offset.dtype == torch.float64 + assert head.left is head.right + assert head.offset is head.other_offset + native.left.data.fill_(4) + trainer._checkpoint_slots["student"].revision += 1 + live.refresh(export_head(trainer, "student", "head")) + output = head(torch.tensor(3.0, dtype=torch.float64)) + assert output.dtype == torch.float64 + assert output.item() == 53 + head.offset.add_(1) + update = live.take_publication() + assert update is not None + assert update.buffers["offset"].dtype == torch.float32 + + +def test_client_explicit_snapshot_supports_nonreentrant_checkpoint(): + trainer, rank = _trainer("student") + native = rank.module("head", TiedHead, checkpoint="student") + collector = CotangentCollector() + live = LiveHead(export_head(trainer, "student", "head"), TiedHead(), collector) + captured = _module(live).snapshot() + x = torch.tensor(3.0, requires_grad=True) + old = checkpoint(captured, x, use_reentrant=False) + native.left.data.fill_(4) + trainer._checkpoint_slots["student"].revision += 1 + live.refresh(export_head(trainer, "student", "head")) + packets = collector.backward(old) + trainer._commit_versioned_gradients(head_gradient_targets(trainer, packets[0])) + torch.testing.assert_close(native.left.grad, torch.tensor(13.0)) + torch.testing.assert_close(x.grad, torch.tensor(4.0)) + + +@pytest.mark.parametrize("client", (False, True)) +def test_buffer_item_and_bitwise_mutations_publish_without_losing_handle(client): + trainer, rank = _trainer("student") + native = rank.buffer("mask", lambda: torch.tensor([1, 2]), checkpoint="student") + live = ( + LiveHead( + export_head(trainer, "student", "mask"), + torch.tensor([1, 2]), + CotangentCollector(), + ) + if client + else None + ) + value = native if live is None else _tensor(live) + value[0] = 4 + original = value + value |= 1 + assert value is original + torch.testing.assert_close(value, torch.tensor([5, 3])) + if live is not None: + update = live.take_publication() + assert update is not None + execute_head_operation(trainer, "head_publish", (update,)) + torch.testing.assert_close(native, torch.tensor([5, 3])) + + +@pytest.mark.parametrize( + "mutation", + ( + lambda parameter: parameter.__setitem__(0, 9), + lambda parameter: parameter.__iadd__(1), + lambda parameter: parameter.requires_grad_(False), + lambda parameter: parameter.data.fill_(9), + ), +) +def test_client_parameter_mutations_fail_before_changing_owned_values(mutation): + trainer, rank = _trainer("student") + rank.parameter("gain", lambda: torch.tensor([2.0]), checkpoint="student") + live = LiveHead( + export_head(trainer, "student", "gain"), + torch.tensor([2.0]), + CotangentCollector(), + ) + with pytest.raises(RuntimeError, match="checkpoint parameters"): + mutation(_tensor(live)) + torch.testing.assert_close(_tensor(live).detach(), torch.tensor([2.0])) + + +def test_head_export_preserves_strict_constructor_policy_for_client(): + from art.trainer_rank._options import ForwardOptions + + trainer, rank = _trainer("student") + setattr(trainer, "_forward_options", ForwardOptions(max_gradient_staleness=0)) + rank.parameter("gain", lambda: torch.tensor(2.0), checkpoint="student") + collector = CotangentCollector() + live = LiveHead( + export_head(trainer, "student", "gain"), torch.tensor(2.0), collector + ) + old = _tensor(live).square() + trainer._checkpoint_slots["student"].revision += 1 + with pytest.raises(RuntimeError, match="staleness"): + head_gradient_targets(trainer, collector.backward(old)[0]) + + +def test_reusing_module_factory_value_does_not_share_checkpoint_storage(): + trainer, rank = _trainer("A", "B") + source = TiedHead() + first = rank.module("head", lambda: source, checkpoint="A") + second = rank.module("head", lambda: source, checkpoint="B") + first.left.data.fill_(4) + assert first(torch.tensor(3.0)).item() == 53 + assert second(torch.tensor(3.0)).item() == 15 + assert source.left.item() == 2 + with trainer._gradient_transaction(): + second(torch.tensor(3.0)).backward() + assert first.left.grad is None + torch.testing.assert_close(second.left.grad, torch.tensor(13.0)) + + +def test_registration_under_no_grad_preserves_authoritative_trainability(): + trainer, rank = _trainer("student") + native = rank.module("head", TiedHead, checkpoint="student") + collector = CotangentCollector() + with torch.no_grad(): + live = LiveHead(export_head(trainer, "student", "head"), TiedHead(), collector) + assert _module(live).left.requires_grad + packets = collector.backward(_module(live)(torch.tensor(3.0))) + trainer._commit_versioned_gradients(head_gradient_targets(trainer, packets[0])) + torch.testing.assert_close(native.left.grad, torch.tensor(13.0)) + with torch.no_grad(): + assert not _module(live)(torch.tensor(3.0)).requires_grad + + +@pytest.mark.parametrize("external", (False, True)) +def test_client_reentrant_checkpoint_rejects_nested_remote_bridges(external): + trainer, rank = _trainer("student") + native = rank.module("head", TiedHead, checkpoint="student") + collector = CotangentCollector() + live = LiveHead( + export_head(trainer, "student", "head"), + TiedHead(None if external else True), + collector, + ) + x = torch.tensor(3.0, requires_grad=True) + loss = ( + checkpoint(_module(live).snapshot(), x, use_reentrant=True) + if external + else _module(live)(x) + ) + with pytest.raises(RuntimeError, match="nested remote backward is unsupported"): + collector.backward(loss) + assert native.left.grad is None + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_cuda_batchnorm_explicit_placement_publishes_buffers_and_preserves_old_graph(): + from art.trainer_rank._tensors import managed_tensor + + trainer, rank = _trainer("student") + trainer.device = torch.device("cuda", 0) + native = rank.module("bn", lambda: torch.nn.BatchNorm1d(2), checkpoint="student") + collector = CotangentCollector() + live = LiveHead( + export_head(trainer, "student", "bn"), torch.nn.BatchNorm1d(2), collector + ) + head = _module(live).to("cuda") + head.eval() + x = torch.tensor([[1.0, 3.0], [2.0, 5.0]], device="cuda", requires_grad=True) + old = head(managed_tensor(x)).sum() + head.train() + head(torch.tensor([[4.0, 2.0], [8.0, 10.0]], device="cuda")) + update = live.take_publication() + assert update is not None + execute_head_operation(trainer, "head_publish", (update,)) + live.refresh(export_head(trainer, "student", "bn")) + assert native.num_batches_tracked.item() == 1 + torch.testing.assert_close( + native.running_mean, torch.tensor([0.6, 0.6], device="cuda") + ) + packets = collector.backward(old) + trainer._commit_versioned_gradients(head_gradient_targets(trainer, packets[0])) + torch.testing.assert_close(x.grad, torch.full_like(x, (1.0 + 1e-5) ** -0.5)) + torch.testing.assert_close( + native.bias.grad, torch.tensor([2.0, 2.0], device="cuda") + ) + + +def test_remote_registration_of_existing_frozen_snapshot_keeps_frozen_parameters(): + trainer, rank = _trainer("snapshot") + trainer._checkpoint_slots["snapshot"].snapshot = True + rank.module("head", TiedHead, checkpoint="snapshot") + source = TiedHead() + state = execute_head_operation( + trainer, "head_register", HeadRegistration("snapshot", "head", "module", source) + ) + assert source.left.requires_grad + assert not state.parameters["left"].requires_grad + live = LiveHead(state, source, CotangentCollector()) + assert not _module(live).left.requires_grad + x = torch.tensor(3.0, requires_grad=True) + with trainer._gradient_transaction(): + _module(live)(x).backward() + torch.testing.assert_close(x.grad, torch.tensor(4.0)) + + +def test_client_reused_module_factory_does_not_share_checkpoint_handles(): + trainer, rank = _trainer("A", "B") + rank.module("head", TiedHead, checkpoint="A") + rank.module("head", TiedHead, checkpoint="B") + source = TiedHead() + first = LiveHead(export_head(trainer, "A", "head"), source, CotangentCollector()) + second = LiveHead(export_head(trainer, "B", "head"), source, CotangentCollector()) + _module(first).offset.add_(2) + assert _module(first)(torch.tensor(3.0)).item() == 17 + assert _module(second)(torch.tensor(3.0)).item() == 15 + assert source.offset.item() == 1 + assert _module(first).left is not _module(second).left + assert first.take_publication() is not None + assert second.take_publication() is None + + +@pytest.mark.parametrize("operation", ("model_first", "parameter_first", "linear")) +def test_managed_model_operand_captures_live_parameter_and_keeps_old_version(operation): + from art.trainer_rank._tensors import detach_tree + + trainer, rank = _trainer("student") + native = rank.parameter( + "weight", lambda: torch.tensor([2.0, 4.0]), checkpoint="student" + ) + collector = CotangentCollector() + live = LiveHead( + export_head(trainer, "student", "weight"), torch.zeros(2), collector + ) + hidden = collector.attach( + detach_tree("model", torch.tensor([3.0, 5.0], requires_grad=True)), managed=True + ) + weight = _tensor(live) + loss = ( + hidden @ weight + if operation == "model_first" + else weight @ hidden + if operation == "parameter_first" + else torch.nn.functional.linear(hidden, weight) + ) + native.data.copy_(torch.tensor([7.0, 11.0])) + trainer._checkpoint_slots["student"].revision += 1 + live.refresh(export_head(trainer, "student", "weight")) + assert (hidden @ weight).item() == 76 + packets = collector.backward(loss) + assert len(packets) == 2 + model = next(packet for packet in packets if packet.handle == "model") + head = next(packet for packet in packets if packet.handle.startswith("head:")) + torch.testing.assert_close(model.gradients[0], torch.tensor([2.0, 4.0])) + trainer._commit_versioned_gradients(head_gradient_targets(trainer, head)) + torch.testing.assert_close(native.grad, torch.tensor([3.0, 5.0])) + + +@pytest.mark.parametrize("client", (False, True)) +def test_module_buffer_reassignment_publishes(client): + class Counter(torch.nn.Module): + def __init__(self): + super().__init__() + self.register_buffer("count", torch.tensor(0)) + + def forward(self, value): + self.count = self.count + 1 + return value + self.count + + trainer, rank = _trainer("student") + native = rank.module("head", Counter, checkpoint="student") + live = ( + LiveHead( + export_head(trainer, "student", "head"), Counter(), CotangentCollector() + ) + if client + else None + ) + head = native if live is None else _module(live) + assert head(torch.tensor(0)).item() == 1 + assert head(torch.tensor(0)).item() == 2 + assert head.count.item() == 2 + if live is not None: + execute_head_operation(trainer, "head_publish", (live.take_publication(),)) + assert native.count.item() == 2 + + +@pytest.mark.parametrize("client", (False, True)) +def test_failed_buffer_shape_change_rolls_back_every_buffer(client): + class Resize(torch.nn.Module): + a: torch.Tensor + b: torch.Tensor + + def __init__(self): + super().__init__() + self.register_buffer("a", torch.zeros(1)) + self.register_buffer("b", torch.zeros(1)) + + def forward(self, value): + self.a.add_(1) + self.b.resize_(2) + return value + + trainer, rank = _trainer("student") + native = rank.module("head", Resize, checkpoint="student") + live = ( + LiveHead( + export_head(trainer, "student", "head"), Resize(), CotangentCollector() + ) + if client + else None + ) + head = native if live is None else _module(live) + before = export_head(trainer, "student", "head").buffer_revision + with pytest.raises(ValueError, match="preserve buffer shape"): + head(torch.tensor(0)) + assert head.a.item() == 0 + assert head.b.shape == (1,) + assert export_head(trainer, "student", "head").buffer_revision == before + if live is not None: + assert live.take_publication() is None + + +@pytest.mark.parametrize("client", (False, True)) +def test_failing_handle_forward_hook_does_not_publish_buffers(client): + trainer, rank = _trainer("student") + native = rank.module("head", lambda: torch.nn.BatchNorm1d(2), checkpoint="student") + live = ( + LiveHead( + export_head(trainer, "student", "head"), + torch.nn.BatchNorm1d(2), + CotangentCollector(), + ) + if client + else None + ) + head = native if live is None else _module(live) + + def fail(*args): + raise RuntimeError("forward hook failed") + + hook = head.register_forward_hook(fail) + with pytest.raises(RuntimeError, match="forward hook failed"): + head(torch.ones(4, 2)) + assert head.num_batches_tracked.item() == 0 + torch.testing.assert_close(head.running_mean, torch.zeros(2)) + if live is not None: + assert live.take_publication() is None + hook.remove() + head(torch.ones(4, 2)) + assert head.num_batches_tracked.item() == 1 + + +@pytest.mark.parametrize("client", (False, True)) +@pytest.mark.parametrize("failure", (None, "pre", "forward", "post", "always")) +@pytest.mark.parametrize("pending_before", (False, True)) +@pytest.mark.parametrize("use_saved_alias", (False, True)) +def test_handle_hooks_share_one_buffer_publication( + client, failure, pending_before, use_saved_alias +): + calls = [] + + class Counter(torch.nn.Module): + count: torch.Tensor + alias: torch.Tensor + + def __init__(self): + super().__init__() + self.register_buffer("count", torch.tensor(0.0)) + self.register_buffer("alias", self.count) + + def forward(self, value): + assert self.count is self.alias + self.count.add_(10) + calls.append("forward") + if failure == "forward": + raise RuntimeError("forward failed") + return value + self.count + + trainer, rank = _trainer("student") + native = rank.module("head", Counter, checkpoint="student") + live = ( + LiveHead( + export_head(trainer, "student", "head"), Counter(), CotangentCollector() + ) + if client + else None + ) + head = native if live is None else _module(live) + initial = 7 if pending_before else 0 + if pending_before: + head.count.add_(initial) + revision = export_head(trainer, "student", "head").buffer_revision + saved_alias = head.count + + def mutate(module, amount, phase): + assert module.count is module.alias + (saved_alias if use_saved_alias else module.count).add_(amount) + calls.append(phase) + if failure == phase: + raise RuntimeError(f"{phase} failed") + + head.register_forward_pre_hook(lambda module, args: mutate(module, 1, "pre")) + head.register_forward_hook(lambda module, args, out: mutate(module, 100, "post")) + head.register_forward_hook( + lambda module, args, out: mutate(module, 1000, "always"), always_call=True + ) + if failure is not None: + with pytest.raises(RuntimeError, match=f"{failure} failed"): + head(torch.tensor(2.0)) + assert calls[-1] == "always" + expected = initial + else: + result = head(torch.tensor(2.0)) + assert calls == ["pre", "forward", "post", "always"] + assert result.item() == initial + 13 + expected = initial + 1111 + assert head.count.item() == expected + assert head.count is head.alias + if live is not None: + update = live.take_publication() + assert (update is not None) == (pending_before or failure is None) + if update is not None: + execute_head_operation(trainer, "head_publish", (update,)) + assert native.count.item() == expected + assert export_head(trainer, "student", "head").buffer_revision == revision + int( + failure is None or (client and pending_before) + ) + + +@pytest.mark.parametrize("client", (False, True)) +def test_handle_hook_parameter_gradients_keep_original_version(client): + trainer, rank = _trainer("student") + native = rank.module("head", TiedHead, checkpoint="student") + collector = CotangentCollector() + live = ( + LiveHead(export_head(trainer, "student", "head"), TiedHead(), collector) + if client + else None + ) + head = native if live is None else _module(live) + parameter = head.left + head.register_forward_hook(lambda module, args, out: out + module.left.square()) + value = torch.tensor(3.0, requires_grad=True) + old = head(value) + native.left.data.fill_(4) + trainer._checkpoint_slots["student"].revision += 1 + if live is not None: + live.refresh(export_head(trainer, "student", "head")) + packets = collector.backward(old) + assert len(packets) == 1 + trainer._commit_versioned_gradients(head_gradient_targets(trainer, packets[0])) + else: + with trainer._gradient_transaction(): + old.backward() + assert head.left is parameter + torch.testing.assert_close(native.left.grad, torch.tensor(17.0)) + torch.testing.assert_close(value.grad, torch.tensor(4.0)) + + +@pytest.mark.parametrize("client", (False, True)) +def test_recursive_handle_hook_does_not_publish_before_outer_failure(client): + trainer, rank = _trainer("student") + native = rank.module("head", TiedHead, checkpoint="student") + live = ( + LiveHead( + export_head(trainer, "student", "head"), TiedHead(), CotangentCollector() + ) + if client + else None + ) + head = native if live is None else _module(live) + + def recurse(module, args): + module.offset.add_(1) + if args[0].item() == 1: + module(torch.tensor(2.0)) + raise RuntimeError("outer call failed") + + head.register_forward_pre_hook(recurse) + with pytest.raises(RuntimeError, match="outer call failed"): + head(torch.tensor(1.0)) + assert head.offset.item() == 1 + assert export_head(trainer, "student", "head").buffer_revision == 0 + if live is not None: + assert live.take_publication() is None + + +@pytest.mark.parametrize("client", (False, True)) +@pytest.mark.parametrize( + "scenario", + ("alias_failure", "reentry_failure", "reentry_success", "successful_inner"), +) +def test_nested_handle_calls_preserve_active_captures(client, scenario): + trainer, rank = _trainer("student") + native = { + name: rank.module(name, TiedHead, checkpoint="student") + for name in ("outer", "inner") + } + collector = CotangentCollector() + live = ( + { + name: LiveHead(export_head(trainer, "student", name), TiedHead(), collector) + for name in native + } + if client + else {} + ) + heads = {name: _module(value) for name, value in live.items()} if client else native + outer, inner = heads["outer"], heads["inner"] + outer_alias = outer.offset + seen = [] + + def outer_hook(module, args): + seen.append(module.offset.item()) + module.offset.add_(10) + if args[0].item() == 0: + inner(torch.tensor(1.0)) + if scenario == "successful_inner": + raise RuntimeError("outer failed") + + def inner_hook(module, args): + module.offset.add_(100) + outer_alias.add_(1) + if scenario.startswith("reentry"): + outer(torch.tensor(2.0)) + if scenario.endswith("failure"): + raise RuntimeError("inner failed") + + outer.register_forward_pre_hook(outer_hook) + inner.register_forward_pre_hook(inner_hook) + if scenario == "reentry_success": + assert outer(torch.tensor(0.0)).item() == 24 + else: + with pytest.raises(RuntimeError, match="failed"): + outer(torch.tensor(0.0)) + assert seen == ([1, 12] if scenario.startswith("reentry") else [1]) + expected = { + "outer": 22 if scenario == "reentry_success" else 1, + "inner": 101 if scenario in ("reentry_success", "successful_inner") else 1, + } + for name, value in heads.items(): + assert value.offset.item() == expected[name] + if client: + update = live[name].take_publication() + assert (update is not None) == (expected[name] != 1) + if update is not None: + execute_head_operation(trainer, "head_publish", (update,)) + assert native[name].offset.item() == expected[name] + assert export_head(trainer, "student", name).buffer_revision == int( + expected[name] != 1 + ) + + +def test_inplace_operation_snapshots_readonly_checkpoint_parameter(): + trainer, rank = _trainer("student") + parameter = rank.parameter( + "weight", lambda: torch.tensor(2.0), checkpoint="student" + ) + input_value = torch.tensor(3.0, requires_grad=True) + original = (input_value * 1).mul_(parameter) + parameter.data.fill_(7) + with trainer._gradient_transaction(): + original.backward() + torch.testing.assert_close(parameter.grad, torch.tensor(3.0)) + torch.testing.assert_close(input_value.grad, torch.tensor(2.0)) + parameter.grad = None + stale = torch.tensor(3.0).mul_(parameter) + trainer._checkpoint_slots["student"].revision += 3 + with ( + pytest.raises(RuntimeError, match="staleness"), + trainer._gradient_transaction(), + ): + stale.backward() + assert parameter.grad is None + + +def _live_buffer_authority_worker(process_rank, init_method, asymmetric=False): + from datetime import timedelta + + import torch.distributed as dist + + from art.trainer_rank._heads import synchronize_head_buffers + + dist.init_process_group( + "gloo", + rank=process_rank, + world_size=2, + init_method=init_method, + timeout=timedelta(seconds=30), + ) + try: + trainer, rank = _trainer("student") + native = rank.module( + "head", lambda: torch.nn.BatchNorm1d(2), checkpoint="student" + ) + if asymmetric: + if process_rank == 1: + if asymmetric == "buffers": + del native._buffers["running_mean"] + else: + del trainer._checkpoint_slots["student"].custom["head"] + with pytest.raises( + (RuntimeError, ValueError), + match="registrations differ|preserve buffer names", + ): + synchronize_head_buffers(trainer) + dist.barrier() + return + live = LiveHead( + export_head(trainer, "student", "head"), + torch.nn.BatchNorm1d(2), + CotangentCollector(), + ) + if process_rank == 1: + _module(live)(torch.ones(4, 2)) + update = live.take_publication() + execute_head_operation( + trainer, "head_publish", () if update is None else (update,) + ) + synchronize_head_buffers(trainer) + live.refresh(export_head(trainer, "student", "head")) + assert native.num_batches_tracked.item() == 0 + assert _module(live).num_batches_tracked.item() == 0 + synchronized_revision = live.state.buffer_revision + synchronize_head_buffers(trainer) + assert ( + export_head(trainer, "student", "head").buffer_revision + == synchronized_revision + ) + if process_rank == 1: + _module(live)(torch.ones(4, 2)) + update = live.take_publication() + execute_head_operation( + trainer, "head_publish", () if update is None else (update,) + ) + assert native.num_batches_tracked.item() == process_rank + finally: + dist.destroy_process_group() + + +def test_distributed_live_buffer_refresh_accepts_dp_zero_authority(tmp_path): + torch.multiprocessing.spawn( + _live_buffer_authority_worker, + args=(f"file://{tmp_path / 'live_heads'}",), + nprocs=2, + join=True, + ) + + +@pytest.mark.parametrize("kind", ("parameter", "buffer")) +def test_inplace_operation_snapshots_readonly_client_tensor(kind): + trainer, rank = _trainer("student") + factory = lambda: torch.tensor(2.0) + native = getattr(rank, kind)("scale", factory, checkpoint="student") + collector = CotangentCollector() + live = LiveHead(export_head(trainer, "student", "scale"), factory(), collector) + x = torch.tensor(3.0, requires_grad=True) + loss = (x * 1).mul_(_tensor(live)) + with torch.no_grad(): + native.fill_(7) + trainer._checkpoint_slots["student"].revision += 1 + live.refresh(export_head(trainer, "student", "scale")) + packets = collector.backward(loss) + torch.testing.assert_close(x.grad, torch.tensor(2.0)) + assert live.take_publication() is None + if kind == "parameter": + trainer._commit_versioned_gradients(head_gradient_targets(trainer, packets[0])) + torch.testing.assert_close(native.grad, torch.tensor(3.0)) + + +def test_logical_callback_reentrant_head_rejects_before_gradient_publication(): + from types import SimpleNamespace + + from art.trainer_rank._heads import logical_register_head + + trainer, _ = _trainer("student") + collector = CotangentCollector() + + def invoke(operation, kind, payload): + assert operation == "head" + return execute_head_operation(trainer, kind, payload) + + view = SimpleNamespace( + _rank=trainer, + _invoke=invoke, + _executor=SimpleNamespace( + state=SimpleNamespace(collector=collector), invoke=invoke + ), + device=torch.device("cpu"), + ) + + def callback(rank): + head = logical_register_head( + rank, "module", "head", lambda: TiedHead(True), checkpoint="student" + ) + loss = head(torch.tensor(3.0, requires_grad=True)) + # The logical executor submits packets only after local collection succeeds. + packets = collector.backward(loss) + trainer._commit_versioned_gradients( + [ + target + for packet in packets + for target in head_gradient_targets(trainer, packet) + ] + ) + + with pytest.raises(RuntimeError, match="nested remote backward is unsupported"): + callback(view) + custom = trainer._checkpoint_slots["student"].custom["head"].value + assert isinstance(custom, torch.nn.Module) + assert all(parameter.grad is None for parameter in custom.parameters()) + assert ( + getattr(trainer, "_logical_head_handles")[ + ("student", "head") + ].take_publication() + is None + ) + + +@pytest.mark.parametrize("mode", ("rank", "zero")) +@pytest.mark.parametrize("kind", ("buffer", "module")) +def test_logical_registration_recovers_after_rejected_buffer_publication(mode, kind): + async def run(): + trainer, rank = _trainer("student") + factory = TiedHead if kind == "module" else lambda: torch.tensor(1.0) + native = getattr(rank, kind)("head", factory, checkpoint="student") + retained = [] + + def buffer(head): + return head.offset if kind == "module" else head + + def register(view): + return getattr(view, kind)("head", factory, checkpoint="student") + + def conflict(view): + old = register(view) + retained.append(old) + buffer(old).add_(1) + # The logical copy now has a publication against the old revision. + buffer(native).add_(5) + + with pytest.raises(RuntimeError, match="changed before publication"): + await run_rank_callback(trainer, conflict, mode=mode) + (old,) = retained + with pytest.raises(RuntimeError, match="publication failed"): + buffer(old).item() + assert buffer(native).item() == 6 + + def recover(view): + fresh = register(view) + assert fresh is not old + assert register(view) is fresh + assert buffer(fresh).item() == 6 + if kind == "module": + assert fresh.left is fresh.right + assert fresh.offset is fresh.other_offset + assert fresh(torch.tensor(3.0)).item() == 20 + buffer(fresh).add_(2) + return fresh + + fresh = (await run_rank_callback(trainer, recover, mode=mode)).value + assert buffer(native).item() == 8 + assert (await run_rank_callback(trainer, register, mode=mode)).value is fresh + assert buffer(fresh).item() == 8 + with pytest.raises(RuntimeError, match="publication failed"): + buffer(old).item() + + asyncio.run(run()) + + +@pytest.mark.parametrize("client", (False, True)) +@pytest.mark.parametrize("grad_enabled", (False, True)) +@pytest.mark.parametrize( + "view", + ( + lambda value: value[1:], + lambda value: value.view(2, 2), + lambda value: value.narrow(0, 1, 2), + ), +) +def test_live_buffer_view_mutation_rejects_without_silent_write( + client, grad_enabled, view +): + trainer, rank = _trainer("student") + native = rank.buffer("stats", lambda: torch.zeros(4), checkpoint="student") + live = ( + LiveHead( + export_head(trainer, "student", "stats"), + torch.zeros(4), + CotangentCollector(), + ) + if client + else None + ) + buffer = native if live is None else _tensor(live) + revision = export_head(trainer, "student", "stats").buffer_revision + with torch.set_grad_enabled(grad_enabled): + snapshot = view(buffer) + with pytest.raises( + RuntimeError, match="Views of live checkpoint buffers are read-only" + ): + snapshot.fill_(7) + copy = snapshot.clone() + copy.fill_(7) + buffer[2:] = 7 + torch.testing.assert_close(buffer, torch.tensor([0.0, 0.0, 7.0, 7.0])) + if live is not None: + update = live.take_publication() + assert update is not None + execute_head_operation(trainer, "head_publish", (update,)) + assert export_head(trainer, "student", "stats").buffer_revision == revision + 1 + + +@pytest.mark.parametrize("client", (False, True)) +def test_functional_batchnorm_no_grad_publishes_buffer_changes(client): + trainer, rank = _trainer("student") + native = rank.buffer("mean", lambda: torch.zeros(2), checkpoint="student") + live = ( + LiveHead( + export_head(trainer, "student", "mean"), + torch.zeros(2), + CotangentCollector(), + ) + if client + else None + ) + mean = native if live is None else _tensor(live) + before = export_head(trainer, "student", "mean").buffer_revision + with torch.no_grad(): + torch.nn.functional.batch_norm( + torch.ones(4, 2), mean, torch.ones(2), training=True + ) + torch.testing.assert_close(mean, torch.full((2,), 0.1)) + if live is not None: + update = live.take_publication() + assert update is not None + execute_head_operation(trainer, "head_publish", (update,)) + assert export_head(trainer, "student", "mean").buffer_revision == before + 1 + + +@pytest.mark.parametrize("client", (False, True)) +def test_live_parameter_metadata_does_not_capture_weights(monkeypatch, client): + trainer, rank = _trainer("student") + native = rank.parameter("weight", lambda: torch.zeros(3, 4), checkpoint="student") + live = ( + LiveHead( + export_head(trainer, "student", "weight"), + torch.zeros(3, 4), + CotangentCollector(), + ) + if client + else None + ) + value = native if live is None else _tensor(live) + + def unexpected(*args, **kwargs): + raise AssertionError("metadata inspection must not capture parameter values") + + monkeypatch.setattr(trainer, "_snapshot_parameter", unexpected) + if live is not None: + monkeypatch.setattr(live, "capture", unexpected) + assert value.size() == (3, 4) + assert value.numel() == 12 + assert value.dim() == 2 + assert value.shape == (3, 4) + assert value.dtype == torch.float32 + assert value.device == torch.device("cpu") + assert value.requires_grad + assert value.is_leaf + assert value.grad_fn is None + assert value.grad is None + + +@pytest.mark.parametrize("client", (False, True)) +@pytest.mark.parametrize("property_name", ("T", "mT", "H", "mH", "real", "imag")) +@pytest.mark.parametrize("complex_dtype", (False, True)) +def test_parameter_tensor_properties_capture_immutable_versions( + client, property_name, complex_dtype +): + initial = torch.tensor([[1 + 2j, 3 - 4j], [5 + 6j, 7 - 8j]]) + if not complex_dtype: + if property_name == "imag": + pytest.skip("imag requires a complex tensor") + initial = initial.real.clone() + trainer, rank = _trainer("student") + native = rank.parameter("weight", lambda: initial.clone(), checkpoint="student") + collector = CotangentCollector() + live = ( + LiveHead(export_head(trainer, "student", "weight"), initial, collector) + if client + else None + ) + value = native if live is None else _tensor(live) + expected = initial.clone().requires_grad_() + getattr(expected, property_name).abs().square().sum().backward() + old = getattr(value, property_name).abs().square().sum() + stale = getattr(value, property_name).abs().square().sum() + native.data.copy_(initial * 3) + trainer._checkpoint_slots["student"].revision += 1 + if live is not None: + live.refresh(export_head(trainer, "student", "weight")) + torch.testing.assert_close( + getattr(value, property_name), getattr(initial * 3, property_name) + ) + if live is None: + with trainer._gradient_transaction(): + old.backward() + else: + packets = collector.backward(old) + assert len(packets) == 1 + trainer._commit_versioned_gradients(head_gradient_targets(trainer, packets[0])) + assert value.grad is None + torch.testing.assert_close(native.grad, expected.grad) + trainer._checkpoint_slots["student"].revision += 2 + with pytest.raises(RuntimeError, match="staleness"): + if live is None: + with trainer._gradient_transaction(): + stale.backward() + else: + head_gradient_targets(trainer, collector.backward(stale)[0]) + torch.testing.assert_close(native.grad, expected.grad) + + +@pytest.mark.parametrize("client", (False, True)) +@pytest.mark.parametrize("property_name", ("T", "mT", "H", "mH", "real", "imag")) +def test_buffer_tensor_properties_reject_unpublished_mutation(client, property_name): + initial = torch.ones(2, 2, dtype=torch.complex64) + trainer, rank = _trainer("student") + native = rank.buffer("stats", lambda: initial.clone(), checkpoint="student") + live = ( + LiveHead( + export_head(trainer, "student", "stats"), initial, CotangentCollector() + ) + if client + else None + ) + value = native if live is None else _tensor(live) + with pytest.raises(RuntimeError, match="Views of live checkpoint buffers"): + getattr(value, property_name).fill_(7) + torch.testing.assert_close(native, initial) + assert export_head(trainer, "student", "stats").buffer_revision == 0 + if live is not None: + assert live.take_publication() is None + + +@pytest.mark.parametrize("asymmetric", (True, "buffers")) +def test_distributed_buffer_registration_mismatch_fails_on_every_rank( + tmp_path, asymmetric +): + torch.multiprocessing.spawn( + _live_buffer_authority_worker, + args=(f"file://{tmp_path / 'mismatched_heads'}", asymmetric), + nprocs=2, + join=True, + ) + + +@pytest.mark.parametrize("client", (False, True)) +def test_module_buffer_view_mutation_rejects_without_publishing(client): + trainer, rank = _trainer("student") + native = rank.module("head", lambda: torch.nn.BatchNorm1d(2), checkpoint="student") + live = ( + LiveHead( + export_head(trainer, "student", "head"), + torch.nn.BatchNorm1d(2), + CotangentCollector(), + ) + if client + else None + ) + head = native if live is None else _module(live) + with pytest.raises( + RuntimeError, match="Views of live checkpoint buffers are read-only" + ): + head.running_mean[:1].fill_(7) + torch.testing.assert_close(head.running_mean, torch.zeros(2)) + if live is not None: + assert live.take_publication() is None + + +@pytest.mark.parametrize("client", (False, True)) +def test_stateful_function_on_buffer_snapshot_view_rejects(client): + trainer, rank = _trainer("student") + native = rank.buffer("mean", lambda: torch.zeros(2), checkpoint="student") + live = ( + LiveHead( + export_head(trainer, "student", "mean"), + torch.zeros(2), + CotangentCollector(), + ) + if client + else None + ) + mean = native if live is None else _tensor(live) + with ( + torch.no_grad(), + pytest.raises( + RuntimeError, match="Stateful operations on buffer snapshot views" + ), + ): + torch.nn.functional.batch_norm( + torch.ones(4, 2), mean[:], torch.ones(2), training=True + ) + torch.testing.assert_close(mean, torch.zeros(2)) + if live is not None: + assert live.take_publication() is None + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +@pytest.mark.parametrize("client", (False, True)) +def test_cuda_live_buffer_views_and_functional_publication(client): + trainer, rank = _trainer("student") + trainer.device = torch.device("cuda", 0) + native = rank.buffer("mean", lambda: torch.zeros(2), checkpoint="student") + live = ( + LiveHead( + export_head(trainer, "student", "mean"), + torch.zeros(2, device="cuda"), + CotangentCollector(), + ) + if client + else None + ) + mean = native if live is None else _tensor(live) + with pytest.raises( + RuntimeError, match="Views of live checkpoint buffers are read-only" + ): + mean[:].fill_(7) + with torch.no_grad(): + torch.nn.functional.batch_norm( + torch.ones(4, 2, device="cuda"), + mean, + torch.ones(2, device="cuda"), + training=True, + ) + if live is not None: + update = live.take_publication() + assert update is not None + execute_head_operation(trainer, "head_publish", (update,)) + torch.testing.assert_close(native, torch.full((2,), 0.1, device="cuda")) + assert export_head(trainer, "student", "mean").buffer_revision == 1 + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_cuda_buffer_sync_stages_cpu_authority_before_comparison(tmp_path): + import torch.distributed as dist + + from art.trainer_rank._heads import synchronize_head_buffers + + dist.init_process_group( + "gloo", rank=0, world_size=1, init_method=f"file://{tmp_path / 'cuda_sync'}" + ) + try: + trainer, rank = _trainer("student") + trainer.device = torch.device("cuda", 0) + buffer = rank.buffer("mean", lambda: torch.ones(2), checkpoint="student") + synchronize_head_buffers(trainer) + torch.testing.assert_close(buffer, torch.ones(2, device="cuda")) + assert export_head(trainer, "student", "mean").buffer_revision == 0 + finally: + dist.destroy_process_group() diff --git a/tests/unit/test_trainer_rank_memory_admission.py b/tests/unit/test_trainer_rank_memory_admission.py new file mode 100644 index 000000000..9a0b087fc --- /dev/null +++ b/tests/unit/test_trainer_rank_memory_admission.py @@ -0,0 +1,564 @@ +from dataclasses import replace +from types import SimpleNamespace +from typing import Any, cast + +import pytest +import torch + +from art.trainer_rank import ( + ForwardInput, + ForwardOptions, + ImportanceSamplingGradientCorrection, + TrainerRank, + _impl, +) + + +@pytest.fixture +def rank(monkeypatch): + class Model(torch.nn.Module): + def __init__(self): + super().__init__() + self.weight = torch.nn.Parameter(torch.zeros((), dtype=torch.bfloat16)) + self.config = SimpleNamespace( + hidden_size=8, num_layers=4, padded_vocab_size=32 + ) + self.decoder = object() + + def _preprocess(self, *args, **kwargs): + return None + + result = TrainerRank( + cast( + Any, + SimpleNamespace( + model=[Model()], + optimizer=None, + provider=SimpleNamespace( + hidden_size=8, num_layers=4, recompute_granularity="full" + ), + model_support_handler=SimpleNamespace(build_gdn_execution_spec=False), + ), + ) + ) + monkeypatch.setattr(result, "_graph_memory_policy_enabled", lambda: True) + monkeypatch.setattr(result, "_available_memory_bytes", lambda: 250) + monkeypatch.setattr(result, "_available_cpu_memory_bytes", lambda: 1_000_000) + monkeypatch.setattr( + result, + "_estimate_group_request_output_bytes", + lambda requests: 10 * len(requests), + ) + monkeypatch.setattr( + result, + "_plan_cost", + lambda plan: _impl._SubforwardCost( + 100 * sum(len(g.items) for g in plan.groups), + 80 * sum(len(g.items) for g in plan.groups), + ), + ) + return result + + +def _requests(options=None): + return [ + ForwardInput( + input_tokens=torch.tensor([i, 10, 11]), hidden_states=True, options=options + ) + for i in range(4) + ] + + +def test_complete_root_can_split_and_replay_where_gpu_retention_refuses(rank): + requests = _requests() + plan, check = rank._plan_admissible_forward( + requests, checkpoint=_impl.Unset, context="test" + ) + assert isinstance(plan, _impl._SplitForwardPlan) + assert plan.subforward_count == 2 + assert sorted(i for indices in plan.request_indices for i in indices) == list( + range(4) + ) + assert check.fits and check.estimated_required_bytes == 240 + assert all(g.memory_placement.backward_state == "replay" for g in plan.groups) + assert all(g.memory_placement.output_device == "model" for g in plan.groups) + with pytest.raises(_impl.TrainerRankMemoryError): + rank._plan_admissible_forward( + _requests(ForwardOptions(backward_state="gpu")), + checkpoint=_impl.Unset, + context="test", + ) + + +def test_auto_cpu_outputs_preserve_saved_state_without_replay(rank): + plan, check = rank._plan_admissible_forward( + _requests(ForwardOptions(output_device="auto")), + checkpoint=_impl.Unset, + context="test", + ) + assert check.fits and check.estimated_required_bytes == 220 + assert all(g.memory_placement.backward_state == "cpu" for g in plan.groups) + assert all(g.memory_placement.output_device == "cpu" for g in plan.groups) + + +@pytest.mark.parametrize( + "transfer_seconds, samples, expected", + [(0.3, 2, "replay"), (0.01, 2, "cpu"), (0.3, 1, "cpu"), (0.0001, 2, "cpu")], +) +def test_measured_fallback_costs_choose_only_after_gpu_refusal( + rank, monkeypatch, transfer_seconds, samples, expected +): + monkeypatch.setattr(rank, "_available_memory_bytes", lambda: 150) + requests = _requests(ForwardOptions(output_device="cpu"))[:2] + children = tuple(rank._plan_flat_forward([request]) for request in requests) + plan = _impl._SplitForwardPlan(children, ((0,), (1,)), 2) + for _ in range(samples): + rank._record_graph_forward_time(children[0], 0.1) + rank._graph_cache = SimpleNamespace( + handles=lambda: (), + transfer_stats=SimpleNamespace( + offload_bytes=140, + restore_bytes=140, + offload_seconds=transfer_seconds, + restore_seconds=transfer_seconds, + ), + ) + selected, check = rank._admit_graph_memory(plan) + assert check.fits + assert {g.memory_placement.backward_state for g in selected.groups} == {expected} + assert check.fallback_costs["preferred"] == expected + assert check.fallback_costs["source"] == ( + "measured_forward_and_transfers" + if samples >= 2 and transfer_seconds >= 0.001 + else "insufficient_samples" + ) + + +def test_forward_timing_requires_matching_shape_and_gpu_retention(rank): + plan = rank._plan_flat_forward(_requests()[:1]) + for seconds in (100.0, 0.1, 0.2, 0.3): + rank._record_graph_forward_time(plan, seconds) + assert next(rank._graph_memory_units(plan))[2].replay_seconds == 0.3 + other = replace(plan, packed_tokens=plan.packed_tokens + 1) + assert next(rank._graph_memory_units(other))[2].replay_seconds is None + group = plan.groups[0] + segment = group.packed.segments[0] + different_tree = replace( + group.packed, + segments=(replace(segment, parent_id=segment.parent_id + 1),), + ) + other = replace(plan, groups=(replace(group, packed=different_tree),)) + assert next(rank._graph_memory_units(other))[2].replay_seconds is None + replay, _ = rank._admit_graph_memory( + rank._plan_flat_forward(_requests(ForwardOptions(backward_state="replay"))[:1]) + ) + rank._record_graph_forward_time(replay, 99.0) + assert next(rank._graph_memory_units(plan))[2].replay_seconds == 0.3 + + +def test_gpu_headroom_does_not_consult_transfer_costs(rank): + class Cache: + @staticmethod + def handles(): + return () + + @property + def transfer_stats(self): + raise AssertionError("GPU headroom path consulted fallback costs") + + rank._graph_cache = Cache() + selected, check = rank._admit_graph_memory(rank._plan_flat_forward(_requests()[:1])) + assert check.fits and check.fallback_costs is None + assert selected.groups[0].memory_placement.backward_state == "gpu" + + +@pytest.mark.parametrize( + "policy, expected", + [ + (ForwardOptions(backward_state="cpu"), "cpu"), + (ForwardOptions(backward_state="replay"), "replay"), + (ForwardOptions(allow_replay=False), "cpu"), + (ForwardOptions(allow_cpu_offload=False), "replay"), + ], +) +def test_fallback_costs_preserve_forced_and_disabled_policies( + rank, monkeypatch, policy, expected +): + monkeypatch.setattr(rank, "_available_memory_bytes", lambda: 150) + requests = _requests(replace(policy, output_device="cpu"))[:2] + plan = _impl._SplitForwardPlan( + tuple(rank._plan_flat_forward([request]) for request in requests), + ((0,), (1,)), + 2, + ) + selected, check = rank._admit_graph_memory(plan) + assert check.fits + assert {g.memory_placement.backward_state for g in selected.groups} == {expected} + + +def test_mixed_policies_split_physical_groups_and_share_estimator_keys(rank): + requests = _requests()[:2] + requests[1] = replace( + requests[1], + options=ForwardOptions(backward_state="replay", output_device="cpu"), + ) + plan = rank._plan_flat_forward(requests) + assert len(plan.groups) == 2 + estimated = rank._estimate_flat_forward(requests, exact=True) + assert estimated[2] == plan.signature + selected, check = rank._admit_graph_memory(plan) + assert check.fits + assert [g.memory_placement.backward_state for g in selected.groups] == [ + "gpu", + "replay", + ] + assert [g.memory_placement.output_device for g in selected.groups] == [ + "model", + "cpu", + ] + assert selected.signature.memory_placement == (("gpu", "model"), ("replay", "cpu")) + assert selected.signature != plan.signature + + +def test_equal_resolved_policies_still_pack_together(rank): + requests = _requests()[:2] + requests[1] = replace(requests[1], options=ForwardOptions(max_gradient_staleness=2)) + assert len(rank._plan_flat_forward(requests).groups) == 1 + + +def test_cpu_shortage_is_not_overwritten_by_fresh_gpu_check(rank, monkeypatch): + monkeypatch.setattr(rank, "_available_cpu_memory_bytes", lambda: 0) + plan = rank._plan_flat_forward(_requests()[:1]) + selected, check = rank._admit_graph_memory(plan) + assert not check.fits and not check.cpu_fits + refusal = _impl._ForwardRefusal(selected, check, "test") + with pytest.raises(_impl.TrainerRankMemoryError, match="per-rank CPU headroom=0"): + rank._recover_admission( + lambda: (selected, check), + lambda value: value, + lambda value, check: (value[0], check), + context="test", + sync_across_dp=True, + ) + assert "CPU retained" in str(refusal.error("test")) + + +def test_retained_weight_versions_stay_in_gpu_budget_even_for_replay(rank, monkeypatch): + monkeypatch.setattr( + rank, "_lora_version_capture_bytes", lambda *args: 200, raising=False + ) + request = _requests(ForwardOptions(backward_state="replay", output_device="cpu"))[ + :1 + ] + _, check = rank._admit_graph_memory(rank._plan_flat_forward(request)) + assert check.estimated_required_bytes == 300 + assert not check.fits + + +def test_gradient_transaction_reserves_one_aggregate_per_slot_across_children( + rank, monkeypatch +): + monkeypatch.setattr(rank, "_lora_gradient_staging_bytes", lambda _ref: 60) + requests = _requests(ForwardOptions(backward_state="replay", output_device="cpu")) + split = _impl._SplitForwardPlan( + tuple(rank._plan_flat_forward([request]) for request in requests), + tuple((i,) for i in range(4)), + 4, + ) + _, check = rank._admit_graph_memory(split) + assert check.estimated_required_bytes == 160 + assert check.fits + + +def test_prior_larger_replay_workspace_survives_new_small_root_admission(rank): + rank._graph_cache = SimpleNamespace( + handles=lambda: ("old",), + state=lambda _handle: SimpleNamespace(restore_workspace_bytes=300), + ) + requests = _requests(ForwardOptions(backward_state="replay", output_device="cpu"))[ + :1 + ] + _, check = rank._admit_graph_memory(rank._plan_flat_forward(requests)) + assert not check.fits + assert check.estimated_required_bytes == 300 + + +def test_prior_checkpoint_staging_deduplicates_old_and_new_graph_targets( + rank, monkeypatch +): + rank._checkpoint_slots["old"] = _impl._CheckpointSlot() + old = rank._slot_ref("old") + monkeypatch.setattr( + rank, "_lora_gradient_staging_bytes", lambda ref: 40 if ref == old else 0 + ) + monkeypatch.setattr(rank, "_lora_version_capture_bytes", lambda *_args: 0) + rank._graph_cache = SimpleNamespace( + handles=lambda: ("old1", "old2"), + state=lambda _handle: SimpleNamespace( + restore_workspace_bytes=150, + checkpoint_versions=(SimpleNamespace(checkpoint="old"),), + ), + ) + requests = _requests(ForwardOptions(backward_state="replay", output_device="cpu"))[ + :1 + ] + plan = rank._plan_flat_forward(requests) + _, check = rank._admit_graph_memory(plan) + assert check.estimated_required_bytes == 190 + plan = replace(plan, groups=(replace(plan.groups[0], slot_ref=old),)) + _, check = rank._admit_graph_memory(plan) + assert check.estimated_required_bytes == 190 + + +@pytest.mark.parametrize("existing_gradient", [False, True]) +def test_outstanding_forwards_reserve_sequential_atomic_gradient_publication( + rank, monkeypatch, existing_gradient +): + parameter = torch.nn.Parameter(torch.ones(16)) + if existing_gradient: + parameter.grad = torch.zeros_like(parameter) + size = parameter.numel() * parameter.element_size() + baseline_grad_bytes = size if existing_gradient else 0 + rank._checkpoint_slots["student"] = _impl._CheckpointSlot(params=(parameter,)) + ref = rank._slot_ref("student") + monkeypatch.setattr(rank, "_iter_slot_parameters", lambda _ref: iter((parameter,))) + monkeypatch.setattr(rank, "_lora_version_capture_bytes", lambda *_args: 0) + monkeypatch.setattr(rank, "_available_memory_bytes", lambda: 100 + 3 * size) + options = ForwardOptions(backward_state="replay", output_device="cpu") + plan = rank._plan_flat_forward(_requests(options)[:1]) + plan = replace(plan, groups=(replace(plan.groups[0], slot_ref=ref),)) + # Both forwards are admitted before either one publishes a gradient. Replay + # and CPU outputs need not release GPU storage between their backwards. + checks = [rank._admit_graph_memory(plan)[1] for _ in range(2)] + reservation = min(check.estimated_required_bytes for check in checks) - 100 + snapshot = rank._snapshot_parameter( + parameter, rank._capture_checkpoint_version("student") + ) + losses = [(snapshot * factor).sum() for factor in (2, 3)] + state = rank._version_state() + publish = state._publish + observed = [] + + def measure(prepared): + batch = state._transaction + assert batch is not None + tensors = [gradient for _, gradient in batch.gradients.values()] + tensors.extend( + tensor + for _, combined, previous in prepared.parameters + for tensor in (combined, previous) + if tensor is not None + ) + storages = { + tensor.untyped_storage().data_ptr(): tensor.untyped_storage().nbytes() + for tensor in tensors + } + live = sum(storages.values()) + observed.append(live) + assert live - baseline_grad_bytes <= reservation + publish(prepared) + + monkeypatch.setattr(state, "_publish", measure) + for loss in losses: + with rank._gradient_transaction(): + loss.backward() + assert observed == ([3 * size, 3 * size] if existing_gradient else [size, 3 * size]) + torch.testing.assert_close(parameter.grad, torch.full_like(parameter, 5)) + + +def test_empty_output_does_not_pin_hidden_storage_and_keeps_autograd(): + hidden = torch.randn(1024, 8, requires_grad=True) + empty = _impl._select_positions(hidden, torch.empty(0, dtype=torch.long)) + assert empty.untyped_storage().nbytes() == 0 + empty.sum().backward() + assert hidden.grad is not None and hidden.grad.count_nonzero() == 0 + + +def test_correction_metadata_and_explicit_prepass_are_budgeted(rank): + request = ForwardInput( + input_tokens=torch.arange(3), + target_tokens=torch.arange(3), + top_k=2, + ) + group = rank._plan_flat_forward([request]).groups[0] + default = _impl._resolved_request_policy(None) + always = _impl._resolved_request_policy( + ForwardOptions( + stale_gradient_corrections=( + ImportanceSamplingGradientCorrection(policy="always"), + ) + ) + ) + assert _impl._correction_state_bytes(group, default) == 3 * 4 + 6 * 12 + assert _impl._correction_state_bytes(group, always) == 3 * 8 + 6 * 16 + assert ( + _impl._correction_state_bytes( + group, + _impl._resolved_request_policy( + ForwardOptions(stale_gradient_corrections=()) + ), + ) + == 6 * 12 + ) + + +@pytest.mark.parametrize("grad_enabled", [False, True]) +@pytest.mark.parametrize("top_k", [0, 4]) +@pytest.mark.parametrize("label_columns", [0, 1, 2]) +@pytest.mark.parametrize( + "corrections", + [None, (), (ImportanceSamplingGradientCorrection(policy="always"),)], + ids=["default", "disabled", "always"], +) +def test_correction_budget_covers_captured_storage_and_aliases( + rank, grad_enabled, top_k, label_columns, corrections +): + from art.trainer_rank._corrections import capture_forward_corrections + from art.trainer_rank._tensors import flatten_tensors + + labels = ( + torch.arange(3 * label_columns).reshape( + (3,) if label_columns == 1 else (3, label_columns) + ) + if label_columns + else None + ) + options = _impl._resolved_request_policy( + ForwardOptions( + stale_gradient_corrections=_impl.Unset + if corrections is None + else corrections + ) + ) + request = ForwardInput( + input_tokens=torch.arange(3), + target_tokens=labels, + top_k=top_k or None, + hidden_states=True, + no_grad=not grad_enabled, + ) + plan = rank._plan_flat_forward([request]) + group = plan.groups[0] + output = _impl.ForwardOutput( + target_logprobs=None + if labels is None + else torch.zeros_like(labels, dtype=torch.float32, requires_grad=grad_enabled), + top_k=_impl.TopK( + torch.zeros(3, top_k, dtype=torch.float32, requires_grad=grad_enabled), + torch.arange(3 * top_k).reshape(3, top_k), + ) + if top_k + else None, + logits=None, + hidden_states=torch.zeros(3, 1, requires_grad=grad_enabled), + ) + tree = {"output": output, "alias": [output]} + tensors, _ = flatten_tensors(tree) + context = capture_forward_corrections(tree, tensors, options) + retained = sum(tensor.untyped_storage().nbytes() for tensor in context.tensors) + gradients = tuple( + torch.ones_like(tensor) if tensor.requires_grad else None for tensor in tensors + ) + staged = 0 + if corrections and corrections[0].policy == "always": + corrected = context.correct(gradients, tensors) + staged = sum( + after.untyped_storage().nbytes() + for before, after in zip(gradients, corrected, strict=True) + if after is not None and after is not before + ) + assert _impl._correction_state_bytes(group, options) == retained + staged + if not grad_enabled: + assert next(rank._graph_memory_units(plan))[2].replay_bytes == 0 + elif output.top_k is not None: + current = list(tensors) + index = next( + i for i, tensor in enumerate(tensors) if tensor is output.top_k.tokens + ) + current[index] = current[index].flip(-1) + with pytest.raises(RuntimeError, match="changed active top-k token identities"): + context.validate_replay(gradients, current) + + +@pytest.mark.parametrize( + "corrections,expected", + [ + ((), 144), + (None, 144), + ((ImportanceSamplingGradientCorrection(policy="always"),), 192), + ], + ids=["disabled", "default", "always"], +) +def test_top_k_context_must_fit_host_budget_before_forward( + rank, monkeypatch, corrections, expected +): + options = ForwardOptions( + backward_state="gpu", + stale_gradient_corrections=_impl.Unset if corrections is None else corrections, + ) + request = ForwardInput(input_tokens=torch.arange(3), top_k=4, options=options) + plan = rank._plan_flat_forward([request]) + required = _impl._snapshot_tensor_bytes(plan.groups[0]) + 64 * 1024 + expected + monkeypatch.setattr(rank, "_available_cpu_memory_bytes", lambda: required - 1) + _, check = rank._admit_graph_memory(plan) + assert not check.fits and not check.cpu_fits + assert check.cpu_required_bytes == required + monkeypatch.setattr(rank, "_available_cpu_memory_bytes", lambda: required) + _, check = rank._admit_graph_memory(plan) + assert check.fits and check.cpu_fits + + +def test_reclamation_respects_captured_forced_policy_and_cpu_capacity( + rank, monkeypatch +): + decisions = [] + states = { + "forced": SimpleNamespace( + retention="gpu", offloadable=False, replayable=False, offload_bytes=100 + ), + "auto": SimpleNamespace( + retention="gpu", offloadable=True, replayable=True, offload_bytes=100 + ), + } + rank._graph_cache = SimpleNamespace( + handles=lambda: tuple(states), + state=states.__getitem__, + offload=lambda handle: decisions.append(("cpu", handle)), + evict=lambda handle: decisions.append(("replay", handle)), + ) + monkeypatch.setattr(torch.cuda, "empty_cache", lambda: None) + check = _impl._MemoryCheck(1000, 250, False, cpu_required_bytes=50) + assert rank._reclaim_graph_memory(check, sync_across_dp=False) + assert decisions == [("cpu", "auto")] + decisions.clear() + monkeypatch.setattr(rank, "_available_cpu_memory_bytes", lambda: 100) + assert rank._reclaim_graph_memory(check, sync_across_dp=False) + assert decisions == [("replay", "auto")] + + +def test_reclamation_transfer_error_is_exchanged_before_propagating(rank, monkeypatch): + error = RuntimeError("CPU offload allocation failed") + exchanges = [] + + def fail(_handle): + raise error + + def exchange(values, **_kwargs): + exchanges.append(values) + return values + + rank._graph_cache = SimpleNamespace( + handles=lambda: ("graph",), + state=lambda _: SimpleNamespace( + retention="gpu", offloadable=True, replayable=True, offload_bytes=100 + ), + offload=fail, + evict=lambda _: None, + ) + monkeypatch.setattr(rank, "_recovery_reduce", exchange) + with pytest.raises(RuntimeError) as caught: + rank._reclaim_graph_memory( + _impl._MemoryCheck(1000, 250, False), sync_across_dp=False + ) + assert caught.value is error + assert exchanges[-1] == [0.0] diff --git a/tests/unit/test_trainer_rank_memory_policy.py b/tests/unit/test_trainer_rank_memory_policy.py new file mode 100644 index 000000000..9eb0ce1a8 --- /dev/null +++ b/tests/unit/test_trainer_rank_memory_policy.py @@ -0,0 +1,306 @@ +from pathlib import Path +from typing import Literal + +import pytest + +from art.trainer_rank._memory_policy import ( + ForwardMemoryCost, + HostMemoryBudget, + MemoryScope, + choose_memory_placement, + choose_output_placements, + host_memory_budget, + local_rank_count, + placement_cost, +) + + +def test_aggregate_outputs_reserve_explicit_model_before_auto(): + assert choose_output_placements( + [(60, "auto"), (60, "model"), (20, "auto"), (100, "cpu")], + gpu_available_bytes=100, + ) == ("cpu", "model", "model", "cpu") + + +def test_aggregate_outputs_refuse_explicit_model_before_any_copy(): + with pytest.raises(MemoryError, match="require 120 GPU bytes"): + choose_output_placements( + [(60, "model"), (60, "model")], gpu_available_bytes=100 + ) + + +def test_aggregate_outputs_fresh_headroom_accounts_previous_waves(): + outputs: list[tuple[int, Literal["auto", "model", "cpu"]]] = [ + (60, "auto"), + (20, "auto"), + ] + assert choose_output_placements(outputs, gpu_available_bytes=100) == ( + "model", + "model", + ) + assert choose_output_placements(outputs, gpu_available_bytes=20) == ("cpu", "model") + + +def test_staged_gradients_accumulate_across_children_beside_restore_workspace(): + placement = placement_cost( + [ForwardMemoryCost(100, 80, 10, gradient_staging_bytes=30)] * 3, + backward_state="cpu", + output_device="cpu", + ) + assert placement.gpu_retained_bytes == 30 + assert placement.gpu_backward_bytes == 90 + assert placement.gpu_required_bytes == 30 + 90 + 90 + + +def _host(tmp_path: Path, *, v2=True, namespace=False): + proc = tmp_path / "proc" + (proc / "self").mkdir(parents=True) + (proc / "meminfo").write_text("MemTotal: 1000 kB\nMemAvailable: 800 kB\n") + mount = tmp_path / "memory" + child = mount / "pod" / "rank" + child.mkdir(parents=True) + root = "/delegated" if namespace else "/" + member = root.rstrip("/") + "/pod/rank" + (proc / "self/cgroup").write_text( + f"0::{member}\n" if v2 else f"2:cpu,memory:{member}\n" + ) + (proc / "self/mountinfo").write_text( + f"30 1 0:29 {root} {mount} rw - " + + ("cgroup2 cgroup rw\n" if v2 else "cgroup cgroup rw,memory\n") + ) + limit_name = "memory.max" if v2 else "memory.limit_in_bytes" + used_name = "memory.current" if v2 else "memory.usage_in_bytes" + for path, limit, used in ( + (mount, 800_000, 100_000), + (child.parent, 500_000, 300_000), + (child, 400_000, 100_000), + ): + (path / limit_name).write_text(str(limit)) + (path / used_name).write_text(str(used)) + return proc, mount, child, limit_name, used_name + + +@pytest.mark.parametrize("v2", [True, False]) +@pytest.mark.parametrize("namespace", [True, False]) +def test_shared_host_and_ancestor_cgroup_budgets(tmp_path, v2, namespace): + proc, _, child, _, _ = _host(tmp_path, v2=v2, namespace=namespace) + budget = host_memory_budget(local_world_size=4, proc_root=proc) + # The pod ancestor binds: (500000 - 300000 - 10% reserve) / 4. + assert budget.available_bytes == 37_500 + assert len(budget.scopes) == 4 + assert min(budget.scopes, key=lambda s: s.per_rank_available_bytes).name == str( + child.parent + ) + assert ( + budget.available_bytes * 4 + == host_memory_budget(local_world_size=1, proc_root=proc).available_bytes + ) + + +def test_fresh_usage_reduces_additional_offload_credit(tmp_path): + proc, _, child, _, used_name = _host(tmp_path) + first = host_memory_budget(local_world_size=2, proc_root=proc) + (child.parent / used_name).write_text("400000") + second = host_memory_budget(local_world_size=2, proc_root=proc) + assert first.available_bytes == 75_000 + assert second.available_bytes == 25_000 + + +def test_topology_cache_observes_membership_and_mount_changes_immediately(tmp_path): + proc, mount, child, limit_name, used_name = _host(tmp_path) + assert ( + host_memory_budget(local_world_size=4, proc_root=proc).available_bytes == 37_500 + ) + moved = child.parent / "moved" + moved.mkdir() + (moved / limit_name).write_text("10000") + (moved / used_name).write_text("0") + (proc / "self/cgroup").write_text("0::/pod/moved\n") + assert ( + host_memory_budget(local_world_size=4, proc_root=proc).available_bytes == 2_250 + ) + remount = tmp_path / "remounted" + replacement = remount / "pod/moved" + replacement.mkdir(parents=True) + (replacement / limit_name).write_text("12000") + (replacement / used_name).write_text("2000") + mountinfo = proc / "self/mountinfo" + mountinfo.write_text(mountinfo.read_text().replace(str(mount), str(remount))) + assert ( + host_memory_budget(local_world_size=4, proc_root=proc).available_bytes == 2_200 + ) + + +@pytest.mark.parametrize("used", [None, "not-a-counter", "500000"]) +def test_missing_or_exhausted_cgroup_usage_grants_no_credit(tmp_path, used): + proc, _, child, _, used_name = _host(tmp_path) + path = child / used_name + if used is None: + path.unlink() + else: + path.write_text(used) + assert host_memory_budget(local_world_size=2, proc_root=proc).available_bytes == 0 + + +def test_unlimited_cgroup_still_obeys_host_and_parent(tmp_path): + proc, mount, child, limit_name, _ = _host(tmp_path) + for path in (mount, child): + (path / limit_name).write_text("max") + budget = host_memory_budget(local_world_size=2, proc_root=proc) + assert len(budget.scopes) == 2 + assert budget.available_bytes == 75_000 + + +def test_unknown_host_budget_fails_closed_and_empty_scopes_grant_nothing(tmp_path): + assert ( + host_memory_budget(local_world_size=1, proc_root=tmp_path).available_bytes == 0 + ) + assert HostMemoryBudget(()).available_bytes == 0 + + +@pytest.mark.parametrize("ranks", [0, -1]) +def test_invalid_rank_count_is_rejected(ranks): + with pytest.raises(ValueError, match="positive"): + host_memory_budget(local_world_size=ranks) + with pytest.raises(ValueError, match="positive"): + _ = MemoryScope("host", 100, 100, ranks).per_rank_available_bytes + + +def test_local_rank_count_does_not_use_gpu_count(monkeypatch): + for name in ("LOCAL_WORLD_SIZE", "OMPI_COMM_WORLD_LOCAL_SIZE", "MPI_LOCALNRANKS"): + monkeypatch.delenv(name, raising=False) + monkeypatch.setenv("WORLD_SIZE", "16") + assert local_rank_count() == 16 + monkeypatch.setenv("LOCAL_WORLD_SIZE", "4") + assert local_rank_count() == 4 + + +ROOT = (ForwardMemoryCost(100, 80, 10, 5),) * 3 + + +@pytest.mark.parametrize( + "state,device,gpu,cpu,retained", + [ + ("gpu", "model", 290, 15, 270), + ("gpu", "cpu", 260, 45, 240), + ("cpu", "model", 150, 225, 60), + ("cpu", "cpu", 120, 255, 30), + ("replay", "model", 130, 15, 30), + ("replay", "cpu", 100, 45, 0), + ], +) +def test_root_cost_keeps_outputs_and_backward_restore_workspace( + state, device, gpu, cpu, retained +): + placement = placement_cost(ROOT, backward_state=state, output_device=device) + assert placement.gpu_required_bytes == gpu + assert placement.cpu_required_bytes == cpu + assert placement.gpu_retained_bytes == retained + + +@pytest.mark.parametrize( + "gpu,cpu,state,device", + [ + (290, 15, "gpu", "model"), + (259, 225, "cpu", "model"), + (130, 224, "replay", "model"), + (129, 45, "replay", "cpu"), + ], +) +def test_admission_prefers_retention_and_admits_previously_refused_roots( + gpu, cpu, state, device +): + placement = choose_memory_placement( + ROOT, gpu_available_bytes=gpu, cpu_available_bytes=cpu, output_device="auto" + ) + assert placement is not None + assert (placement.backward_state, placement.output_device) == (state, device) + assert placement.gpu_required_bytes <= gpu + assert placement.cpu_required_bytes <= cpu + + +@pytest.mark.parametrize("gpu,cpu", [(99, 1000), (259, 14), (129, 44)]) +def test_policy_never_admits_without_cpu_capacity_or_one_child_workspace(gpu, cpu): + assert ( + choose_memory_placement( + ROOT, gpu_available_bytes=gpu, cpu_available_bytes=cpu, output_device="auto" + ) + is None + ) + + +def test_disable_levels_and_explicit_output_opt_out_are_authoritative(): + assert ( + choose_memory_placement( + ROOT, + gpu_available_bytes=129, + cpu_available_bytes=1000, + allow_cpu_offload=False, + output_device="model", + ) + is None + ) + assert ( + choose_memory_placement( + ROOT, + gpu_available_bytes=130, + cpu_available_bytes=224, + allow_replay=False, + ) + is None + ) + forced = choose_memory_placement( + ROOT, + gpu_available_bytes=1000, + cpu_available_bytes=1000, + backward_state="replay", + output_device="cpu", + ) + assert forced is not None + assert (forced.backward_state, forced.output_device) == ("replay", "cpu") + + +@pytest.mark.parametrize( + "kwargs", + [ + {"backward_state": "cpu", "allow_cpu_offload": False}, + {"backward_state": "replay", "allow_replay": False}, + {"backward_state": "typo"}, + {"output_device": "typo"}, + ], +) +def test_invalid_policy_raises_before_admission(kwargs): + with pytest.raises(ValueError): + choose_memory_placement( + ROOT, gpu_available_bytes=1000, cpu_available_bytes=1000, **kwargs + ) + + +def test_no_grad_cpu_outputs_release_gpu_storage_without_replay(): + costs = (ForwardMemoryCost(100, 80, 80, backward_required=False),) * 3 + placement = choose_memory_placement( + costs, + gpu_available_bytes=100, + cpu_available_bytes=240, + allow_cpu_offload=False, + allow_replay=False, + output_device="auto", + ) + assert placement is not None + assert (placement.backward_state, placement.output_device) == ("gpu", "cpu") + assert placement.gpu_required_bytes == 100 + assert placement.cpu_required_bytes == 240 + + +def test_cost_is_independent_of_child_execution_order(): + costs = [*ROOT, ForwardMemoryCost(300, 60, 5, 10)] + for state in ("gpu", "cpu", "replay"): + assert placement_cost(costs, backward_state=state, output_device="cpu") == ( + placement_cost(costs[::-1], backward_state=state, output_device="cpu") + ) + + +def test_current_probability_correction_reserves_a_forward_beside_old_graphs(): + cost = ForwardMemoryCost(100, 80, 10, correction_workspace_bytes=100) + placement = placement_cost((cost,) * 3, backward_state="gpu", output_device="model") + assert placement.gpu_required_bytes == 370 diff --git a/tests/unit/test_trainer_rank_memory_policy_cuda.py b/tests/unit/test_trainer_rank_memory_policy_cuda.py new file mode 100644 index 000000000..6548dd440 --- /dev/null +++ b/tests/unit/test_trainer_rank_memory_policy_cuda.py @@ -0,0 +1,213 @@ +"""Real allocator/host measurements; run only on a reserved validation GPU.""" + +from dataclasses import asdict +import gc +import json +import os +from pathlib import Path +from types import SimpleNamespace + +import pytest +import torch + +from art.trainer_rank import TrainerRank +from art.trainer_rank._graphs import GraphCache +from art.trainer_rank._impl import _CheckpointSlot +from art.trainer_rank._memory_policy import ( + ForwardMemoryCost, + choose_memory_placement, + host_memory_budget, + placement_cost, +) + +pytestmark = pytest.mark.skipif( + os.environ.get("ART_GRAPH_GPU_TEST") != "1", reason="requires a reserved GPU" +) + + +@pytest.mark.parametrize("existing_gradient", [False, True]) +def test_sequential_gradient_publications_fit_original_admission_reserve( + monkeypatch, existing_gradient +): + trainer = TrainerRank.__new__(TrainerRank) + parameter = torch.nn.Parameter(torch.ones(4 * 1024**2, device="cuda")) + source = torch.ones_like(parameter) + if existing_gradient: + parameter.grad = torch.zeros_like(parameter) + trainer.device = parameter.device + trainer.runtime = SimpleNamespace(model=[], optimizer=None) + trainer._checkpoint_slots = {"student": _CheckpointSlot(params=(parameter,))} + monkeypatch.setattr( + trainer, "_iter_slot_parameters", lambda _ref: iter((parameter,)) + ) + ref = trainer._slot_ref("student") + # Two pending forwards observe the same pre-backward gradient state. + reserved = min(trainer._lora_gradient_staging_bytes(ref) for _ in range(2)) + state = trainer._version_state() + version = state.capture("student") + torch.cuda.synchronize() + baseline = torch.cuda.memory_allocated() + torch.cuda.reset_peak_memory_stats() + for _ in range(2): + with trainer._gradient_transaction(): + state.accumulate(((version, 2, parameter, source),)) + torch.cuda.synchronize() + peak = torch.cuda.max_memory_allocated() - baseline + assert peak <= reserved + assert peak == parameter.numel() * parameter.element_size() * ( + 2 if existing_gradient else 3 + ) + print( + "GRADIENT_RESERVATION=" + + json.dumps( + { + "existing_gradient": existing_gradient, + "reserved_bytes": reserved, + "peak_bytes": peak, + } + ) + ) + torch.testing.assert_close(parameter.grad, source * 2) + + +def _rss(): + for line in Path("/proc/self/status").read_text().splitlines(): + if line.startswith("VmRSS:"): + return int(line.split()[1]) * 1024 + return 0 + + +@pytest.fixture(scope="module") +def workload(): + torch.manual_seed(1729) + inputs = torch.randn(4096, 1024) + weight = torch.nn.Parameter(torch.randn(1024, device="cuda")) + + def execute(value): + return (value.to("cuda").sin() * weight,) + + # Warm kernels/gradient allocation before measuring the same real workload. + execute(inputs)[0].sum().backward() + weight.grad = None + gc.collect() + torch.cuda.empty_cache() + torch.cuda.synchronize() + baseline = torch.cuda.memory_allocated() + torch.cuda.reset_peak_memory_stats() + (output,) = execute(inputs) + torch.cuda.synchronize() + retained = torch.cuda.memory_allocated() - baseline + output_bytes = output.numel() * output.element_size() + output.backward(torch.ones_like(output)) + torch.cuda.synchronize() + peak = torch.cuda.max_memory_allocated() - baseline + assert weight.grad is not None + expected = weight.grad.detach().clone() + weight.grad = None + del output + cost = ForwardMemoryCost( + peak_bytes=int(peak * 1.1), + retained_bytes=int(retained * 1.1), + output_bytes=output_bytes, + replay_bytes=inputs.numel() * inputs.element_size() + 65536, + ) + return inputs, weight, execute, expected, cost + + +def _run_root(workload, state, device): + inputs, weight, execute, expected, cost = workload + cache = GraphCache() + weight.grad = None + gc.collect() + torch.cuda.empty_cache() + torch.cuda.synchronize() + baseline, rss_before = torch.cuda.memory_allocated(), _rss() + torch.cuda.reset_peak_memory_stats() + storage = weight.untyped_storage().data_ptr() + handles, outputs = [], [] + for _ in range(4): + handle, (output,) = cache.run( + execute, + inputs, + retention=state, + output_device="cpu" if device == "cpu" else "cuda", + cuda_devices=[torch.cuda.current_device()], + keep_on_device=lambda tensor: ( + tensor.untyped_storage().data_ptr() == storage + ), + ) + handles.append(handle) + outputs.append(output) + torch.cuda.synchronize() + forward_retained = torch.cuda.memory_allocated() - baseline + records = [cache.state(handle) for handle in handles] + rss_retained = _rss() - rss_before + caller_bytes = sum(output.numel() * output.element_size() for output in outputs) + # Stream cotangents from CPU so this measures runtime restoration workspace, + # independently of arbitrary caller loss graphs/temporary CUDA tensors. + cache.backward_many( + tuple( + (handle, (torch.ones(output.shape, dtype=output.dtype),)) + for handle, output in zip(handles, outputs, strict=True) + ) + ) + torch.cuda.synchronize() + peak = torch.cuda.max_memory_allocated() - baseline + torch.testing.assert_close(weight.grad, expected * 4, rtol=2e-5, atol=2e-4) + assert not cache.handles() + planned = placement_cost((cost,) * 4, backward_state=state, output_device=device) + measurement = { + "state": state, + "output_device": device, + "gpu_forward_retained_bytes": forward_retained, + "gpu_peak_bytes": peak, + "rss_forward_delta_bytes": rss_retained, + "cache_gpu_bytes": sum(record.gpu_bytes for record in records), + "cache_cpu_bytes": sum(record.cpu_bytes for record in records), + "caller_output_bytes": caller_bytes, + "planned": asdict(planned), + "host_budget_two_ranks": asdict(host_memory_budget(local_world_size=2)), + } + print("MEMORY_MEASUREMENT=" + json.dumps(measurement, sort_keys=True)) + assert peak <= planned.gpu_required_bytes + 8 * 1024**2 + if device == "cpu" and state == "replay": + assert forward_retained < 1024**2 + if state in ("cpu", "replay"): + assert ( + sum(record.cpu_bytes for record in records) + >= cost.replay_bytes * 4 - 4 * 65536 + ) + del outputs + return measurement + + +@pytest.mark.parametrize("state", ["gpu", "cpu", "replay"]) +@pytest.mark.parametrize("device", ["model", "cpu"]) +def test_real_root_memory_and_gradient_policy(workload, state, device): + _run_root(workload, state, device) + + +def test_real_root_admitted_by_replay_with_cpu_outputs(workload): + *_, cost = workload + budget = cost.peak_bytes + 8 * 1024**2 + cpu_budget = host_memory_budget(local_world_size=2).available_bytes + assert ( + choose_memory_placement( + (cost,) * 4, + gpu_available_bytes=budget, + cpu_available_bytes=cpu_budget, + backward_state="gpu", + output_device="model", + ) + is None + ) + chosen = choose_memory_placement( + (cost,) * 4, + gpu_available_bytes=budget, + cpu_available_bytes=cpu_budget, + output_device="auto", + ) + assert chosen is not None + assert (chosen.backward_state, chosen.output_device) == ("replay", "cpu") + measured = _run_root(workload, chosen.backward_state, chosen.output_device) + assert measured["gpu_peak_bytes"] <= budget diff --git a/tests/unit/test_trainer_rank_memory_recovery_distributed.py b/tests/unit/test_trainer_rank_memory_recovery_distributed.py new file mode 100644 index 000000000..e54a5decb --- /dev/null +++ b/tests/unit/test_trainer_rank_memory_recovery_distributed.py @@ -0,0 +1,194 @@ +"""Real collectives around injected transfer failure; no CUDA memory claims.""" + +from datetime import timedelta +import json +from pathlib import Path +from types import SimpleNamespace + +import pytest +import torch +import torch.distributed as dist +import torch.multiprocessing as mp + +from art.trainer_rank import _impl +from art.trainer_rank._memory_policy import ForwardMemoryCost +from art.trainer_rank._options import resolve_forward_options + + +def _reclaim_worker(index: int, directory: str, sync_across_dp: bool) -> None: + dist.init_process_group( + "gloo", + init_method=f"file://{directory}/rendezvous", + rank=index, + world_size=2, + timeout=timedelta(seconds=30), + ) + try: + rank = object.__new__(_impl.TrainerRank) + rank.device = torch.device("cpu") + rank._graph_memory_policy_enabled = lambda: True + rank._available_cpu_memory_bytes = lambda: 1000 + rank._forward_memory_group = lambda: dist.group.WORLD + setattr(torch.cuda, "empty_cache", lambda: None) + + def offload(_handle): + if index == 0: + raise RuntimeError("injected CPU allocation failure") + + rank._graph_cache = SimpleNamespace( + handles=lambda: (f"rank-local-{index}",), + state=lambda _: SimpleNamespace( + retention="gpu", offloadable=True, replayable=True, offload_bytes=100 + ), + offload=offload, + evict=lambda _: None, + ) + try: + rank._reclaim_graph_memory( + _impl._MemoryCheck(2000, 1000, False), sync_across_dp=sync_across_dp + ) + except RuntimeError as error: + result = str(error) + else: + raise AssertionError("Failed graph transfer was incorrectly accepted") + gathered = [None, None] + # Proves both participants left reclamation at a matching boundary. + dist.all_gather_object(gathered, result) + if index == 0: + Path(directory, "results.json").write_text(json.dumps(gathered)) + finally: + dist.destroy_process_group() + + +@pytest.mark.parametrize("sync_across_dp", [False, True]) +def test_failed_offload_does_not_strand_a_physical_peer(tmp_path, sync_across_dp): + mp.spawn(_reclaim_worker, args=(str(tmp_path), sync_across_dp), nprocs=2, join=True) + assert json.loads((tmp_path / "results.json").read_text()) == [ + "injected CPU allocation failure", + "Graph reclamation failed on another physical rank", + ] + + +def _fallback_worker(index: int, directory: str, sync_across_dp: bool) -> None: + dist.init_process_group( + "gloo", + init_method=f"file://{directory}/rendezvous", + rank=index, + world_size=2, + timeout=timedelta(seconds=30), + ) + try: + rank = object.__new__(_impl.TrainerRank) + rank.device = torch.device("cpu") + rank._forward_memory_group = lambda: dist.group.WORLD + # Locally rank 0 prefers CPU (.2 versus .5 seconds), while rank 1 + # prefers replay (2 versus .1). Both must use the same global order. + rank._graph_cache = SimpleNamespace( + transfer_stats=SimpleNamespace( + offload_bytes=100, + restore_bytes=100, + offload_seconds=0.1 if index == 0 else 1.0, + restore_seconds=0.1 if index == 0 else 1.0, + ) + ) + cost = ForwardMemoryCost( + 110, 110, 10, replay_seconds=0.5 if index == 0 else 0.1 + ) + units = [(0, (0,), cost, resolve_forward_options())] + candidates = list( + rank._graph_memory_candidates(units, sync_across_dp=sync_across_dp) + ) + assert candidates[2][0] == "replay" + assert candidates[2][2] == { + "source": "measured_forward_and_transfers", + "cpu_extra_seconds": 2.0, + "replay_extra_seconds": 0.5, + "preferred": "replay", + } + # One cold peer keeps every participant conservative. + if index == 1: + rank._graph_cache.transfer_stats.restore_bytes = 0 + candidates = list( + rank._graph_memory_candidates(units, sync_across_dp=sync_across_dp) + ) + assert candidates[2][0] == "cpu" + assert candidates[2][2]["source"] == "insufficient_samples" + finally: + dist.destroy_process_group() + + +@pytest.mark.parametrize("sync_across_dp", [False, True]) +def test_measured_fallback_order_agrees_across_physical_peers(tmp_path, sync_across_dp): + mp.spawn( + _fallback_worker, args=(str(tmp_path), sync_across_dp), nprocs=2, join=True + ) + + +def _head_failure_worker(index: int, directory: str, failure: str) -> None: + from test_trainer_rank_custom_tensors import _trainer + + from art.trainer_rank._commands import _Executor + from art.trainer_rank._heads import LiveHead, export_head + from art.trainer_rank._tensors import CotangentCollector + + dist.init_process_group( + "gloo", + init_method=f"file://{directory}/rendezvous", + rank=index, + world_size=2, + timeout=timedelta(seconds=20), + ) + try: + trainer, api = _trainer("student") + parameter = api.parameter("head", lambda: torch.ones(4), checkpoint="student") + parameter.grad = torch.ones_like(parameter) + collector = CotangentCollector() + live = LiveHead( + export_head(trainer, "student", "head"), torch.ones(4), collector + ) + assert isinstance(live.value, torch.Tensor) + packets = collector.backward( + torch.stack([(live.value * scale).sum() for scale in (2, 3)]).sum() + ) + cache = trainer._forward_graph_cache() + setattr( + cache, + "backward_many", + lambda *_a, **_k: pytest.fail("peer entered model backward"), + ) + commit, convert = trainer._commit_versioned_gradients, torch.Tensor.to + calls = 0 + + def stage(gradients): + nonlocal calls + calls += 1 + if index == 0 and failure == "stage" and calls == 2: + raise MemoryError("injected head stage allocation") + return commit(gradients) + + source_ids = {id(packet.gradients[0]) for packet in packets} + + def copy(tensor, *args, **kwargs): + if index == 0 and failure == "copy" and id(tensor) in source_ids: + raise MemoryError("injected head copy allocation") + return convert(tensor, *args, **kwargs) + + setattr(trainer, "_commit_versioned_gradients", stage) + setattr(torch.Tensor, "to", copy) + try: + with pytest.raises((MemoryError, RuntimeError), match="head .* allocation"): + _Executor(trainer, "zero")._backward(packets, retain_graph=False) + finally: + setattr(torch.Tensor, "to", convert) + torch.testing.assert_close(parameter.grad, torch.ones_like(parameter)) + assert trainer._version_state()._transaction is None + dist.barrier() + finally: + dist.destroy_process_group() + + +@pytest.mark.parametrize("failure", ["copy", "stage"]) +def test_remote_head_allocation_failure_is_coordinated_before_model_backward( + tmp_path, failure +): + mp.spawn(_head_failure_worker, args=(str(tmp_path), failure), nprocs=2, join=True) diff --git a/tests/unit/test_trainer_rank_options.py b/tests/unit/test_trainer_rank_options.py new file mode 100644 index 000000000..70b573604 --- /dev/null +++ b/tests/unit/test_trainer_rank_options.py @@ -0,0 +1,252 @@ +from __future__ import annotations + +from copy import copy, deepcopy +from dataclasses import FrozenInstanceError +import pickle +from typing import Any, cast, get_type_hints + +import cloudpickle +import pytest +import torch + +from art.trainer_rank import ( + ForwardInput, + ForwardOptions, + ResolvedForwardOptions, + Unset, + resolve_forward_options, +) +from art.trainer_rank import ( + ImportanceSamplingGradientCorrection as Correction, +) +from art.trainer_rank._corrections import ( + correct_logprob_cotangent, + importance_weights, +) + + +@pytest.mark.parametrize( + "public_type", [ForwardOptions, ResolvedForwardOptions, Correction] +) +def test_public_option_annotations_resolve(public_type: type) -> None: + hints = get_type_hints(public_type) + assert hints + if public_type is ForwardOptions: + assert hints["max_gradient_staleness"] == int | type(Unset) + + +def test_options_resolve_per_field_and_preserve_explicit_overrides() -> None: + constructor = ForwardOptions(max_gradient_staleness=8, output_device="cpu") + method = ForwardOptions(max_gradient_staleness=4, allow_cpu_offload=False) + input = ForwardOptions(max_gradient_staleness=0, stale_gradient_corrections=[]) + resolved = resolve_forward_options(constructor, method, input) + assert resolved == ResolvedForwardOptions( + max_gradient_staleness=0, + stale_gradient_corrections=(), + allow_cpu_offload=False, + output_device="cpu", + ) + assert resolve_forward_options().stale_gradient_corrections == (Correction(),) + assert resolve_forward_options().max_gradient_staleness == 2 + + +def test_options_snapshot_corrections_and_replace_inherited_collection() -> None: + corrections = [Correction(clip_high=2)] + options = ForwardOptions(stale_gradient_corrections=corrections) + corrections.clear() + resolved = resolve_forward_options( + options, ForwardOptions(stale_gradient_corrections=[Correction(clip_high=3)]) + ) + assert options.stale_gradient_corrections == (Correction(clip_high=2),) + assert resolved.stale_gradient_corrections == (Correction(clip_high=3),) + with pytest.raises(FrozenInstanceError): + setattr(resolved, "max_gradient_staleness", 0) + with pytest.raises(FrozenInstanceError): + setattr(options, "output_device", "model") + + +@pytest.mark.parametrize( + "roundtrip", + [ + copy, + deepcopy, + lambda x: pickle.loads(pickle.dumps(x)), + lambda x: cloudpickle.loads(cloudpickle.dumps(x)), + ], +) +def test_unset_identity_survives_transport(roundtrip) -> None: + assert roundtrip(Unset) is Unset + options = roundtrip(ForwardOptions(max_gradient_staleness=0)) + assert options.output_device is Unset + assert resolve_forward_options(options).max_gradient_staleness == 0 + request = roundtrip(ForwardInput(input_tokens=torch.tensor([1]), options=options)) + assert request.checkpoint is Unset + assert request.options == options + + +@pytest.mark.parametrize( + "kwargs", + [ + {"max_gradient_staleness": -1}, + {"max_gradient_staleness": True}, + {"max_gradient_staleness": 1.5}, + {"allow_cpu_offload": 0}, + {"allow_replay": None}, + {"backward_state": "disk"}, + {"output_device": "cuda"}, + {"stale_gradient_corrections": [Correction(), Correction()]}, + ], +) +def test_invalid_options_fail_before_submission(kwargs) -> None: + with pytest.raises(ValueError): + ForwardOptions(**kwargs) + + +@pytest.mark.parametrize( + "state,flag", [("cpu", "allow_cpu_offload"), ("replay", "allow_replay")] +) +def test_cross_level_conflicts_are_checked_after_resolution(state, flag) -> None: + constructor = ForwardOptions(**cast(dict[str, Any], {flag: False})) + method = ForwardOptions(backward_state=state) + with pytest.raises(ValueError, match="requires"): + resolve_forward_options(constructor, method) + # An input can override either side of the inherited conflict. + assert ( + resolve_forward_options( + constructor, method, ForwardOptions(**cast(dict[str, Any], {flag: True})) + ).backward_state + == state + ) + + +@pytest.mark.parametrize( + "kwargs", + [ + {"clip_low": -1}, + {"clip_low": 6}, + {"clip_high": float("inf")}, + {"clip_low": float("nan")}, + {"policy": "sometimes"}, + ], +) +def test_invalid_correction(kwargs) -> None: + with pytest.raises(ValueError): + Correction(**kwargs) + + +def test_score_function_weighting_matches_current_categorical_expectation() -> None: + # Enumerate all actions: E_old[(p_new/p_old) A grad log p_new]. + logits = torch.tensor([0.2, -0.3, 0.7], dtype=torch.float64, requires_grad=True) + old = torch.tensor([0.5, 0.3, 0.2], dtype=torch.float64) + rewards = torch.tensor([1.0, -0.5, 2.0], dtype=torch.float64) + current_logprobs = logits.log_softmax(-1) + weights = importance_weights(old.log(), current_logprobs, Correction(clip_high=100)) + assert not weights.requires_grad + weighted_score = (old * weights * rewards * current_logprobs).sum() + actual = torch.autograd.grad(weighted_score, logits, retain_graph=True)[0] + expected = torch.autograd.grad((current_logprobs.exp() * rewards).sum(), logits)[0] + torch.testing.assert_close(actual, expected) + + +@pytest.mark.parametrize( + "dtype", [torch.float16, torch.bfloat16, torch.float32, torch.float64] +) +def test_ratio_clipping_is_stable_in_low_precision(dtype) -> None: + original = torch.tensor([-10000, -1, -1, -10000], dtype=dtype, requires_grad=True) + current = torch.tensor( + [-1, -10000, -float("inf"), -10000], dtype=dtype, requires_grad=True + ) + weights = importance_weights( + original, current, Correction(clip_low=0.1, clip_high=5) + ) + torch.testing.assert_close(weights, weights.new_tensor([5, 0.1, 0.1, 1])) + assert weights.dtype == (torch.float64 if dtype == torch.float64 else torch.float32) + assert not weights.requires_grad + assert torch.isfinite(weights).all() + assert torch.equal( + importance_weights(original, current, Correction(clip_high=0)), + torch.zeros_like(weights), + ) + + +def test_top_k_uses_original_ids_and_full_distribution_probabilities() -> None: + old = torch.tensor([[0.4, 0.3]], dtype=torch.float64) + current = torch.tensor([[0.1, 0.6]], dtype=torch.float64) + tokens = torch.tensor([[2, 0]]) + cotangent = torch.tensor([[3.0, -2.0]], dtype=torch.float64) + result = correct_logprob_cotangent( + cotangent, + original_logprobs=old.log(), + current_logprobs=current.log(), + correction=Correction(), + original_tokens=tokens, + current_tokens=tokens, + ) + torch.testing.assert_close( + result, torch.tensor([[0.75, -4.0]], dtype=torch.float64) + ) + with pytest.raises(ValueError, match="same token IDs"): + correct_logprob_cotangent( + cotangent, + original_logprobs=old.log(), + current_logprobs=current.log(), + correction=Correction(), + original_tokens=tokens, + current_tokens=tokens.flip(-1), + ) + + +def test_unavailable_correction_follows_policy_without_forward() -> None: + grad = torch.ones(2) + kwargs: dict[str, Any] = dict( + original_logprobs=torch.zeros(2), current_logprobs=None + ) + assert correct_logprob_cotangent(grad, correction=Correction(), **kwargs) is grad + with pytest.raises(RuntimeError, match="requires current logprobs"): + correct_logprob_cotangent( + grad, correction=Correction(policy="always"), **kwargs + ) + + +@pytest.mark.parametrize( + "original,current", + [ + ([0.0], [float("nan")]), + ([0.0], [float("inf")]), + ([float("-inf")], [-1.0]), + ([float("-inf")], [float("-inf")]), + ], +) +def test_undefined_ratios_and_missing_support_raise(original, current) -> None: + with pytest.raises(ValueError): + importance_weights(torch.tensor(original), torch.tensor(current), Correction()) + + +def test_correction_rejects_implicit_broadcasting() -> None: + with pytest.raises(ValueError, match="shapes must match"): + importance_weights(torch.zeros(2, 1), torch.zeros(2), Correction()) + with pytest.raises(ValueError, match="shapes must match"): + correct_logprob_cotangent( + torch.ones(2, 1), + original_logprobs=torch.zeros(2), + current_logprobs=torch.zeros(2), + correction=Correction(), + ) + + +@pytest.mark.parametrize("bound", [1e-300, 1e300]) +def test_extreme_finite_bounds_preserve_representable_weights(bound: float) -> None: + weights = importance_weights( + torch.tensor([-1.0]), + torch.tensor([-1.0]), + Correction(clip_low=bound, clip_high=bound), + ) + assert weights.dtype == torch.float64 + assert weights.item() == bound + + +def test_resolved_options_reject_unset_and_unsupported_corrections() -> None: + with pytest.raises(ValueError, match="cannot contain Unset"): + ResolvedForwardOptions(max_gradient_staleness=cast(Any, Unset)) + with pytest.raises(TypeError, match="unsupported stale gradient correction"): + ForwardOptions(stale_gradient_corrections=cast(Any, [object()])) diff --git a/tests/unit/test_trainer_rank_output_memory.py b/tests/unit/test_trainer_rank_output_memory.py new file mode 100644 index 000000000..adb5a5d80 --- /dev/null +++ b/tests/unit/test_trainer_rank_output_memory.py @@ -0,0 +1,104 @@ +"""Logical placement must preserve native pending-backward reservations.""" + +from types import SimpleNamespace + +import pytest +from test_trainer_rank_memory_admission import _requests, rank # noqa: F401 +import torch + +from art.trainer_rank import ForwardInput, ForwardOptions, ForwardOutput, _impl +from art.trainer_rank._commands import _Executor, _OutputPacket, _view +from art.trainer_rank._tensors import detach_tree + + +def _output(size=80, *, policy="auto", handle="new"): + return ( + _OutputPacket( + detach_tree(handle, ForwardOutput(None, None, None, torch.ones(size // 4))), + (False,), + False, + ), + ForwardInput( + input_tokens=torch.tensor([1]), options=ForwardOptions(output_device=policy) + ), + ) + + +def _pending(rank, monkeypatch, *states): + monkeypatch.setattr( + rank, + "_graph_cache", + SimpleNamespace( + handles=lambda: tuple(range(len(states))), state=states.__getitem__ + ), + raising=False, + ) + + +@pytest.mark.parametrize("policy", ["auto", "model", "cpu"]) +def test_logical_copy_preserves_pending_restore(rank, monkeypatch, policy): + _pending(rank, monkeypatch, SimpleNamespace(restore_workspace_bytes=100)) + monkeypatch.setattr(rank, "_available_memory_bytes", lambda: 120) + view = _view(_Executor(rank, "zero")) + released = [] + monkeypatch.setattr(view, "_invoke", lambda *args: released.append(args)) + if policy == "model": + with pytest.raises(MemoryError, match="only 20 bytes"): + view._place_outputs([_output(policy=policy)]) + assert released == [("release", ("new",))] + else: + assert view._place_outputs([_output(policy=policy)])[0].cpu == (True,) + assert released == [] + + +def test_logical_copy_reserves_distinct_checkpoints_and_standalone_heads( + rank, monkeypatch +): + parameters = [torch.nn.Parameter(torch.ones(size)) for size in (16, 8, 4, 64)] + parameters[1].grad = torch.ones_like(parameters[1]) + rank._tag_custom_parameters(parameters[2:3]) + for name, parameter in zip(("first", "second", "head_only", "unused"), parameters): + rank._checkpoint_slots[name] = _impl._CheckpointSlot(params=(parameter,)) + _pending( + rank, + monkeypatch, + *( + SimpleNamespace( + restore_workspace_bytes=workspace, + checkpoint_versions=(rank._capture_checkpoint_version(name),), + ) + for name, workspace in (("first", 80), ("first", 100), ("second", 60)) + ), + ) + assert rank._pending_backward_memory() == (100, 192 + 64 + 48) + assert rank._pending_backward_memory(exclude_staging=("first",)) == (100, 112) + free = 484 + monkeypatch.setattr(rank, "_available_memory_bytes", lambda: free) + view = _view(_Executor(rank, "zero")) + planned = view._place_outputs([_output(handle="a"), _output(handle="b")]) + assert [output.cpu for output in planned] == [(False,), (True,)] + free -= 80 # First wave's GPU output is now live. + assert view._place_outputs([_output()])[0].cpu == (True,) + parameters[1].grad = None # Staging is recomputed, not cached with prior placement. + assert rank._pending_backward_memory() == (100, 192 + 96 + 48) + + +def test_client_cpu_transport_does_not_consume_worker_gpu_reserve(rank, monkeypatch): + view = _view(_Executor(rank, "zero")) + view._transport_handles = [] + monkeypatch.setattr( + rank, "_available_memory_bytes", lambda: pytest.fail("GPU query") + ) + assert view._place_outputs([_output(policy="model")])[0].cpu == (True,) + assert view._transport_handles == ["new"] + + +def test_native_admission_keeps_standalone_head_staging(rank, monkeypatch): + rank._checkpoint_slots["head_only"] = _impl._CheckpointSlot() + rank.parameter("head", lambda: torch.ones(1), checkpoint="head_only") + plan = rank._plan_flat_forward( + _requests(ForwardOptions(backward_state="replay", output_device="cpu"))[:1] + ) + monkeypatch.setattr(rank, "_available_memory_bytes", lambda: 111) + _, check = rank._admit_graph_memory(plan) + assert not check.fits and check.estimated_required_bytes == 112 diff --git a/tests/unit/test_trainer_rank_output_memory_cuda.py b/tests/unit/test_trainer_rank_output_memory_cuda.py new file mode 100644 index 000000000..9a18be164 --- /dev/null +++ b/tests/unit/test_trainer_rank_output_memory_cuda.py @@ -0,0 +1,86 @@ +"""Reserved-GPU canary for logical copies beside an existing replay backward.""" + +import gc +import json +import os + +import pytest +from test_trainer_rank_custom_tensors import _trainer +from test_trainer_rank_output_memory import _output +import torch + +from art.trainer_rank._commands import _Executor, _view +from art.trainer_rank._tensors import detach_tree + +pytestmark = pytest.mark.skipif( + os.environ.get("ART_GRAPH_GPU_TEST") != "1", reason="reserved GPU required" +) + + +def test_logical_outputs_leave_existing_replay_backward_admissible(monkeypatch): + torch.set_num_threads(2) + trainer, api = _trainer("student") + trainer.device = torch.device("cuda") + weight = api.parameter("weight", lambda: torch.ones(()), checkpoint="student") + cache = trainer._forward_graph_cache() + inputs = torch.ones(4 * 1024**2, device="cuda") + workspace = 64 * 1024**2 + handle, outputs = cache.run( + lambda values: ((values.sin() * weight).sum(),), + inputs, + retention="replay", + output_device="cpu", + execution_peak_bytes=workspace, + checkpoint_versions=(trainer._capture_checkpoint_version("student"),), + cuda_devices=[torch.cuda.current_device()], + ) + loss = trainer._forward_cotangent_collector().attach(detach_tree(handle, outputs))[ + 0 + ] + view = _view(_Executor(trainer, "zero")) + released = [] + monkeypatch.setattr(view, "_invoke", lambda *args: released.append(args)) + gc.collect() + torch.cuda.synchronize() + baseline = torch.cuda.memory_allocated() + capacity = 80 * 1024**2 + limit = baseline + capacity + monkeypatch.setattr( + trainer, + "_available_memory_bytes", + lambda: limit - torch.cuda.memory_allocated(), + ) + large = 32 * 1024**2 + with pytest.raises(MemoryError): + view._place_outputs([_output(large, policy="model")]) + assert released == [("release", ("new",))] + assert cache.handles() == (handle,) + auto = view._attach(view._place_outputs([_output(large)])[0]) + assert auto.hidden_states.device.type == "cpu" + model = view._attach(view._place_outputs([_output(8 * 1024**2, policy="model")])[0]) + assert model.hidden_states.device.type == "cuda" + torch.cuda.reset_peak_memory_stats() + trainer.backward(loss) + torch.cuda.synchronize() + peak = torch.cuda.max_memory_allocated() - baseline + assert peak <= capacity + assert weight.grad is not None + torch.testing.assert_close( + weight.grad.cpu(), + inputs.numel() * torch.sin(torch.tensor(1.0)), + rtol=1e-5, + atol=0, + ) + assert not cache.handles() + print( + "OUTPUT_BACKWARD_RESERVE=" + + json.dumps( + dict( + available_bytes=capacity, + restore_bytes=workspace, + rejected_copy_bytes=large, + admitted_copy_bytes=8 * 1024**2, + peak_bytes=peak, + ) + ) + ) diff --git a/tests/unit/test_trainer_rank_parameter_hooks.py b/tests/unit/test_trainer_rank_parameter_hooks.py new file mode 100644 index 000000000..bd49e19b4 --- /dev/null +++ b/tests/unit/test_trainer_rank_parameter_hooks.py @@ -0,0 +1,290 @@ +from __future__ import annotations + +import gc +import json + +import pytest +from test_trainer_rank_custom_tensors import _trainer +import torch + +from art.trainer_rank import ModuleHandle +from art.trainer_rank._heads import LiveHead, export_head, head_gradient_targets +from art.trainer_rank._tensors import CotangentCollector + + +def setup(client): + trainer, rank = _trainer("student") + native = rank.parameter("p", lambda: torch.tensor(2.0), checkpoint="student") + collector = CotangentCollector() + live = LiveHead(export_head(trainer, "student", "p"), torch.tensor(2.0), collector) + parameter = live.value if client else native + assert isinstance(parameter, torch.Tensor) + + def backward(loss, **kwargs): + if client: + packets = collector.backward(loss, **kwargs) + with trainer._gradient_transaction(): + for packet in packets: + trainer._commit_versioned_gradients( + head_gradient_targets(trainer, packet) + ) + return packets + trainer.backward(loss, **kwargs) + + return trainer, native, parameter, live, collector, backward + + +@pytest.mark.parametrize("client", (False, True)) +def test_live_parameter_hook_masks_real_gradient(client): + _, native, parameter, _, _, backward = setup(client) + called = [] + handle = parameter.register_hook( + lambda gradient: called.append(gradient.item()) or gradient * 0 + ) + backward(parameter.square()) + assert called == [4] + assert native.grad.item() == 0 + handle.remove() + backward(parameter.square()) + assert native.grad.item() == 4 + + +@pytest.mark.parametrize("client", (False, True)) +def test_hooks_aggregate_uses_and_old_versions_before_existing_grad(client): + trainer, native, parameter, live, _, backward = setup(client) + old = parameter.square() + parameter * 3 + native.data.fill_(4) + trainer._checkpoint_slots["student"].revision += 1 + if client: + live.refresh(export_head(trainer, "student", "p")) + live.max_gradient_staleness = 0 + new = parameter.square() + native.grad = torch.tensor(5.0) + seen = [] + first = parameter.register_hook( + lambda gradient: seen.append(gradient.item()) or gradient.square() + ) + second = parameter.register_hook(lambda gradient: gradient + 1) + packets = backward(old + new) + assert seen == [15] + assert native.grad.item() == 231 + if client: + combined = [ + json.loads(packet.handle[5:]) + for packet in packets + if any(g is not None for g in packet.gradients) + ] + assert len(combined) == 1 + assert combined[0]["revision"] == 0 + assert combined[0]["max_gradient_staleness"] == 1 + first.remove() + second.remove() + trainer._checkpoint_slots["student"].revision += 1 if client else 2 + with pytest.raises(RuntimeError, match="staleness"): + trainer._version_state().validate_accumulated(["student"]) + + +@pytest.mark.parametrize("client", (False, True)) +def test_hook_removal_applies_to_retained_graph(client): + _, native, parameter, _, _, backward = setup(client) + seen = [] + loss = parameter.square() + handle = parameter.register_hook( + lambda gradient: seen.append(gradient.item()) or gradient * 2 + ) + backward(loss, retain_graph=True) + handle.remove() + backward(loss) + assert seen == [4] + assert native.grad.item() == 12 + + +@pytest.mark.parametrize("client", (False, True)) +@pytest.mark.parametrize("bad_result", (False, True)) +def test_hook_failure_leaves_all_authoritative_gradients_unchanged(client, bad_result): + trainer, native, parameter, _, collector, backward = setup(client) + rank = trainer + other = rank.parameter("q", lambda: torch.tensor(3.0), checkpoint="student") + qlive = LiveHead(export_head(trainer, "student", "q"), torch.tensor(3.0), collector) + q = qlive.value if client else other + native.grad, other.grad = torch.tensor(5.0), torch.tensor(6.0) + parameter.register_hook(lambda gradient: gradient * 0) + + def fail(gradient): + if bad_result: + return torch.ones(2) + raise RuntimeError("user hook failed") + + q.register_hook(fail) + with pytest.raises(RuntimeError, match="hook"): + backward(parameter.square() + q.square()) + assert native.grad.item() == 5 + assert other.grad.item() == 6 + + +@pytest.mark.parametrize("client", (False, True)) +def test_post_accumulate_hook_rejected_at_registration(client): + _, _, parameter, _, _, _ = setup(client) + with pytest.raises(RuntimeError, match="do not support register_post_accumulate"): + parameter.register_post_accumulate_grad_hook(lambda parameter: None) + + +def test_head_hook_registries_follow_graph_lifetime(): + _, _, parameter, _, collector, _ = setup(True) + loss = parameter.square() + assert len(collector._head_hooks) == 1 + del loss + gc.collect() + assert not collector._head_hooks + + +@pytest.mark.parametrize("client", (False, True)) +@pytest.mark.parametrize("sparse_first", (False, True)) +@pytest.mark.parametrize("mixed", (False, True)) +def test_sparse_hook_gradients_preserve_layout_and_mix_with_dense( + client, sparse_first, mixed +): + trainer, rank = _trainer("student") + values = torch.arange(1.0, 7).reshape(3, 2) + native = rank.parameter("p", lambda: values.clone(), checkpoint="student") + collector = CotangentCollector() + live = LiveHead(export_head(trainer, "student", "p"), values, collector) + parameter = live.value if client else native + assert isinstance(parameter, torch.Tensor) + seen = [] + parameter.register_hook(lambda gradient: seen.append(gradient.layout) or gradient) + indices = torch.tensor([0, 2, 0]) + sparse = lambda: torch.nn.functional.embedding( + indices, parameter, sparse=True + ).sum() + dense = lambda: parameter.square().sum() + loss = ( + (sparse() + dense() if sparse_first else dense() + sparse()) + if mixed + else sparse() + ) + native.grad = torch.ones_like(values) + + def backward(): + if client: + packets = collector.backward(loss) + with trainer._gradient_transaction(): + for packet in packets: + trainer._commit_versioned_gradients( + head_gradient_targets(trainer, packet) + ) + else: + trainer.backward(loss) + + if not mixed: + with pytest.raises(ValueError, match="layout"): + backward() + torch.testing.assert_close(native.grad, torch.ones_like(values)) + if client: + assert seen == [torch.sparse_coo] + else: + backward() + assert seen == [torch.strided] + expected = 1 + 2 * values + torch.tensor([[2.0, 2.0], [0.0, 0.0], [1.0, 1.0]]) + torch.testing.assert_close(native.grad, expected) + + +@pytest.mark.parametrize("client", (False, True)) +def test_tied_module_parameter_hook_sums_all_calls(client): + from test_trainer_rank_live_heads import TiedHead + + trainer, rank = _trainer("student") + native = rank.module("head", TiedHead, checkpoint="student") + collector = CotangentCollector() + live = LiveHead(export_head(trainer, "student", "head"), TiedHead(), collector) + head = live.value if client else native + assert isinstance(head, ModuleHandle) + seen = [] + head.left.register_hook( + lambda gradient: seen.append(gradient.item()) or gradient.clamp(max=10) + ) + loss = head(torch.tensor(3.0)) + head(torch.tensor(1.0)) + if client: + with trainer._gradient_transaction(): + for packet in collector.backward(loss): + trainer._commit_versioned_gradients( + head_gradient_targets(trainer, packet) + ) + else: + trainer.backward(loss) + assert seen == [18] + assert native.left.grad.item() == 10 + + +@pytest.mark.parametrize("client", (False, True)) +def test_hook_on_direct_root_and_removed_before_first_backward(client): + _, native, parameter, _, _, backward = setup(client) + called = [] + handle = parameter.register_hook( + lambda gradient: called.append(gradient.item()) or gradient * 3 + ) + backward(parameter) + assert native.grad.item() == 3 + assert called == [1] + loss = parameter.square() + handle.remove() + backward(loss) + assert native.grad.item() == 7 + assert called == [1] + + +def test_hook_registry_survives_release_during_inline_backward(monkeypatch): + _, native, parameter, _, collector, backward = setup(True) + parameter.register_hook(lambda gradient: gradient * 0) + record = collector._record + + def release_while_recording(*args): + record(*args) + collector._head_hooks.clear() + + monkeypatch.setattr(collector, "_record", release_while_recording) + backward(parameter.square()) + assert native.grad.item() == 0 + + +@pytest.mark.parametrize("client", (False, True)) +def test_hook_cannot_mutate_authoritative_parameter(client): + trainer, native, parameter, _, _, backward = setup(client) + + def mutate(gradient): + parameter.add_(1) + return gradient + + parameter.register_hook(mutate) + with pytest.raises( + RuntimeError, match="mutate checkpoint parameters|only be changed" + ): + backward(parameter.square()) + assert native.item() == 2 + assert native.grad is None + assert trainer._checkpoint_slots["student"].revision == 0 + + +def test_later_hook_failure_discards_already_prepared_gradient(): + trainer, native, parameter, _, _, _ = setup(False) + other = trainer.parameter("q", lambda: torch.tensor(3.0), checkpoint="student") + native.grad, other.grad = torch.tensor(5.0), torch.tensor(6.0) + called = [] + parameter.register_hook(lambda gradient: called.append("p") or gradient * 0) + + def fail(gradient): + called.append("q") + raise RuntimeError("second hook failed") + + other.register_hook(fail) + version = trainer._capture_checkpoint_version("student") + with pytest.raises(RuntimeError, match="second hook failed"): + trainer._commit_versioned_gradients( + [ + (version, 2, native, torch.tensor(4.0)), + (version, 2, other, torch.tensor(6.0)), + ] + ) + assert called == ["p", "q"] + assert native.grad.item() == 5 + assert other.grad.item() == 6 diff --git a/tests/unit/test_trainer_rank_parameter_no_grad.py b/tests/unit/test_trainer_rank_parameter_no_grad.py new file mode 100644 index 000000000..73053308e --- /dev/null +++ b/tests/unit/test_trainer_rank_parameter_no_grad.py @@ -0,0 +1,100 @@ +from __future__ import annotations + +import asyncio + +import pytest +from test_trainer_rank_custom_tensors import _trainer +from test_trainer_rank_live_heads import _step +import torch + +from art.trainer_rank._commands import run_rank_callback +from art.trainer_rank._heads import LiveHead, export_head +from art.trainer_rank._tensors import CotangentCollector + + +def _read(parameter, operation): + if operation == "detach": + return parameter.detach() + if operation == "view": + return parameter.view(2, 2) + if operation == "slice": + return parameter[:, :1] + if operation == "transpose": + return parameter.T + return parameter.to(device=parameter.device, dtype=parameter.dtype) + + +@pytest.mark.parametrize("surface", ("native", "client", "zero", "rank")) +@pytest.mark.parametrize("operation", ("detach", "view", "slice", "transpose", "to")) +def test_no_grad_parameter_reads_preserve_saved_loss_after_optimizer_step( + monkeypatch, surface, operation +): + trainer, rank = _trainer("student") + initial = torch.tensor([[2.0, 3.0], [4.0, 5.0]]) + native = rank.parameter("p", lambda: initial.clone(), checkpoint="student") + collector = CotangentCollector() + live = LiveHead(export_head(trainer, "student", "p"), initial, collector) + + def capture(parameter): + with torch.no_grad(): + saved = _read(parameter, operation) + assert not saved.requires_grad + value = torch.ones_like(saved, requires_grad=True) + return saved, value, (value * saved).sum() + + def logical_capture(view): + parameter = view.parameter("p", lambda: initial.clone(), checkpoint="student") + return capture(parameter) + + if surface in {"zero", "rank"}: + saved, value, loss = asyncio.run( + run_rank_callback(trainer, logical_capture, mode=surface) + ).value + else: + saved, value, loss = capture(native if surface == "native" else live.value) + + trainer.backward(native.square().sum()) + _step(trainer, monkeypatch) + assert trainer._checkpoint_slots["student"].revision == 1 + assert not torch.equal(native, initial) + live.refresh(export_head(trainer, "student", "p")) + if surface in {"zero", "rank"}: + asyncio.run( + run_rank_callback(trainer, lambda view: view.backward(loss), mode=surface) + ) + elif surface == "client": + assert collector.backward(loss) == () + else: + trainer.backward(loss) + torch.testing.assert_close(saved, _read(initial, operation)) + torch.testing.assert_close(value.grad, _read(initial, operation)) + assert native.grad is None + + +@pytest.mark.parametrize("requires_grad", (False, True)) +def test_native_no_grad_detached_read_is_private_and_does_not_capture_gradients( + monkeypatch, requires_grad +): + trainer, rank = _trainer("student") + + def factory(): + module = torch.nn.Module() + module.register_parameter( + "p", torch.nn.Parameter(torch.tensor(2.0), requires_grad=requires_grad) + ) + return module + + parameter = rank.module("head", factory, checkpoint="student").p + assert parameter.requires_grad is requires_grad + + def forbidden_snapshot(*args, **kwargs): + raise AssertionError("no-grad reads must not create backward snapshots") + + monkeypatch.setattr(trainer, "_snapshot_parameter", forbidden_snapshot) + with torch.no_grad(): + saved = parameter.detach() + saved.add_(5) + assert saved.item() == 7 + assert parameter.item() == 2 + assert parameter.grad is None + assert trainer._checkpoint_slots["student"].revision == 0 diff --git a/tests/unit/test_trainer_rank_recompute_memory.py b/tests/unit/test_trainer_rank_recompute_memory.py index ea145e4d4..47fc8e5f5 100644 --- a/tests/unit/test_trainer_rank_recompute_memory.py +++ b/tests/unit/test_trainer_rank_recompute_memory.py @@ -75,7 +75,7 @@ def test_reported_cold_request_is_refused_before_execution( assert not rank._memory_check(_plan(rank)).fits assert rank._memory_check(_plan(rank, tokens=1024)).fits with pytest.raises(TrainerRankMemoryError): - rank.dp_rank_forward( + rank.forward( [ForwardInput(input_tokens=torch.arange(32710), hidden_states=True)] ) diff --git a/tests/unit/test_trainer_rank_recovery_slots.py b/tests/unit/test_trainer_rank_recovery_slots.py index abcbbd889..20d9662b9 100644 --- a/tests/unit/test_trainer_rank_recovery_slots.py +++ b/tests/unit/test_trainer_rank_recovery_slots.py @@ -45,7 +45,7 @@ def search(actual, **kwargs): rank._snapshot_planning_telemetry = lambda *args: None rank._try_cache_recovery = lambda *args, **kwargs: True result = rank._plan_admissible_forward( - requests, checkpoint=checkpoint, context="dp_rank_forward" + requests, checkpoint=checkpoint, context="forward" ) assert result == fit and not results assert events == ["ensure"] + ["search"] * search_count @@ -71,7 +71,7 @@ def forbidden(*args, **kwargs): rank._recover_admission = forbidden rank._find_admissible_forward = forbidden with pytest.raises(error_type) as captured: - rank._plan_admissible_forward([], checkpoint=None, context="dp_rank_forward") + rank._plan_admissible_forward([], checkpoint=None, context="forward") assert captured.value is error assert error.__cause__ is cause and error.__context__ is context assert error.__suppress_context__ and events == ["ensure"] @@ -80,6 +80,7 @@ def forbidden(*args, **kwargs): @pytest.mark.parametrize("ensure_slots", (None, False, True)) def test_direct_search_keeps_default_setup(ensure_slots): rank = TrainerRank.__new__(TrainerRank) + rank.device = _impl.torch.device("cpu") events = [] plan, check = object(), _impl._MemoryCheck(80, 200, True) rank._ensure_checkpoint_slots_for = lambda *a, **kw: events.append("ensure") diff --git a/tests/unit/test_trainer_rank_recovery_slots_distributed.py b/tests/unit/test_trainer_rank_recovery_slots_distributed.py index 3326663cb..6192a1dea 100644 --- a/tests/unit/test_trainer_rank_recovery_slots_distributed.py +++ b/tests/unit/test_trainer_rank_recovery_slots_distributed.py @@ -134,9 +134,7 @@ def ensure(values): ) error = barrier_error = None try: - rank._plan_admissible_forward( - [request], checkpoint=None, context="dp_rank_forward" - ) + rank._plan_admissible_forward([request], checkpoint=None, context="forward") except BaseException as exc: error = {"type": type(exc).__name__, "message": str(exc)} try: diff --git a/tests/unit/test_trainer_rank_release_completion.py b/tests/unit/test_trainer_rank_release_completion.py new file mode 100644 index 000000000..81d765bd4 --- /dev/null +++ b/tests/unit/test_trainer_rank_release_completion.py @@ -0,0 +1,88 @@ +"""Pending callback cleanup remains observable and ordered after failure.""" + +import asyncio +from typing import Any + +import pytest +from test_trainer_rank_commands import _Rank + +from art.trainer_rank._commands import _Executor, _Release + + +def test_completed_release_is_finalized_before_queued_done_callback(): + async def run(): + rank: Any = _Rank() + executor = _Executor(rank, "zero") + state = executor.state + state.graphs["zero:old:dp:0"] = (rank.weight,) + state.released.add("zero:old:dp:0") + completed = asyncio.get_running_loop().create_future() + release = state.pending_release = _Release(completed, [tuple(state.released)]) + completed.set_result(None) + completed.add_done_callback(lambda _: executor._finish_release(release)) + await executor._join_release() + assert not state.graphs and not state.released + assert state.pending_release is None + # The already queued callback cannot alter the next release's ownership. + state.graphs["zero:new:dp:0"] = (rank.weight,) + await asyncio.sleep(0) + assert tuple(state.graphs) == ("zero:new:dp:0",) + + asyncio.run(run()) + + +@pytest.mark.parametrize("cancelled", [False, True]) +def test_background_cleanup_failure_is_reported_and_blocks_next_entry(cancelled): + async def run(): + rank: Any = _Rank() + executor = _Executor(rank, "zero") + loop = asyncio.get_running_loop() + reports = [] + loop.set_exception_handler(lambda _loop, context: reports.append(context)) + completed = loop.create_future() + release = executor.state.pending_release = _Release(completed, [()]) + completed.add_done_callback(lambda _: executor._finish_release(release)) + if cancelled: + completed.cancel() + else: + completed.set_exception(RuntimeError("injected release transport error")) + with pytest.raises( + RuntimeError, match="Callback release reconciliation failed" + ): + await asyncio.wait_for(executor._join_release(), 1) + assert len(reports) == 1 + with pytest.raises( + RuntimeError, match="Callback release reconciliation failed" + ): + await executor.reconcile_releases() + assert executor.state.pending_release is None + + asyncio.run(run()) + + +def test_success_in_unrelated_exception_handler_still_awaits_cleanup(monkeypatch): + async def run(): + rank: Any = _Rank() + executor = _Executor(rank, "zero") + started, finish = asyncio.Event(), asyncio.Event() + + async def reconcile(**kwargs): + started.set() + await finish.wait() + + monkeypatch.setattr(executor, "reconcile_releases", reconcile) + + async def callback(): + try: + raise ValueError("unrelated handled exception") + except ValueError: + async with executor.release_on_exit(): + pass + + pending = asyncio.create_task(callback()) + await started.wait() + assert not pending.done() + finish.set() + await pending + + asyncio.run(run()) diff --git a/tests/unit/test_trainer_rank_release_lifetime.py b/tests/unit/test_trainer_rank_release_lifetime.py new file mode 100644 index 000000000..cad94400b --- /dev/null +++ b/tests/unit/test_trainer_rank_release_lifetime.py @@ -0,0 +1,122 @@ +"""Dead caller graphs release native captures across logical callback modes.""" + +import asyncio +from functools import partial +from typing import Any +import weakref + +import pytest +from test_trainer_rank_commands import _input, _Rank +from test_trainer_rank_versions import _trainer +import torch + +from art.trainer_rank import ForwardInput, ForwardOutput, TrainerRankSlotStateError +from art.trainer_rank._commands import run_rank_callback +from art.trainer_rank._graphs import GraphCache +from art.trainer_rank._tensors import CotangentCollector, TensorPacket, flatten_tensors + + +class _CachedRank(_Rank): + def __init__(self): + super().__init__() + self.native, self.weight = _trainer() + self.ref = self.native._slot_ref("student") + self.cache = GraphCache() + self.collector = CotangentCollector() + self.snapshots = [] + + def forward(self, tree, **kwargs): + if not isinstance(tree, ForwardInput): + return super().forward(tree, **kwargs) + version = self.native._capture_checkpoint_version("student") + + def execute(tokens): + snapshot = self.native._snapshot_parameter(self.weight, version) + self.snapshots.append(weakref.ref(snapshot)) + return (snapshot * tokens,) + + handle, tensors = self.cache.run( + execute, tree.input_tokens.float(), checkpoint_versions=(version,) + ) + _, spec = flatten_tensors(ForwardOutput(None, None, None, tensors[0])) + packet = TensorPacket( + handle, spec, tensors, tuple(tensor.requires_grad for tensor in tensors) + ) + output = self.collector.attach( + packet, on_release=partial(self.cache.release, handle) + ) + # Compose the actual slot replacement guard with the retained inner + # cache bridge, outside the cache's saved-variable offload hooks. + (output,) = self.native._track_slot_graph_outputs(self.ref, [output]) + return output + + def _forward_cotangent_collector(self): + return self.collector + + def _forward_graph_cache(self): + return self.cache + + def _forward_memory_group(self): + return None + + def _gradient_transaction(self, **kwargs): + return self.native._gradient_transaction(**kwargs) + + +def _run(rank, callback, mode): + return asyncio.run(run_rank_callback(rank, callback, mode=mode)).value + + +def _can_replace(rank): + rank.native._guard_slot_can_load(rank.native._slot_ref("student")) + + +def _released(rank): + assert not rank._rank_command_state.graphs + assert not rank._rank_command_state.released + assert not rank.cache.handles() + assert all(reference() is None for reference in rank.snapshots) + _can_replace(rank) + + +@pytest.mark.parametrize("mode", ["zero", "rank"]) +def test_callback_local_output_releases_physical_graph_and_checkpoint_capture(mode): + rank: Any = _CachedRank() + assert ( + _run(rank, lambda view: view.forward(_input(3)).hidden_states.item(), mode) == 6 + ) + _released(rank) + + +@pytest.mark.parametrize("mode", ["zero", "rank"]) +def test_output_dropped_after_stop_releases_before_other_mode_callback(mode): + rank: Any = _CachedRank() + output = _run(rank, lambda view: view.forward(_input(3)), mode) + caller = weakref.ref(output.hidden_states) + with pytest.raises(TrainerRankSlotStateError, match="live backward graph"): + _can_replace(rank) + del output + assert caller() is None + assert rank._rank_command_state.released and rank.cache.handles() + # The native checkpoint replacement guard runs before any view operation. + _run(rank, lambda _view: _can_replace(rank), "rank" if mode == "zero" else "zero") + _released(rank) + + +@pytest.mark.parametrize("mode", ["zero", "rank"]) +def test_live_dependent_loss_survives_other_mode_and_backwards_later(mode): + rank: Any = _CachedRank() + output = _run(rank, lambda view: view.forward(_input(3)), mode) + caller = weakref.ref(output.hidden_states) + loss = output.hidden_states.sum() + del output + assert caller() is None # The tensor wrapper is gone; its autograd graph lives. + _run(rank, lambda view: view.zero_grad(), "rank" if mode == "zero" else "zero") + assert rank.cache.handles() + with pytest.raises(TrainerRankSlotStateError, match="live backward graph"): + _can_replace(rank) + _run(rank, lambda view: view.backward(loss), mode) + torch.testing.assert_close(rank.weight.grad, torch.tensor(3.0, dtype=torch.float64)) + del loss + _run(rank, lambda view: view.zero_grad(), "rank" if mode == "zero" else "zero") + _released(rank) diff --git a/tests/unit/test_trainer_rank_resident_memory.py b/tests/unit/test_trainer_rank_resident_memory.py new file mode 100644 index 000000000..a4a6108c1 --- /dev/null +++ b/tests/unit/test_trainer_rank_resident_memory.py @@ -0,0 +1,49 @@ +"""CPU placement accounts for contexts which retain their original tensor storage.""" + +from dataclasses import replace + +from test_trainer_rank_memory_admission import _requests, rank # noqa: F401 + +from art.trainer_rank import ForwardOptions +from art.trainer_rank._memory_policy import ForwardMemoryCost, placement_cost + + +def test_partial_offload_retains_residual_for_every_complete_root_child(): + cost = ForwardMemoryCost(100, 80, 10, cpu_resident_bytes=40) + cpu = placement_cost([cost] * 3, backward_state="cpu", output_device="cpu") + replay = placement_cost([cost] * 3, backward_state="replay", output_device="cpu") + assert cpu.gpu_retained_bytes == 120 + assert cpu.gpu_required_bytes == 180 + assert replay.gpu_retained_bytes == 0 and replay.gpu_required_bytes == 100 + + +def test_cp_residency_requires_matching_layout_and_keeps_cold_bound(rank, monkeypatch): + monkeypatch.setattr(rank, "_topology_key", lambda: (1, 1, 2, 1)) + plan = rank._plan_flat_forward( + _requests(ForwardOptions(backward_state="cpu", output_device="cpu"))[:1] + ) + assert tuple(rank._graph_memory_units(plan))[0][2].cpu_resident_bytes == 80 + group = plan.groups[0] + rank._graph_residency = {rank._graph_residency_key(group): 40} + measured = tuple(rank._graph_memory_units(plan))[0][2] + assert 40 <= measured.cpu_resident_bytes < 80 + other = replace( + plan, + groups=(replace(group, packed=replace(group.packed, segments=())),), + ) + assert tuple(rank._graph_memory_units(other))[0][2].cpu_resident_bytes == 80 + branched = replace( + group, + packed=replace( + group.packed, + segments=tuple( + replace(segment, parent_id=segment.parent_id + 1) + for segment in group.packed.segments + ), + ), + ) + assert rank._graph_residency_key(branched) != rank._graph_residency_key(group) + rank._graph_residency = {rank._graph_residency_key(group): 120} + underestimated = tuple(rank._graph_memory_units(plan))[0][2] + assert underestimated.retained_bytes >= 120 + assert underestimated.peak_bytes == underestimated.retained_bytes + 20 diff --git a/tests/unit/test_trainer_rank_resident_memory_cuda.py b/tests/unit/test_trainer_rank_resident_memory_cuda.py new file mode 100644 index 000000000..9434bfca2 --- /dev/null +++ b/tests/unit/test_trainer_rank_resident_memory_cuda.py @@ -0,0 +1,60 @@ +"""Raw autograd context storage is nonoffloadable and must never be restored twice.""" + +import gc +import os + +import pytest +import torch + +from art._tensor_residency import record_resident_tensors +from art.trainer_rank._graphs import GraphCache + +pytestmark = pytest.mark.skipif( + os.environ.get("ART_GRAPH_GPU_TEST") != "1", reason="reserved GPU required" +) + + +class _RawContext(torch.autograd.Function): + @staticmethod + def forward(ctx, value): + ctx.raw = value * 2 + ctx.save_for_backward(ctx.raw) + record_resident_tensors((ctx.raw, ctx.raw[1:])) + return ctx.raw.sum() + + @staticmethod + def backward(ctx, *grad_outputs): + (gradient,) = grad_outputs + (saved,) = ctx.saved_tensors + assert ( + saved.untyped_storage().data_ptr() == ctx.raw.untyped_storage().data_ptr() + ) + return gradient.expand_as(saved) * 2 + + +@pytest.mark.parametrize("retention", ["gpu", "cpu"]) +def test_raw_context_residency_and_saved_alias_restore(retention): + cache = GraphCache() + inputs = torch.ones(4 * 1024**2, device="cuda", requires_grad=True) + handle, (output,) = cache.run( + lambda value: (_RawContext.apply(value),), + inputs, + retention=retention, + output_device="cpu", + ) + state = cache.state(handle) + assert state.non_offloadable_bytes == inputs.numel() * 4 + 4 + cache.offload(handle) + offloaded = cache.state(handle) + assert offloaded.gpu_bytes == state.non_offloadable_bytes + assert offloaded.cpu_bytes == offloaded.replay_bytes + cache.backward(handle, (torch.ones_like(output),), retain_graph=True) + gc.collect() + torch.cuda.synchronize() + before = torch.cuda.memory_allocated() + cache.evict(handle) + torch.cuda.synchronize() + assert before - torch.cuda.memory_allocated() >= inputs.numel() * 4 + assert cache.state(handle).gpu_bytes == 0 + cache.backward(handle, (torch.ones_like(output),)) + assert not cache.handles() diff --git a/tests/unit/test_trainer_rank_slot_graph_lifetime.py b/tests/unit/test_trainer_rank_slot_graph_lifetime.py new file mode 100644 index 000000000..ec872a982 --- /dev/null +++ b/tests/unit/test_trainer_rank_slot_graph_lifetime.py @@ -0,0 +1,168 @@ +from __future__ import annotations + +from contextlib import nullcontext +import gc +from typing import Literal + +import pytest +from test_trainer_rank_custom_tensors import _trainer +import torch + +from art.megatron.context_parallel.types import ParallelTopology +from art.megatron.prefix_tree_packing import prefix_tree_pack +from art.trainer_rank import ForwardInput, ForwardOptions, ForwardOutput +from art.trainer_rank._impl import ( + TrainerRankSlotStateError, + _ForwardGroupPlan, + _ForwardItem, +) +from art.trainer_rank._tensors import ManagedTensor + + +@pytest.fixture(params=["model", "cpu"]) +def output_device(request): + return request.param + + +@pytest.fixture(params=["none", "detach", "cpu"], autouse=True) +def ambient_hooks(request): + saved = [] + + def pack(tensor): + saved.append(tensor.dtype) + return tensor.detach() + + context = ( + torch.autograd.graph.saved_tensors_hooks(pack, lambda tensor: tensor) + if request.param == "detach" + else torch.autograd.graph.save_on_cpu() + if request.param == "cpu" + else nullcontext() + ) + with context: + yield saved if request.param == "detach" else None + + +def _forward(monkeypatch, retention, output_device, *, grad_enabled=True): + trainer, _ = _trainer("student") + ref = trainer._slot_ref("student") + weight = torch.nn.Parameter(torch.tensor(2.0)) + tokens = torch.tensor([1, 2]) + request = ForwardInput( + input_tokens=tokens, + hidden_states=True, + logits=True, + options=ForwardOptions(backward_state=retention, output_device=output_device), + ) + group = _ForwardGroupPlan( + ref, + grad_enabled, + (0,), + (_ForwardItem(request, tokens, None),), + prefix_tree_pack((tokens,), max_depth=0), + ) + monkeypatch.setattr(trainer, "_topology", lambda: ParallelTopology()) + monkeypatch.setattr(trainer, "_configure_hybridep", lambda *args, **kwargs: None) + monkeypatch.setattr(trainer, "_prepare_packed_forward", lambda packed: None) + monkeypatch.setattr(trainer, "_hybridep_graph_tracking", True, raising=False) + monkeypatch.setattr( + trainer, + "_forward_packed", + lambda items, prepared: [ + ForwardOutput( + target_logprobs=None, + top_k=None, + hidden_states=weight.square(), + logits=weight.pow(3), + ) + ], + ) + output = trainer._execute_graph_group(group)[0] + return trainer, ref, weight, output + + +@pytest.mark.parametrize("retention", ["gpu", "cpu", "replay"]) +def test_native_cached_slot_guard_tracks_retained_and_consumed_graph( + monkeypatch, retention: Literal["gpu", "cpu", "replay"], output_device +): + trainer, ref, weight, output = _forward(monkeypatch, retention, output_device) + loss = output.hidden_states + assert loss is not None + assert isinstance(loss, ManagedTensor) == (output_device == "cpu") + assert trainer._has_live_slot_graph(ref) + assert trainer._has_live_hybridep_graphs() + with pytest.raises(TrainerRankSlotStateError, match="live backward graph"): + trainer._guard_slot_can_load(ref) + handle = trainer._forward_graph_cache().handles()[0] + trainer._forward_graph_cache().evict(handle) + assert trainer._has_live_slot_graph(ref) + for _ in range(2): + trainer.backward(loss, retain_graph=True) + assert trainer._has_live_slot_graph(ref) + assert trainer._has_live_hybridep_graphs() + trainer.backward(loss) + torch.testing.assert_close(weight.grad, torch.tensor(12.0)) + # Keep both returned outputs alive: the unused logits cannot be consumed + # after final backward releases their shared physical graph. + assert output.logits is not None + assert not trainer._forward_graph_cache().handles() + assert not trainer._has_live_slot_graph(ref) + assert not trainer._has_live_hybridep_graphs() + trainer._guard_slot_can_load(ref) + + +@pytest.mark.parametrize("retention", ["gpu", "cpu", "replay"]) +def test_native_cached_slot_guard_follows_dependent_loss( + monkeypatch, retention, output_device +): + trainer, ref, _, output = _forward(monkeypatch, retention, output_device) + assert output.hidden_states is not None + loss = output.hidden_states.square() + del output + gc.collect() + assert trainer._has_live_slot_graph(ref) + del loss + gc.collect() + assert not trainer._has_live_slot_graph(ref) + assert not trainer._has_live_hybridep_graphs() + assert not trainer._forward_graph_cache().handles() + trainer._guard_slot_can_load(ref) + + +@pytest.mark.parametrize("retention", ["gpu", "cpu", "replay"]) +def test_native_no_grad_forward_has_no_slot_graph_guard( + monkeypatch, retention, output_device +): + trainer, ref, _, output = _forward( + monkeypatch, retention, output_device, grad_enabled=False + ) + assert output.hidden_states is not None and not output.hidden_states.requires_grad + assert not trainer._has_live_slot_graph(ref) + assert not trainer._has_live_hybridep_graphs() + assert not trainer._forward_graph_cache().handles() + trainer._guard_slot_can_load(ref) + + +@pytest.mark.parametrize("kind", ["parameter", "module"]) +def test_custom_slot_guard_preserves_markers_and_ambient_activation_hooks( + kind, ambient_hooks +): + trainer, rank = _trainer("student") + ref = trainer._slot_ref("student") + if kind == "parameter": + parameter = rank.parameter("p", lambda: torch.tensor(2.0), checkpoint="student") + loss = parameter.square() + else: + head = rank.module("head", lambda: torch.nn.Linear(1, 1), checkpoint="student") + loss = head(torch.ones(1, 1, requires_grad=True)).square().sum() + with pytest.raises(TrainerRankSlotStateError, match="live backward graph"): + trainer._guard_slot_can_load(ref) + if ambient_hooks is not None: + assert torch.float32 in ambient_hooks + assert torch.bool not in ambient_hooks + trainer.backward(loss, retain_graph=True) + assert trainer._has_live_slot_graph(ref) + trainer.backward(loss) + trainer.zero_grad() + assert not trainer._has_live_slot_graph(ref) + trainer._guard_slot_can_load(ref) diff --git a/tests/unit/test_trainer_rank_split.py b/tests/unit/test_trainer_rank_split.py index 359700def..7e2aa48bd 100644 --- a/tests/unit/test_trainer_rank_split.py +++ b/tests/unit/test_trainer_rank_split.py @@ -3,7 +3,7 @@ Written before the implementation (test-first, as for the automatic planner) and expected to FAIL on the pre-split tree. Contract, as agreed: -- ``dp_rank_forward`` should try not to raise when splitting the call into +- ``forward`` should try not to raise when splitting the call into sequential subforwards would make execution feasible. The split ladder is bounded and deterministic: the fewest subforwards that fit, cutting the requests in prefix-local depth-first order so most sharing stays inside one @@ -23,7 +23,7 @@ would return them. - Refusing is acceptable when the ladder is exhausted (a single request alone cannot fit) — confident refusal over expensive search. -- The same machinery applies inside ``forward_micro_batches`` when even the +- The same machinery applies inside ``forward_batches`` when even the minimum wave cannot fit unsplit. - Telemetry reports ``subforward_count`` (``last_forward_telemetry`` and ``MicroBatchStats``); it is 1 for unsplit calls. @@ -48,7 +48,6 @@ TrainerRank, TrainerRankMemoryError, TrainerRankPartialExecutionError, - TrainerRankSlotStateError, ) from art.trainer_rank._impl import ( Unset, @@ -151,7 +150,7 @@ def _rank(monkeypatch: pytest.MonkeyPatch) -> TrainerRank: return rank -def test_dp_rank_forward_splits_instead_of_raising( +def test_forward_splits_instead_of_raising( monkeypatch: pytest.MonkeyPatch, ) -> None: rank = _rank(monkeypatch) @@ -163,7 +162,7 @@ def test_dp_rank_forward_splits_instead_of_raising( monkeypatch.setattr(rank, "_retained_memory_bytes", lambda *_args, **_kwargs: 0) _packed_budget(monkeypatch, rank, 20) - outputs = rank.dp_rank_forward(inputs) + outputs = rank.forward(inputs) assert [int(output.target_logprobs.item()) for output in outputs] == [0, 1, 2, 3] assert len(executed) == 2 @@ -185,7 +184,7 @@ def test_unsplit_call_reports_a_single_subforward( _recording_executor(monkeypatch, rank) _packed_budget(monkeypatch, rank, 1_000) - rank.dp_rank_forward([_request(marker) for marker in range(4)]) + rank.forward([_request(marker) for marker in range(4)]) telemetry = rank.last_forward_telemetry() assert telemetry["subforward_count"] == 1 @@ -205,7 +204,7 @@ def test_split_outputs_preserve_nested_caller_order( monkeypatch.setattr(rank, "_retained_memory_bytes", lambda *_args, **_kwargs: 0) _packed_budget(monkeypatch, rank, 20) - outputs = rank.dp_rank_forward(nested) + outputs = rank.forward(nested) assert [ [int(output.target_logprobs.item()) for output in group] for group in outputs @@ -236,7 +235,7 @@ def plan(requests, **kwargs): monkeypatch.setattr(rank, "_available_memory_bytes", lambda: 9) with pytest.raises(TrainerRankMemoryError) as exc_info: - rank.dp_rank_forward(inputs) + rank.forward(inputs) assert exc_info.value.predicted_peak_bytes > exc_info.value.usable_limit_bytes assert "smaller" in exc_info.value.suggestion @@ -267,7 +266,7 @@ def test_split_admission_accounts_for_live_graphs_cumulatively( _packed_budget(monkeypatch, rank, 25) with pytest.raises(TrainerRankMemoryError): - rank.dp_rank_forward([_request(marker) for marker in range(4)]) + rank.forward([_request(marker) for marker in range(4)]) assert executed == [] @@ -291,13 +290,13 @@ def test_split_admission_uses_a_retained_profile_when_available( lambda *_args, **kwargs: int(kwargs["required"] * 0.1), ) - outputs = rank.dp_rank_forward([_request(marker) for marker in range(4)]) + outputs = rank.forward([_request(marker) for marker in range(4)]) assert len(outputs) == 4 assert len(executed) == 2 -def test_forward_micro_batches_splits_the_minimum_wave( +def test_forward_batches_splits_the_minimum_wave( monkeypatch: pytest.MonkeyPatch, ) -> None: rank = _rank(monkeypatch) @@ -308,7 +307,7 @@ def test_forward_micro_batches_splits_the_minimum_wave( items = [[_request(marker) for marker in range(4)]] _packed_budget(monkeypatch, rank, 20) - batches = list(rank.forward_micro_batches(items)) + batches = list(rank.forward_batches(items)) assert len(batches) == 1 assert batches[0].stats.global_count == 1 @@ -334,10 +333,10 @@ def partition() -> list[tuple[int, ...]]: for plan in executed ] - rank.dp_rank_forward(inputs) + rank.forward(inputs) first = partition() executed.clear() - rank.dp_rank_forward(inputs) + rank.forward(inputs) second = partition() assert first == second @@ -365,7 +364,7 @@ def ensure(names): monkeypatch.setattr(rank, "_ensure_checkpoint_slots", ensure) _packed_budget(monkeypatch, rank, 30) - rank.dp_rank_forward([_request(marker) for marker in range(8)]) + rank.forward([_request(marker) for marker in range(8)]) assert rank.last_forward_telemetry()["subforward_count"] == 4 assert ensured == 1 @@ -399,11 +398,11 @@ def test_retained_profile_is_trusted_only_near_its_observed_scale( ) if expect_split: - rank.dp_rank_forward(inputs) + rank.forward(inputs) assert len(executed) == 2 else: with pytest.raises(TrainerRankMemoryError): - rank.dp_rank_forward(inputs) + rank.forward(inputs) assert executed == [] @@ -860,7 +859,7 @@ def run(plan: _FlatForwardPlan, **_kwargs: object) -> tuple[list, None]: monkeypatch.setattr(rank, "_run_flat_plan_with_memory_tracking", run) with pytest.raises(TrainerRankPartialExecutionError) as exc_info: - rank.dp_rank_forward([_request(marker) for marker in range(4)]) + rank.forward([_request(marker) for marker in range(4)]) message = str(exc_info.value) assert f"subforward {failing_ordinal + 1} of 2 failed during execution" in message @@ -888,7 +887,7 @@ def budget() -> int: "_execute_flat_plan", lambda plan: [ForwardOutput(None, None, None, None)] * plan.request_count, ) - rank.dp_rank_forward([_request(9, length=5)]) + rank.forward([_request(9, length=5)]) assert rank.last_forward_telemetry()["predicted_peak_bytes"] == 5 oom = torch.cuda.OutOfMemoryError("injected forward allocation failure") executed = 0 @@ -905,9 +904,9 @@ def run(plan): inputs = [tuple(_request(marker) for marker in range(4 if split else 1))] with pytest.raises(TrainerRankMemoryError) as caught: if micro_batches: - next(rank.forward_micro_batches(inputs)) + next(rank.forward_batches(inputs)) else: - rank.dp_rank_forward(inputs) + rank.forward(inputs) error = caught.value assert isinstance(error, TrainerRankPartialExecutionError) == split @@ -931,10 +930,10 @@ def test_micro_batch_refusal_replaces_previous_admission_telemetry( _recording_executor(monkeypatch, rank) available = 100 _packed_budget(monkeypatch, rank, lambda: available) - rank.dp_rank_forward([_request(0, length=5)]) + rank.forward([_request(0, length=5)]) available = 1 with pytest.raises(TrainerRankMemoryError) as caught: - next(rank.forward_micro_batches([_request(1)])) + next(rank.forward_batches([_request(1)])) telemetry = rank.last_forward_telemetry() assert telemetry["predicted_peak_bytes"] == caught.value.predicted_peak_bytes == 10 assert telemetry["usable_limit_bytes"] == caught.value.usable_limit_bytes == 1 @@ -949,9 +948,7 @@ class _SlotRef: def test_split_subforwards_track_independent_slot_graphs( monkeypatch: pytest.MonkeyPatch, ) -> None: - """Two subforwards on one slot carry independent slot-graph sentinels: - releasing the first subforward's graph keeps slot load/step blocked until - the second is released too.""" + """Consuming one child keeps the other child's cached graph available.""" rank = _rank(monkeypatch) monkeypatch.setattr(rank, "_retained_memory_bytes", lambda *_args, **_kwargs: 0) @@ -960,7 +957,9 @@ def test_split_subforwards_track_independent_slot_graphs( monkeypatch.setattr(rank, "_slot_ref", lambda name: _SlotRef(name)) monkeypatch.setattr(rank, "_resolve_slot_ref", lambda request, **_kwargs: ref) monkeypatch.setattr(rank, "_validate_hybridep_topology", lambda: None) - monkeypatch.setattr(rank, "_topology", lambda: object()) + topology = object() + monkeypatch.setattr(rank, "_topology", lambda: topology) + monkeypatch.setattr(rank, "_capture_lora_version", lambda *_args, **_kwargs: None) monkeypatch.setattr(rank, "_configure_hybridep", lambda *_args, **_kwargs: None) monkeypatch.setattr(rank, "_prepare_packed_forward", lambda _packed: None) @@ -977,23 +976,27 @@ def forward(items: object, _prepared: object) -> list[ForwardOutput]: monkeypatch.setattr(rank, "_forward_packed", forward) lora = ModuleType("art.megatron.lora") - cast(Any, lora).use_lora_slot = lambda _slot: nullcontext() + cast(Any, lora).use_lora_slot = lambda _slot, **_kwargs: nullcontext() monkeypatch.setitem(sys.modules, "art.megatron.lora", lora) - outputs = rank.dp_rank_forward([_request(marker) for marker in range(4)]) + outputs = rank.forward([_request(marker) for marker in range(4)]) first, second = rank.last_forward_telemetry()["subforward_request_indices"] def loss(indices: tuple[int, ...]) -> torch.Tensor: return torch.stack([outputs[index].target_logprobs for index in indices]).sum() - with pytest.raises(TrainerRankSlotStateError, match="live backward graph"): - rank._guard_slot_can_load(ref) - loss(first).backward() - with pytest.raises(TrainerRankSlotStateError, match="live backward graph"): - rank._guard_slot_can_load(ref) - with pytest.raises(TrainerRankSlotStateError, match="Cannot optim_step"): - rank._guard_checkpoint_can_step("teacher") - loss(second).backward() + def backward(indices: tuple[int, ...]) -> None: + packets = rank._forward_cotangent_collector().backward(loss(indices)) + rank._forward_graph_cache().backward_many( + [(packet.handle, packet.gradients) for packet in packets] + ) + + cache = rank._forward_graph_cache() + assert len(cache.handles()) == 2 + backward(first) + assert len(cache.handles()) == 1 + backward(second) + assert cache.handles() == () rank._guard_slot_can_load(ref) rank._guard_checkpoint_can_step("teacher") diff --git a/tests/unit/test_trainer_rank_split_peak.py b/tests/unit/test_trainer_rank_split_peak.py index dbfb21d54..bbac43557 100644 --- a/tests/unit/test_trainer_rank_split_peak.py +++ b/tests/unit/test_trainer_rank_split_peak.py @@ -170,6 +170,8 @@ def test_profile_order_change_cannot_drop_completed_split_floor(monkeypatch): def _counter_split(monkeypatch): rank = _rank() + # This executor injects allocator counters without creating cached graphs. + monkeypatch.setattr(rank, "_graph_memory_policy_enabled", lambda: False) monkeypatch.setattr(rank, "_dp_rank_and_size", lambda: (0, 1)) _native_slot_fields(monkeypatch, rank) requests = _requests() @@ -217,7 +219,7 @@ def test_completed_iterator_preserves_caller_peak_for_next_admission(monkeypatch monkeypatch.setattr(torch.cuda, "mem_get_info", lambda _: (1_000_000, 1_000_000)) releases = [] monkeypatch.setattr(torch.cuda, "empty_cache", lambda: releases.append(True)) - iterator = rank.forward_micro_batches([requests], yield_empty=True) + iterator = rank.forward_batches([requests], yield_empty=True) batch = next(iterator) assert batch.stats.subforward_count == counters["executed"] == 2 assert counters["resets"] == [100, 600] @@ -228,7 +230,7 @@ def test_completed_iterator_preserves_caller_peak_for_next_admission(monkeypatch del batch counters["allocated"] = 100 with pytest.raises(tr.TrainerRankMemoryError): - next(rank.forward_micro_batches([requests], yield_empty=True)) + next(rank.forward_batches([requests], yield_empty=True)) assert counters["executed"] == 2 assert releases == [] @@ -237,7 +239,7 @@ def test_completed_iterator_preserves_caller_peak_for_next_admission(monkeypatch @pytest.mark.parametrize("termination", ["throw", "close"]) def test_incomplete_caller_does_not_learn_split_peak(monkeypatch, termination): rank, requests, counters = _counter_split(monkeypatch) - iterator = rank.forward_micro_batches([requests], yield_empty=True) + iterator = rank.forward_batches([requests], yield_empty=True) batch = next(iterator) assert batch.stats.subforward_count == counters["executed"] == 2 children = dict(rank._memory_profiles) @@ -259,7 +261,7 @@ def test_partial_forward_does_not_learn_split_peak(monkeypatch): rank, requests, counters = _counter_split(monkeypatch) original = torch.cuda.OutOfMemoryError("second split child allocation") counters.update(fail_at=2, error=original) - iterator = rank.forward_micro_batches([requests], yield_empty=True) + iterator = rank.forward_batches([requests], yield_empty=True) with pytest.raises(tr.TrainerRankPartialExecutionError) as caught: next(iterator) assert "1 of 2 completed" in str(caught.value) diff --git a/tests/unit/test_trainer_rank_tensors.py b/tests/unit/test_trainer_rank_tensors.py new file mode 100644 index 000000000..ebc754e37 --- /dev/null +++ b/tests/unit/test_trainer_rank_tensors.py @@ -0,0 +1,1009 @@ +from __future__ import annotations + +from collections import OrderedDict, namedtuple +from dataclasses import dataclass, field +import pickle + +import pytest +import torch + +from art.trainer_rank import ForwardOutput, TopK +from art.trainer_rank._tensors import ( + CotangentCollector, + ManagedTensor, + detach_tree, + flatten_tensors, + managed_tensor, + managed_tree, + unflatten_tensors, +) + + +@dataclass(frozen=True, slots=True) +class NestedOutput: + values: object + label: str = "root" + metadata: int = field(default=7, init=False) + + +def test_tree_preserves_nested_types_metadata_and_tensor_aliases(): + value = torch.tensor([1.0, 2.0], requires_grad=True) + tokens = torch.tensor([3, 4]) + Pair = namedtuple("Pair", "first second") + tree = OrderedDict( + outer=[NestedOutput((value, None)), Pair(value, tokens)], + size=torch.Size((2, 3)), + empty=[], + ) + tensors, spec = flatten_tensors(tree) + assert len(tensors) == 2 + restored = unflatten_tensors(spec, tuple(t + 1 for t in tensors)) + assert isinstance(restored, OrderedDict) + assert isinstance(restored["outer"][0], NestedOutput) + assert isinstance(restored["outer"][1], Pair) + assert restored["outer"][0].metadata == 7 + assert restored["outer"][0].values[0] is restored["outer"][1].first + assert restored["outer"][0].values[1] is None + assert restored["size"] == torch.Size((2, 3)) + assert restored["empty"] == [] + + +def test_forward_packet_pickle_and_detached_storage(): + parameter = torch.tensor([2.0, 3.0], requires_grad=True) + physical = ForwardOutput( + parameter.square(), TopK(parameter + 1, torch.tensor([0, 1])), None, None + ) + packet = detach_tree("forward:1", {"output": physical}) + assert all( + tensor.grad_fn is None and not tensor.requires_grad for tensor in packet.tensors + ) + packet = pickle.loads(pickle.dumps(packet)) + collector = CotangentCollector() + output = collector.attach(packet)["output"] + assert output.target_logprobs.requires_grad + assert not output.top_k.tokens.requires_grad + assert output.logits is None and output.hidden_states is None + packet.tensors[0].zero_() + torch.testing.assert_close(physical.target_logprobs, torch.tensor([4.0, 9.0])) + + +def _model(x, weight): + hidden = torch.tanh(x @ weight) + return ForwardOutput( + hidden.log_softmax(-1), + TopK(hidden[:, :1], torch.zeros((2, 1), dtype=torch.long)), + None, + hidden, + ) + + +def _loss(outputs, head): + a, b = outputs + return ((a.hidden_states @ head) - (b.hidden_states @ head)).square().sum() + ( + a.target_logprobs[:, 0] - b.target_logprobs[:, 1] + ).square().sum() + + +@pytest.mark.parametrize("managed", [False, True]) +def test_coupled_forwards_and_local_head_match_connected_reference(managed): + generator = torch.Generator().manual_seed(12) + weight = torch.randn( + 3, 4, dtype=torch.float64, generator=generator, requires_grad=True + ) + inputs = [ + torch.randn(2, 3, dtype=torch.float64, generator=generator) for _ in range(2) + ] + head = torch.randn( + 4, 1, dtype=torch.float64, generator=generator, requires_grad=True + ) + physical = [_model(x, weight) for x in inputs] + collector = CotangentCollector() + outputs = [ + collector.attach(detach_tree(str(i), value), managed=managed) + for i, value in enumerate(physical) + ] + packets = collector.backward(_loss(outputs, head)) + assert weight.grad is None + assert [packet.handle for packet in packets] == ["0", "1"] + actual_outputs, cotangents = [], [] + for packet, value in zip(packets, physical, strict=True): + leaves, _ = flatten_tensors(value) + assert packet.gradients[1] is None # unused top-k logprobs + assert packet.gradients[2] is None # integer token IDs + for leaf, grad in zip(leaves, packet.gradients, strict=True): + if grad is not None: + actual_outputs.append(leaf) + cotangents.append(grad) + torch.autograd.backward(actual_outputs, cotangents) + reference_weight = weight.detach().clone().requires_grad_() + reference_head = head.detach().clone().requires_grad_() + _loss([_model(x, reference_weight) for x in inputs], reference_head).backward() + torch.testing.assert_close(weight.grad, reference_weight.grad) + torch.testing.assert_close(head.grad, reference_head.grad) + + +def test_aliases_unused_forwards_and_explicit_repeated_backward(): + collector = CotangentCollector() + value = torch.tensor([2.0, 3.0], requires_grad=True) + output = collector.attach(detach_tree("used", [value, value])) + collector.attach(detach_tree("unused", value)) + assert output[0] is output[1] + loss = (output[0] * output[1]).sum() + first = collector.backward(loss, retain_graph=True) + second = collector.backward(loss) + assert len(first) == len(second) == 1 + torch.testing.assert_close(first[0].gradients[0], 2 * value) + torch.testing.assert_close(second[0].gradients[0], first[0].gradients[0]) + with pytest.raises(RuntimeError, match="second time"): + collector.backward(loss) + + +def test_tuple_backward_and_multiple_snapshots_of_same_handle(): + collector = CotangentCollector() + packet = detach_tree("head:1", torch.tensor([2.0, 3.0], requires_grad=True)) + a, b = collector.attach(packet), collector.attach(packet) + gradients = collector.backward((a, b), (torch.ones(2), torch.full((2,), 2.0))) + assert len(gradients) == 1 + torch.testing.assert_close(gradients[0].gradients[0], torch.full((2,), 3.0)) + + +@pytest.mark.parametrize("managed", [False, True]) +def test_local_failure_discards_already_collected_remote_cotangents(managed): + collector = CotangentCollector() + bridge = collector.attach( + detach_tree("first", torch.tensor(3.0, requires_grad=True)) + ) + value = managed_tensor(bridge) if managed else bridge + seen = [] + + def fail_after_collection(*args): + seen.append(bool(collector._pending)) + raise RuntimeError("local loss failed") + + hook = bridge.grad_fn.register_hook(fail_after_collection) + with pytest.raises(RuntimeError, match="local loss failed"): + collector.backward(value, retain_graph=True) + assert seen == [True] + assert collector._pending is None + hook.remove() + recovered = collector.backward(value) + assert len(recovered) == 1 + torch.testing.assert_close(recovered[0].gradients[0], torch.tensor(1.0)) + + +def test_unscoped_backward_is_rejected(): + collector = CotangentCollector() + value = collector.attach( + detach_tree("forward", torch.tensor(3.0, requires_grad=True)) + ) + with pytest.raises(RuntimeError, match="trainer.backward"): + value.backward() + + +def test_empty_and_nondifferentiable_trees(): + collector = CotangentCollector() + tree = [None, {"tokens": torch.tensor([1, 2]), "labels": []}] + result = collector.attach(detach_tree("empty", tree)) + assert result[0] is None + assert not result[1]["tokens"].requires_grad + assert result[1]["labels"] == [] + assert collector.backward(torch.tensor(2.0, requires_grad=True)) == () + + +@pytest.mark.parametrize("reverse", [False, True]) +def test_managed_cpu_arithmetic_keeps_original_gradient_paths(reverse): + a = torch.tensor([2.0, 3.0], requires_grad=True) + b = torch.tensor([5.0, 7.0], requires_grad=True) + managed = managed_tree({"a": a})["a"] + result = b * managed if reverse else managed * b + assert isinstance(result, ManagedTensor) + result.sum().backward() + torch.testing.assert_close(a.grad, b.detach()) + torch.testing.assert_close(b.grad, a.detach()) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") +@pytest.mark.parametrize("managed_device", ["cpu", "cuda"]) +@pytest.mark.parametrize("reverse", [False, True]) +def test_managed_cross_device_arithmetic_and_original_gradients( + managed_device, reverse +): + other_device = "cuda" if managed_device == "cpu" else "cpu" + a = torch.tensor([2.0, 3.0], device=managed_device, requires_grad=True) + b = torch.tensor([5.0, 7.0], device=other_device, requires_grad=True) + managed = managed_tensor(a) + result = b * managed if reverse else managed * b + assert result.device.type == managed_device + result.sum().backward() + assert a.grad is not None and b.grad is not None + assert a.grad.device == a.device and b.grad.device == b.device + torch.testing.assert_close(a.grad.cpu(), b.detach().cpu()) + torch.testing.assert_close(b.grad.cpu(), a.detach().cpu()) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") +@pytest.mark.parametrize("input_device", ["cpu", "cuda"]) +def test_managed_linear_cpu_cuda_matches_explicit_copies(input_device): + other_device = "cuda" if input_device == "cpu" else "cpu" + torch.manual_seed(7) + inputs = torch.randn( + 2, 3, device=input_device, dtype=torch.float64, requires_grad=True + ) + head = torch.nn.Linear(3, 4, device=other_device, dtype=torch.float64) + result = head(managed_tensor(inputs)).square().sum() + assert result.device.type == input_device + result.backward() + expected_input = inputs.detach().cpu().requires_grad_() + expected_weight = head.weight.detach().cpu().requires_grad_() + expected_bias = head.bias.detach().cpu().requires_grad_() + torch.nn.functional.linear( + expected_input, expected_weight, expected_bias + ).square().sum().backward() + for actual, expected in [ + (inputs, expected_input), + (head.weight, expected_weight), + (head.bias, expected_bias), + ]: + assert actual.grad is not None + assert actual.grad.device == actual.device + torch.testing.assert_close(actual.grad.cpu(), expected.grad) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") +def test_managed_nested_operands_and_mixed_device_mutation_rejection(): + a = torch.tensor([2.0], requires_grad=True) + b = torch.tensor([3.0], device="cuda", requires_grad=True) + result = torch.cat([managed_tensor(a), b]) + assert result.device.type == "cpu" + result.sum().backward() + assert a.grad is not None and b.grad is not None + assert a.grad.item() == b.grad.item() == 1 + with pytest.raises(RuntimeError, match="mutation"): + managed_tensor(a.detach()).add_(b.detach()) + with pytest.raises(RuntimeError, match="mutation"): + torch.add(managed_tensor(a.detach()), b.detach(), out=torch.empty_like(b)) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") +@pytest.mark.parametrize("reverse", [False, True]) +def test_conflicting_managed_devices_choose_cpu_in_either_direction(reverse): + cpu = torch.tensor([2.0], requires_grad=True) + cuda = torch.tensor([3.0], device="cuda", requires_grad=True) + a, b = managed_tensor(cpu), managed_tensor(cuda) + result = b * a if reverse else a * b + assert result.device.type == "cpu" + result.sum().backward() + assert cpu.grad is not None and cuda.grad is not None + assert cpu.grad.device.type == "cpu" and cpu.grad.item() == 3 + assert cuda.grad.device.type == "cuda" and cuda.grad.item() == 2 + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") +def test_cpu_output_bridge_gpu_head_and_remote_model_cotangents(): + weight = torch.tensor( + [[2.0], [3.0]], device="cuda", dtype=torch.float64, requires_grad=True + ) + physical = weight.square() + collector = CotangentCollector() + output = collector.attach( + detach_tree("gpu-model", physical, device="cpu"), managed=True + ) + head = torch.nn.Linear(1, 1, bias=False, device="cuda", dtype=torch.float64) + with torch.no_grad(): + head.weight.fill_(4) + loss = head(output).sum() + assert loss.device.type == "cpu" + (packet,) = collector.backward(loss) + assert packet.gradients[0] is not None + assert packet.gradients[0].device.type == "cpu" + assert weight.grad is None + torch.autograd.backward(physical, packet.gradients[0].to(physical.device)) + torch.testing.assert_close(weight.grad, 8 * weight.detach()) + torch.testing.assert_close( + head.weight.grad, torch.tensor([[13.0]], device="cuda", dtype=torch.float64) + ) + + +@pytest.mark.parametrize("managed", [False, True]) +def test_release_follows_dependent_loss_not_temporary_output(managed): + import gc + import weakref + + collector = CotangentCollector() + released = [] + value = collector.attach( + detach_tree("released", torch.tensor(2.0, requires_grad=True)), + managed=managed, + on_release=lambda: released.append("released"), + ) + original = weakref.ref(value) + loss = value + 1 # addition does not save its input Python wrapper + del value + gc.collect() + assert original() is None + assert not released + collector.backward(loss) + assert not released + del loss + gc.collect() + assert released == ["released"] + gc.collect() + assert released == ["released"] + + +def test_release_follows_retained_graph_and_all_output_branches(): + import gc + + collector = CotangentCollector() + released = [] + outputs = collector.attach( + detach_tree( + "retained", + [ + torch.tensor(2.0, requires_grad=True), + torch.tensor(3.0, requires_grad=True), + ], + ), + on_release=lambda: released.append(True), + ) + first, second = [output.square() for output in outputs] + del outputs + collector.backward(first, retain_graph=True) + collector.backward(first, retain_graph=True) + del first + gc.collect() + assert not released + collector.backward(second) + del second + gc.collect() + assert released == [True] + + +def test_release_dropped_and_nondifferentiable_packets(): + import gc + + collector = CotangentCollector() + released = [] + output = collector.attach( + detach_tree("unused", torch.tensor(2.0, requires_grad=True)), + on_release=lambda: released.append("unused"), + ) + del output + gc.collect() + assert released == ["unused"] + output = collector.attach( + detach_tree("frozen", torch.tensor(3.0)), + on_release=lambda: released.append("frozen"), + ) + assert released == ["unused", "frozen"] + assert output.item() == 3 + with torch.no_grad(): + collector.attach( + detach_tree("disabled", torch.tensor(4.0, requires_grad=True)), + on_release=lambda: released.append("disabled"), + ) + gc.collect() + assert released == ["unused", "frozen", "disabled"] + + +@pytest.mark.parametrize("use_grad", [False, True]) +def test_managed_backward_preserves_existing_graph_under_no_grad(use_grad): + source = torch.tensor([2.0, 3.0], requires_grad=True) + managed = managed_tensor(source) + loss = managed.square().sum() + with torch.no_grad(): + assert not (managed * 2).requires_grad + if use_grad: + (gradient,) = torch.autograd.grad(loss, source) + else: + loss.backward() + gradient = source.grad + torch.testing.assert_close(gradient, 2 * source.detach()) + + +def test_managed_autograd_grad_targets_original_tensor_and_higher_derivatives(): + source = torch.tensor([2.0, 3.0], requires_grad=True) + managed = managed_tensor(source) + loss = managed.pow(3).sum() + (gradient,) = torch.autograd.grad(loss, managed, create_graph=True) + (second,) = torch.autograd.grad(gradient.sum(), managed) + torch.testing.assert_close(gradient, 3 * source.detach().square()) + torch.testing.assert_close(second, 6 * source.detach()) + + +def test_managed_hooks_and_retained_gradients_use_original_tensor(): + source = torch.tensor([2.0, 3.0], requires_grad=True) + managed = managed_tensor(source) + calls = [] + managed.register_hook(lambda gradient: calls.append(gradient.clone())) + managed.retain_grad() + managed.square().sum().backward() + assert len(calls) == 1 + torch.testing.assert_close(calls[0], 2 * source.detach()) + torch.testing.assert_close(managed.grad, source.grad) + + +def test_local_autograd_grad_of_managed_bridge_output_does_not_commit(): + collector = CotangentCollector() + value = collector.attach( + detach_tree("local-grad", torch.tensor(3.0, requires_grad=True)), managed=True + ) + loss = value.square() + (gradient,) = torch.autograd.grad(loss, value, retain_graph=True) + assert gradient.item() == 6 and collector._pending is None + (packet,) = collector.backward(loss) + assert packet.gradients[0] is not None + assert packet.gradients[0].item() == 6 + + +@pytest.mark.parametrize("managed", [False, True]) +def test_output_hooks_change_or_reject_collected_cotangents(managed): + collector = CotangentCollector() + value = collector.attach( + detach_tree("hook", torch.tensor(3.0, requires_grad=True)), managed=managed + ) + hook = value.register_hook(lambda gradient: gradient * 0) + (packet,) = collector.backward(value.square(), retain_graph=True) + torch.testing.assert_close(packet.gradients[0], torch.tensor(0.0)) + hook.remove() + + def reject(gradient): + raise RuntimeError("hook rejects loss") + + value.register_hook(reject) + with pytest.raises(RuntimeError, match="hook rejects loss"): + collector.backward(value.square()) + assert collector._pending is None + + +@pytest.mark.parametrize("managed", [False, True]) +def test_unrelated_failed_backward_cannot_enter_an_active_collection(managed): + from concurrent.futures import ThreadPoolExecutor + from threading import Event + + collector = CotangentCollector() + a = collector.attach( + detach_tree("a", torch.tensor(2.0, requires_grad=True)), managed=managed + ) + b = collector.attach(detach_tree("b", torch.tensor(3.0, requires_grad=True))) + paused, resume = Event(), Event() + + def pause(gradient): + paused.set() + assert resume.wait(10) + return gradient + + def fail_foreign(*args): + raise RuntimeError("foreign backward failed after collection") + + a.register_hook(pause) + b.grad_fn.register_hook(fail_foreign) + with ThreadPoolExecutor(1) as pool: + future = pool.submit(collector.backward, a) + try: + assert paused.wait(10) + with pytest.raises(RuntimeError, match="owning trainer.backward"): + b.backward() + with pytest.raises(RuntimeError, match="already active"): + collector.backward(b) + finally: + resume.set() + packets = future.result(timeout=10) + assert [packet.handle for packet in packets] == ["a"] + torch.testing.assert_close(packets[0].gradients[0], torch.tensor(1.0)) + + +def test_nested_remote_backward_is_rejected_without_polluting_parent_task(): + collector = CotangentCollector() + a = collector.attach(detach_tree("a", torch.tensor(2.0, requires_grad=True))) + b = collector.attach(detach_tree("b", torch.tensor(3.0, requires_grad=True))) + rejected = [] + + def nested(gradient): + with pytest.raises(RuntimeError, match="nested remote backward"): + b.backward() + rejected.append(True) + return gradient + + a.register_hook(nested) + packets = collector.backward(a) + assert rejected == [True] + assert [packet.handle for packet in packets] == ["a"] + + +@pytest.mark.parametrize("use_reentrant", [False, True]) +@pytest.mark.parametrize("managed", [False, True]) +def test_local_checkpoint_recomputation_keeps_collection_task(use_reentrant, managed): + from torch.utils.checkpoint import checkpoint + + collector = CotangentCollector() + physical = torch.tensor([2.0, 3.0], dtype=torch.float64, requires_grad=True) + value = collector.attach(detach_tree("checkpoint", physical), managed=managed) + head = torch.tensor([0.3, -0.2], dtype=torch.float64, requires_grad=True) + output = checkpoint(lambda x: (x * head).sin(), value, use_reentrant=use_reentrant) + (packet,) = collector.backward(output.sum()) + torch.testing.assert_close( + packet.gradients[0], (physical.detach() * head.detach()).cos() * head.detach() + ) + torch.testing.assert_close( + head.grad, (physical.detach() * head.detach()).cos() * physical.detach() + ) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") +@pytest.mark.parametrize("destination", ["cpu", "cuda"]) +def test_cross_device_assignment_rejects_without_discarding_writes(destination): + source_device = "cuda" if destination == "cpu" else "cpu" + dst = torch.zeros(2, device=destination) + src = managed_tensor(torch.tensor([4.0, 5.0], device=source_device)) + with pytest.raises(RuntimeError, match="__setitem__.*cpu.*cuda.*explicit"): + dst[:] = src + torch.testing.assert_close(dst, torch.zeros_like(dst)) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") +def test_mixed_device_stateful_modules_require_explicit_placement(): + inputs = managed_tensor(torch.tensor([[1.0, 3.0], [2.0, 6.0]])) + batch_norm = torch.nn.BatchNorm1d(2, device="cuda") + assert batch_norm.running_mean is not None and batch_norm.running_var is not None + mean, variance = batch_norm.running_mean.clone(), batch_norm.running_var.clone() + with pytest.raises(RuntimeError, match="batch_norm.*explicit"): + batch_norm(inputs) + torch.testing.assert_close(batch_norm.running_mean, mean) + torch.testing.assert_close(batch_norm.running_var, variance) + weight = torch.full((2, 2), 4.0, device="cuda") + tokens = managed_tensor(torch.tensor([0, 1])) + with pytest.raises(RuntimeError, match="embedding.*explicit"): + torch.nn.functional.embedding(tokens, weight, max_norm=1) + torch.testing.assert_close(weight, torch.full_like(weight, 4.0)) + + +def test_same_device_managed_batch_norm_preserves_buffer_updates(): + inputs = torch.tensor([[1.0, 3.0], [2.0, 6.0]], requires_grad=True) + actual, reference = torch.nn.BatchNorm1d(2), torch.nn.BatchNorm1d(2) + result = actual(managed_tensor(inputs)) + expected = reference(inputs) + torch.testing.assert_close(result, expected) + for name, buffer in actual.named_buffers(): + torch.testing.assert_close(buffer, dict(reference.named_buffers())[name]) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") +def test_mixed_device_common_loss_and_indexing_are_supported(): + source = torch.tensor([[1.0, 2.0], [3.0, 4.0]], requires_grad=True) + target = torch.zeros((2, 2), device="cuda", requires_grad=True) + output = managed_tensor(source) + loss = torch.nn.functional.mse_loss(output, target) + assert loss.device.type == "cpu" + loss.backward() + assert target.grad is not None and source.grad is not None + torch.testing.assert_close(source.grad, source.detach() / 2) + torch.testing.assert_close(target.grad.cpu(), -source.detach() / 2) + selected = torch.index_select(output, 0, torch.tensor([1], device="cuda")) + torch.testing.assert_close(selected, source[1:]) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") +def test_collection_task_spans_cpu_and_cuda_bridge_nodes(): + collector = CotangentCollector() + source = [ + torch.tensor([2.0, 3.0], device=device, requires_grad=True) + for device in ("cpu", "cuda") + ] + outputs = [ + collector.attach(detach_tree(str(i), value), managed=True) + for i, value in enumerate(source) + ] + losses = tuple(value.square().sum() for value in outputs) + for retain_graph in (True, False): + packets = collector.backward(losses, retain_graph=retain_graph) + assert len(packets) == 2 + for packet, value in zip(packets, source, strict=True): + assert packet.gradients[0] is not None + assert packet.gradients[0].device == value.device + torch.testing.assert_close(packet.gradients[0], 2 * value.detach()) + + +def test_managed_requires_grad_mutates_original_identity(): + value = managed_tensor(torch.tensor([2.0, 3.0])) + returned = value.requires_grad_() + assert returned is value and value.requires_grad + value.square().sum().backward() + torch.testing.assert_close(value.grad, torch.tensor([4.0, 6.0])) + + +@pytest.mark.parametrize("different_count", [False, True]) +def test_duplicate_handles_reject_incompatible_output_signatures(different_count): + collector = CotangentCollector() + a = collector.attach(detach_tree("duplicate", torch.ones(2, requires_grad=True))) + source = ( + [torch.ones(2, requires_grad=True), torch.ones(2, requires_grad=True)] + if different_count + else torch.ones(3, requires_grad=True) + ) + b = collector.attach(detach_tree("duplicate", source)) + other_loss = sum(value.sum() for value in b) if different_count else b.sum() + with pytest.raises(ValueError, match="Incompatible output signatures.*duplicate"): + collector.backward(a.sum() + other_loss) + assert collector._pending is None and not collector._signatures + + +@pytest.mark.parametrize("managed", [False, True]) +@pytest.mark.parametrize("under_no_grad", [False, True]) +def test_bridged_outputs_require_clone_before_inplace_writes(managed, under_no_grad): + source = torch.tensor([2.0, 3.0], requires_grad=True) + collector = CotangentCollector() + output = collector.attach(detach_tree("readonly", source), managed=managed) + if under_no_grad: + with pytest.raises(RuntimeError, match="view.*modified|modified.*view"): + with torch.no_grad(): + output.mul_(2) + collector.backward(output.sum()) + else: + with pytest.raises(RuntimeError, match="view.*modified|modified.*view"): + output.mul_(2) + torch.testing.assert_close(source, torch.tensor([2.0, 3.0])) + writable = collector.attach( + detach_tree("writable", source), managed=managed + ).clone() + writable.mul_(2) + (packet,) = collector.backward(writable.sum()) + torch.testing.assert_close(packet.gradients[0], torch.full((2,), 2.0)) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") +@pytest.mark.parametrize("source_device", ["cpu", "cuda"]) +@pytest.mark.parametrize("method", ["to", "type_as"]) +def test_explicit_tensor_transfer_specs_preserve_destination_and_gradient( + source_device, method +): + destination = "cuda" if source_device == "cpu" else "cpu" + source = torch.tensor([2.0, 3.0], device=source_device, requires_grad=True) + spec = torch.empty(0, device=destination, dtype=torch.float64) + result = getattr(managed_tensor(source), method)(spec) + assert result.device == spec.device and result.dtype == spec.dtype + result.sum().backward() + torch.testing.assert_close(source.grad, torch.ones_like(source)) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") +@pytest.mark.parametrize( + "operator", ["__iand__", "__ior__", "__ixor__", "__ilshift__", "__irshift__"] +) +def test_mixed_device_bitwise_inplace_dunders_reject_without_rebinding(operator): + target = torch.ones(2, dtype=torch.int64, device="cuda") + source = managed_tensor(torch.ones(2, dtype=torch.int64)) + with pytest.raises(RuntimeError, match="mutation.*explicit"): + getattr(target, operator)(source) + torch.testing.assert_close(target, torch.ones_like(target)) + assert target.device.type == "cuda" and type(target) is torch.Tensor + + +@pytest.mark.skipif(torch.cuda.device_count() < 2, reason="Two CUDA devices required") +def test_managed_ambiguous_cuda_devices_require_explicit_transfer(): + source = managed_tensor(torch.ones(2, device="cuda:0")) + other = torch.ones(2, device="cuda:1") + with pytest.raises(RuntimeError, match="explicit move between accelerator devices"): + source + other + moved = source.to(other) + assert moved.device == other.device + + +@pytest.mark.parametrize( + "container", ["set", "object", "tensor_key", "default_factory"] +) +def test_output_packets_reject_opaque_tensor_bearing_metadata(container): + from collections import defaultdict + from types import SimpleNamespace + + source = torch.tensor(3.0, requires_grad=True) + tree = ( + {source} + if container == "set" + else SimpleNamespace(value=source) + if container == "object" + else defaultdict(lambda: source) + if container == "default_factory" + else {source: "key"} + ) + with pytest.raises(TypeError, match="Unsupported output tree metadata"): + detach_tree("opaque", tree) + + +def test_graph_release_callback_does_not_run_at_interpreter_shutdown(): + import subprocess + import sys + + result = subprocess.run( + [ + sys.executable, + "-c", + """ +import torch +from art.trainer_rank._tensors import CotangentCollector, detach_tree +collector = CotangentCollector() +value = collector.attach( + detach_tree('alive-at-exit', torch.tensor(2., requires_grad=True)), + on_release=lambda: print('UNEXPECTED_RELEASE_AT_EXIT', flush=True), +) +print('OUTPUT_REMAINS_ALIVE', flush=True) +""", + ], + check=True, + capture_output=True, + text=True, + timeout=60, + ) + assert "OUTPUT_REMAINS_ALIVE" in result.stdout + assert "UNEXPECTED_RELEASE_AT_EXIT" not in result.stdout + + +def test_real_microbatch_packet_roundtrip_preserves_unset_and_backward(): + from art.trainer_rank import ForwardInput, MicroBatch, MicroBatchStats, Unset + + source = torch.tensor([2.0, 3.0], requires_grad=True) + batch = MicroBatch( + inputs=[ForwardInput(input_tokens=torch.tensor([1, 2]))], + outputs=[ForwardOutput(source, None, None, None)], + indices=[2], + stats=MicroBatchStats(0, 3, 3, 1, 2, 2, 0, 0, 0, False), + ) + packet = pickle.loads(pickle.dumps(detach_tree("microbatch", batch))) + collector = CotangentCollector() + attached = collector.attach(packet) + assert isinstance(attached, MicroBatch) + assert attached.inputs[0].checkpoint is Unset + assert attached.select(["a", "b", "c"]) == ["c"] + assert attached.stats == batch.stats + (gradient,) = collector.backward(attached.outputs[0].target_logprobs.square().sum()) + assert gradient.gradients[0] is None # input token IDs + torch.testing.assert_close(gradient.gradients[1], 2 * source.detach()) + + +@pytest.mark.parametrize("managed", [False, True]) +@pytest.mark.parametrize("device", ["cpu", "cuda"]) +def test_cloned_output_inplace_preserves_original_hooks(managed, device): + if device == "cuda" and not torch.cuda.is_available(): + pytest.skip("CUDA required") + collector = CotangentCollector() + value = collector.attach( + detach_tree("inplace-hook", torch.ones(2, device=device, requires_grad=True)), + managed=managed, + ).clone() + calls = [] + + def zero(gradient): + calls.append(gradient.clone()) + return gradient * 0 + + value.register_hook(zero) + assert value.mul_(2) is value + (packet,) = collector.backward(value.sum()) + assert len(calls) == 1 + torch.testing.assert_close(calls[0], torch.full_like(value, 2)) + torch.testing.assert_close(packet.gradients[0], torch.zeros_like(value)) + + +@pytest.mark.parametrize("managed", [False, True]) +def test_inplace_transpose_updates_shape_and_gradient_on_original(managed): + collector = CotangentCollector() + source = torch.arange(6.0).reshape(2, 3).requires_grad_() + value = collector.attach(detach_tree("transpose", source), managed=managed).clone() + assert value.transpose_(0, 1) is value + assert value.shape == (3, 2) + weights = torch.arange(1.0, 7.0).reshape(3, 2) + (packet,) = collector.backward((value * weights).sum()) + torch.testing.assert_close(packet.gradients[0], weights.T) + + +@pytest.mark.parametrize("managed", [False, True]) +@pytest.mark.parametrize("device", ["cpu", "cuda"]) +@pytest.mark.parametrize( + "conversion", ["device", "tensor", "dtype", "type_as", "host_or_cuda", "contiguous"] +) +def test_noop_conversion_preserves_identity_gradient_queries_and_later_hooks( + managed, device, conversion +): + if device == "cuda" and not torch.cuda.is_available(): + pytest.skip("CUDA required") + collector = CotangentCollector() + source = torch.tensor([2.0, 3.0], device=device, requires_grad=True) + value = collector.attach(detach_tree("noop", source), managed=managed) + loss = value.square().sum() + with torch.no_grad(): + converted = ( + value.to(value.device) + if conversion == "device" + else value.to(source) + if conversion == "tensor" + else value.to(value.dtype) + if conversion == "dtype" + else value.type_as(source) + if conversion == "type_as" + else (value.cpu() if device == "cpu" else value.cuda()) + if conversion == "host_or_cuda" + else value.contiguous() + ) + assert converted is value + hook = converted.register_hook(lambda gradient: gradient * 0) + (gradient,) = torch.autograd.grad(loss, converted, retain_graph=True) + torch.testing.assert_close(gradient, torch.zeros_like(source)) + hook.remove() + (gradient,) = torch.autograd.grad(loss, converted, retain_graph=True) + torch.testing.assert_close(gradient, 2 * source.detach()) + (packet,) = collector.backward(loss) + torch.testing.assert_close(packet.gradients[0], 2 * source.detach()) + + +def test_same_device_out_preserves_ordinary_output_identity(): + source = managed_tensor(torch.tensor([2.0, 3.0])) + destination = torch.empty(2) + result = torch.add(source, 1, out=destination) + assert result is destination and type(result) is torch.Tensor + torch.testing.assert_close(destination, torch.tensor([3.0, 4.0])) + + +@pytest.mark.parametrize("reverse", [False, True]) +def test_managed_operands_defer_to_other_tensor_snapshot_dispatch(reverse): + collector = CotangentCollector() + source = torch.tensor([2.0, 3.0], requires_grad=True) + managed = collector.attach(detach_tree("model", source), managed=True) + captures = [] + + class SnapshotParameter(torch.Tensor): + @classmethod + def __torch_function__(cls, func, types, args=(), kwargs=None): + tensors, spec = flatten_tensors((args, kwargs or {})) + snapshot = collector.attach( + detach_tree("head", torch.tensor([5.0, 7.0], requires_grad=True)) + ) + captures.append(True) + args, kwargs = unflatten_tensors( + spec, + tuple( + snapshot if isinstance(tensor, cls) else tensor + for tensor in tensors + ), + ) + return func(*args, **kwargs) + + # A live proxy's storage is deliberately stale; only its dispatch supplies + # the current captured head value, as the real client parameter handle does. + proxy = torch.Tensor._make_subclass(SnapshotParameter, torch.zeros(2)) + loss = (proxy * managed if reverse else managed * proxy).sum() + assert captures == [True] + packets = {packet.handle: packet for packet in collector.backward(loss)} + torch.testing.assert_close(packets["model"].gradients[0], torch.tensor([5.0, 7.0])) + torch.testing.assert_close(packets["head"].gradients[0], source.detach()) + + +@pytest.mark.parametrize("fail", [False, True]) +def test_flatten_releases_tensor_references_without_cyclic_gc(fail): + import gc + import weakref + + @dataclass + class Broken: + value: int = 0 + + def __getattribute__(self, name): + if name == "value": + raise RuntimeError("broken dataclass field") + return object.__getattribute__(self, name) + + was_enabled = gc.isenabled() + gc.disable() + try: + source = torch.ones(2, requires_grad=True) + reference = weakref.ref(source) + tree = [source, Broken()] if fail else [source] + if fail: + with pytest.raises(RuntimeError, match="broken dataclass field"): + flatten_tensors(tree) + else: + leaves, spec = flatten_tensors(tree) + del leaves, spec + del tree, source + assert reference() is None + finally: + if was_enabled: + gc.enable() + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") +@pytest.mark.parametrize("device", ["cpu", "cuda"]) +@pytest.mark.parametrize("attribute", ["T", "mT", "H", "mH", "real", "imag"]) +def test_tensor_properties_keep_managed_placement_and_original_gradients( + device, attribute +): + source = torch.tensor( + [[1 + 2j, 3 - 1j], [2 - 1j, 4 + 3j]], + dtype=torch.complex128, + device=device, + requires_grad=True, + ) + view = getattr(managed_tensor(source), attribute) + assert isinstance(view, ManagedTensor) + operand = torch.full( + (2, 3), + 2.0, + dtype=view.dtype, + device="cuda" if device == "cpu" else "cpu", + requires_grad=True, + ) + result = view @ operand + assert result.device == source.device + result.abs().square().sum().backward() + expected_source = source.detach().cpu().requires_grad_() + expected_operand = operand.detach().cpu().requires_grad_() + ( + getattr(expected_source, attribute) @ expected_operand + ).abs().square().sum().backward() + assert source.grad is not None and operand.grad is not None + assert source.grad.device == source.device and operand.grad.device == operand.device + torch.testing.assert_close(source.grad.cpu(), expected_source.grad) + torch.testing.assert_close(operand.grad.cpu(), expected_operand.grad) + + +@pytest.mark.parametrize("device", ["cpu", "cuda"]) +@pytest.mark.parametrize( + "layouts", [("dense", "sparse"), ("sparse", "dense"), ("sparse", "sparse")] +) +def test_duplicate_handle_sparse_cotangents_accumulate_in_any_order(device, layouts): + if device == "cuda" and not torch.cuda.is_available(): + pytest.skip("CUDA required") + collector = CotangentCollector() + packet = detach_tree("mixed-grad", torch.ones(4, device=device, requires_grad=True)) + outputs = [collector.attach(packet) for _ in layouts] + indices = torch.tensor([0, 2], device=device) + losses = [ + output.sum() + if layout == "dense" + else output.gather(0, indices, sparse_grad=True).sum() + for output, layout in zip(outputs, layouts, strict=True) + ] + (collected,) = collector.backward(sum(losses)) + gradient = collected.gradients[0] + assert gradient is not None + if layouts == ("sparse", "sparse"): + assert gradient.is_sparse + expected = torch.tensor([2.0, 0.0, 2.0, 0.0], device=device) + else: + expected = torch.tensor([2.0, 1.0, 2.0, 1.0], device=device) + torch.testing.assert_close(gradient.to_dense(), expected) + + +def test_transformers_dataclass_mapping_packet_preserves_fields_entries_and_aliases(): + from typing import cast + + from transformers.modeling_outputs import BaseModelOutput + + source = cast(torch.FloatTensor, torch.tensor([2.0, 3.0], requires_grad=True)) + physical = BaseModelOutput( + last_hidden_state=source, + hidden_states=(source, cast(torch.FloatTensor, source * 2)), + ) + physical["extra"] = source * 3 + packet = pickle.loads(pickle.dumps(detach_tree("mapping", physical))) + collector = CotangentCollector() + output = collector.attach(packet) + assert isinstance(output, BaseModelOutput) + assert list(output) == list(physical) + assert output.last_hidden_state is output["last_hidden_state"] is output[0] + assert output.hidden_states[0] is output.last_hidden_state + assert output.attentions is None and "attentions" not in output + assert output.to_tuple()[-1] is output["extra"] + (gradients,) = collector.backward( + output.last_hidden_state.sum() + output["extra"].sum() + ) + assert gradients.gradients[1] is None + leaves, _ = flatten_tensors(physical) + selected = [ + (value, gradient) + for value, gradient in zip(leaves, gradients.gradients, strict=True) + if gradient is not None + ] + torch.autograd.backward( + [value for value, _ in selected], [gradient for _, gradient in selected] + ) + torch.testing.assert_close(source.grad, torch.full_like(source, 4)) diff --git a/tests/unit/test_trainer_rank_validation.py b/tests/unit/test_trainer_rank_validation.py index d6ed71048..9e9127be6 100644 --- a/tests/unit/test_trainer_rank_validation.py +++ b/tests/unit/test_trainer_rank_validation.py @@ -3,7 +3,7 @@ import asyncio from collections.abc import Iterable from concurrent.futures import ThreadPoolExecutor -from dataclasses import dataclass, replace +from dataclasses import dataclass, fields, replace from datetime import timedelta import gc from importlib.util import find_spec @@ -263,7 +263,7 @@ def test_forward_input_distinguishes_unset_and_base_checkpoint( assert request.checkpoint is expected -def test_dp_rank_forward_rejects_unloaded_explicit_checkpoint( +def test_forward_rejects_unloaded_explicit_checkpoint( monkeypatch: pytest.MonkeyPatch, ) -> None: trainer = TrainerRank(_runtime()) @@ -275,11 +275,11 @@ def test_dp_rank_forward_rejects_unloaded_explicit_checkpoint( ) with pytest.raises(TrainerRankSlotStateError, match="unloaded.*'typo'"): - trainer.dp_rank_forward([request]) + trainer.forward([request]) @pytest.mark.parametrize("checkpoint", (None, "student")) -def test_dp_rank_forward_accepts_base_or_loaded_explicit_checkpoint( +def test_forward_accepts_base_or_loaded_explicit_checkpoint( monkeypatch: pytest.MonkeyPatch, checkpoint: str | None, ) -> None: @@ -292,7 +292,7 @@ def test_dp_rank_forward_accepts_base_or_loaded_explicit_checkpoint( checkpoint=checkpoint, ) - output = trainer.dp_rank_forward([request]) + output = trainer.forward([request]) assert isinstance(output[0], ForwardOutput) @@ -325,7 +325,7 @@ def execute(plan: object, **_kwargs: object) -> list[ForwardOutput]: ), ] - trainer.dp_rank_forward(inputs, checkpoint="method") + trainer.forward(inputs, checkpoint="method") assert seen == ["method", "request", None] @@ -337,7 +337,7 @@ def test_forward_method_checkpoint_rejects_unloaded_name( _stub_forward(monkeypatch, trainer) with pytest.raises(TrainerRankSlotStateError, match="unloaded.*'typo'"): - trainer.dp_rank_forward([_target_request(1)], checkpoint="typo") + trainer.forward([_target_request(1)], checkpoint="typo") @pytest.mark.parametrize( @@ -349,7 +349,7 @@ def test_forward_method_checkpoint_rejects_unloaded_name( (False, False, True), ), ) -def test_dp_rank_forward_grad_mode( +def test_forward_grad_mode( monkeypatch: pytest.MonkeyPatch, ambient_grad: bool, no_grad: bool | None, @@ -364,12 +364,12 @@ def execute(plan: object, **_kwargs: object) -> list[ForwardOutput]: _stub_forward(monkeypatch, trainer, execute) with torch.set_grad_enabled(ambient_grad): - trainer.dp_rank_forward([_target_request(1)], no_grad=no_grad) + trainer.forward([_target_request(1)], no_grad=no_grad) assert seen == [expected] -@pytest.mark.parametrize("api", ("dp_rank_forward", "forward_micro_batches")) +@pytest.mark.parametrize("api", ("forward", "forward_batches")) def test_forward_input_overrides_grad_mode_by_group( monkeypatch: pytest.MonkeyPatch, api: str, @@ -390,10 +390,10 @@ def execute(plan: object, **_kwargs: object) -> list[ForwardOutput]: ) for token, no_grad in ((1, True), (2, False)) ] - if api == "dp_rank_forward": - trainer.dp_rank_forward(inputs) + if api == "forward": + trainer.forward(inputs) else: - list(trainer.forward_micro_batches(inputs)) + list(trainer.forward_batches(inputs)) assert seen == [False, True] @@ -424,24 +424,16 @@ def test_forward_groups_execute_in_their_selected_grad_modes( monkeypatch.setattr(trainer, "_validate_hybridep_topology", lambda: None) monkeypatch.setattr(trainer, "_topology", lambda: object()) monkeypatch.setattr(trainer, "_configure_hybridep", lambda *_args, **_kwargs: None) - monkeypatch.setattr(trainer, "_prepare_packed_forward", lambda _packed: None) - class UseLoRASlot: - def __enter__(self) -> None: - pass - - def __exit__(self, *_args: object) -> None: - pass - - lora = ModuleType("art.megatron.lora") - cast(Any, lora).use_lora_slot = lambda _slot: UseLoRASlot() - monkeypatch.setitem(sys.modules, "art.megatron.lora", lora) - - def forward(_items: object, _prepared: object) -> list[ForwardOutput]: + def forward(group: Any) -> list[ForwardOutput]: seen.append(torch.is_grad_enabled()) - return [ForwardOutput(None, None, None, None)] + return [ + ForwardOutput( + None, None, None, None, group.slot_ref.name, not group.grad_enabled + ) + ] - monkeypatch.setattr(trainer, "_forward_packed", forward) + monkeypatch.setattr(trainer, "_execute_graph_group", forward) outputs = cast(Any, trainer)._execute_flat_plan(plan) assert seen == [False, True] @@ -451,7 +443,7 @@ def forward(_items: object, _prepared: object) -> list[ForwardOutput]: ] -def test_forward_micro_batches_keeps_grad_mode_across_iteration( +def test_forward_batches_keeps_grad_mode_across_iteration( monkeypatch: pytest.MonkeyPatch, ) -> None: trainer = TrainerRank(_runtime()) @@ -462,7 +454,7 @@ def execute(plan: object, **_kwargs: object) -> list[ForwardOutput]: return _empty_outputs(plan) _stub_forward(monkeypatch, trainer, execute, profiled=True) - batches = trainer.forward_micro_batches( + batches = trainer.forward_batches( [_target_request(index) for index in range(3)], no_grad=True ) assert torch.is_grad_enabled() @@ -473,7 +465,7 @@ def execute(plan: object, **_kwargs: object) -> list[ForwardOutput]: assert torch.is_grad_enabled() -def test_forward_micro_batches_uses_method_checkpoint_fallback( +def test_forward_batches_uses_method_checkpoint_fallback( monkeypatch: pytest.MonkeyPatch, ) -> None: trainer = TrainerRank(_runtime()) @@ -487,7 +479,7 @@ def execute(plan: object, **_kwargs: object) -> list[ForwardOutput]: _stub_forward(monkeypatch, trainer, execute, profiled=True) - list(trainer.forward_micro_batches([_target_request(1)], checkpoint="teacher")) + list(trainer.forward_batches([_target_request(1)], checkpoint="teacher")) assert seen == ["teacher"] @@ -509,7 +501,7 @@ def test_forward_input_preserves_public_runtime_shape() -> None: ) def test_trainer_rank_rejects_removed_planner_knobs(knob: str) -> None: with pytest.raises(TypeError): - TrainerRank(_runtime(), **{knob: 1}) + TrainerRank(_runtime(), **cast(dict[str, Any], {knob: 1})) @pytest.mark.skipif(find_spec("megatron") is None, reason="requires Megatron") @@ -569,9 +561,9 @@ def test_hybridep_validates_topology_for_empty_forward( if dp > 1: with pytest.raises(NotImplementedError, match="DP=1"): - trainer.dp_rank_forward([]) + trainer.forward([]) else: - assert trainer.dp_rank_forward([]) == [] + assert trainer.forward([]) == [] def test_no_grad_groups_keep_the_fallback_score_and_their_own_cache_key() -> None: @@ -803,7 +795,7 @@ def test_snapshot_disposal_is_not_public() -> None: "parameter", "buffer", "forward", - "forward_micro_batches", + "forward_batches", "optim_step", "save", "export_lora", @@ -833,10 +825,10 @@ def install(target: TrainerRank, _source: object, name: str) -> None: trainer.buffer("mean", lambda: torch.zeros(1), checkpoint="student") elif consumer == "forward": _stub_forward(monkeypatch, trainer) - trainer.dp_rank_forward([_target_request(1)], checkpoint="student") - elif consumer == "forward_micro_batches": + trainer.forward([_target_request(1)], checkpoint="student") + elif consumer == "forward_batches": _stub_forward(monkeypatch, trainer, profiled=True) - list(trainer.forward_micro_batches([_target_request(1)], checkpoint="student")) + list(trainer.forward_batches([_target_request(1)], checkpoint="student")) elif consumer == "optim_step": with pytest.raises(TrainerRankSlotStateError, match="no gradients"): trainer.optim_step( @@ -1423,10 +1415,17 @@ def test_forward_snapshot_is_independent_and_forward_only( trainer.optim_step(params=AdamParams(learning_rate=1e-3), checkpoints=["saved"]) with pytest.raises(TrainerRankSlotStateError, match="load over forward-only"): trainer._guard_slot_can_load(saved) + with pytest.raises(TrainerRankSlotStateError, match="not a forward-only"): + trainer._discard_snapshot_checkpoint("student") + version = trainer._capture_checkpoint_version("saved") trainer._discard_snapshot_checkpoint("saved") assert "saved" not in trainer._checkpoint_slots assert lora._slot(saved) is None + assert trainer.snapshot_checkpoint("student", "saved") + assert trainer._capture_checkpoint_version("saved").generation > version.generation + with pytest.raises(TrainerRankSlotStateError, match="replaced"): + trainer._validate_checkpoint_version(version) def test_prepared_snapshot_loads_forward_only_without_replacing_slots( @@ -2438,6 +2437,7 @@ def make_trainer() -> TrainerRank: original.save_checkpoint(str(output), "student") assert not list(tmp_path.glob(".exact.snapshot-*")) assert not (tmp_path / ".exact.reserved").exists() + prepared = prepare_checkpoint(str(output)) assert prepared.manifest is not None assert validate_checkpoint(output) == prepared.manifest @@ -2477,6 +2477,43 @@ def make_trainer() -> TrainerRank: assert not list(tmp_path.glob(".exact.snapshot-*")) assert not (tmp_path / ".exact.reserved").exists() + old_version = restored._capture_checkpoint_version("student") + assert old_version.revision == 1 + checkpoint_module.load_checkpoint(restored, prepared, "student") + replacement = restored._capture_checkpoint_version("student") + assert replacement.generation == old_version.generation + 1 + assert replacement.revision == old_version.revision + 1 + with pytest.raises(TrainerRankSlotStateError, match="replaced"): + restored._validate_checkpoint_version(old_version) + restored._validate_checkpoint_version(replacement, 0) + for parameter in restored._checkpoint_slots["student"].params: + parameter.grad = torch.full_like(parameter, -0.125) + restored.optim_step(params=adam) + restored._validate_checkpoint_version(replacement, 1) + with pytest.raises(TrainerRankSlotStateError, match="gradient staleness 1"): + restored._validate_checkpoint_version(replacement, 0) + before_failure = restored._capture_checkpoint_version("student") + + with monkeypatch.context() as context: + + def fail_commit(*_args: object) -> None: + raise RuntimeError("injected generation commit failure") + + context.setattr(checkpoint_module, "_commit_slot", fail_commit) + with pytest.raises(RuntimeError, match="generation commit failure"): + checkpoint_module.load_checkpoint(restored, prepared, "student") + assert restored._capture_checkpoint_version("student") == before_failure + reserved_generation = restored._version_state().generation + assert reserved_generation > replacement.generation + checkpoint_module.load_checkpoint(restored, prepared, "student") + assert ( + restored._capture_checkpoint_version("student").revision + == before_failure.revision + 1 + ) + assert ( + restored._capture_checkpoint_version("student").generation > reserved_generation + ) + def test_trainer_rank_default_forward_uses_explicit_base_slot() -> None: trainer = TrainerRank(_runtime()) @@ -2845,9 +2882,7 @@ def test_optim_step_implicitly_ignores_resident_forward_snapshot( trainer._set_default_slot(_slot_ref("student")) _stub_forward(monkeypatch, trainer, profiled=True) list( - trainer.forward_micro_batches( - [_target_request(1)], checkpoint="saved", no_grad=True - ) + trainer.forward_batches([_target_request(1)], checkpoint="saved", no_grad=True) ) monkeypatch.setattr( trainer, @@ -3255,13 +3290,13 @@ def test_optim_step_live_graph_error_is_collective(tmp_path: Path) -> None: pytest.fail("collective live-graph policy test hung") -def test_dp_rank_forward_preserves_nested_shape_for_inactive_requests() -> None: +def test_forward_preserves_nested_shape_for_inactive_requests() -> None: trainer = TrainerRank(_runtime()) trainer._default_slot_ref = _slot_ref("teacher") request_a = ForwardInput(input_tokens=torch.tensor([1])) request_b = ForwardInput(input_tokens=torch.tensor([2])) - outputs = trainer.dp_rank_forward([[request_a], [request_b]], no_grad=True) + outputs = trainer.forward([[request_a], [request_b]], no_grad=True) assert len(outputs) == 2 assert len(outputs[0]) == 1 @@ -3271,11 +3306,11 @@ def test_dp_rank_forward_preserves_nested_shape_for_inactive_requests() -> None: assert outputs[1][0].checkpoint == "teacher" assert outputs[0][0].no_grad assert outputs[1][0].no_grad - assert not hasattr(trainer, "forward") + assert not hasattr(trainer, "dp_rank_forward") assert not hasattr(trainer, "micro_batches") -def test_dp_rank_forward_supports_arbitrary_nested_depth( +def test_forward_supports_arbitrary_nested_depth( monkeypatch: pytest.MonkeyPatch, ) -> None: trainer = TrainerRank(_runtime()) @@ -3285,7 +3320,7 @@ def test_dp_rank_forward_supports_arbitrary_nested_depth( [[[[[_target_request(3), _target_request(5)]]]]], ] - outputs = cast(Any, trainer).dp_rank_forward(nested) + outputs = cast(Any, trainer).forward(nested) assert _output_shape(outputs) == [ [[[[["output"]]]]], @@ -3295,7 +3330,7 @@ def test_dp_rank_forward_supports_arbitrary_nested_depth( @pytest.mark.parametrize("yield_empty", [False, True]) -def test_forward_micro_batches_uses_deterministic_dp_windows( +def test_forward_batches_uses_deterministic_dp_windows( monkeypatch: pytest.MonkeyPatch, yield_empty: bool, ) -> None: @@ -3303,7 +3338,7 @@ def test_forward_micro_batches_uses_deterministic_dp_windows( _stub_forward(monkeypatch, trainer, dp=(1, 2)) batches = list( - trainer.forward_micro_batches( + trainer.forward_batches( [_target_request(i) for i in range(5)], yield_empty=yield_empty ) ) @@ -3319,7 +3354,7 @@ def test_forward_micro_batches_uses_deterministic_dp_windows( @pytest.mark.parametrize( "operation", [ - "dp_reduce", + "reduce", "optim_step", "parameter", "module", @@ -3332,8 +3367,8 @@ def test_forward_micro_batches_uses_deterministic_dp_windows( "finish_checkpoint_save", "abort_checkpoint_save", "export_lora", - "dp_rank_forward", - "forward_micro_batches", + "forward", + "forward_batches", ], ) def test_skipped_forward_wave_rejects_collectives_before_backend( @@ -3341,7 +3376,7 @@ def test_skipped_forward_wave_rejects_collectives_before_backend( ) -> None: trainer = TrainerRank(_runtime()) _stub_forward(monkeypatch, trainer, dp=(0, 2)) - batches = trainer.forward_micro_batches([_target_request(1)]) + batches = trainer.forward_batches([_target_request(1)]) next(batches) assert trainer._skipped_forward_waves @@ -3351,7 +3386,7 @@ def unexpected(*_args: object, **_kwargs: object) -> Never: monkeypatch.setattr(dist, "all_reduce", unexpected) monkeypatch.setattr(trainer, "_checkpoint_group", unexpected) calls = { - "dp_reduce": lambda: trainer.dp_reduce(torch.tensor(1)), + "reduce": lambda: trainer.reduce(torch.tensor(1)), "optim_step": lambda: trainer.optim_step(params=AdamParams(learning_rate=1e-3)), "parameter": lambda: trainer.parameter("p", unexpected), "module": lambda: trainer.module("m", unexpected), @@ -3364,9 +3399,9 @@ def unexpected(*_args: object, **_kwargs: object) -> Never: "finish_checkpoint_save": lambda: trainer.finish_checkpoint_save("/unused"), "abort_checkpoint_save": lambda: trainer.abort_checkpoint_save("/unused"), "export_lora": lambda: trainer.export_lora("/unused"), - "dp_rank_forward": lambda: trainer.dp_rank_forward([_target_request(1)]), - "forward_micro_batches": lambda: next( - trainer.forward_micro_batches([_target_request(1)], yield_empty=True) + "forward": lambda: trainer.forward([_target_request(1)]), + "forward_batches": lambda: next( + trainer.forward_batches([_target_request(1)], yield_empty=True) ), } with pytest.raises(RuntimeError, match="yield_empty=False skips"): @@ -3381,7 +3416,7 @@ def test_skipped_forward_wave_cleans_up_retained_iterator( ) -> None: trainer = TrainerRank(_runtime()) _stub_forward(monkeypatch, trainer, dp=(0, 2)) - batches = trainer.forward_micro_batches([_target_request(1)], no_grad=True) + batches = trainer.forward_batches([_target_request(1)], no_grad=True) for _batch in batches: break assert torch.is_grad_enabled() @@ -3408,9 +3443,9 @@ def test_skipped_forward_wave_cannot_resume_another_iterator( ) -> None: trainer = TrainerRank(_runtime()) _stub_forward(monkeypatch, trainer, dp=(0, 2)) - outer = trainer.forward_micro_batches([_target_request(i) for i in range(4)]) + outer = trainer.forward_batches([_target_request(i) for i in range(4)]) next(outer) # Every rank participates in this wave, so nesting is allowed. - inner = trainer.forward_micro_batches([_target_request(1)]) + inner = trainer.forward_batches([_target_request(1)]) next(inner) with pytest.raises(RuntimeError, match="yield_empty=False skips"): next(outer) @@ -3464,9 +3499,9 @@ def execute(plan, **_kwargs): monkeypatch.setattr(trainer, "_forward_memory_group", lambda: None) requests = [_target_request(i) for i in range(count)] batches = ( - trainer.forward_micro_batches(requests) + trainer.forward_batches(requests) if mode is None - else trainer.forward_micro_batches(requests, yield_empty=mode) + else trainer.forward_batches(requests, yield_empty=mode) ) local_indices: list[int] = [] total_loss = torch.tensor(0.0) @@ -3474,11 +3509,11 @@ def execute(plan, **_kwargs): local_indices.extend(batch.indices) participation = torch.tensor(len(batch.outputs)) if mode is True or batch.stats.global_count >= world_size: - trainer.dp_reduce(participation) + trainer.reduce(participation) assert participation.item() == batch.stats.global_count else: with pytest.raises(RuntimeError, match="skips"): - trainer.dp_reduce(participation) + trainer.reduce(participation) loss = torch.tensor(0.0) for output in batch.outputs: loss = loss + output.target_logprobs.sum() @@ -3494,22 +3529,22 @@ def execute(plan, **_kwargs): if parameter.grad is None else parameter.grad ) - trainer.dp_reduce(gradient) - trainer.dp_reduce(total_loss) + trainer.reduce(gradient) + trainer.reduce(total_loss) assert gradient.item() == count**2 assert total_loss.item() == 2 * count**2 with pytest.raises(ValueError, match="yield_empty setting"): - list(trainer.forward_micro_batches(requests, yield_empty=rank == 0)) + list(trainer.forward_batches(requests, yield_empty=rank == 0)) completed = torch.tensor(1) - trainer.dp_reduce(completed) + trainer.reduce(completed) assert completed.item() == world_size finally: dist.destroy_process_group() @pytest.mark.parametrize("world_size", [2, 4]) -def test_forward_micro_batches_yield_modes_collectively( +def test_forward_batches_yield_modes_collectively( tmp_path: Path, world_size: int ) -> None: context = mp.spawn( @@ -3527,7 +3562,7 @@ def test_forward_micro_batches_yield_modes_collectively( pytest.fail("forward yield modes test hung") -def test_forward_micro_batches_syncs_fit_decision_across_dp( +def test_forward_batches_syncs_fit_decision_across_dp( monkeypatch: pytest.MonkeyPatch, ) -> None: trainer = TrainerRank(_runtime()) @@ -3543,13 +3578,13 @@ def memory_check(required: int, *, sync_across_dp: bool = False) -> _MemoryCheck ) monkeypatch.setattr(trainer, "_memory_check_required", memory_check) - next(iter(trainer.forward_micro_batches([_target_request(i) for i in range(6)]))) + next(iter(trainer.forward_batches([_target_request(i) for i in range(6)]))) assert sync_flags assert all(sync_flags) -def test_forward_micro_batches_supports_arbitrary_nested_depth( +def test_forward_batches_supports_arbitrary_nested_depth( monkeypatch: pytest.MonkeyPatch, ) -> None: trainer = TrainerRank(_runtime()) @@ -3560,9 +3595,9 @@ def test_forward_micro_batches_supports_arbitrary_nested_depth( ] nested = [(child for child in item) for item in expected] - batches = list(cast(Any, trainer).forward_micro_batches(nested)) + batches = list(cast(Any, trainer).forward_batches(nested)) - assert batches[0].inputs == expected + _assert_nested_tensors_equal(batches[0].inputs, expected) assert _output_shape(batches[0].outputs) == [ [[[[["output"]]]]], [[[[["output", "output"]]]]], @@ -3570,7 +3605,7 @@ def test_forward_micro_batches_supports_arbitrary_nested_depth( assert _output_values(batches[0].outputs) == [0, 1, 2] -def test_forward_micro_batches_ramps_after_first_success( +def test_forward_batches_ramps_after_first_success( monkeypatch: pytest.MonkeyPatch, ) -> None: trainer = TrainerRank(_runtime()) @@ -3586,9 +3621,7 @@ def run(plan, **_kwargs): _stub_forward(monkeypatch, trainer, run) - batches = list( - trainer.forward_micro_batches([_target_request(i) for i in range(8)]) - ) + batches = list(trainer.forward_batches([_target_request(i) for i in range(8)])) assert batches[0].stats.global_count == 1 assert batches[0].stats.cold_start @@ -3596,7 +3629,7 @@ def run(plan, **_kwargs): assert not batches[1].stats.cold_start -def test_forward_micro_batches_profiles_caller_peak_after_yield( +def test_forward_batches_profiles_caller_peak_after_yield( monkeypatch: pytest.MonkeyPatch, ) -> None: trainer = TrainerRank(_runtime()) @@ -3616,7 +3649,7 @@ def test_forward_micro_batches_profiles_caller_peak_after_yield( ), ) - batches = trainer.forward_micro_batches([_target_request(1)]) + batches = trainer.forward_batches([_target_request(1)]) next(batches) assert profiles == [] @@ -3627,7 +3660,7 @@ def test_forward_micro_batches_profiles_caller_peak_after_yield( @pytest.mark.parametrize("no_grad", [False, True]) @pytest.mark.parametrize("retain_previous", [False, True]) -def test_forward_micro_batches_releases_completed_wave_before_planning( +def test_forward_batches_releases_completed_wave_before_planning( monkeypatch: pytest.MonkeyPatch, no_grad: bool, retain_previous: bool ) -> None: trainer = TrainerRank(_runtime()) @@ -3657,7 +3690,7 @@ def profile(*_args): profiled.append(tensors[-1]() is not None) monkeypatch.setattr(trainer, "_update_peak_memory_profile", profile) - batches = trainer.forward_micro_batches( + batches = trainer.forward_batches( [_target_request(1), _target_request(3)], no_grad=no_grad ) first = next(batches) @@ -3691,7 +3724,7 @@ def test_memory_profiles_distinguish_grad_mode() -> None: assert grad_signature != no_grad_signature -def test_forward_micro_batches_does_not_overtrust_tiny_memory_profile( +def test_forward_batches_does_not_overtrust_tiny_memory_profile( monkeypatch: pytest.MonkeyPatch, ) -> None: trainer = TrainerRank(_runtime()) @@ -3709,7 +3742,7 @@ def test_forward_micro_batches_does_not_overtrust_tiny_memory_profile( assert candidate.plan.packed_tokens == 16 -def test_forward_micro_batches_tail_does_not_reset_stable_window( +def test_forward_batches_tail_does_not_reset_stable_window( monkeypatch: pytest.MonkeyPatch, ) -> None: trainer = TrainerRank(_runtime()) @@ -3729,15 +3762,13 @@ def test_forward_micro_batches_tail_does_not_reset_stable_window( fits=required <= 128, ), ) - batches = list( - trainer.forward_micro_batches([_target_request(i) for i in range(130)]) - ) + batches = list(trainer.forward_batches([_target_request(i) for i in range(130)])) assert [batch.stats.global_count for batch in batches] == [64, 64, 2] assert trainer._last_global_micro_batch_size == 64 -def test_forward_micro_batches_raises_when_smallest_batch_will_not_fit( +def test_forward_batches_raises_when_smallest_batch_will_not_fit( monkeypatch: pytest.MonkeyPatch, ) -> None: trainer = TrainerRank(_runtime()) @@ -3757,10 +3788,10 @@ def test_forward_micro_batches_raises_when_smallest_batch_will_not_fit( ), ) with pytest.raises(TrainerRankMemoryError, match="smallest DP microbatch"): - next(iter(trainer.forward_micro_batches([_target_request(1)]))) + next(iter(trainer.forward_batches([_target_request(1)]))) -def test_forward_micro_batches_rejects_mismatched_replicated_counts( +def test_forward_batches_rejects_mismatched_replicated_counts( monkeypatch: pytest.MonkeyPatch, ) -> None: trainer = TrainerRank(_runtime()) @@ -3777,7 +3808,7 @@ def gather(output, value): monkeypatch.setattr(trainer_rank.dist, "all_gather_object", gather) with pytest.raises(ValueError, match="same top-level input count"): - list(trainer.forward_micro_batches([_target_request(1)])) + list(trainer.forward_batches([_target_request(1)])) monkeypatch.setattr(trainer_rank.dist, "is_initialized", lambda: False) _stub_forward(monkeypatch, trainer, dp=(1, 2)) @@ -3785,7 +3816,7 @@ def gather(output, value): input_tokens=torch.tensor([1, 2]), target_tokens=torch.tensor([1, 2, 3]) ) with pytest.raises(ValueError, match="target_tokens"): - next(iter(trainer.forward_micro_batches([invalid, _target_request(1)]))) + next(iter(trainer.forward_batches([invalid, _target_request(1)]))) def test_forward_plan_estimates_output_memory_for_request_combo() -> None: @@ -3851,6 +3882,12 @@ def _assert_nested_tensors_equal(actual: object, expected: object) -> None: if isinstance(expected, torch.Tensor): assert isinstance(actual, torch.Tensor) torch.testing.assert_close(actual, expected, atol=0, rtol=0) + elif isinstance(expected, ForwardInput): + assert isinstance(actual, ForwardInput) + for field in fields(expected): + _assert_nested_tensors_equal( + getattr(actual, field.name), getattr(expected, field.name) + ) elif isinstance(expected, dict): assert isinstance(actual, dict) and actual.keys() == expected.keys() actual_dict = cast(dict[Any, object], actual) diff --git a/tests/unit/test_trainer_rank_versions.py b/tests/unit/test_trainer_rank_versions.py new file mode 100644 index 000000000..36076c420 --- /dev/null +++ b/tests/unit/test_trainer_rank_versions.py @@ -0,0 +1,539 @@ +from __future__ import annotations + +from contextlib import nullcontext +from datetime import timedelta +import gc +from pathlib import Path +import time +from types import SimpleNamespace +import weakref + +import pytest +import torch +import torch.distributed as dist +import torch.multiprocessing as mp +from torch.utils.checkpoint import checkpoint + +from art.trainer_rank import TrainerRank, TrainerRankSlotStateError +from art.trainer_rank._impl import _CheckpointSlot + + +def _trainer() -> tuple[TrainerRank, torch.nn.Parameter]: + trainer = TrainerRank.__new__(TrainerRank) + parameter = torch.nn.Parameter(torch.tensor(2.0, dtype=torch.float64)) + trainer._checkpoint_slots = {"student": _CheckpointSlot(params=(parameter,))} + trainer.runtime = SimpleNamespace(model=[], optimizer=None) + return trainer, parameter + + +@pytest.mark.parametrize("reentrant", (False, True)) +def test_snapshot_recompute_routes_original_gradient_after_update( + reentrant: bool, +) -> None: + trainer, current = _trainer() + version = trainer._capture_checkpoint_version("student") + old = trainer._snapshot_parameter(current, version) + x = torch.tensor(3.0, dtype=torch.float64, requires_grad=True) + loss = checkpoint(lambda x: x * old.square(), x, use_reentrant=reentrant) + with torch.no_grad(): + current.fill_(3) + trainer._checkpoint_slots["student"].revision += 1 + with trainer._gradient_transaction(): + loss.backward() + torch.testing.assert_close(current.grad, torch.tensor(12.0, dtype=torch.float64)) + torch.testing.assert_close(x.grad, torch.tensor(4.0, dtype=torch.float64)) + assert current.item() == 3 + + +def test_coupled_versions_accumulate_and_repeated_backward_routes_once() -> None: + trainer, current = _trainer() + old = trainer._snapshot_parameter( + current, trainer._capture_checkpoint_version("student") + ) + with torch.no_grad(): + current.fill_(3) + trainer._checkpoint_slots["student"].revision += 1 + new = trainer._snapshot_parameter( + current, trainer._capture_checkpoint_version("student") + ) + loss = (old.square() - new).square() + with trainer._gradient_transaction(): + loss.backward(retain_graph=True) + assert current.grad is not None + assert current.grad.item() == 6 + with trainer._gradient_transaction(): + loss.backward() + assert current.grad is not None + assert current.grad.item() == 12 + assert old.grad is None and new.grad is None + + +def test_stale_backward_preserves_existing_current_gradient() -> None: + trainer, current = _trainer() + old = trainer._snapshot_parameter( + current, trainer._capture_checkpoint_version("student"), 0 + ) + loss = old.square() + current.grad = torch.tensor(7.0, dtype=torch.float64) + trainer._checkpoint_slots["student"].revision += 1 + with pytest.raises(TrainerRankSlotStateError, match="staleness 1"): + with trainer._gradient_transaction(): + loss.backward() + assert current.grad.item() == 7 + + +def test_failed_backward_discards_staged_gradients_and_releases_batch() -> None: + trainer, current = _trainer() + old = trainer._snapshot_parameter( + current, trainer._capture_checkpoint_version("student") + ) + + def fail(_gradient: torch.Tensor) -> None: + raise RuntimeError("local autograd failed") + + old.register_hook(fail) + with pytest.raises(RuntimeError, match="local autograd failed"): + with trainer._gradient_transaction(): + old.square().backward() + assert current.grad is None + assert trainer._version_state()._transaction is None + + +@pytest.mark.parametrize("reentrant", (False, True)) +@pytest.mark.parametrize("transaction", (False, True)) +def test_snapshot_backward_requires_atomic_scope_for_outer_failure( + reentrant: bool, transaction: bool +) -> None: + trainer, p = _trainer() + q = torch.nn.Parameter(torch.tensor(3.0, dtype=torch.float64)) + trainer._checkpoint_slots["student"].params += (q,) + version = trainer._capture_checkpoint_version("student") + old_p = trainer._snapshot_parameter(p, version) + old_q = trainer._snapshot_parameter(q, version, 0) + loss = checkpoint(lambda x: x * old_p.square(), old_q, use_reentrant=reentrant) + p.grad, q.grad = torch.full_like(p, 7), torch.full_like(q, 11) + previous_p, previous_q = p.grad, q.grad + trainer._checkpoint_slots["student"].revision += 1 + message = "staleness 1" if transaction else "requires TrainerRank.backward" + with pytest.raises(RuntimeError, match=message): + with trainer._gradient_transaction() if transaction else nullcontext(): + loss.backward() + assert p.grad is previous_p and p.grad.item() == 7 + assert q.grad is previous_q and q.grad.item() == 11 + assert trainer._version_state()._origins == {} + assert trainer._version_state()._transaction is None + + +def test_reentrant_backward_transaction_rolls_back_completed_nested_task() -> None: + trainer, current = _trainer() + old = trainer._snapshot_parameter( + current, trainer._capture_checkpoint_version("student") + ) + x = torch.tensor(3.0, dtype=torch.float64, requires_grad=True) + loss = checkpoint(lambda x: x * old.square(), x, use_reentrant=True) + with pytest.raises(RuntimeError, match="later backward failed"): + with trainer._gradient_transaction(): + loss.backward() + assert current.grad is None + raise RuntimeError("later backward failed") + assert current.grad is None + + +def test_cotangent_batch_preflights_all_versions_before_any_mutation() -> None: + trainer, current = _trainer() + version = trainer._capture_checkpoint_version("student") + trainer._checkpoint_slots["student"].revision += 1 + now = trainer._capture_checkpoint_version("student") + with pytest.raises(TrainerRankSlotStateError, match="staleness"): + trainer._commit_versioned_gradients( + [ + (now, 0, current, torch.ones_like(current)), + (version, 0, current, torch.ones_like(current)), + ] + ) + assert current.grad is None + + +def test_replacement_invalidates_origin_and_old_target() -> None: + trainer, current = _trainer() + version = trainer._capture_checkpoint_version("student") + snapshot = trainer._snapshot_parameter(current, version) + new = torch.nn.Parameter(torch.tensor(5.0, dtype=torch.float64)) + trainer._checkpoint_slots["student"] = _CheckpointSlot(params=(new,), generation=1) + with pytest.raises(TrainerRankSlotStateError, match="replaced"): + with trainer._gradient_transaction(): + snapshot.square().backward() + assert current.grad is None and new.grad is None + + +def test_replaced_target_identity_rejects_entire_gradient_batch() -> None: + trainer, current = _trainer() + version = trainer._capture_checkpoint_version("student") + replacement = torch.nn.Parameter(current.detach().clone()) + trainer._checkpoint_slots["student"].params = (replacement,) + with pytest.raises(TrainerRankSlotStateError, match="target was replaced"): + trainer._commit_versioned_gradients( + [ + (version, 2, replacement, torch.ones_like(replacement)), + (version, 2, current, torch.ones_like(current)), + ] + ) + assert current.grad is None and replacement.grad is None + + +def test_accumulated_origin_is_checked_before_optimizer_mutation() -> None: + trainer, current = _trainer() + old = trainer._snapshot_parameter( + current, trainer._capture_checkpoint_version("student"), 0 + ) + with trainer._gradient_transaction(): + old.square().backward() + trainer._checkpoint_slots["student"].revision += 1 + with pytest.raises(TrainerRankSlotStateError, match="staleness"): + trainer._dynamic_optim_step(["student"], params={}, scale_grads={}) + assert current.grad is not None + assert current.item() == 2 and current.grad.item() == 4 + trainer.zero_grad() + trainer._version_state().validate_accumulated(["student"]) + + +def test_snapshot_lifetime_follows_graph_references() -> None: + trainer, current = _trainer() + snapshot = trainer._snapshot_parameter( + current, trainer._capture_checkpoint_version("student") + ) + reference = weakref.ref(snapshot) + loss = snapshot.square() + del snapshot + gc.collect() + assert reference() is not None + with trainer._gradient_transaction(): + loss.backward() + del loss + gc.collect() + assert reference() is None + + +def test_replay_capture_does_not_reset_origin_age() -> None: + trainer, current = _trainer() + origin = trainer._capture_checkpoint_version("student") + trainer._checkpoint_slots["student"].revision = 2 + replay = trainer._snapshot_parameter(current, origin, 2) + trainer._checkpoint_slots["student"].revision = 3 + with pytest.raises(TrainerRankSlotStateError, match="staleness 3"): + with trainer._gradient_transaction(): + replay.square().backward() + assert current.grad is None + + +def test_collective_preflight_failure_does_not_commit_local_gradients() -> None: + trainer, current = _trainer() + old = trainer._snapshot_parameter( + current, trainer._capture_checkpoint_version("student") + ) + + def other_rank_failed(validate) -> None: + validate() + assert current.grad is None + raise RuntimeError("another rank rejected stale gradients") + + with pytest.raises(RuntimeError, match="another rank"): + with trainer._gradient_transaction(before_commit=other_rank_failed): + old.square().backward() + assert current.grad is None + assert not trainer._version_state()._origins + + +def test_gradient_dtype_is_validated_before_any_accumulator_is_published() -> None: + trainer, current = _trainer() + other = torch.nn.Parameter(torch.tensor(1.0, dtype=torch.float64)) + trainer._checkpoint_slots["student"].params += (other,) + version = trainer._capture_checkpoint_version("student") + with pytest.raises(ValueError, match="dtype"): + trainer._commit_versioned_gradients( + [ + (version, 2, current, torch.ones_like(current)), + (version, 2, other, torch.tensor(2.0, dtype=torch.float32)), + ] + ) + assert current.grad is None and other.grad is None + + +def test_nested_transaction_still_participates_in_collective_preflight() -> None: + trainer, current = _trainer() + old = trainer._snapshot_parameter( + current, trainer._capture_checkpoint_version("student") + ) + calls = [] + + def collective(validate) -> None: + validate() + calls.append(True) + assert current.grad is None + + with trainer._gradient_transaction(): + with trainer._gradient_transaction(before_commit=collective): + old.square().backward() + assert calls == [True] + assert current.grad is None + assert current.grad is not None + assert current.grad.item() == 4 + + +def test_caught_nested_backward_failure_invalidates_whole_transaction() -> None: + trainer, current = _trainer() + old = trainer._snapshot_parameter( + current, trainer._capture_checkpoint_version("student") + ) + with pytest.raises(RuntimeError, match="nested gradient transaction failed"): + with trainer._gradient_transaction(): + with pytest.raises(RuntimeError, match="failed"): + with trainer._gradient_transaction(): + old.square().backward() + raise RuntimeError("failed") + assert current.grad is None + + +def test_explicit_head_cotangents_join_model_gradient_transaction() -> None: + trainer, current = _trainer() + version = trainer._capture_checkpoint_version("student") + old = trainer._snapshot_parameter(current, version) + with pytest.raises(RuntimeError, match="model backward failed"): + with trainer._gradient_transaction(): + trainer._commit_versioned_gradients( + [(version, 2, current, torch.ones_like(current))] + ) + old.square().backward() + assert current.grad is None + raise RuntimeError("model backward failed") + assert current.grad is None + with trainer._gradient_transaction(): + trainer._commit_versioned_gradients( + [(version, 2, current, torch.ones_like(current))] + ) + old.square().backward() + assert current.grad.item() == 5 + + +def _divergent_version_worker(rank: int, rendezvous: str) -> None: + dist.init_process_group( + "gloo", + init_method=rendezvous, + rank=rank, + world_size=2, + timeout=timedelta(seconds=30), + ) + try: + trainer, current = _trainer() + version = trainer._capture_checkpoint_version("student") + if rank == 0: + trainer._commit_versioned_gradients( + [(version, 0, current, torch.full_like(current, 3))] + ) + else: + current.grad = torch.full_like(current, 3) + trainer._checkpoint_slots["student"].revision = 1 + with pytest.raises(RuntimeError, match="staleness|Another rank failed"): + trainer._dynamic_optim_step(["student"], params={}, scale_grads={}) + assert current.grad is not None + assert current.item() == 2 and current.grad.item() == 3 + completed = torch.tensor(1) + dist.all_reduce(completed) + assert completed.item() == 2 + finally: + dist.destroy_process_group() + + +def test_divergent_optimizer_provenance_rejects_collectively_before_mutation( + tmp_path: Path, +) -> None: + processes = mp.spawn( + _divergent_version_worker, + args=(f"file://{tmp_path / 'versions'}",), + nprocs=2, + join=False, + ) + deadline = time.monotonic() + 90 + while time.monotonic() < deadline: + if processes.join(timeout=1): + return + for process in processes.processes: + process.terminate() + pytest.fail("Divergent version preflight did not complete collectively") + + +def test_transaction_coalesces_many_children_but_keeps_each_origin() -> None: + trainer, current = _trainer() + old = trainer._snapshot_parameter( + current, trainer._capture_checkpoint_version("student") + ) + trainer._checkpoint_slots["student"].revision = 1 + new = trainer._snapshot_parameter( + current, trainer._capture_checkpoint_version("student") + ) + current.grad = torch.full_like(current, 7) + pointer = None + with trainer._gradient_transaction(): + for _ in range(50): + old.square().backward() + new.sum().backward() + batch = trainer._version_state()._transaction + assert batch is not None + assert len(batch.gradients) == 1 and len(batch.origins) == 2 + staged = batch.gradients[id(current)][1] + if pointer is None: + pointer = staged.data_ptr() + assert staged.data_ptr() == pointer + assert current.grad.item() == 7 + assert current.grad.item() == 257 + assert len(trainer._version_state()._origins["student"]) == 2 + assert batch.gradients == {} and batch.origins == set() + + +def test_retained_failed_transaction_traceback_releases_staging() -> None: + trainer, current = _trainer() + old = trainer._snapshot_parameter( + current, trainer._capture_checkpoint_version("student") + ) + saved_error = None + reference = None + try: + with trainer._gradient_transaction(): + old.square().backward() + batch = trainer._version_state()._transaction + assert batch is not None + reference = weakref.ref(batch.gradients[id(current)][1]) + raise RuntimeError("failed after staged child") + except RuntimeError as exc: + saved_error = exc + assert saved_error is not None and saved_error.__traceback__ is not None + gc.collect() + assert reference is not None and reference() is None + assert current.grad is None and trainer._version_state()._origins == {} + + +def test_csr_cotangent_rejected_before_any_gradient_publication() -> None: + trainer, _ = _trainer() + first, second = [torch.nn.Parameter(torch.ones(2, 2)) for _ in range(2)] + first.grad = torch.full_like(first, 7) + previous = first.grad + trainer._checkpoint_slots["student"].params = (first, second) + version = trainer._capture_checkpoint_version("student") + with pytest.raises(ValueError, match="strided"): + trainer._commit_versioned_gradients( + [ + (version, 2, first, torch.ones_like(first)), + (version, 2, second, torch.ones_like(second).to_sparse_csr()), + ] + ) + assert first.grad is previous and torch.all(first.grad == 7) + assert second.grad is None and trainer._version_state()._origins == {} + + +def test_gradient_assignment_failure_rolls_back_earlier_publication() -> None: + class RejectGradient(torch.nn.Parameter): + def __setattr__(self, name: str, value: object) -> None: + if name == "grad": + raise RuntimeError("injected gradient publication failure") + super().__setattr__(name, value) + + trainer, first = _trainer() + second = RejectGradient(torch.ones_like(first)) + trainer._checkpoint_slots["student"].params += (second,) + first.grad = torch.full_like(first, 7) + previous = first.grad + version = trainer._capture_checkpoint_version("student") + with pytest.raises(RuntimeError, match="publication failure"): + trainer._commit_versioned_gradients( + [ + (version, 2, first, torch.ones_like(first)), + (version, 2, second, torch.ones_like(second)), + ] + ) + assert first.grad is previous and first.grad.item() == 7 + assert second.grad is None and trainer._version_state()._origins == {} + + +def _transaction_exit_failure_worker(rank: int, rendezvous: str, nested: bool) -> None: + dist.init_process_group( + "gloo", + init_method=rendezvous, + rank=rank, + world_size=2, + timeout=timedelta(seconds=15), + ) + try: + trainer, current = _trainer() + version = trainer._capture_checkpoint_version("student") + trainer._commit_versioned_gradients( + [(version, 2, current, torch.full_like(current, 7))] + ) + previous = current.grad + origins = trainer._version_state()._origins.copy() + snapshot = trainer._snapshot_parameter(current, version) + calls = [] + + def coordinate(validate) -> None: + calls.append(True) + error = None + try: + validate() + except BaseException as exc: + error = exc + errors = [None, None] + dist.all_gather_object(errors, None if error is None else str(error)) + if error is not None: + raise error + if any(errors): + raise RuntimeError(f"another rank failed: {errors}") + + saved_error = None + reference = None + try: + with trainer._gradient_transaction() if nested else nullcontext(): + with trainer._gradient_transaction(before_commit=coordinate): + snapshot.square().backward() + batch = trainer._version_state()._transaction + assert batch is not None + reference = weakref.ref(batch.gradients[id(current)][1]) + if rank == 0: + raise ValueError("injected second-child replay failure") + snapshot.sum().backward() + except (ValueError, RuntimeError) as exc: + saved_error = exc + assert saved_error is not None and "second-child replay failure" in str( + saved_error + ) + if rank == 0: + assert isinstance(saved_error, ValueError) + assert calls == [True] + assert current.grad is previous and current.grad is not None + assert current.grad.item() == 7 + assert trainer._version_state()._origins == origins + gc.collect() + assert reference is not None and reference() is None + completed = torch.tensor(1) + dist.all_reduce(completed) + assert completed.item() == 2 + finally: + dist.destroy_process_group() + + +@pytest.mark.parametrize("nested", [False, True]) +def test_local_replay_failure_uses_same_collective_exit_phase_on_every_rank( + tmp_path: Path, + nested: bool, +) -> None: + processes = mp.spawn( + _transaction_exit_failure_worker, + args=(f"file://{tmp_path / 'exit'}", nested), + nprocs=2, + join=False, + ) + deadline = time.monotonic() + 60 + while time.monotonic() < deadline: + if processes.join(timeout=1): + return + for process in processes.processes: + process.terminate() + pytest.fail("Transaction success/failure exit phases did not match") diff --git a/tests/unit/test_trainer_rank_weird_shapes.py b/tests/unit/test_trainer_rank_weird_shapes.py index 8db3ef72e..8c6ace7aa 100644 --- a/tests/unit/test_trainer_rank_weird_shapes.py +++ b/tests/unit/test_trainer_rank_weird_shapes.py @@ -237,7 +237,7 @@ def test_planner_handles_vineppo_nested_shape_and_request_mix() -> None: ) -def test_forward_micro_batches_preserves_nested_vineppo_groups( +def test_forward_batches_preserves_nested_vineppo_groups( monkeypatch: pytest.MonkeyPatch, ) -> None: rank = TrainerRank(_runtime()) @@ -260,7 +260,7 @@ def test_forward_micro_batches_preserves_nested_vineppo_groups( ) groups = _vineppo_like_inputs() - micro_batches = list(rank.forward_micro_batches(groups)) + micro_batches = list(rank.forward_batches(groups)) assert [batch.indices for batch in micro_batches] == [(0, 1, 2, 3)] assert micro_batches[0].select(groups) == groups @@ -271,7 +271,7 @@ def test_forward_micro_batches_preserves_nested_vineppo_groups( ) -def test_forward_micro_batches_prewarms_next_wave_during_yield( +def test_forward_batches_prewarms_next_wave_during_yield( monkeypatch: pytest.MonkeyPatch, ) -> None: rank = TrainerRank(_runtime()) @@ -290,7 +290,7 @@ def test_forward_micro_batches_prewarms_next_wave_during_yield( ), ) - generator = rank.forward_micro_batches(inputs) + generator = rank.forward_batches(inputs) first = next(generator) assert first.stats.global_count == 4 @@ -363,7 +363,7 @@ def test_speculative_planning_uses_immutable_snapshots( original_rows = tuple(row.clone() for row in _rows(inputs[4:8])) original_key = rank._layout_cache_key(original_rows) - generator = rank.forward_micro_batches(inputs) + generator = rank.forward_batches(inputs) next(generator) # The caller mutates its (aliased) input tensors while suspended. for request in inputs[4:8]: @@ -388,7 +388,7 @@ def test_speculative_planning_warms_this_dp_ranks_local_slice( # Local budget of 4 items per rank -> global waves of 8 at DP2. rank = _prewarmed_rank(monkeypatch, inputs, inputs[0:8:2], dp=(0, 2)) - generator = rank.forward_micro_batches(inputs) + generator = rank.forward_batches(inputs) first = next(generator) assert first.stats.global_count == 8 future = rank._speculative_planning_future @@ -432,7 +432,7 @@ def test_width_search_lets_prefix_sharing_widen_the_wave( # (4,002); the wave must still take both requests. _set_packed_token_budget(monkeypatch, rank, 2_400) - batches = list(rank.forward_micro_batches(inputs)) + batches = list(rank.forward_batches(inputs)) assert [batch.stats.global_count for batch in batches] == [2] @@ -500,13 +500,13 @@ def plans(shared_len: int) -> tuple[TrainerRank, list, int, int]: # the width-3 one does. _set_packed_token_budget(monkeypatch, rank, (two + three) // 2) - batches = list(rank.forward_micro_batches(inputs)) + batches = list(rank.forward_batches(inputs)) assert [batch.stats.global_count for batch in batches] == [3] assert batches[0].stats.packed_tokens <= (two + three) // 2 -def test_dp_rank_forward_falls_back_to_memory_minimal_layout_before_refusing( +def test_forward_falls_back_to_memory_minimal_layout_before_refusing( monkeypatch: pytest.MonkeyPatch, ) -> None: shared = tuple(range(10_000, 10_040)) @@ -528,7 +528,7 @@ def test_dp_rank_forward_falls_back_to_memory_minimal_layout_before_refusing( # Cost-optimal layout declines sharing (82 tokens); full sharing (42) fits. _set_packed_token_budget(monkeypatch, rank, 60) - outputs = rank.dp_rank_forward(inputs) + outputs = rank.forward(inputs) assert len(outputs) == 2 assert executed == [42] @@ -666,24 +666,24 @@ def test_profiled_steady_state_keeps_the_wide_shared_wave( ) _set_packed_token_budget(monkeypatch, rank, 1_100) - batches = list(rank.forward_micro_batches(inputs)) + batches = list(rank.forward_batches(inputs)) assert [batch.stats.global_count for batch in batches] == [16] assert not batches[0].stats.cold_start -def test_forward_micro_batches_telemetry_reports_hidden_speculation( +def test_forward_batches_telemetry_reports_hidden_speculation( monkeypatch: pytest.MonkeyPatch, ) -> None: inputs = _unshared_requests(8) rank = _prewarmed_rank(monkeypatch, inputs, inputs[:4]) - list(rank.forward_micro_batches(inputs)) + list(rank.forward_batches(inputs)) telemetry = rank.last_forward_telemetry() assert telemetry["planning_ms"] > 0.0 assert "speculative_planning_ms" in telemetry -@pytest.mark.parametrize("api", ("dp_rank_forward", "forward_micro_batches")) +@pytest.mark.parametrize("api", ("forward", "forward_batches")) def test_forward_preserves_caller_owned_nested_input_tensors( api: str, monkeypatch: pytest.MonkeyPatch, @@ -710,10 +710,10 @@ def test_forward_preserves_caller_owned_nested_input_tensors( for _request, inputs, targets in tensors ] - if api == "dp_rank_forward": - rank.dp_rank_forward(groups) + if api == "forward": + rank.forward(groups) else: - list(rank.forward_micro_batches(groups)) + list(rank.forward_batches(groups)) for (request, inputs, targets), (expected_inputs, expected_targets) in zip( tensors, snapshots, strict=True @@ -855,7 +855,7 @@ def test_adaptive_planner_grows_stable_window_to_largest_aligned_fit( assert candidate.rejected_candidates <= 2 -def test_forward_micro_batches_shrinks_when_memory_budget_drops( +def test_forward_batches_shrinks_when_memory_budget_drops( monkeypatch: pytest.MonkeyPatch, ) -> None: rank = TrainerRank(_runtime()) @@ -889,7 +889,7 @@ def run(plan, **_kwargs): _set_packed_token_budget(monkeypatch, rank, lambda: available["packed_tokens"]) monkeypatch.setattr(rank, "_run_flat_plan_with_memory_tracking", run) - batches = list(rank.forward_micro_batches(inputs)) + batches = list(rank.forward_batches(inputs)) assert [batch.stats.global_count for batch in batches] == [8, 3, 3] assert [batch.stats.available_bytes for batch in batches] == [ @@ -949,13 +949,13 @@ def slot_ref(name: str | None) -> SlotRef | None: } -@pytest.mark.parametrize("api", ("dp_rank_forward", "forward_micro_batches")) +@pytest.mark.parametrize("api", ("forward", "forward_batches")) def test_forward_raises_before_expected_oom_with_actionable_context( api: str, monkeypatch: pytest.MonkeyPatch, ) -> None: rank = TrainerRank(_runtime()) - if api == "dp_rank_forward": + if api == "forward": monkeypatch.setattr( rank, "_memory_check", @@ -977,9 +977,9 @@ def test_forward_raises_before_expected_oom_with_actionable_context( with pytest.raises(TrainerRankMemoryError) as exc_info: ( - rank.dp_rank_forward(request) - if api == "dp_rank_forward" - else next(iter(rank.forward_micro_batches(request))) + rank.forward(request) + if api == "forward" + else next(iter(rank.forward_batches(request))) ) message = str(exc_info.value) diff --git a/uv.lock b/uv.lock index 9e654d262..6e7f3a4ea 100644 --- a/uv.lock +++ b/uv.lock @@ -4330,6 +4330,7 @@ source = { editable = "." } dependencies = [ { name = "aiohttp" }, { name = "anthropic" }, + { name = "cloudpickle" }, { name = "litellm" }, { name = "nest-asyncio" }, { name = "numpy" }, @@ -4504,6 +4505,7 @@ requires-dist = [ { name = "awscli", marker = "extra == 'backend-cu130'", specifier = ">=1.38.1" }, { name = "bitsandbytes", marker = "extra == 'backend'", specifier = ">=0.45.2,!=0.50.0" }, { name = "bitsandbytes", marker = "extra == 'backend-cu130'", specifier = ">=0.45.2,!=0.50.0" }, + { name = "cloudpickle", specifier = ">=3.1.1" }, { name = "datrie", marker = "extra == 'tinker'", specifier = ">=0.8.3" }, { name = "duckdb", marker = "extra == 'backend'", specifier = ">=1.0.0" }, { name = "duckdb", marker = "extra == 'backend-cu130'", specifier = ">=1.0.0" }, From 407f8871e2813fa7311976cd3cf343dbff0d9410 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 18 Sep 2026 00:56:19 +0000 Subject: [PATCH 002/150] Use native topology type in graph cache GPU fixtures --- .../integration/megatron/lora/test_trainer_v1_graph_cache.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/tests/integration/megatron/lora/test_trainer_v1_graph_cache.py b/tests/integration/megatron/lora/test_trainer_v1_graph_cache.py index 92d2975cb..fe419ea69 100644 --- a/tests/integration/megatron/lora/test_trainer_v1_graph_cache.py +++ b/tests/integration/megatron/lora/test_trainer_v1_graph_cache.py @@ -8,6 +8,7 @@ torch = pytest.importorskip("torch") pytest.importorskip("megatron.core") +from art.megatron.context_parallel.types import ParallelTopology # noqa: E402 from art.megatron.lora import LoRA, LoRASlotRef, use_lora_slot # noqa: E402 from art.megatron.prefix_tree_packing import prefix_tree_pack # noqa: E402 from art.trainer_rank import ( # noqa: E402 @@ -52,7 +53,7 @@ def test_group_cache_routes_old_gradients_after_optimizer_update( references = tuple( value.detach().clone().requires_grad_() for value in originals ) - monkeypatch.setattr(trainer, "_topology", lambda: (1, 1, 1, 1)) + monkeypatch.setattr(trainer, "_topology", lambda: ParallelTopology()) monkeypatch.setattr( trainer, "_prepare_packed_forward", @@ -152,7 +153,7 @@ def test_native_stale_logprob_correction_keeps_original_gradient_age(mode, monke historical = [value.detach().clone().requires_grad_() for value in parameters] tokens = torch.arange(12) x = tokens.to(device).float().reshape(-1, 4) / 13 - monkeypatch.setattr(trainer, "_topology", lambda: (1, 1, 1, 1)) + monkeypatch.setattr(trainer, "_topology", lambda: ParallelTopology()) monkeypatch.setattr( trainer, "_prepare_packed_forward", From 3770f95b3f173f4bc588279d24eba7893b4ce71e Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 18 Sep 2026 01:12:37 +0000 Subject: [PATCH 003/150] Model independent forward groups separately in graph lifetime test --- tests/unit/test_trainer_rank_validation.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/tests/unit/test_trainer_rank_validation.py b/tests/unit/test_trainer_rank_validation.py index 9e9127be6..f5fb2a140 100644 --- a/tests/unit/test_trainer_rank_validation.py +++ b/tests/unit/test_trainer_rank_validation.py @@ -3137,7 +3137,9 @@ 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("teacher") - first, second = _tracked_targets(trainer, ref, 2, 3) + # Each physical forward group has its own tracking call and cache lifetime. + first = _tracked_targets(trainer, ref, 2)[0] + second = _tracked_targets(trainer, ref, 3)[0] first.sum().backward() with pytest.raises(TrainerRankSlotStateError, match="live backward graph"): From 00a3f6f45bf7cf21192359aa96c5f51b02e1d920 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 18 Sep 2026 01:14:26 +0000 Subject: [PATCH 004/150] test: declare topology for split graph lifetime fixture --- tests/unit/test_trainer_rank_split.py | 11 ++++++++++- 1 file changed, 10 insertions(+), 1 deletion(-) diff --git a/tests/unit/test_trainer_rank_split.py b/tests/unit/test_trainer_rank_split.py index 7e2aa48bd..4a31c35e8 100644 --- a/tests/unit/test_trainer_rank_split.py +++ b/tests/unit/test_trainer_rank_split.py @@ -48,6 +48,7 @@ TrainerRank, TrainerRankMemoryError, TrainerRankPartialExecutionError, + TrainerRankSlotStateError, ) from art.trainer_rank._impl import ( Unset, @@ -957,7 +958,7 @@ def test_split_subforwards_track_independent_slot_graphs( monkeypatch.setattr(rank, "_slot_ref", lambda name: _SlotRef(name)) monkeypatch.setattr(rank, "_resolve_slot_ref", lambda request, **_kwargs: ref) monkeypatch.setattr(rank, "_validate_hybridep_topology", lambda: None) - topology = object() + topology = SimpleNamespace(tp=1, cp=1, dp=1, pp=1, sp=False) monkeypatch.setattr(rank, "_topology", lambda: topology) monkeypatch.setattr(rank, "_capture_lora_version", lambda *_args, **_kwargs: None) monkeypatch.setattr(rank, "_configure_hybridep", lambda *_args, **_kwargs: None) @@ -991,10 +992,18 @@ def backward(indices: tuple[int, ...]) -> None: [(packet.handle, packet.gradients) for packet in packets] ) + def assert_pending_graph() -> None: + with pytest.raises(TrainerRankSlotStateError, match="live backward graph"): + rank._guard_slot_can_load(ref) + with pytest.raises(TrainerRankSlotStateError, match="not been backpropagated"): + rank._guard_checkpoint_can_step("teacher") + cache = rank._forward_graph_cache() assert len(cache.handles()) == 2 + assert_pending_graph() backward(first) assert len(cache.handles()) == 1 + assert_pending_graph() backward(second) assert cache.handles() == () rank._guard_slot_can_load(ref) From 6d62ce2cd3e6f3daf5a92c9d13bb6f525b2f36b8 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 19 Sep 2026 20:11:09 +0000 Subject: [PATCH 005/150] refactor: remove unused memory placement selector --- src/art/trainer_rank/_memory_policy.py | 48 ---------- .../cp_attn/test_cpu_offload_residency.py | 32 ++----- tests/unit/test_trainer_rank_memory_policy.py | 89 +------------------ .../test_trainer_rank_memory_policy_cuda.py | 27 ++---- 4 files changed, 14 insertions(+), 182 deletions(-) diff --git a/src/art/trainer_rank/_memory_policy.py b/src/art/trainer_rank/_memory_policy.py index 8af88108c..e7809f22b 100644 --- a/src/art/trainer_rank/_memory_policy.py +++ b/src/art/trainer_rank/_memory_policy.py @@ -302,51 +302,3 @@ def placement_cost( staging, max((cost.peak_bytes for cost in costs), default=0), ) - - -def choose_memory_placement( - costs: Sequence[ForwardMemoryCost], - *, - gpu_available_bytes: int, - cpu_available_bytes: int, - backward_state: Literal["auto", "gpu", "cpu", "replay"] = "auto", - output_device: Literal["auto", "model", "cpu"] = "model", - allow_cpu_offload: bool = True, - allow_replay: bool = True, -) -> MemoryPlacement | None: - """Prefer retained GPU, then CPU saved state, then replay; O(children).""" - if backward_state not in ("auto", "gpu", "cpu", "replay"): - raise ValueError(f"Unknown backward_state: {backward_state}") - if output_device not in ("auto", "model", "cpu"): - raise ValueError(f"Unknown output_device: {output_device}") - if backward_state == "cpu" and not allow_cpu_offload: - raise ValueError("CPU backward state conflicts with allow_cpu_offload=False") - if backward_state == "replay" and not allow_replay: - raise ValueError("Replay backward state conflicts with allow_replay=False") - states: tuple[Retention, ...] = ( - tuple( - state - for state, enabled in ( - ("gpu", True), - ("cpu", allow_cpu_offload), - ("replay", allow_replay), - ) - if enabled - ) - if backward_state == "auto" - else (backward_state,) - ) - devices: tuple[OutputDevice, ...] = ( - ("model", "cpu") if output_device == "auto" else (output_device,) - ) - for state in states: - for device in devices: - placement = placement_cost( - costs, backward_state=state, output_device=device - ) - if ( - placement.gpu_required_bytes <= gpu_available_bytes - and placement.cpu_required_bytes <= cpu_available_bytes - ): - return placement - return None diff --git a/tests/integration/megatron/cp_attn/test_cpu_offload_residency.py b/tests/integration/megatron/cp_attn/test_cpu_offload_residency.py index ed7b1910f..dcda63aad 100644 --- a/tests/integration/megatron/cp_attn/test_cpu_offload_residency.py +++ b/tests/integration/megatron/cp_attn/test_cpu_offload_residency.py @@ -29,7 +29,6 @@ from art.trainer_rank._graphs import GraphCache # noqa: E402 from art.trainer_rank._memory_policy import ( # noqa: E402 ForwardMemoryCost, - choose_memory_placement, placement_cost, ) @@ -204,33 +203,16 @@ def run(retention, count=1): ) cap = old.gpu_required_bytes + resident assert old.gpu_required_bytes <= cap < corrected.gpu_required_bytes - assert ( - choose_memory_placement( - [cost] * count, - gpu_available_bytes=cap, - cpu_available_bytes=1 << 40, - backward_state="cpu", - output_device="cpu", - ) - is None - ) - replay_plan = choose_memory_placement( - [cost] * count, - gpu_available_bytes=cap, - cpu_available_bytes=1 << 40, - backward_state="replay", - output_device="cpu", + replay_plan = placement_cost( + [cost] * count, backward_state="replay", output_device="cpu" ) - assert replay_plan is not None + assert replay_plan.gpu_required_bytes <= cap + assert replay_plan.cpu_required_bytes <= 1 << 40 replay = run("replay", count) assert replay["peak"] <= cap - partial_plan = choose_memory_placement( - [cost] * count, - gpu_available_bytes=corrected.gpu_required_bytes, - cpu_available_bytes=1 << 40, - output_device="cpu", - ) - assert partial_plan is not None and partial_plan.backward_state == "cpu" + retained = placement_cost([cost] * count, backward_state="gpu", output_device="cpu") + assert corrected.gpu_required_bytes < retained.gpu_required_bytes + assert corrected.cpu_required_bytes <= 1 << 40 partial = run("cpu", count) assert cap < partial["peak"] <= corrected.gpu_required_bytes print( diff --git a/tests/unit/test_trainer_rank_memory_policy.py b/tests/unit/test_trainer_rank_memory_policy.py index 9eb0ce1a8..6ae2d7fce 100644 --- a/tests/unit/test_trainer_rank_memory_policy.py +++ b/tests/unit/test_trainer_rank_memory_policy.py @@ -7,7 +7,6 @@ ForwardMemoryCost, HostMemoryBudget, MemoryScope, - choose_memory_placement, choose_output_placements, host_memory_budget, local_rank_count, @@ -198,95 +197,9 @@ def test_root_cost_keeps_outputs_and_backward_restore_workspace( assert placement.gpu_retained_bytes == retained -@pytest.mark.parametrize( - "gpu,cpu,state,device", - [ - (290, 15, "gpu", "model"), - (259, 225, "cpu", "model"), - (130, 224, "replay", "model"), - (129, 45, "replay", "cpu"), - ], -) -def test_admission_prefers_retention_and_admits_previously_refused_roots( - gpu, cpu, state, device -): - placement = choose_memory_placement( - ROOT, gpu_available_bytes=gpu, cpu_available_bytes=cpu, output_device="auto" - ) - assert placement is not None - assert (placement.backward_state, placement.output_device) == (state, device) - assert placement.gpu_required_bytes <= gpu - assert placement.cpu_required_bytes <= cpu - - -@pytest.mark.parametrize("gpu,cpu", [(99, 1000), (259, 14), (129, 44)]) -def test_policy_never_admits_without_cpu_capacity_or_one_child_workspace(gpu, cpu): - assert ( - choose_memory_placement( - ROOT, gpu_available_bytes=gpu, cpu_available_bytes=cpu, output_device="auto" - ) - is None - ) - - -def test_disable_levels_and_explicit_output_opt_out_are_authoritative(): - assert ( - choose_memory_placement( - ROOT, - gpu_available_bytes=129, - cpu_available_bytes=1000, - allow_cpu_offload=False, - output_device="model", - ) - is None - ) - assert ( - choose_memory_placement( - ROOT, - gpu_available_bytes=130, - cpu_available_bytes=224, - allow_replay=False, - ) - is None - ) - forced = choose_memory_placement( - ROOT, - gpu_available_bytes=1000, - cpu_available_bytes=1000, - backward_state="replay", - output_device="cpu", - ) - assert forced is not None - assert (forced.backward_state, forced.output_device) == ("replay", "cpu") - - -@pytest.mark.parametrize( - "kwargs", - [ - {"backward_state": "cpu", "allow_cpu_offload": False}, - {"backward_state": "replay", "allow_replay": False}, - {"backward_state": "typo"}, - {"output_device": "typo"}, - ], -) -def test_invalid_policy_raises_before_admission(kwargs): - with pytest.raises(ValueError): - choose_memory_placement( - ROOT, gpu_available_bytes=1000, cpu_available_bytes=1000, **kwargs - ) - - def test_no_grad_cpu_outputs_release_gpu_storage_without_replay(): costs = (ForwardMemoryCost(100, 80, 80, backward_required=False),) * 3 - placement = choose_memory_placement( - costs, - gpu_available_bytes=100, - cpu_available_bytes=240, - allow_cpu_offload=False, - allow_replay=False, - output_device="auto", - ) - assert placement is not None + placement = placement_cost(costs, backward_state="gpu", output_device="cpu") assert (placement.backward_state, placement.output_device) == ("gpu", "cpu") assert placement.gpu_required_bytes == 100 assert placement.cpu_required_bytes == 240 diff --git a/tests/unit/test_trainer_rank_memory_policy_cuda.py b/tests/unit/test_trainer_rank_memory_policy_cuda.py index 6548dd440..de9acbefb 100644 --- a/tests/unit/test_trainer_rank_memory_policy_cuda.py +++ b/tests/unit/test_trainer_rank_memory_policy_cuda.py @@ -15,7 +15,6 @@ from art.trainer_rank._impl import _CheckpointSlot from art.trainer_rank._memory_policy import ( ForwardMemoryCost, - choose_memory_placement, host_memory_budget, placement_cost, ) @@ -187,27 +186,13 @@ def test_real_root_memory_and_gradient_policy(workload, state, device): _run_root(workload, state, device) -def test_real_root_admitted_by_replay_with_cpu_outputs(workload): +def test_real_root_replay_with_cpu_outputs_fits_budget(workload): *_, cost = workload budget = cost.peak_bytes + 8 * 1024**2 cpu_budget = host_memory_budget(local_world_size=2).available_bytes - assert ( - choose_memory_placement( - (cost,) * 4, - gpu_available_bytes=budget, - cpu_available_bytes=cpu_budget, - backward_state="gpu", - output_device="model", - ) - is None - ) - chosen = choose_memory_placement( - (cost,) * 4, - gpu_available_bytes=budget, - cpu_available_bytes=cpu_budget, - output_device="auto", - ) - assert chosen is not None - assert (chosen.backward_state, chosen.output_device) == ("replay", "cpu") - measured = _run_root(workload, chosen.backward_state, chosen.output_device) + retained = placement_cost((cost,) * 4, backward_state="gpu", output_device="model") + replay = placement_cost((cost,) * 4, backward_state="replay", output_device="cpu") + assert replay.gpu_required_bytes <= budget < retained.gpu_required_bytes + assert replay.cpu_required_bytes <= cpu_budget + measured = _run_root(workload, "replay", "cpu") assert measured["gpu_peak_bytes"] <= budget From 5475043d11c8c01f56eda86669fa5ba47d47e00d Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 19 Sep 2026 20:16:10 +0000 Subject: [PATCH 006/150] Stamp trainer head registration and export observations --- src/art/trainer_rank/_heads.py | 32 +++++++++++++++++----- tests/unit/test_trainer_rank_commands.py | 2 +- tests/unit/test_trainer_rank_live_heads.py | 4 +-- 3 files changed, 28 insertions(+), 10 deletions(-) diff --git a/src/art/trainer_rank/_heads.py b/src/art/trainer_rank/_heads.py index c27ee7d33..d4f725d04 100644 --- a/src/art/trainer_rank/_heads.py +++ b/src/art/trainer_rank/_heads.py @@ -521,6 +521,18 @@ class HeadState: max_gradient_staleness: int = 2 +@dataclass(frozen=True) +class HeadObservation: + sequence: int + state: HeadState | None + + +def _observe_head(trainer: TrainerRank, state: HeadState | None) -> HeadObservation: + sequence = getattr(trainer, "_head_observation_sequence", 0) + 1 + setattr(trainer, "_head_observation_sequence", sequence) + return HeadObservation(sequence, state) + + @dataclass(frozen=True) class HeadBufferUpdate: version: Any @@ -597,13 +609,18 @@ def execute_head_operation( lambda: deepcopy(registration.value), checkpoint=checkpoint, ) - return export_head(trainer, checkpoint, registration.name) + return _observe_head( + trainer, export_head(trainer, checkpoint, registration.name) + ) if kind == "head_export": return tuple( - export_head(trainer, checkpoint, name) - if checkpoint in trainer._checkpoint_slots - and name in trainer._checkpoint_slots[checkpoint].custom - else None + _observe_head( + trainer, + export_head(trainer, checkpoint, name) + if checkpoint in trainer._checkpoint_slots + and name in trainer._checkpoint_slots[checkpoint].custom + else None, + ) for checkpoint, name in payload ) if kind == "head_publish": @@ -1316,7 +1333,7 @@ def logical_register_head( value = factory() state = view._invoke( "head", "head_register", HeadRegistration(checkpoint, name, kind, value) - ) + ).state else: value = deepcopy(view._rank._checkpoint_slots[checkpoint].custom[name].value) value = ( @@ -1356,7 +1373,8 @@ def refresh_logical_heads(view: Any) -> None: keys = tuple(key for key, head in registry.items() if not head.invalid) if keys: states = view._executor.invoke("head", "head_export", keys) - for key, state in zip(keys, states, strict=True): + for key, observation in zip(keys, states, strict=True): + state = observation.state if state is None: registry[key].invalid = True else: diff --git a/tests/unit/test_trainer_rank_commands.py b/tests/unit/test_trainer_rank_commands.py index 9cb010cc5..4dbc50d45 100644 --- a/tests/unit/test_trainer_rank_commands.py +++ b/tests/unit/test_trainer_rank_commands.py @@ -636,7 +636,7 @@ def factory(): def register(view): parameter = view.parameter("gain", factory, checkpoint="student") assert view.parameter("gain", factory, checkpoint="student") is parameter - return view._invoke("head", "head_export", (("student", "gain"),))[0] + return view._invoke("head", "head_export", (("student", "gain"),))[0].state state = asyncio.run(run_rank_callback(trainer, register, mode="zero")).value assert calls == [True] diff --git a/tests/unit/test_trainer_rank_live_heads.py b/tests/unit/test_trainer_rank_live_heads.py index 7ff5a3354..dd6eb4006 100644 --- a/tests/unit/test_trainer_rank_live_heads.py +++ b/tests/unit/test_trainer_rank_live_heads.py @@ -215,7 +215,7 @@ def test_registration_materializes_factory_state_and_persistent_buffers(): source = TiedHead() state = execute_head_operation( trainer, "head_register", HeadRegistration("student", "head", "module", source) - ) + ).state assert list(state.parameters) == ["left"] assert list(state.buffers) == ["offset"] registered = trainer._checkpoint_slots["student"].custom["head"].value @@ -579,7 +579,7 @@ def test_remote_registration_of_existing_frozen_snapshot_keeps_frozen_parameters source = TiedHead() state = execute_head_operation( trainer, "head_register", HeadRegistration("snapshot", "head", "module", source) - ) + ).state assert source.left.requires_grad assert not state.parameters["left"].requires_grad live = LiveHead(state, source, CotangentCollector()) From 0746fca6e420386c25db81d1ca3146275c2bef3e Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 19 Sep 2026 20:21:57 +0000 Subject: [PATCH 007/150] fix: remap trainer payloads and compact operation retirement --- src/art/trainer_rank/_commands.py | 15 +- src/art/trainer_rank/_operations.py | 140 +++++++++++---- src/art/trainer_rank/_transport.py | 19 ++ tests/unit/test_trainer_driver_transport.py | 32 +++- tests/unit/test_trainer_operations.py | 190 ++++++++++++++++++-- 5 files changed, 336 insertions(+), 60 deletions(-) create mode 100644 src/art/trainer_rank/_transport.py diff --git a/src/art/trainer_rank/_commands.py b/src/art/trainer_rank/_commands.py index 315adf5b8..e2bda3a8d 100644 --- a/src/art/trainer_rank/_commands.py +++ b/src/art/trainer_rank/_commands.py @@ -8,7 +8,6 @@ from dataclasses import dataclass, field, replace from functools import partial import inspect -from io import BytesIO from typing import TYPE_CHECKING, Any, Literal, TypeVar, cast import weakref @@ -16,7 +15,7 @@ import torch import torch.distributed as dist -from . import _impl +from . import _impl, _transport if TYPE_CHECKING: from . import TrainerRank @@ -60,11 +59,7 @@ class _Command: def _encode_command(command: _Command) -> bytes: - # Torch owns storage serialization so nested modules and tensor views keep - # their shared storages. Cloudpickle still supports callback-local objects. - stream = BytesIO() - torch.save(command, stream, pickle_module=cloudpickle) - return stream.getvalue() + return _transport.encode(command) @dataclass(frozen=True) @@ -274,11 +269,7 @@ def _broadcast(self, command: _Command | None) -> _Command: def _decode(self, payload: Any) -> _Command: error, decoded = None, None try: - # A leader's CUDA ordinal is not a peer's local model device. Decode - # storages on CPU; native command handlers place them on their rank. - decoded = torch.load( - BytesIO(payload), map_location="cpu", weights_only=False - ) + decoded = _transport.decode(payload) except Exception as exc: error = f"Command deserialization failed: {exc}" failures = self._gather(error) diff --git a/src/art/trainer_rank/_operations.py b/src/art/trainer_rank/_operations.py index 770747c0f..5af7df142 100644 --- a/src/art/trainer_rank/_operations.py +++ b/src/art/trainer_rank/_operations.py @@ -4,30 +4,101 @@ import asyncio from copy import copy -from dataclasses import dataclass +from dataclasses import dataclass, field import hashlib import inspect +import secrets from typing import Any import cloudpickle +from . import _transport + +OperationId = tuple[str, int] + @dataclass(frozen=True) class TrainerOperation: - id: str + id: OperationId kind: str payload: bytes @classmethod - def capture(cls, id: str, kind: str, payload: Any) -> TrainerOperation: + def capture(cls, id: OperationId, kind: str, payload: Any) -> TrainerOperation: """Freeze arguments before asynchronous submission can observe mutations.""" - return cls(id, kind, cloudpickle.dumps(payload)) + return cls(id, kind, _transport.encode(payload)) + + +@dataclass +class OperationSequence: + """Client IDs and cumulative settlement, including work never admitted.""" + + session: str = field(default_factory=lambda: secrets.token_hex(16)) + issued: int = 0 + pending: set[int] = field(default_factory=set) + abandoned: set[int] = field(default_factory=set) + + def next(self) -> OperationId: + self.issued += 1 + self.pending.add(self.issued) + return self.session, self.issued + + def acknowledge(self, *ids: OperationId, abandon: bool = False) -> TrainerOperation: + for session, sequence in ids: + if session != self.session: + raise ValueError("trainer operation belongs to another client session") + self.pending.discard(sequence) + if abandon: + self.abandoned.add(sequence) + return TrainerOperation.capture( + (self.session, self.issued), "acknowledge", (self.pending, self.abandoned) + ) @dataclass class _Outcome: fingerprint: tuple[str, bytes] - completion: asyncio.Future[Any] | None + completion: asyncio.Future[Any] + + +@dataclass +class _Acknowledged: + through: int = 0 + pending: set[int] = field(default_factory=set) + + +@dataclass +class _Ledger: + outcomes: dict[OperationId, _Outcome] = field(default_factory=dict) + acknowledged: dict[str, _Acknowledged] = field(default_factory=dict) + + def retired(self, id: OperationId) -> bool: + session, sequence = id + state = self.acknowledged.get(session) + return ( + state is not None + and sequence <= state.through + and sequence not in state.pending + ) + + def acknowledge(self, id: OperationId, pending: set[int]) -> None: + session, through = id + if through < 1 or any( + sequence < 1 or sequence > through for sequence in pending + ): + raise ValueError("invalid trainer acknowledgement sequence") + state = self.acknowledged.setdefault(session, _Acknowledged()) + # Snapshots can arrive out of order. Old snapshots may retire more IDs, + # but must never resurrect an already retired operation. + state.pending = { + sequence + for sequence in state.pending + if sequence > through or sequence in pending + } | {sequence for sequence in pending if sequence > state.through} + state.through = max(state.through, through) + for operation_id, outcome in tuple(self.outcomes.items()): + if self.retired(operation_id) and outcome.completion.done(): + del self.outcomes[operation_id] @dataclass(frozen=True) @@ -56,16 +127,28 @@ def capture(cls, error: BaseException) -> _Failure: class OperationResultReleasedError(RuntimeError): - """The operation completed, but its acknowledged result is no longer retained.""" + """The client retired this operation; its result is no longer available.""" -def _ledger(rank_zero: Any) -> dict[str, _Outcome]: +def _ledger(rank_zero: Any) -> _Ledger: owner = rank_zero._rank if not hasattr(owner, "_operation_outcomes"): - owner._operation_outcomes = {} + owner._operation_outcomes = _Ledger() return owner._operation_outcomes +async def _abandon(rank_zero: Any, outcome: _Outcome) -> None: + result = await asyncio.shield(outcome.completion) + if result is None or isinstance(result, _Failure): + return + if outcome.fingerprint[0] == "batches_open": + cleanup = rank_zero.close_forward_batches(result) + else: + cleanup = rank_zero.release_forward((result.handle,)) + if inspect.isawaitable(cleanup): + await cleanup + + async def execute_operation(rank_zero: Any, operation: TrainerOperation) -> Any: """Apply at most once, including concurrent retries and failed mutations. @@ -73,31 +156,30 @@ async def execute_operation(rank_zero: Any, operation: TrainerOperation) -> Any: session; callers must not transparently replay updates on a replacement. """ ledger = _ledger(rank_zero) - payload = cloudpickle.loads(operation.payload) if operation.kind == "acknowledge": - for operation_id in payload: - if (outcome := ledger.get(operation_id)) is not None: - if outcome.completion is not None and not outcome.completion.done(): - raise RuntimeError( - "cannot acknowledge an unfinished trainer operation" - ) - outcome.completion = None + pending, abandoned = _transport.decode(operation.payload) + for sequence in abandoned: + outcome = ledger.outcomes.get((operation.id[0], sequence)) + if outcome is not None: + await _abandon(rank_zero, outcome) + ledger.acknowledge(operation.id, set(pending)) return None + if ledger.retired(operation.id): + raise OperationResultReleasedError(operation.id) fingerprint = (operation.kind, hashlib.sha256(operation.payload).digest()) - if (outcome := ledger.get(operation.id)) is not None: + if (outcome := ledger.outcomes.get(operation.id)) is not None: if outcome.fingerprint != fingerprint: raise ValueError( "trainer operation identity was reused with different arguments" ) - if outcome.completion is None: - raise OperationResultReleasedError(operation.id) result = await asyncio.shield(outcome.completion) if isinstance(result, _Failure): raise cloudpickle.loads(result.payload) from None return result completion = asyncio.get_running_loop().create_future() - ledger[operation.id] = _Outcome(fingerprint, completion) + ledger.outcomes[operation.id] = _Outcome(fingerprint, completion) try: + payload = _transport.decode(operation.payload) if operation.kind in ("forward", "batches_next"): # Only driver replies use CPU transport views. The physical command # receives the original policy, and native callback views are intact. @@ -124,16 +206,9 @@ async def execute_operation(rank_zero: Any, operation: TrainerOperation) -> Any: result = rank_zero.open_forward_batches(**payload) elif operation.kind == "batches_close": result = rank_zero.close_forward_batches(payload["handle"]) - pending = ledger.get(payload.get("pending_operation")) - if ( - pending is not None - and pending.completion is not None - and pending.completion.done() - ): - if ( - packet := pending.completion.result() - ) is not None and not isinstance(packet, _Failure): - rank_zero.release_forward((packet.handle,)) + pending = ledger.outcomes.get(payload.get("pending_operation")) + if pending is not None: + await _abandon(rank_zero, pending) elif operation.kind.startswith("head_"): from ._heads import execute_head_operation @@ -148,3 +223,8 @@ async def execute_operation(rank_zero: Any, operation: TrainerOperation) -> Any: else: completion.set_result(result) return result + finally: + # A cancellation can retire an operation before its remote completion. + # Fence retries immediately, then drop the result once execution settles. + if ledger.retired(operation.id): + ledger.outcomes.pop(operation.id, None) diff --git a/src/art/trainer_rank/_transport.py b/src/art/trainer_rank/_transport.py new file mode 100644 index 000000000..7705cdcc8 --- /dev/null +++ b/src/art/trainer_rank/_transport.py @@ -0,0 +1,19 @@ +"""Storage-aware snapshots shared by client operations and physical commands.""" + +from io import BytesIO +from typing import Any + +import cloudpickle +import torch + + +def encode(value: Any) -> bytes: + # Torch preserves shared storages; cloudpickle supports callback-local types. + stream = BytesIO() + torch.save(value, stream, pickle_module=cloudpickle) + return stream.getvalue() + + +def decode(payload: bytes) -> Any: + # Sender CUDA ordinals are not receiver devices. Native handlers place tensors. + return torch.load(BytesIO(payload), map_location="cpu", weights_only=False) diff --git a/tests/unit/test_trainer_driver_transport.py b/tests/unit/test_trainer_driver_transport.py index 2b9938af1..97292eb16 100644 --- a/tests/unit/test_trainer_driver_transport.py +++ b/tests/unit/test_trainer_driver_transport.py @@ -5,7 +5,7 @@ import asyncio from dataclasses import replace import gc -from typing import Any +from typing import Any, cast import weakref import pytest @@ -48,7 +48,7 @@ def forward(self, tree, **kwargs): def _operation(view, kind, payload, identity=None): return execute_operation( - view, TrainerOperation.capture(identity or str(id(payload)), kind, payload) + view, TrainerOperation.capture((identity or str(id(payload)), 1), kind, payload) ) @@ -253,3 +253,31 @@ async def run(): finally: if enabled: gc.enable() + + +@pytest.mark.parametrize("kind", ["forward", "batches_open"]) +def test_abandoned_reply_releases_native_graphs_and_iterators(kind): + async def run(): + rank = cast(Any, _TransportRank()) + view = _view(_Executor(rank, "zero")) + state = rank._rank_command_state + operation = TrainerOperation.capture( + ("client", 1), + kind, + {"inputs": _input(3) if kind == "forward" else [_input(3)]}, + ) + await execute_operation(view, operation) + if kind == "forward": + assert state.graphs and state.exports + else: + assert state.iterators and state.batch_inputs + acknowledgement = TrainerOperation.capture( + operation.id, "acknowledge", ((), (1,)) + ) + await execute_operation(view, acknowledgement) + await execute_operation(view, acknowledgement) + assert not state.graphs and not state.exports + assert not state.iterators and not state.batch_inputs + assert not rank._operation_outcomes.outcomes + + asyncio.run(run()) diff --git a/tests/unit/test_trainer_operations.py b/tests/unit/test_trainer_operations.py index 781928e99..2b31eccdd 100644 --- a/tests/unit/test_trainer_operations.py +++ b/tests/unit/test_trainer_operations.py @@ -22,7 +22,7 @@ async def run(): optim_step=lambda **kwargs: calls.append(kwargs) or {"step": len(calls)}, ) operation = TrainerOperation.capture( - "step-1", "optim_step", {"params": {"lr": 1}} + ("client", 1), "optim_step", {"params": {"lr": 1}} ) assert await execute_operation(rank, operation) == {"step": 1} assert await execute_operation(rank, operation) == {"step": 1} @@ -30,10 +30,12 @@ async def run(): with pytest.raises(ValueError, match="different arguments"): await execute_operation( rank, - TrainerOperation.capture("step-1", "optim_step", {"params": {"lr": 2}}), + TrainerOperation.capture( + ("client", 1), "optim_step", {"params": {"lr": 2}} + ), ) await execute_operation( - rank, TrainerOperation.capture("ack", "acknowledge", ["step-1"]) + rank, TrainerOperation.capture(("client", 1), "acknowledge", ((), ())) ) with pytest.raises(OperationResultReleasedError): await execute_operation(rank, operation) @@ -56,13 +58,19 @@ def backward_packets(**kwargs): rank = SimpleNamespace( _rank=SimpleNamespace(), backward_packets=backward_packets ) - operation = TrainerOperation.capture("backward-1", "backward", {"packets": ()}) + operation = TrainerOperation.capture(("client", 1), "backward", {"packets": ()}) for _ in range(2): with pytest.raises(TrainerRankSlotStateError) as error: await execute_operation(rank, operation) assert type(error.value) is type(stale) assert str(error.value) == str(stale) assert len(calls) == 1 + await execute_operation( + rank, TrainerOperation.capture(operation.id, "acknowledge", ((), ())) + ) + with pytest.raises(OperationResultReleasedError): + await execute_operation(rank, operation) + assert len(calls) == 1 asyncio.run(run()) @@ -80,7 +88,7 @@ async def optim_step(): return calls rank = SimpleNamespace(_rank=SimpleNamespace(), optim_step=optim_step) - operation = TrainerOperation.capture("step-1", "optim_step", {}) + operation = TrainerOperation.capture(("client", 1), "optim_step", {}) original = asyncio.create_task(execute_operation(rank, operation)) await entered.wait() retry = asyncio.create_task(execute_operation(rank, operation)) @@ -98,7 +106,9 @@ async def optim_step(): def test_operation_captures_tensor_arguments_at_submission(): async def run(): source = torch.tensor([2.0]) - operation = TrainerOperation.capture("forward-1", "forward", {"inputs": source}) + operation = TrainerOperation.capture( + ("client", 1), "forward", {"inputs": source} + ) source.add_(10) rank = SimpleNamespace( _rank=SimpleNamespace(), @@ -112,6 +122,137 @@ async def run(): asyncio.run(run()) +@pytest.mark.parametrize("cuda_tag", [False, True]) +def test_operation_codec_remaps_storages_and_preserves_aliases(monkeypatch, cuda_tag): + from test_trainer_command_transport import _check_payload, _payload + + async def run(): + source = _payload("cpu") + with monkeypatch.context() as capture: + if cuda_tag: + capture.setattr( + torch.serialization, + "_package_registry", + [ + (0, lambda storage: "cuda:7", lambda storage, location: None), + *torch.serialization._package_registry, + ], + ) + operation = TrainerOperation.capture( + ("client", 1), "optim_step", {"value": source} + ) + monkeypatch.setattr(torch.cuda, "is_available", lambda: False) + rank = SimpleNamespace(_rank=SimpleNamespace(), optim_step=lambda value: value) + result = await execute_operation(rank, operation) + _check_payload(result) + assert result.base.untyped_storage() is not source.base.untyped_storage() + + asyncio.run(run()) + + +def test_acknowledgement_bounds_history_including_unadmitted_holes(): + async def run(): + calls = [] + rank = SimpleNamespace( + _rank=SimpleNamespace(), optim_step=lambda: calls.append(1) + ) + # ID 1 is still pending. Odd IDs after it were cancelled before admission. + for sequence in range(2, 2002, 2): + await execute_operation( + rank, TrainerOperation.capture(("client", sequence), "optim_step", {}) + ) + await execute_operation( + rank, + TrainerOperation.capture( + ("client", sequence), "acknowledge", ((1,), ()) + ), + ) + ledger = rank._rank._operation_outcomes + assert not ledger.outcomes + assert len(ledger.acknowledged) == 1 + state = ledger.acknowledged["client"] + assert state.through == 2000 and state.pending == {1} + for sequence in (2, 3, 1999, 2000): + with pytest.raises(OperationResultReleasedError): + await execute_operation( + rank, + TrainerOperation.capture(("client", sequence), "optim_step", {}), + ) + await execute_operation( + rank, TrainerOperation.capture(("client", 1), "optim_step", {}) + ) + await execute_operation( + rank, TrainerOperation.capture(("client", 2000), "acknowledge", ((), ())) + ) + assert not state.pending and not ledger.outcomes + assert len(calls) == 1001 + + asyncio.run(run()) + + +def test_out_of_order_acknowledgements_never_resurrect_ids_or_retire_other_sessions(): + async def run(): + calls = [] + rank = SimpleNamespace( + _rank=SimpleNamespace(), optim_step=lambda: calls.append(1) + ) + for through, pending in ( + (5, (1, 3, 4)), + (3, (1,)), + (6, (1, 4, 6)), + (5, (1, 2, 3)), + ): + await execute_operation( + rank, + TrainerOperation.capture( + ("client", through), "acknowledge", (pending, ()) + ), + ) + state = rank._rank._operation_outcomes.acknowledged["client"] + assert state.through == 6 and state.pending == {1, 6} + for sequence in (2, 3, 4, 5): + with pytest.raises(OperationResultReleasedError): + await execute_operation( + rank, + TrainerOperation.capture(("client", sequence), "optim_step", {}), + ) + for identity in (("client", 1), ("client", 6), ("client", 7), ("other", 3)): + operation = TrainerOperation.capture(identity, "optim_step", {}) + await execute_operation(rank, operation) + await execute_operation(rank, operation) + assert len(calls) == 4 + + asyncio.run(run()) + + +def test_retiring_running_update_fences_retries_until_completion_is_dropped(): + async def run(): + entered, release = asyncio.Event(), asyncio.Event() + calls = [] + + async def optim_step(): + calls.append(1) + entered.set() + await release.wait() + raise ValueError("failed after mutation") + + rank = SimpleNamespace(_rank=SimpleNamespace(), optim_step=optim_step) + operation = TrainerOperation.capture(("client", 1), "optim_step", {}) + original = asyncio.create_task(execute_operation(rank, operation)) + await entered.wait() + await execute_operation( + rank, TrainerOperation.capture(operation.id, "acknowledge", ((), ())) + ) + with pytest.raises(OperationResultReleasedError): + await execute_operation(rank, operation) + release.set() + with pytest.raises(ValueError, match="failed after mutation"): + await original + assert len(calls) == 1 and not rank._rank._operation_outcomes.outcomes + + asyncio.run(run()) + + def test_batch_pulls_and_close_are_identified_without_advancing_twice(): async def run(): events = [] @@ -123,18 +264,24 @@ async def run(): close_forward_batches=lambda handle: events.append(("close", handle)), release_forward=lambda handles: events.append(("release", tuple(handles))), ) - opening = TrainerOperation.capture("open", "batches_open", {"inputs": []}) + opening = TrainerOperation.capture( + ("client", 1), "batches_open", {"inputs": []} + ) assert await execute_operation(rank, opening) == "iterator" assert await execute_operation(rank, opening) == "iterator" next_wave = TrainerOperation.capture( - "next", "batches_next", {"handle": "iterator"} + ("client", 2), "batches_next", {"handle": "iterator"} ) first = await execute_operation(rank, next_wave) assert await execute_operation(rank, next_wave) is first + # Other results can be acknowledged while this wave's reply is lost. + await execute_operation( + rank, TrainerOperation.capture(("client", 3), "acknowledge", ((2, 3), ())) + ) close = TrainerOperation.capture( - "close", + ("client", 3), "batches_close", - {"handle": "iterator", "pending_operation": "next"}, + {"handle": "iterator", "pending_operation": next_wave.id}, ) await execute_operation(rank, close) await execute_operation(rank, close) @@ -144,6 +291,12 @@ async def run(): ("close", "iterator"), ("release", ("packet",)), ] + await execute_operation( + rank, TrainerOperation.capture(close.id, "acknowledge", ((), ())) + ) + assert not rank._rank._operation_outcomes.outcomes + with pytest.raises(OperationResultReleasedError): + await execute_operation(rank, next_wave) asyncio.run(run()) @@ -175,7 +328,7 @@ def forward(): forward=forward, export_forward=lambda output: output, ) - operation = TrainerOperation.capture("failed", "forward", {}) + operation = TrainerOperation.capture(("client", 1), "forward", {}) for _ in range(3): try: await execute_operation(rank, operation) @@ -190,7 +343,10 @@ def forward(): gc.collect() assert references[0]() is None assert len(references) == 1 - assert rank._rank._operation_outcomes["failed"].completion.exception() is None + assert ( + rank._rank._operation_outcomes.outcomes[operation.id].completion.exception() + is None + ) asyncio.run(run()) @@ -208,7 +364,7 @@ async def optim_step(): raise ValueError("update rejected") rank = SimpleNamespace(_rank=SimpleNamespace(), optim_step=optim_step) - operation = TrainerOperation.capture("failed", "optim_step", {}) + operation = TrainerOperation.capture(("client", 1), "optim_step", {}) original = asyncio.create_task(execute_operation(rank, operation)) await entered.wait() retry = asyncio.create_task(execute_operation(rank, operation)) @@ -267,7 +423,9 @@ def fail(*args, **kwargs): else: setattr(view, "_submit_backward", fail) operation = TrainerOperation.capture( - "backward", "backward", {"packets": packets, "retain_graph": retain_graph} + ("client", 1), + "backward", + {"packets": packets, "retain_graph": retain_graph}, ) for _ in range(2): try: @@ -288,7 +446,7 @@ def fail(*args, **kwargs): await execute_operation( view, TrainerOperation.capture( - "retry-intentionally", "backward", {"packets": packets} + ("client", 2), "backward", {"packets": packets} ), ) gc.collect() @@ -314,7 +472,7 @@ async def run(): await execute_operation( view, TrainerOperation.capture( - "invalid", + ("client", 1), "backward", {"packets": (CotangentPacket(handle, gradients),)}, ), From d5b4fb284e7fb76e549a27740f435c6fca029925 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 19 Sep 2026 20:21:45 +0000 Subject: [PATCH 008/150] Share trainer v1 test and development setup --- dev/trainer_v1_acceptance.py | 65 ++++---------- dev/trainer_v1_benchmark.py | 63 ++++---------- dev/trainer_v1_checkpoint.py | 24 ++---- dev/trainer_v1_facade_benchmark.py | 26 ++---- dev/trainer_v1_head_memory.py | 27 ++---- dev/trainer_v1_memory.py | 52 ++++------- dev/trainer_v1_support.py | 86 +++++++++++++++++++ .../test_trainer_rank_callback_cleanup.py | 33 ++----- .../test_trainer_rank_callback_failures.py | 34 ++------ .../test_trainer_rank_callback_lifecycle.py | 30 ++----- tests/unit/test_trainer_rank_graph_order.py | 13 +-- .../test_trainer_rank_memory_admission.py | 29 +------ ...rainer_rank_memory_recovery_distributed.py | 23 +---- tests/unit/test_trainer_rank_split_peak.py | 31 +------ tests/unit/test_trainer_rank_versions.py | 52 ++--------- tests/unit/trainer_rank_test_support.py | 74 ++++++++++++++++ 16 files changed, 263 insertions(+), 399 deletions(-) create mode 100644 dev/trainer_v1_support.py create mode 100644 tests/unit/trainer_rank_test_support.py diff --git a/dev/trainer_v1_acceptance.py b/dev/trainer_v1_acceptance.py index cc0604637..50080e724 100644 --- a/dev/trainer_v1_acceptance.py +++ b/dev/trainer_v1_acceptance.py @@ -15,7 +15,6 @@ import os from pathlib import Path import resource -import subprocess import sys import time import traceback @@ -23,9 +22,16 @@ import torch import torch.distributed as dist from trainer_rank_diag import rank0_checked -from trainer_rank_support import load_random_checkpoints +from trainer_v1_support import ( + _cached_backward, + _leaves, + build_rank, + load_checkpoint, + nccl_group, + source_commits, +) -from art.trainer_rank import ForwardInput, TrainerRank +from art.trainer_rank import ForwardInput def _requests(checkpoint, offset, lengths, options=None): @@ -45,12 +51,6 @@ def _requests(checkpoint, offset, lengths, options=None): return [[leaves[0], *leaves[1:2]], leaves[2:]] if len(leaves) > 1 else leaves -def _leaves(tree): - if isinstance(tree, (list, tuple)): - return [leaf for item in tree for leaf in _leaves(item)] - return [tree] - - def _loss_and_cotangents(first, second): a = torch.cat([leaf.target_logprobs for leaf in _leaves(first)]) b = torch.cat([leaf.target_logprobs for leaf in _leaves(second)]) @@ -68,14 +68,6 @@ def _loss_and_cotangents(first, second): return loss, tensors, gradients -def _cached_backward(rank, loss): - with rank._gradient_transaction(): - packets = rank._forward_cotangent_collector().backward(loss) - rank._forward_graph_cache().backward_many( - [(packet.handle, packet.gradients) for packet in packets] - ) - - def _canonical_gradients(rank, checkpoint): """Reduce once, then gather canonical LoRA shards without altering weights.""" from art.megatron.lora import LoRA @@ -189,33 +181,21 @@ def main(): device = int(os.environ["LOCAL_RANK"]) if args.reverse_devices: device = int(os.environ["LOCAL_WORLD_SIZE"]) - 1 - device - torch.cuda.set_device(device) - dist.init_process_group("nccl") - try: + with nccl_group(device): from megatron.core import parallel_state as ps - from art.megatron.train import build_training_runtime - if args.mode in ("oracle", "retained") and dist.get_world_size() != 1: raise ValueError( "The mathematical global-loss oracle requires one physical rank" ) - torch.manual_seed(90217) - runtime = build_training_runtime( - model_identifier=args.model, - provider_configure=(lambda p: setattr(p, "num_layers", args.layers)) - if args.layers - else None, + physical = build_rank( + args.model, + layers=args.layers or None, print_env=dist.get_rank() == 0, ) - for chunk in runtime.model: - chunk.eval() - physical = TrainerRank(runtime) if args.mode == "control" and ps.get_data_parallel_world_size() != 1: raise ValueError("A physical topology control requires DP=1") - (checkpoint,) = load_random_checkpoints( - runtime, physical, 1, base_model=args.model, lora_rank=2 - ) + checkpoint = load_checkpoint(physical, args.model) options = None if args.mode == "retained": from art.trainer_rank import ForwardOptions @@ -230,7 +210,7 @@ def main(): second = _requests(checkpoint, 113, [19], options) callbacks = 0 cache_before_backward = None - hidden_size = runtime.provider.hidden_size + hidden_size = physical.runtime.provider.hidden_size def head_factory(): head = torch.nn.Linear( @@ -472,18 +452,7 @@ def finish(): }, "device": torch.cuda.get_device_name(), "torch": torch.__version__, - "source_commit": subprocess.check_output( - ["git", "rev-parse", "HEAD"], - cwd=Path(sys.modules["art.trainer_rank"].__file__) - .resolve() - .parents[3], - text=True, - ).strip(), - "harness_commit": subprocess.check_output( - ["git", "rev-parse", "HEAD"], - cwd=Path(__file__).resolve().parent.parent, - text=True, - ).strip(), + **source_commits(), "measurements": measurements, "comparison": comparison, "retention": args.retention if args.mode == "retained" else None, @@ -501,8 +470,6 @@ def finish(): print(json.dumps(metadata), flush=True) rank0_checked("trainer v1 global-loss acceptance", finish) - finally: - dist.destroy_process_group() if __name__ == "__main__": diff --git a/dev/trainer_v1_benchmark.py b/dev/trainer_v1_benchmark.py index f716d45d2..eedd3564c 100644 --- a/dev/trainer_v1_benchmark.py +++ b/dev/trainer_v1_benchmark.py @@ -16,18 +16,23 @@ from pathlib import Path import resource import statistics -import subprocess -import sys import time import weakref from dotenv import load_dotenv import torch import torch.distributed as dist -from trainer_rank_support import load_random_checkpoints -from trainer_v1_acceptance import _cached_backward, _leaves +from trainer_v1_support import ( + _cached_backward, + _leaves, + build_rank, + deterministic_kernels, + load_checkpoint, + nccl_group, + source_commits, +) -from art.trainer_rank import AdamParams, ForwardInput, TrainerRank +from art.trainer_rank import AdamParams, ForwardInput def main(): @@ -53,39 +58,16 @@ def main(): parser.error("Delayed graphs require v1; use gpu as the delayed oracle") load_dotenv(".env") if args.deterministic: - os.environ["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8" - torch.use_deterministic_algorithms(True) - import art.megatron.flex_attn.compiled as flex - - flex._FORCED_FLEX_BACKEND = "TRITON" - flex._FORCED_FLEX_KERNEL_OPTIONS = {"BACKEND": "TRITON"} - flex.dense_compiled_flex_attention = flex.triton_dense_compiled_flex_attention - flex.sparse_compiled_flex_attention = flex.triton_sparse_compiled_flex_attention + deterministic_kernels() for axis in ("TENSOR_MODEL", "CONTEXT", "DATA", "PIPELINE_MODEL"): os.environ[f"ART_MEGATRON_{axis}_PARALLEL_SIZE"] = "1" os.environ["ART_TRAINER_RANK_TEST_HOOKS"] = "1" os.environ["ART_TRAINER_RANK_TEST_ANCHOR"] = "no_sharing" - torch.cuda.set_device(int(os.environ["LOCAL_RANK"])) - dist.init_process_group("nccl") - try: + with nccl_group(): if dist.get_world_size() != 1: raise ValueError("This paired benchmark requires one physical rank") - from art.megatron.train import build_training_runtime - - torch.manual_seed(90217) - runtime = build_training_runtime( - model_identifier=args.model, - provider_configure=(lambda p: setattr(p, "num_layers", args.layers)) - if args.layers - else None, - print_env=False, - ) - for chunk in runtime.model: - chunk.eval() - rank = TrainerRank(runtime) - (checkpoint,) = load_random_checkpoints( - runtime, rank, 1, base_model=args.model, lora_rank=2 - ) + rank = build_rank(args.model, layers=args.layers or None) + checkpoint = load_checkpoint(rank, args.model) forward = getattr(rank, "forward", None) or rank.dp_rank_forward options = None if args.mode != "baseline": @@ -279,20 +261,11 @@ def step(value): "arguments": { k: str(v) if isinstance(v, Path) else v for k, v in vars(args).items() }, - "source_commit": subprocess.check_output( - ["git", "rev-parse", "HEAD"], - cwd=Path(sys.modules["art.trainer_rank"].__file__).resolve().parents[3], - text=True, - ).strip(), - "harness_commit": subprocess.check_output( - ["git", "rev-parse", "HEAD"], - cwd=Path(__file__).resolve().parent.parent, - text=True, - ).strip(), + **source_commits(), "torch": torch.__version__, "device": torch.cuda.get_device_name(), - "layers": runtime.provider.num_layers, - "dtype": str(next(runtime.model[0].parameters()).dtype), + "layers": rank.runtime.provider.num_layers, + "dtype": str(next(rank.runtime.model[0].parameters()).dtype), "initial_fixed_objective": initial_eval, "final_fixed_objective": final_eval, "median_tokens_per_second": statistics.median( @@ -322,8 +295,6 @@ def step(value): raise AssertionError( f"Fixed objective did not improve: {initial_eval} -> {final_eval}" ) - finally: - dist.destroy_process_group() if __name__ == "__main__": diff --git a/dev/trainer_v1_checkpoint.py b/dev/trainer_v1_checkpoint.py index 528d6239b..a006dbcb0 100644 --- a/dev/trainer_v1_checkpoint.py +++ b/dev/trainer_v1_checkpoint.py @@ -5,8 +5,9 @@ from pathlib import Path from dotenv import load_dotenv -import torch -import torch.distributed as dist +from trainer_v1_support import build_rank, load_checkpoint, nccl_group + +from art.trainer_rank import validate_checkpoint def main(): @@ -17,25 +18,12 @@ def main(): load_dotenv(".env") for axis in ("TENSOR_MODEL", "CONTEXT", "DATA", "PIPELINE_MODEL"): os.environ[f"ART_MEGATRON_{axis}_PARALLEL_SIZE"] = "1" - torch.cuda.set_device(int(os.environ["LOCAL_RANK"])) - dist.init_process_group("nccl") - try: - from trainer_rank_support import load_random_checkpoints - - from art.megatron.train import build_training_runtime - from art.trainer_rank import TrainerRank, validate_checkpoint - - torch.manual_seed(90217) - runtime = build_training_runtime(model_identifier=args.model, print_env=False) - rank = TrainerRank(runtime) - (checkpoint,) = load_random_checkpoints( - runtime, rank, 1, base_model=args.model, lora_rank=2 - ) + with nccl_group(): + rank = build_rank(args.model, eval_mode=False) + checkpoint = load_checkpoint(rank, args.model) rank.save_checkpoint(str(args.output), checkpoint) assert validate_checkpoint(args.output) is not None print(f"native checkpoint saved: {args.output}", flush=True) - finally: - dist.destroy_process_group() if __name__ == "__main__": diff --git a/dev/trainer_v1_facade_benchmark.py b/dev/trainer_v1_facade_benchmark.py index 282c39192..9b1629253 100644 --- a/dev/trainer_v1_facade_benchmark.py +++ b/dev/trainer_v1_facade_benchmark.py @@ -15,9 +15,9 @@ from dotenv import load_dotenv import torch import torch.distributed as dist -from trainer_rank_support import load_random_checkpoints +from trainer_v1_support import build_rank, load_checkpoint, nccl_group -from art.trainer_rank import ForwardInput, ForwardOptions, TrainerRank, _commands +from art.trainer_rank import ForwardInput, ForwardOptions, _commands def main(): @@ -29,9 +29,7 @@ def main(): parser.add_argument("--profile", action="store_true") args = parser.parse_args() load_dotenv(".env") - torch.cuda.set_device(int(os.environ["LOCAL_RANK"])) - dist.init_process_group("nccl") - try: + with nccl_group(): world = dist.get_world_size() for axis in ("TENSOR_MODEL", "CONTEXT", "DATA", "PIPELINE_MODEL"): os.environ[f"ART_MEGATRON_{axis}_PARALLEL_SIZE"] = str( @@ -39,20 +37,8 @@ def main(): ) os.environ["ART_TRAINER_RANK_TEST_HOOKS"] = "1" os.environ["ART_TRAINER_RANK_TEST_ANCHOR"] = "no_sharing" - from art.megatron.train import build_training_runtime - - torch.manual_seed(90217) - runtime = build_training_runtime( - model_identifier="Qwen/Qwen3-0.6B", - provider_configure=lambda p: setattr(p, "num_layers", args.layers), - print_env=False, - ) - for chunk in runtime.model: - chunk.eval() - physical = TrainerRank(runtime) - (checkpoint,) = load_random_checkpoints( - runtime, physical, 1, base_model="Qwen/Qwen3-0.6B", lora_rank=2 - ) + physical = build_rank("Qwen/Qwen3-0.6B", layers=args.layers) + checkpoint = load_checkpoint(physical, "Qwen/Qwen3-0.6B") options = ForwardOptions( backward_state="gpu", output_device="model", stale_gradient_corrections=() ) @@ -227,8 +213,6 @@ def measure(mode, repetitions): "cumulative" ).print_stats(100) print("FACADE_SUMMARY=" + json.dumps(summary), flush=True) - finally: - dist.destroy_process_group() if __name__ == "__main__": diff --git a/dev/trainer_v1_head_memory.py b/dev/trainer_v1_head_memory.py index 3f20d52b2..7cbdefdce 100644 --- a/dev/trainer_v1_head_memory.py +++ b/dev/trainer_v1_head_memory.py @@ -9,14 +9,12 @@ from dotenv import load_dotenv import torch -import torch.distributed as dist from torch.multiprocessing.reductions import StorageWeakRef -from trainer_rank_support import load_random_checkpoints +from trainer_v1_support import build_rank, load_checkpoint, nccl_group from art.trainer_rank import ( ForwardInput, ForwardOptions, - TrainerRank, TrainerRankMemoryError, run_rank_callback, ) @@ -28,23 +26,16 @@ def main(): args = parser.parse_args() load_dotenv(".env") torch.set_num_threads(2) - torch.cuda.set_device(int(os.environ.get("LOCAL_RANK", "0"))) - dist.init_process_group("nccl") - try: - from art.megatron.train import build_training_runtime - + with nccl_group(int(os.environ.get("LOCAL_RANK", "0"))): os.environ["ART_TRAINER_RANK_TEST_HOOKS"] = "1" - torch.manual_seed(716) - runtime = build_training_runtime( - model_identifier="Qwen/Qwen3-0.6B", + rank = build_rank( + "Qwen/Qwen3-0.6B", + seed=716, + layers=2, + eval_mode=False, model_initialization="random", - provider_configure=lambda provider: setattr(provider, "num_layers", 2), - print_env=False, - ) - rank = TrainerRank(runtime) - (checkpoint,) = load_random_checkpoints( - runtime, rank, 1, base_model="Qwen/Qwen3-0.6B", lora_rank=2 ) + checkpoint = load_checkpoint(rank, "Qwen/Qwen3-0.6B") request = ForwardInput( input_tokens=torch.arange(32), hidden_states=True, @@ -173,8 +164,6 @@ def callback(view): args.output.parent.mkdir(parents=True, exist_ok=True) args.output.write_text(json.dumps(rows, indent=2)) print("HEAD_MEMORY=" + json.dumps(rows), flush=True) - finally: - dist.destroy_process_group() if __name__ == "__main__": diff --git a/dev/trainer_v1_memory.py b/dev/trainer_v1_memory.py index db0065a1c..cb7b36eaf 100644 --- a/dev/trainer_v1_memory.py +++ b/dev/trainer_v1_memory.py @@ -13,33 +13,29 @@ from dotenv import load_dotenv import torch import torch.distributed as dist -from trainer_rank_support import load_random_checkpoints +from trainer_v1_support import ( + _cached_backward, + _leaves, + build_rank, + deterministic_kernels, + load_checkpoint, + nccl_group, +) from art.trainer_rank import ( ForwardInput, ForwardOptions, - TrainerRank, TrainerRankMemoryError, ) from art.trainer_rank._impl import _SplitForwardPlan from art.trainer_rank._memory_policy import host_memory_budget, placement_cost -def _leaves(tree): - if isinstance(tree, (list, tuple)): - return [leaf for child in tree for leaf in _leaves(child)] - return [tree] - - def _backward(rank, outputs): loss = sum( output.hidden_states.float().mean() for output in _leaves(outputs) ).square() - with rank._gradient_transaction(): - packets = rank._forward_cotangent_collector().backward(loss) - rank._forward_graph_cache().backward_many( - [(p.handle, p.gradients) for p in packets] - ) + _cached_backward(rank, loss) return float(loss.detach().cpu()) @@ -54,22 +50,12 @@ def main(): args = parser.parse_args() load_dotenv(".env") if args.deterministic: - os.environ["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8" - torch.use_deterministic_algorithms(True) - import art.megatron.flex_attn.compiled as flex - - flex._FORCED_FLEX_BACKEND = "TRITON" - flex._FORCED_FLEX_KERNEL_OPTIONS = {"BACKEND": "TRITON"} - flex.dense_compiled_flex_attention = flex.triton_dense_compiled_flex_attention - flex.sparse_compiled_flex_attention = flex.triton_sparse_compiled_flex_attention - torch.cuda.set_device(int(os.environ.get("LOCAL_RANK", "0"))) - dist.init_process_group("nccl") - try: + deterministic_kernels() + with nccl_group(int(os.environ.get("LOCAL_RANK", "0"))): if dist.get_world_size() != 1: raise ValueError("This admission oracle uses one physical rank") os.environ["ART_TRAINER_RANK_TEST_HOOKS"] = "1" os.environ["ART_TRAINER_RANK_TEST_ANCHOR"] = "no_sharing" - from art.megatron.train import build_training_runtime def configure(provider): provider.num_layers = args.layers @@ -78,19 +64,13 @@ def configure(provider): provider.recompute_num_layers = None provider.recompute_modules = [] - torch.manual_seed(913) - runtime = build_training_runtime( - model_identifier=args.model, + rank = build_rank( + args.model, + seed=913, model_initialization="random", provider_configure=configure, - print_env=False, - ) - for chunk in runtime.model: - chunk.eval() - rank = TrainerRank(runtime) - (checkpoint,) = load_random_checkpoints( - runtime, rank, 1, base_model=args.model, lora_rank=2 ) + checkpoint = load_checkpoint(rank, args.model) forward = getattr(rank, "forward", None) or rank.dp_rank_forward batches = getattr(rank, "forward_batches", None) or rank.forward_micro_batches options = lambda state, device: ForwardOptions( @@ -375,8 +355,6 @@ def fixed_split(requests, *, checkpoint, **_kwargs): default=str, ) ) - finally: - dist.destroy_process_group() if __name__ == "__main__": diff --git a/dev/trainer_v1_support.py b/dev/trainer_v1_support.py new file mode 100644 index 000000000..aeefd127f --- /dev/null +++ b/dev/trainer_v1_support.py @@ -0,0 +1,86 @@ +"""Shared setup for the native trainer-v1 validation programs.""" + +from contextlib import contextmanager +import os +from pathlib import Path +import subprocess +import sys + +import torch +import torch.distributed as dist +from trainer_rank_support import load_random_checkpoints + +from art.trainer_rank import TrainerRank + + +@contextmanager +def nccl_group(local_rank=None): + torch.cuda.set_device( + int(os.environ["LOCAL_RANK"]) if local_rank is None else local_rank + ) + dist.init_process_group("nccl") + try: + yield + finally: + dist.destroy_process_group() + + +def build_rank(model, *, layers=None, seed=90217, eval_mode=True, **runtime_options): + from art.megatron.train import build_training_runtime + + torch.manual_seed(seed) + runtime_options.setdefault("print_env", False) + runtime_options.setdefault( + "provider_configure", + (lambda p: setattr(p, "num_layers", layers)) if layers is not None else None, + ) + runtime = build_training_runtime(model_identifier=model, **runtime_options) + if eval_mode: + for chunk in runtime.model: + chunk.eval() + return TrainerRank(runtime) + + +def load_checkpoint(rank, model): + (checkpoint,) = load_random_checkpoints( + rank.runtime, rank, 1, base_model=model, lora_rank=2 + ) + return checkpoint + + +def deterministic_kernels(): + os.environ["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8" + torch.use_deterministic_algorithms(True) + import art.megatron.flex_attn.compiled as flex + + setattr(flex, "_FORCED_FLEX_BACKEND", "TRITON") + flex._FORCED_FLEX_KERNEL_OPTIONS = {"BACKEND": "TRITON"} + flex.dense_compiled_flex_attention = flex.triton_dense_compiled_flex_attention + flex.sparse_compiled_flex_attention = flex.triton_sparse_compiled_flex_attention + + +def source_commits(): + source_file = sys.modules["art.trainer_rank"].__file__ + assert source_file is not None + source = Path(source_file).resolve().parents[3] + harness = Path(__file__).resolve().parent.parent + return { + f"{name}_commit": subprocess.check_output( + ["git", "rev-parse", "HEAD"], cwd=directory, text=True + ).strip() + for name, directory in (("source", source), ("harness", harness)) + } + + +def _leaves(tree): + if isinstance(tree, (list, tuple)): + return [leaf for child in tree for leaf in _leaves(child)] + return [tree] + + +def _cached_backward(rank, loss): + with rank._gradient_transaction(): + packets = rank._forward_cotangent_collector().backward(loss) + rank._forward_graph_cache().backward_many( + [(packet.handle, packet.gradients) for packet in packets] + ) diff --git a/tests/unit/test_trainer_rank_callback_cleanup.py b/tests/unit/test_trainer_rank_callback_cleanup.py index 44a607663..2b93a422b 100644 --- a/tests/unit/test_trainer_rank_callback_cleanup.py +++ b/tests/unit/test_trainer_rank_callback_cleanup.py @@ -3,10 +3,7 @@ from __future__ import annotations import asyncio -from datetime import timedelta import gc -import sys -from types import ModuleType, SimpleNamespace from typing import Any import pytest @@ -14,6 +11,7 @@ import torch import torch.distributed as dist import torch.multiprocessing as mp +from trainer_rank_test_support import gloo_group, megatron_topology from art.trainer_rank import run_rank_callback, run_rank_callback_stream @@ -149,30 +147,11 @@ async def suspended(view): def _worker(physical: int, rendezvous: str) -> None: torch.set_num_threads(1) - dist.init_process_group( - "gloo", - init_method=f"file://{rendezvous}", - rank=physical, - world_size=4, - timeout=timedelta(seconds=20), - ) - try: - groups = [dist.new_group([0, 1]), dist.new_group([2, 3])] - dp, tp = divmod(physical, 2) - ps = SimpleNamespace( - get_tensor_model_parallel_rank=lambda: tp, - get_context_parallel_rank=lambda: 0, - get_data_parallel_rank=lambda: dp, - get_data_parallel_world_size=lambda: 2, - get_tensor_and_context_parallel_group=lambda **kwargs: groups[dp], - ) - megatron, core = ModuleType("megatron"), ModuleType("megatron.core") - setattr(core, "parallel_state", ps) - setattr(megatron, "core", core) - sys.modules.update({"megatron": megatron, "megatron.core": core}) - asyncio.run(_cases(_Rank(dp, 2), physical)) - finally: - dist.destroy_process_group() + with ( + gloo_group(physical, f"file://{rendezvous}", world_size=4, timeout=20), + megatron_topology(physical, dp_size=2, tp_size=2), + ): + asyncio.run(_cases(_Rank(physical // 2, 2), physical)) def test_gloo_dp2_tp2_callback_cleanup_across_modes(tmp_path): diff --git a/tests/unit/test_trainer_rank_callback_failures.py b/tests/unit/test_trainer_rank_callback_failures.py index 578224489..70af8ab06 100644 --- a/tests/unit/test_trainer_rank_callback_failures.py +++ b/tests/unit/test_trainer_rank_callback_failures.py @@ -3,18 +3,16 @@ from __future__ import annotations import asyncio -from datetime import timedelta +from contextlib import closing import gc from multiprocessing.connection import Connection -import sys -from types import ModuleType, SimpleNamespace from typing import Any import pytest from test_trainer_rank_commands import _input, _loss_tree, _Rank import torch -import torch.distributed as dist import torch.multiprocessing as mp +from trainer_rank_test_support import gloo_group, megatron_topology from art.trainer_rank import run_rank_callback, run_rank_callback_stream @@ -111,30 +109,12 @@ async def execute(): def _worker(physical: int, rendezvous: str, connection: Connection, kind: str) -> None: torch.set_num_threads(1) - dist.init_process_group( - "gloo", - init_method=f"file://{rendezvous}", - rank=physical, - world_size=2, - timeout=timedelta(seconds=12), - ) - try: - groups = [dist.new_group([0]), dist.new_group([1])] - ps = SimpleNamespace( - get_tensor_model_parallel_rank=lambda: 0, - get_context_parallel_rank=lambda: 0, - get_data_parallel_rank=lambda: physical, - get_data_parallel_world_size=lambda: 2, - get_tensor_and_context_parallel_group=lambda **kwargs: groups[physical], - ) - megatron, core = ModuleType("megatron"), ModuleType("megatron.core") - setattr(core, "parallel_state", ps) - setattr(megatron, "core", core) - sys.modules.update({"megatron": megatron, "megatron.core": core}) + with ( + gloo_group(physical, f"file://{rendezvous}", timeout=12), + closing(connection), + megatron_topology(physical, dp_size=2, tp_size=1), + ): asyncio.run(_serve(_Rank(physical, 2), connection, kind)) - finally: - connection.close() - dist.destroy_process_group() async def _gather(executions): diff --git a/tests/unit/test_trainer_rank_callback_lifecycle.py b/tests/unit/test_trainer_rank_callback_lifecycle.py index 4d2f51b2e..211cceb35 100644 --- a/tests/unit/test_trainer_rank_callback_lifecycle.py +++ b/tests/unit/test_trainer_rank_callback_lifecycle.py @@ -1,16 +1,14 @@ """Checkpoint scopes and delayed cleanup stay within their callback session.""" import asyncio -from datetime import timedelta import gc -import sys -from types import ModuleType, SimpleNamespace from typing import Any import pytest from test_trainer_rank_commands import _input, _Rank import torch.distributed as dist import torch.multiprocessing as mp +from trainer_rank_test_support import gloo_group, megatron_topology from art.trainer_rank import run_rank_callback @@ -95,26 +93,10 @@ def following(view): def _lifecycle_worker(physical, rendezvous): - dist.init_process_group( - "gloo", - init_method=f"file://{rendezvous}", - rank=physical, - world_size=2, - timeout=timedelta(seconds=15), - ) - try: - megatron, core = ModuleType("megatron"), ModuleType("megatron.core") - setattr( - core, - "parallel_state", - SimpleNamespace( - get_tensor_model_parallel_rank=lambda: physical, - get_context_parallel_rank=lambda: 0, - get_tensor_and_context_parallel_group=lambda: dist.group.WORLD, - ), - ) - setattr(megatron, "core", core) - sys.modules.update({"megatron": megatron, "megatron.core": core}) + with ( + gloo_group(physical, f"file://{rendezvous}", timeout=15), + megatron_topology(physical, dp_size=1, tp_size=2), + ): for mode in ("zero", "rank"): rank: Any = _CheckpointRank() held = [] @@ -174,8 +156,6 @@ def mismatched(view): rank.zero_grad = zero_grad asyncio.run(run_rank_callback(rank, following, mode=mode)) dist.barrier() - finally: - dist.destroy_process_group() def test_delayed_callback_cleanup_and_checkpoint_scopes_leave_gloo_reusable(tmp_path): diff --git a/tests/unit/test_trainer_rank_graph_order.py b/tests/unit/test_trainer_rank_graph_order.py index 866654135..95110ccc7 100644 --- a/tests/unit/test_trainer_rank_graph_order.py +++ b/tests/unit/test_trainer_rank_graph_order.py @@ -1,12 +1,12 @@ """Physical peers must enter cached/replayed backward collectives identically.""" -from datetime import timedelta from types import SimpleNamespace import pytest import torch import torch.distributed as dist import torch.multiprocessing as mp +from trainer_rank_test_support import gloo_group from art.trainer_rank import TrainerRank, _graphs from art.trainer_rank._impl import _CheckpointSlot @@ -32,14 +32,7 @@ def backward(ctx, *gradients): def _worker(rank, rendezvous, fail_replay): - dist.init_process_group( - "gloo", - init_method=rendezvous, - rank=rank, - world_size=2, - timeout=timedelta(seconds=30), - ) - try: + with gloo_group(rank, rendezvous): names = iter(("z", "a") if rank == 0 else ("a", "z")) setattr(_graphs, "uuid4", lambda: SimpleNamespace(hex=next(names))) cache = _graphs.GraphCache() @@ -103,8 +96,6 @@ def backward(): completed = torch.tensor(1) dist.all_reduce(completed) assert completed.item() == 2 - finally: - dist.destroy_process_group() @pytest.mark.parametrize("fail_replay", [False, True]) diff --git a/tests/unit/test_trainer_rank_memory_admission.py b/tests/unit/test_trainer_rank_memory_admission.py index 9a0b087fc..98b592870 100644 --- a/tests/unit/test_trainer_rank_memory_admission.py +++ b/tests/unit/test_trainer_rank_memory_admission.py @@ -1,46 +1,21 @@ from dataclasses import replace from types import SimpleNamespace -from typing import Any, cast import pytest +from test_trainer_rank_active_memory import _rank import torch from art.trainer_rank import ( ForwardInput, ForwardOptions, ImportanceSamplingGradientCorrection, - TrainerRank, _impl, ) @pytest.fixture def rank(monkeypatch): - class Model(torch.nn.Module): - def __init__(self): - super().__init__() - self.weight = torch.nn.Parameter(torch.zeros((), dtype=torch.bfloat16)) - self.config = SimpleNamespace( - hidden_size=8, num_layers=4, padded_vocab_size=32 - ) - self.decoder = object() - - def _preprocess(self, *args, **kwargs): - return None - - result = TrainerRank( - cast( - Any, - SimpleNamespace( - model=[Model()], - optimizer=None, - provider=SimpleNamespace( - hidden_size=8, num_layers=4, recompute_granularity="full" - ), - model_support_handler=SimpleNamespace(build_gdn_execution_spec=False), - ), - ) - ) + result = _rank() monkeypatch.setattr(result, "_graph_memory_policy_enabled", lambda: True) monkeypatch.setattr(result, "_available_memory_bytes", lambda: 250) monkeypatch.setattr(result, "_available_cpu_memory_bytes", lambda: 1_000_000) diff --git a/tests/unit/test_trainer_rank_memory_recovery_distributed.py b/tests/unit/test_trainer_rank_memory_recovery_distributed.py index e54a5decb..1ce028c86 100644 --- a/tests/unit/test_trainer_rank_memory_recovery_distributed.py +++ b/tests/unit/test_trainer_rank_memory_recovery_distributed.py @@ -9,6 +9,7 @@ import torch import torch.distributed as dist import torch.multiprocessing as mp +from trainer_rank_test_support import gloo_group from art.trainer_rank import _impl from art.trainer_rank._memory_policy import ForwardMemoryCost @@ -16,14 +17,7 @@ def _reclaim_worker(index: int, directory: str, sync_across_dp: bool) -> None: - dist.init_process_group( - "gloo", - init_method=f"file://{directory}/rendezvous", - rank=index, - world_size=2, - timeout=timedelta(seconds=30), - ) - try: + with gloo_group(index, f"file://{directory}/rendezvous"): rank = object.__new__(_impl.TrainerRank) rank.device = torch.device("cpu") rank._graph_memory_policy_enabled = lambda: True @@ -56,8 +50,6 @@ def offload(_handle): dist.all_gather_object(gathered, result) if index == 0: Path(directory, "results.json").write_text(json.dumps(gathered)) - finally: - dist.destroy_process_group() @pytest.mark.parametrize("sync_across_dp", [False, True]) @@ -70,14 +62,7 @@ def test_failed_offload_does_not_strand_a_physical_peer(tmp_path, sync_across_dp def _fallback_worker(index: int, directory: str, sync_across_dp: bool) -> None: - dist.init_process_group( - "gloo", - init_method=f"file://{directory}/rendezvous", - rank=index, - world_size=2, - timeout=timedelta(seconds=30), - ) - try: + with gloo_group(index, f"file://{directory}/rendezvous"): rank = object.__new__(_impl.TrainerRank) rank.device = torch.device("cpu") rank._forward_memory_group = lambda: dist.group.WORLD @@ -113,8 +98,6 @@ def _fallback_worker(index: int, directory: str, sync_across_dp: bool) -> None: ) assert candidates[2][0] == "cpu" assert candidates[2][2]["source"] == "insufficient_samples" - finally: - dist.destroy_process_group() @pytest.mark.parametrize("sync_across_dp", [False, True]) diff --git a/tests/unit/test_trainer_rank_split_peak.py b/tests/unit/test_trainer_rank_split_peak.py index bbac43557..6442f4831 100644 --- a/tests/unit/test_trainer_rank_split_peak.py +++ b/tests/unit/test_trainer_rank_split_peak.py @@ -4,42 +4,15 @@ from contextlib import nullcontext from dataclasses import replace from itertools import permutations -from types import SimpleNamespace -from typing import Any, cast +from typing import Any import pytest +from test_trainer_rank_active_memory import _rank import torch from art.trainer_rank import _impl as tr -class _Model(torch.nn.Module): - def __init__(self): - super().__init__() - self.weight = torch.nn.Parameter(torch.zeros((), dtype=torch.bfloat16)) - self.config = SimpleNamespace(hidden_size=8, num_layers=4, padded_vocab_size=32) - self.decoder = object() - - def _preprocess(self, *args, **kwargs): - return None - - -def _rank(): - return tr.TrainerRank( - cast( - Any, - SimpleNamespace( - model=[_Model()], - optimizer=None, - provider=SimpleNamespace( - hidden_size=8, num_layers=4, recompute_granularity="full" - ), - model_support_handler=SimpleNamespace(build_gdn_execution_spec=False), - ), - ) - ) - - def _requests(count=2, length=100): return [ tr.ForwardInput( diff --git a/tests/unit/test_trainer_rank_versions.py b/tests/unit/test_trainer_rank_versions.py index 36076c420..1caba182f 100644 --- a/tests/unit/test_trainer_rank_versions.py +++ b/tests/unit/test_trainer_rank_versions.py @@ -1,18 +1,16 @@ from __future__ import annotations from contextlib import nullcontext -from datetime import timedelta import gc from pathlib import Path -import time from types import SimpleNamespace import weakref import pytest import torch import torch.distributed as dist -import torch.multiprocessing as mp from torch.utils.checkpoint import checkpoint +from trainer_rank_test_support import gloo_group, spawn_and_join from art.trainer_rank import TrainerRank, TrainerRankSlotStateError from art.trainer_rank._impl import _CheckpointSlot @@ -316,14 +314,7 @@ def test_explicit_head_cotangents_join_model_gradient_transaction() -> None: def _divergent_version_worker(rank: int, rendezvous: str) -> None: - dist.init_process_group( - "gloo", - init_method=rendezvous, - rank=rank, - world_size=2, - timeout=timedelta(seconds=30), - ) - try: + with gloo_group(rank, rendezvous): trainer, current = _trainer() version = trainer._capture_checkpoint_version("student") if rank == 0: @@ -340,26 +331,17 @@ def _divergent_version_worker(rank: int, rendezvous: str) -> None: completed = torch.tensor(1) dist.all_reduce(completed) assert completed.item() == 2 - finally: - dist.destroy_process_group() def test_divergent_optimizer_provenance_rejects_collectively_before_mutation( tmp_path: Path, ) -> None: - processes = mp.spawn( + spawn_and_join( _divergent_version_worker, args=(f"file://{tmp_path / 'versions'}",), - nprocs=2, - join=False, + timeout=90, + failure="Divergent version preflight did not complete collectively", ) - deadline = time.monotonic() + 90 - while time.monotonic() < deadline: - if processes.join(timeout=1): - return - for process in processes.processes: - process.terminate() - pytest.fail("Divergent version preflight did not complete collectively") def test_transaction_coalesces_many_children_but_keeps_each_origin() -> None: @@ -455,14 +437,7 @@ def __setattr__(self, name: str, value: object) -> None: def _transaction_exit_failure_worker(rank: int, rendezvous: str, nested: bool) -> None: - dist.init_process_group( - "gloo", - init_method=rendezvous, - rank=rank, - world_size=2, - timeout=timedelta(seconds=15), - ) - try: + with gloo_group(rank, rendezvous, timeout=15): trainer, current = _trainer() version = trainer._capture_checkpoint_version("student") trainer._commit_versioned_gradients( @@ -515,8 +490,6 @@ def coordinate(validate) -> None: completed = torch.tensor(1) dist.all_reduce(completed) assert completed.item() == 2 - finally: - dist.destroy_process_group() @pytest.mark.parametrize("nested", [False, True]) @@ -524,16 +497,9 @@ def test_local_replay_failure_uses_same_collective_exit_phase_on_every_rank( tmp_path: Path, nested: bool, ) -> None: - processes = mp.spawn( + spawn_and_join( _transaction_exit_failure_worker, args=(f"file://{tmp_path / 'exit'}", nested), - nprocs=2, - join=False, + timeout=60, + failure="Transaction success/failure exit phases did not match", ) - deadline = time.monotonic() + 60 - while time.monotonic() < deadline: - if processes.join(timeout=1): - return - for process in processes.processes: - process.terminate() - pytest.fail("Transaction success/failure exit phases did not match") diff --git a/tests/unit/trainer_rank_test_support.py b/tests/unit/trainer_rank_test_support.py new file mode 100644 index 000000000..97911ba17 --- /dev/null +++ b/tests/unit/trainer_rank_test_support.py @@ -0,0 +1,74 @@ +"""Real Gloo groups and lightweight topology for trainer-rank contract tests.""" + +from contextlib import contextmanager +from datetime import timedelta +import sys +import time +from types import ModuleType, SimpleNamespace +from unittest.mock import patch + +import pytest +import torch.distributed as dist +import torch.multiprocessing as mp + + +@contextmanager +def gloo_group(rank, rendezvous, *, world_size=2, timeout=30): + dist.init_process_group( + "gloo", + init_method=rendezvous, + rank=rank, + world_size=world_size, + timeout=timedelta(seconds=timeout), + ) + try: + yield + finally: + dist.destroy_process_group() + + +@contextmanager +def megatron_topology(physical, *, dp_size, tp_size): + """Install just the callback topology, with each real TP group created in order.""" + assert dist.get_world_size() == dp_size * tp_size + groups = ( + [dist.group.WORLD] + if dp_size == 1 + else [ + dist.new_group(list(range(dp * tp_size, (dp + 1) * tp_size))) + for dp in range(dp_size) + ] + ) + dp, tp = divmod(physical, tp_size) + megatron, core = ModuleType("megatron"), ModuleType("megatron.core") + setattr( + core, + "parallel_state", + SimpleNamespace( + get_tensor_model_parallel_rank=lambda: tp, + get_context_parallel_rank=lambda: 0, + get_data_parallel_rank=lambda: dp, + get_data_parallel_world_size=lambda: dp_size, + get_tensor_and_context_parallel_group=lambda **kwargs: groups[dp], + ), + ) + setattr(megatron, "core", core) + with patch.dict(sys.modules, {"megatron": megatron, "megatron.core": core}): + yield + + +def spawn_and_join(worker, args, *, timeout, failure, nprocs=2): + """Bound a collective test while preserving spawned-worker tracebacks.""" + processes = mp.spawn(worker, args=args, nprocs=nprocs, join=False) + try: + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + if processes.join(timeout=1): + return + pytest.fail(failure) + finally: + for process in processes.processes: + if process.is_alive(): + process.terminate() + for process in processes.processes: + process.join(timeout=5) From 5a8e49111cd59d436671a134b3c69c2faf33d642 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 19 Sep 2026 20:24:16 +0000 Subject: [PATCH 009/150] Fix logical head completion and lookup boundaries --- src/art/trainer_rank/_commands.py | 32 ++++- src/art/trainer_rank/_heads.py | 25 +++- .../unit/test_trainer_rank_head_boundaries.py | 123 ++++++++++++++++++ .../test_trainer_rank_release_completion.py | 17 ++- 4 files changed, 186 insertions(+), 11 deletions(-) create mode 100644 tests/unit/test_trainer_rank_head_boundaries.py diff --git a/src/art/trainer_rank/_commands.py b/src/art/trainer_rank/_commands.py index e2bda3a8d..433966d95 100644 --- a/src/art/trainer_rank/_commands.py +++ b/src/art/trainer_rank/_commands.py @@ -156,7 +156,23 @@ def _start_release(self) -> None: raise RuntimeError(self.state.release_error) if self.state.pending_release is not None: raise RuntimeError("Previous callback release is still pending") - pending = tuple(self.state.released) + synchronize_heads = ( + self.stopped + and self.mode == "rank" + and self.rank._dp_rank_and_size()[1] > 1 + and any( + custom.kind == "buffer" + or ( + custom.kind == "module" + and next(cast(torch.nn.Module, custom.value).buffers(), None) + is not None + ) + for slot in getattr(self.rank, "_checkpoint_slots", {}).values() + for custom in slot.custom.values() + ) + ) + # Piggyback presence so empty replicas still join a needed reconciliation. + pending = (tuple(self.state.released), synchronize_heads) gathered: list[Any] = [pending] loop = asyncio.get_running_loop() if self.distributed: @@ -187,10 +203,16 @@ def _finish_release(self, release: _Release) -> None: self.state.pending_release = None try: release.completed.result() - handles = {handle for values in release.gathered for handle in values} + handles = {handle for values, _ in release.gathered for handle in values} for handle in handles: self.state.graphs.pop(handle, None) self.state.released.difference_update(handles) + if any(synchronize for _, synchronize in release.gathered): + from ._heads import synchronize_head_buffers + + # Every DP session has stopped. Unequal callbacks/yields cannot + # enter this WORLD collective early, including after failure. + synchronize_head_buffers(self.rank) except BaseException as error: self.state.release_error = ( f"Callback release reconciliation failed: {error}" @@ -512,7 +534,11 @@ def _dispatch(self, command: _Command) -> Any: if self.mode == "zero" and args[0] == "head_export": synchronize_head_buffers(self.rank) return execute_head_operation( - self.rank, *args, coordinate=self._coordinated_preflight, **kwargs + self.rank, + *args, + coordinate=self._coordinated_preflight, + local_lookup=self.mode == "rank", + **kwargs, ) if op == "reduce_value": tensor = args[0].to(self.rank.device) diff --git a/src/art/trainer_rank/_heads.py b/src/art/trainer_rank/_heads.py index d4f725d04..0c859760e 100644 --- a/src/art/trainer_rank/_heads.py +++ b/src/art/trainer_rank/_heads.py @@ -587,13 +587,30 @@ def execute_head_operation( payload: Any, *, coordinate: Callable[[Callable[[], None]], None] | None = None, + local_lookup: bool = False, ) -> Any: """Execute on every physical rank inside the owning trainer operation queue.""" if hasattr(trainer, "_rank"): return trainer._invoke("head", kind, payload) if kind == "head_lookup": + from ._impl import Unset + checkpoint, name = payload - checkpoint = trainer._resolve_custom_checkpoint(checkpoint) + if local_lookup and checkpoint is Unset: + ref = ( + trainer._slot_stack[-1] + if trainer._slot_stack + else trainer._default_slot_ref + ) + checkpoint = None if ref is None else ref.name + # Reopening a loaded head is DP-local; lazy loading is still global. + if not local_lookup or checkpoint is None: + checkpoint = trainer._resolve_custom_checkpoint(checkpoint) + elif checkpoint not in trainer._checkpoint_slots: + raise trainer._slot_state_error( + f"Load checkpoint {checkpoint!r} across all ranks before a logical-DP head lookup" + ) + assert checkpoint is not None return checkpoint, export_head( trainer, checkpoint, name ) if name in trainer._checkpoint_slots[checkpoint].custom else None @@ -1319,15 +1336,13 @@ def logical_register_head( flush_logical_heads(view) checkpoint, state = view._invoke("head", "head_lookup", (checkpoint, name)) + if state is not None and state.kind != kind: + raise ValueError(f"Checkpoint object {name!r} is already a {state.kind}") registry = _logical_heads(view) current = registry.get((checkpoint, name)) if current is not None and not current.invalid and state is not None: current.refresh(state) if not current.invalid: - if state.kind != kind: - raise ValueError( - f"Checkpoint object {name!r} is already a {state.kind}" - ) return current.value if state is None: value = factory() diff --git a/tests/unit/test_trainer_rank_head_boundaries.py b/tests/unit/test_trainer_rank_head_boundaries.py new file mode 100644 index 000000000..1abba88f6 --- /dev/null +++ b/tests/unit/test_trainer_rank_head_boundaries.py @@ -0,0 +1,123 @@ +"""Logical head registration enforces the native object kind contract.""" + +import asyncio + +import pytest +from test_trainer_rank_custom_tensors import _trainer +import torch + +from art.trainer_rank import ModuleHandle, run_rank_callback + + +@pytest.mark.parametrize("mode", ("rank", "zero")) +@pytest.mark.parametrize("cached", (False, True)) +@pytest.mark.parametrize("kind", ("buffer", "parameter", "module")) +def test_logical_registration_checks_kind_before_cached_or_native_lookup( + mode, cached, kind +): + native, api = _trainer("student") + factory = torch.nn.Identity if kind == "module" else lambda: torch.tensor(2.0) + original = getattr(api, kind)("head", factory, checkpoint="student") + calls = [] + + def unexpected_factory(): + calls.append(True) + raise AssertionError("existing head must not call its factory") + + def callback(view): + if cached: + getattr(view, kind)("head", unexpected_factory, checkpoint="student") + for wrong in {"buffer", "parameter", "module"} - {kind}: + with pytest.raises(ValueError, match=f"already a {kind}"): + getattr(view, wrong)("head", unexpected_factory, checkpoint="student") + reopened = getattr(view, kind)("head", unexpected_factory, checkpoint="student") + assert isinstance(reopened, ModuleHandle if kind == "module" else torch.Tensor) + if kind != "module": + assert reopened.item() == original.item() == 2 + assert reopened.requires_grad == (kind == "parameter") + assert ( + getattr(view, kind)("head", unexpected_factory, checkpoint="student") + is reopened + ) + + asyncio.run(run_rank_callback(native, callback, mode=mode)) + assert not calls + + +@pytest.mark.parametrize("selection", ("explicit", "default", "pushed")) +def test_loaded_logical_lookup_resolves_locally(monkeypatch, selection): + native, api = _trainer("student", "teacher") + for name, value in (("student", 2.0), ("teacher", 3.0)): + api.buffer("head", lambda: torch.tensor(value), checkpoint=name) + native._default_slot_ref = native._slot_ref("student") + if selection == "pushed": + native._slot_stack.append(native._slot_ref("teacher")) + + def unexpected_load(_names): + raise AssertionError( + "lookup of a loaded slot must not coordinate global loading" + ) + + monkeypatch.setattr(native, "_ensure_checkpoint_slots", unexpected_load) + + def callback(view): + options = {"checkpoint": "student"} if selection == "explicit" else {} + head = view.buffer("head", lambda: None, **options) + assert head.item() == (3 if selection == "pushed" else 2) + + asyncio.run(run_rank_callback(native, callback, mode="rank")) + + +@pytest.mark.parametrize("mode", ("rank", "zero")) +@pytest.mark.parametrize("prefetched", (False, True)) +def test_unloaded_logical_lookup_requires_global_loading(monkeypatch, mode, prefetched): + native, api = _trainer("pending") + api.buffer("head", lambda: torch.tensor(2.0), checkpoint="pending") + pending = native._checkpoint_slots.pop("pending") + loaded = [] + if prefetched: + native._checkpoint_prefetch_sources["pending"] = "/test/prefetched" + + def load(name): + loaded.append(name) + native._checkpoint_slots[name] = pending + + monkeypatch.setattr(native, "_load_registered_checkpoint", load) + + def callback(view): + if mode == "zero" and prefetched: + assert view.buffer("head", lambda: None, checkpoint="pending").item() == 2 + else: + message = ( + "Load checkpoint .* across all ranks" + if mode == "rank" + else "unloaded checkpoint" + ) + with pytest.raises(RuntimeError, match=message): + view.buffer("head", lambda: None, checkpoint="pending") + with pytest.raises(RuntimeError, match="require a loaded named checkpoint"): + view.buffer("head", lambda: None, checkpoint=None) + + asyncio.run(run_rank_callback(native, callback, mode=mode)) + assert loaded == (["pending"] if mode == "zero" and prefetched else []) + + +@pytest.mark.parametrize("mode", ("rank", "zero")) +@pytest.mark.parametrize("dp_size", (1, 2)) +@pytest.mark.parametrize("kind", (None, "parameter", "module", "buffer")) +def test_only_multi_dp_callbacks_with_buffers_add_reconciliation( + monkeypatch, mode, dp_size, kind +): + native, api = _trainer("student") + monkeypatch.setattr(native, "_dp_rank_and_size", lambda: (0, dp_size)) + if kind is not None: + factory = torch.nn.Identity if kind == "module" else lambda: torch.tensor(2.0) + getattr(api, kind)("head", factory, checkpoint="student") + synchronized = [] + monkeypatch.setattr( + "art.trainer_rank._heads.synchronize_head_buffers", synchronized.append + ) + asyncio.run(run_rank_callback(native, lambda _: None, mode=mode)) + assert synchronized == ( + [native] if mode == "rank" and dp_size > 1 and kind == "buffer" else [] + ) diff --git a/tests/unit/test_trainer_rank_release_completion.py b/tests/unit/test_trainer_rank_release_completion.py index 81d765bd4..3b2a52689 100644 --- a/tests/unit/test_trainer_rank_release_completion.py +++ b/tests/unit/test_trainer_rank_release_completion.py @@ -9,7 +9,15 @@ from art.trainer_rank._commands import _Executor, _Release -def test_completed_release_is_finalized_before_queued_done_callback(): +@pytest.mark.parametrize("peer_buffers", [False, True]) +def test_completed_release_is_finalized_before_queued_done_callback( + monkeypatch, peer_buffers +): + synchronized = [] + monkeypatch.setattr( + "art.trainer_rank._heads.synchronize_head_buffers", synchronized.append + ) + async def run(): rank: Any = _Rank() executor = _Executor(rank, "zero") @@ -17,7 +25,9 @@ async def run(): state.graphs["zero:old:dp:0"] = (rank.weight,) state.released.add("zero:old:dp:0") completed = asyncio.get_running_loop().create_future() - release = state.pending_release = _Release(completed, [tuple(state.released)]) + release = state.pending_release = _Release( + completed, [(tuple(state.released), False), ((), peer_buffers)] + ) completed.set_result(None) completed.add_done_callback(lambda _: executor._finish_release(release)) await executor._join_release() @@ -27,6 +37,7 @@ async def run(): state.graphs["zero:new:dp:0"] = (rank.weight,) await asyncio.sleep(0) assert tuple(state.graphs) == ("zero:new:dp:0",) + assert synchronized == ([rank] if peer_buffers else []) asyncio.run(run()) @@ -40,7 +51,7 @@ async def run(): reports = [] loop.set_exception_handler(lambda _loop, context: reports.append(context)) completed = loop.create_future() - release = executor.state.pending_release = _Release(completed, [()]) + release = executor.state.pending_release = _Release(completed, [((), False)]) completed.add_done_callback(lambda _: executor._finish_release(release)) if cancelled: completed.cancel() From 4c89eb408e03d82f8440cf1bffb33b8de8c8c046 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 19 Sep 2026 20:53:11 +0000 Subject: [PATCH 010/150] Expose callback release fence for checkpoint entry --- src/art/trainer_rank/_commands.py | 51 ++++++++++++------- .../test_trainer_rank_release_completion.py | 38 +++++++++++--- 2 files changed, 63 insertions(+), 26 deletions(-) diff --git a/src/art/trainer_rank/_commands.py b/src/art/trainer_rank/_commands.py index 433966d95..1afb121a0 100644 --- a/src/art/trainer_rank/_commands.py +++ b/src/art/trainer_rank/_commands.py @@ -73,6 +73,7 @@ class _OutputPacket: class _Release: completed: asyncio.Future[None] gathered: list[Any] + finish: Callable[[_Release], None] @dataclass @@ -111,6 +112,33 @@ def rank_callback_leader(rank: TrainerRank, *, mode: Mode = "rank") -> bool: ) +async def join_rank_callback_release( + rank: TrainerRank, +) -> asyncio.CancelledError | None: + """Settle prior callback cleanup, returning any deferred cancellation.""" + state: _State | None = getattr(rank, "_rank_command_state", None) + if state is None: + return None + cancelled = None + if release := state.pending_release: + while True: + try: + await asyncio.shield(release.completed) + break + except asyncio.CancelledError as error: + if release.completed.cancelled(): + break + cancelled = error + except Exception: + break # Finalization records and reports the terminal error. + # A done callback may still be queued; ownership must be settled + # synchronously before this actor can execute another command. + release.finish(release) + if state.release_error is not None: + raise RuntimeError(state.release_error) + return cancelled + + class _Executor: def __init__(self, rank: TrainerRank, mode: Mode) -> None: from ._tensors import CotangentCollector @@ -192,7 +220,9 @@ def _start_release(self) -> None: else: completed = loop.create_future() completed.set_result(None) - release = self.state.pending_release = _Release(completed, gathered) + release = self.state.pending_release = _Release( + completed, gathered, self._finish_release + ) # The callback's exception can reach its controller before every DP # sibling exits. Keep ownership until their matching cleanup completes. completed.add_done_callback(lambda _: self._finish_release(release)) @@ -222,24 +252,7 @@ def _finish_release(self, release: _Release) -> None: ) async def _join_release(self) -> asyncio.CancelledError | None: - cancelled = None - if release := self.state.pending_release: - while True: - try: - await asyncio.shield(release.completed) - break - except asyncio.CancelledError as error: - if release.completed.cancelled(): - break - cancelled = error - except Exception: - break # Finalization records and reports the terminal error. - # A done callback may still be queued; ownership must be settled - # synchronously before this actor can execute another command. - self._finish_release(release) - if self.state.release_error is not None: - raise RuntimeError(self.state.release_error) - return cancelled + return await join_rank_callback_release(self.rank) async def reconcile_releases( self, *, defer_cancellation: bool = False diff --git a/tests/unit/test_trainer_rank_release_completion.py b/tests/unit/test_trainer_rank_release_completion.py index 3b2a52689..234dbe02b 100644 --- a/tests/unit/test_trainer_rank_release_completion.py +++ b/tests/unit/test_trainer_rank_release_completion.py @@ -1,12 +1,12 @@ """Pending callback cleanup remains observable and ordered after failure.""" import asyncio -from typing import Any +from typing import Any, cast import pytest from test_trainer_rank_commands import _Rank -from art.trainer_rank._commands import _Executor, _Release +from art.trainer_rank._commands import _Executor, _Release, join_rank_callback_release @pytest.mark.parametrize("peer_buffers", [False, True]) @@ -19,14 +19,16 @@ def test_completed_release_is_finalized_before_queued_done_callback( ) async def run(): - rank: Any = _Rank() + rank = cast(Any, _Rank()) executor = _Executor(rank, "zero") state = executor.state state.graphs["zero:old:dp:0"] = (rank.weight,) state.released.add("zero:old:dp:0") completed = asyncio.get_running_loop().create_future() release = state.pending_release = _Release( - completed, [(tuple(state.released), False), ((), peer_buffers)] + completed, + [(tuple(state.released), False), ((), peer_buffers)], + executor._finish_release, ) completed.set_result(None) completed.add_done_callback(lambda _: executor._finish_release(release)) @@ -45,13 +47,15 @@ async def run(): @pytest.mark.parametrize("cancelled", [False, True]) def test_background_cleanup_failure_is_reported_and_blocks_next_entry(cancelled): async def run(): - rank: Any = _Rank() + rank = cast(Any, _Rank()) executor = _Executor(rank, "zero") loop = asyncio.get_running_loop() reports = [] loop.set_exception_handler(lambda _loop, context: reports.append(context)) completed = loop.create_future() - release = executor.state.pending_release = _Release(completed, [((), False)]) + release = executor.state.pending_release = _Release( + completed, [((), False)], executor._finish_release + ) completed.add_done_callback(lambda _: executor._finish_release(release)) if cancelled: completed.cancel() @@ -73,7 +77,7 @@ async def run(): def test_success_in_unrelated_exception_handler_still_awaits_cleanup(monkeypatch): async def run(): - rank: Any = _Rank() + rank = cast(Any, _Rank()) executor = _Executor(rank, "zero") started, finish = asyncio.Event(), asyncio.Event() @@ -97,3 +101,23 @@ async def callback(): await pending asyncio.run(run()) + + +def test_checkpoint_fence_defers_cancellation_without_cancelling_release(): + async def run(): + rank = cast(Any, _Rank()) + executor = _Executor(rank, "zero") + completed = asyncio.get_running_loop().create_future() + executor.state.pending_release = _Release( + completed, [((), False)], executor._finish_release + ) + pending = asyncio.create_task(join_rank_callback_release(rank)) + await asyncio.sleep(0) + pending.cancel() + await asyncio.sleep(0) + assert not pending.done() and not completed.cancelled() + completed.set_result(None) + assert isinstance(await pending, asyncio.CancelledError) + assert executor.state.pending_release is None + + asyncio.run(run()) From 04334c30369e8a38edebd7a2dd37b7efe45ae275 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 19 Sep 2026 20:58:08 +0000 Subject: [PATCH 011/150] Separate abandoned batch pull cleanup from iterator close --- src/art/trainer_rank/_operations.py | 3 --- tests/unit/test_trainer_operations.py | 8 ++++++-- 2 files changed, 6 insertions(+), 5 deletions(-) diff --git a/src/art/trainer_rank/_operations.py b/src/art/trainer_rank/_operations.py index 5af7df142..4ad76edbd 100644 --- a/src/art/trainer_rank/_operations.py +++ b/src/art/trainer_rank/_operations.py @@ -206,9 +206,6 @@ async def execute_operation(rank_zero: Any, operation: TrainerOperation) -> Any: result = rank_zero.open_forward_batches(**payload) elif operation.kind == "batches_close": result = rank_zero.close_forward_batches(payload["handle"]) - pending = ledger.outcomes.get(payload.get("pending_operation")) - if pending is not None: - await _abandon(rank_zero, pending) elif operation.kind.startswith("head_"): from ._heads import execute_head_operation diff --git a/tests/unit/test_trainer_operations.py b/tests/unit/test_trainer_operations.py index 2b31eccdd..a6ca3d5f5 100644 --- a/tests/unit/test_trainer_operations.py +++ b/tests/unit/test_trainer_operations.py @@ -278,18 +278,22 @@ async def run(): await execute_operation( rank, TrainerOperation.capture(("client", 3), "acknowledge", ((2, 3), ())) ) + # A lost pull is abandoned independently of closing its iterator. + await execute_operation( + rank, TrainerOperation.capture(("client", 3), "acknowledge", ((3,), (2,))) + ) close = TrainerOperation.capture( ("client", 3), "batches_close", - {"handle": "iterator", "pending_operation": next_wave.id}, + {"handle": "iterator"}, ) await execute_operation(rank, close) await execute_operation(rank, close) assert events == [ "open", "next", - ("close", "iterator"), ("release", ("packet",)), + ("close", "iterator"), ] await execute_operation( rank, TrainerOperation.capture(close.id, "acknowledge", ((), ())) From bbb81f36a1460afc6c8fd9c85323614e5f7f34ad Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 20 Sep 2026 01:02:01 +0000 Subject: [PATCH 012/150] Deduplicate live-head test construction --- tests/unit/test_trainer_rank_live_heads.py | 211 +++++---------------- 1 file changed, 43 insertions(+), 168 deletions(-) diff --git a/tests/unit/test_trainer_rank_live_heads.py b/tests/unit/test_trainer_rank_live_heads.py index dd6eb4006..712a220bf 100644 --- a/tests/unit/test_trainer_rank_live_heads.py +++ b/tests/unit/test_trainer_rank_live_heads.py @@ -46,6 +46,16 @@ def compute(x): ) +def _live_head( + trainer, name, source, collector=None, *, checkpoint="student" +) -> LiveHead: + return LiveHead( + export_head(trainer, checkpoint, name), + source, + CotangentCollector() if collector is None else collector, + ) + + def _module(live: LiveHead) -> ModuleHandle: assert isinstance(live.value, ModuleHandle) return live.value @@ -136,7 +146,7 @@ def test_client_tied_head_old_backward_after_native_refresh(monkeypatch): trainer, rank = _trainer("student") native = rank.module("head", TiedHead, checkpoint="student") collector = CotangentCollector() - live = LiveHead(export_head(trainer, "student", "head"), TiedHead(), collector) + live = _live_head(trainer, "head", TiedHead(), collector) old_input = torch.tensor(3.0, requires_grad=True) old = _module(live)(old_input) with trainer._gradient_transaction(): @@ -156,9 +166,7 @@ def test_live_parameter_reuses_handle_and_earlier_capture_retains_version(): trainer, rank = _trainer("student") parameter = rank.parameter("gain", lambda: torch.tensor(2.0), checkpoint="student") collector = CotangentCollector() - live = LiveHead( - export_head(trainer, "student", "gain"), torch.tensor(2.0), collector - ) + live = _live_head(trainer, "gain", torch.tensor(2.0), collector) handle = _tensor(live) old = handle.square() * 3 parameter.data.fill_(4) @@ -174,11 +182,7 @@ def test_live_parameter_reuses_handle_and_earlier_capture_retains_version(): def test_client_buffer_publication_conflict_is_atomic(): trainer, rank = _trainer("student") native = rank.module("bn", lambda: torch.nn.BatchNorm1d(2), checkpoint="student") - live = LiveHead( - export_head(trainer, "student", "bn"), - torch.nn.BatchNorm1d(2), - CotangentCollector(), - ) + live = _live_head(trainer, "bn", torch.nn.BatchNorm1d(2)) _module(live)(torch.ones(4, 2)) update = live.take_publication() assert update is not None @@ -340,11 +344,7 @@ def test_cuda_checkpoint_head_across_completed_optimizer_update(monkeypatch, cli trainer.device = torch.device("cuda", 0) native = rank.module("head", lambda: TiedHead(False), checkpoint="student") collector = CotangentCollector() - live = ( - LiveHead(export_head(trainer, "student", "head"), TiedHead(False), collector) - if client - else None - ) + live = _live_head(trainer, "head", TiedHead(False), collector) if client else None head = native if live is None else _module(live) original_input = torch.tensor(3.0, device="cuda", requires_grad=True) old = head(managed_tensor(original_input) if client else original_input) @@ -366,9 +366,7 @@ def test_client_buffer_reads_keep_old_graph_and_replacement_invalidates_handle() trainer, rank = _trainer("student") native = rank.buffer("scale", lambda: torch.tensor(2.0), checkpoint="student") collector = CotangentCollector() - live = LiveHead( - export_head(trainer, "student", "scale"), torch.tensor(2.0), collector - ) + live = _live_head(trainer, "scale", torch.tensor(2.0), collector) x = torch.tensor(3.0, requires_grad=True) old = x * _tensor(live) native.fill_(4) @@ -388,9 +386,7 @@ def test_client_buffer_reads_keep_old_graph_and_replacement_invalidates_handle() def test_client_module_explicit_dtype_move_retains_ties_and_live_parameters(): trainer, rank = _trainer("student") native = rank.module("head", TiedHead, checkpoint="student") - live = LiveHead( - export_head(trainer, "student", "head"), TiedHead(), CotangentCollector() - ) + live = _live_head(trainer, "head", TiedHead()) head = _module(live).to("cpu").to(dtype=torch.float64) assert head.left.dtype == torch.float64 assert head.offset.dtype == torch.float64 @@ -412,7 +408,7 @@ def test_client_explicit_snapshot_supports_nonreentrant_checkpoint(): trainer, rank = _trainer("student") native = rank.module("head", TiedHead, checkpoint="student") collector = CotangentCollector() - live = LiveHead(export_head(trainer, "student", "head"), TiedHead(), collector) + live = _live_head(trainer, "head", TiedHead(), collector) captured = _module(live).snapshot() x = torch.tensor(3.0, requires_grad=True) old = checkpoint(captured, x, use_reentrant=False) @@ -429,15 +425,7 @@ def test_client_explicit_snapshot_supports_nonreentrant_checkpoint(): def test_buffer_item_and_bitwise_mutations_publish_without_losing_handle(client): trainer, rank = _trainer("student") native = rank.buffer("mask", lambda: torch.tensor([1, 2]), checkpoint="student") - live = ( - LiveHead( - export_head(trainer, "student", "mask"), - torch.tensor([1, 2]), - CotangentCollector(), - ) - if client - else None - ) + live = _live_head(trainer, "mask", torch.tensor([1, 2])) if client else None value = native if live is None else _tensor(live) value[0] = 4 original = value @@ -463,11 +451,7 @@ def test_buffer_item_and_bitwise_mutations_publish_without_losing_handle(client) def test_client_parameter_mutations_fail_before_changing_owned_values(mutation): trainer, rank = _trainer("student") rank.parameter("gain", lambda: torch.tensor([2.0]), checkpoint="student") - live = LiveHead( - export_head(trainer, "student", "gain"), - torch.tensor([2.0]), - CotangentCollector(), - ) + live = _live_head(trainer, "gain", torch.tensor([2.0])) with pytest.raises(RuntimeError, match="checkpoint parameters"): mutation(_tensor(live)) torch.testing.assert_close(_tensor(live).detach(), torch.tensor([2.0])) @@ -480,9 +464,7 @@ def test_head_export_preserves_strict_constructor_policy_for_client(): setattr(trainer, "_forward_options", ForwardOptions(max_gradient_staleness=0)) rank.parameter("gain", lambda: torch.tensor(2.0), checkpoint="student") collector = CotangentCollector() - live = LiveHead( - export_head(trainer, "student", "gain"), torch.tensor(2.0), collector - ) + live = _live_head(trainer, "gain", torch.tensor(2.0), collector) old = _tensor(live).square() trainer._checkpoint_slots["student"].revision += 1 with pytest.raises(RuntimeError, match="staleness"): @@ -509,7 +491,7 @@ def test_registration_under_no_grad_preserves_authoritative_trainability(): native = rank.module("head", TiedHead, checkpoint="student") collector = CotangentCollector() with torch.no_grad(): - live = LiveHead(export_head(trainer, "student", "head"), TiedHead(), collector) + live = _live_head(trainer, "head", TiedHead(), collector) assert _module(live).left.requires_grad packets = collector.backward(_module(live)(torch.tensor(3.0))) trainer._commit_versioned_gradients(head_gradient_targets(trainer, packets[0])) @@ -523,11 +505,7 @@ def test_client_reentrant_checkpoint_rejects_nested_remote_bridges(external): trainer, rank = _trainer("student") native = rank.module("head", TiedHead, checkpoint="student") collector = CotangentCollector() - live = LiveHead( - export_head(trainer, "student", "head"), - TiedHead(None if external else True), - collector, - ) + live = _live_head(trainer, "head", TiedHead(None if external else True), collector) x = torch.tensor(3.0, requires_grad=True) loss = ( checkpoint(_module(live).snapshot(), x, use_reentrant=True) @@ -547,9 +525,7 @@ def test_cuda_batchnorm_explicit_placement_publishes_buffers_and_preserves_old_g trainer.device = torch.device("cuda", 0) native = rank.module("bn", lambda: torch.nn.BatchNorm1d(2), checkpoint="student") collector = CotangentCollector() - live = LiveHead( - export_head(trainer, "student", "bn"), torch.nn.BatchNorm1d(2), collector - ) + live = _live_head(trainer, "bn", torch.nn.BatchNorm1d(2), collector) head = _module(live).to("cuda") head.eval() x = torch.tensor([[1.0, 3.0], [2.0, 5.0]], device="cuda", requires_grad=True) @@ -595,8 +571,8 @@ def test_client_reused_module_factory_does_not_share_checkpoint_handles(): rank.module("head", TiedHead, checkpoint="A") rank.module("head", TiedHead, checkpoint="B") source = TiedHead() - first = LiveHead(export_head(trainer, "A", "head"), source, CotangentCollector()) - second = LiveHead(export_head(trainer, "B", "head"), source, CotangentCollector()) + first = _live_head(trainer, "head", source, checkpoint="A") + second = _live_head(trainer, "head", source, checkpoint="B") _module(first).offset.add_(2) assert _module(first)(torch.tensor(3.0)).item() == 17 assert _module(second)(torch.tensor(3.0)).item() == 15 @@ -615,9 +591,7 @@ def test_managed_model_operand_captures_live_parameter_and_keeps_old_version(ope "weight", lambda: torch.tensor([2.0, 4.0]), checkpoint="student" ) collector = CotangentCollector() - live = LiveHead( - export_head(trainer, "student", "weight"), torch.zeros(2), collector - ) + live = _live_head(trainer, "weight", torch.zeros(2), collector) hidden = collector.attach( detach_tree("model", torch.tensor([3.0, 5.0], requires_grad=True)), managed=True ) @@ -655,13 +629,7 @@ def forward(self, value): trainer, rank = _trainer("student") native = rank.module("head", Counter, checkpoint="student") - live = ( - LiveHead( - export_head(trainer, "student", "head"), Counter(), CotangentCollector() - ) - if client - else None - ) + live = _live_head(trainer, "head", Counter()) if client else None head = native if live is None else _module(live) assert head(torch.tensor(0)).item() == 1 assert head(torch.tensor(0)).item() == 2 @@ -689,13 +657,7 @@ def forward(self, value): trainer, rank = _trainer("student") native = rank.module("head", Resize, checkpoint="student") - live = ( - LiveHead( - export_head(trainer, "student", "head"), Resize(), CotangentCollector() - ) - if client - else None - ) + live = _live_head(trainer, "head", Resize()) if client else None head = native if live is None else _module(live) before = export_head(trainer, "student", "head").buffer_revision with pytest.raises(ValueError, match="preserve buffer shape"): @@ -711,15 +673,7 @@ def forward(self, value): def test_failing_handle_forward_hook_does_not_publish_buffers(client): trainer, rank = _trainer("student") native = rank.module("head", lambda: torch.nn.BatchNorm1d(2), checkpoint="student") - live = ( - LiveHead( - export_head(trainer, "student", "head"), - torch.nn.BatchNorm1d(2), - CotangentCollector(), - ) - if client - else None - ) + live = _live_head(trainer, "head", torch.nn.BatchNorm1d(2)) if client else None head = native if live is None else _module(live) def fail(*args): @@ -765,13 +719,7 @@ def forward(self, value): trainer, rank = _trainer("student") native = rank.module("head", Counter, checkpoint="student") - live = ( - LiveHead( - export_head(trainer, "student", "head"), Counter(), CotangentCollector() - ) - if client - else None - ) + live = _live_head(trainer, "head", Counter()) if client else None head = native if live is None else _module(live) initial = 7 if pending_before else 0 if pending_before: @@ -819,11 +767,7 @@ def test_handle_hook_parameter_gradients_keep_original_version(client): trainer, rank = _trainer("student") native = rank.module("head", TiedHead, checkpoint="student") collector = CotangentCollector() - live = ( - LiveHead(export_head(trainer, "student", "head"), TiedHead(), collector) - if client - else None - ) + live = _live_head(trainer, "head", TiedHead(), collector) if client else None head = native if live is None else _module(live) parameter = head.left head.register_forward_hook(lambda module, args, out: out + module.left.square()) @@ -848,13 +792,7 @@ def test_handle_hook_parameter_gradients_keep_original_version(client): def test_recursive_handle_hook_does_not_publish_before_outer_failure(client): trainer, rank = _trainer("student") native = rank.module("head", TiedHead, checkpoint="student") - live = ( - LiveHead( - export_head(trainer, "student", "head"), TiedHead(), CotangentCollector() - ) - if client - else None - ) + live = _live_head(trainer, "head", TiedHead()) if client else None head = native if live is None else _module(live) def recurse(module, args): @@ -885,10 +823,7 @@ def test_nested_handle_calls_preserve_active_captures(client, scenario): } collector = CotangentCollector() live = ( - { - name: LiveHead(export_head(trainer, "student", name), TiedHead(), collector) - for name in native - } + {name: _live_head(trainer, name, TiedHead(), collector) for name in native} if client else {} ) @@ -993,11 +928,7 @@ def _live_buffer_authority_worker(process_rank, init_method, asymmetric=False): synchronize_head_buffers(trainer) dist.barrier() return - live = LiveHead( - export_head(trainer, "student", "head"), - torch.nn.BatchNorm1d(2), - CotangentCollector(), - ) + live = _live_head(trainer, "head", torch.nn.BatchNorm1d(2)) if process_rank == 1: _module(live)(torch.ones(4, 2)) update = live.take_publication() @@ -1040,7 +971,7 @@ def test_inplace_operation_snapshots_readonly_client_tensor(kind): factory = lambda: torch.tensor(2.0) native = getattr(rank, kind)("scale", factory, checkpoint="student") collector = CotangentCollector() - live = LiveHead(export_head(trainer, "student", "scale"), factory(), collector) + live = _live_head(trainer, "scale", factory(), collector) x = torch.tensor(3.0, requires_grad=True) loss = (x * 1).mul_(_tensor(live)) with torch.no_grad(): @@ -1170,15 +1101,7 @@ def test_live_buffer_view_mutation_rejects_without_silent_write( ): trainer, rank = _trainer("student") native = rank.buffer("stats", lambda: torch.zeros(4), checkpoint="student") - live = ( - LiveHead( - export_head(trainer, "student", "stats"), - torch.zeros(4), - CotangentCollector(), - ) - if client - else None - ) + live = _live_head(trainer, "stats", torch.zeros(4)) if client else None buffer = native if live is None else _tensor(live) revision = export_head(trainer, "student", "stats").buffer_revision with torch.set_grad_enabled(grad_enabled): @@ -1202,15 +1125,7 @@ def test_live_buffer_view_mutation_rejects_without_silent_write( def test_functional_batchnorm_no_grad_publishes_buffer_changes(client): trainer, rank = _trainer("student") native = rank.buffer("mean", lambda: torch.zeros(2), checkpoint="student") - live = ( - LiveHead( - export_head(trainer, "student", "mean"), - torch.zeros(2), - CotangentCollector(), - ) - if client - else None - ) + live = _live_head(trainer, "mean", torch.zeros(2)) if client else None mean = native if live is None else _tensor(live) before = export_head(trainer, "student", "mean").buffer_revision with torch.no_grad(): @@ -1229,15 +1144,7 @@ def test_functional_batchnorm_no_grad_publishes_buffer_changes(client): def test_live_parameter_metadata_does_not_capture_weights(monkeypatch, client): trainer, rank = _trainer("student") native = rank.parameter("weight", lambda: torch.zeros(3, 4), checkpoint="student") - live = ( - LiveHead( - export_head(trainer, "student", "weight"), - torch.zeros(3, 4), - CotangentCollector(), - ) - if client - else None - ) + live = _live_head(trainer, "weight", torch.zeros(3, 4)) if client else None value = native if live is None else _tensor(live) def unexpected(*args, **kwargs): @@ -1272,11 +1179,7 @@ def test_parameter_tensor_properties_capture_immutable_versions( trainer, rank = _trainer("student") native = rank.parameter("weight", lambda: initial.clone(), checkpoint="student") collector = CotangentCollector() - live = ( - LiveHead(export_head(trainer, "student", "weight"), initial, collector) - if client - else None - ) + live = _live_head(trainer, "weight", initial, collector) if client else None value = native if live is None else _tensor(live) expected = initial.clone().requires_grad_() getattr(expected, property_name).abs().square().sum().backward() @@ -1314,13 +1217,7 @@ def test_buffer_tensor_properties_reject_unpublished_mutation(client, property_n initial = torch.ones(2, 2, dtype=torch.complex64) trainer, rank = _trainer("student") native = rank.buffer("stats", lambda: initial.clone(), checkpoint="student") - live = ( - LiveHead( - export_head(trainer, "student", "stats"), initial, CotangentCollector() - ) - if client - else None - ) + live = _live_head(trainer, "stats", initial) if client else None value = native if live is None else _tensor(live) with pytest.raises(RuntimeError, match="Views of live checkpoint buffers"): getattr(value, property_name).fill_(7) @@ -1346,15 +1243,7 @@ def test_distributed_buffer_registration_mismatch_fails_on_every_rank( def test_module_buffer_view_mutation_rejects_without_publishing(client): trainer, rank = _trainer("student") native = rank.module("head", lambda: torch.nn.BatchNorm1d(2), checkpoint="student") - live = ( - LiveHead( - export_head(trainer, "student", "head"), - torch.nn.BatchNorm1d(2), - CotangentCollector(), - ) - if client - else None - ) + live = _live_head(trainer, "head", torch.nn.BatchNorm1d(2)) if client else None head = native if live is None else _module(live) with pytest.raises( RuntimeError, match="Views of live checkpoint buffers are read-only" @@ -1369,15 +1258,7 @@ def test_module_buffer_view_mutation_rejects_without_publishing(client): def test_stateful_function_on_buffer_snapshot_view_rejects(client): trainer, rank = _trainer("student") native = rank.buffer("mean", lambda: torch.zeros(2), checkpoint="student") - live = ( - LiveHead( - export_head(trainer, "student", "mean"), - torch.zeros(2), - CotangentCollector(), - ) - if client - else None - ) + live = _live_head(trainer, "mean", torch.zeros(2)) if client else None mean = native if live is None else _tensor(live) with ( torch.no_grad(), @@ -1400,13 +1281,7 @@ def test_cuda_live_buffer_views_and_functional_publication(client): trainer.device = torch.device("cuda", 0) native = rank.buffer("mean", lambda: torch.zeros(2), checkpoint="student") live = ( - LiveHead( - export_head(trainer, "student", "mean"), - torch.zeros(2, device="cuda"), - CotangentCollector(), - ) - if client - else None + _live_head(trainer, "mean", torch.zeros(2, device="cuda")) if client else None ) mean = native if live is None else _tensor(live) with pytest.raises( From e92db78f74e913106cccfb513afde8e51e246c47 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 20 Sep 2026 01:03:35 +0000 Subject: [PATCH 013/150] Reuse tensor argument container traversal --- src/art/trainer_rank/_heads.py | 30 ++++++++++++++---------------- src/art/trainer_rank/_impl.py | 10 ++-------- 2 files changed, 16 insertions(+), 24 deletions(-) diff --git a/src/art/trainer_rank/_heads.py b/src/art/trainer_rank/_heads.py index 0c859760e..b02cc82fb 100644 --- a/src/art/trainer_rank/_heads.py +++ b/src/art/trainer_rank/_heads.py @@ -163,6 +163,18 @@ def tensor_mutation_targets( } +def _map_tensor_arguments(fn: Callable[[Any], Any], value: Any) -> Any: + """Map one argument container, preserving namedtuples and opaque values.""" + if isinstance(value, tuple): + items = tuple(fn(item) for item in value) + return type(value)(*items) if hasattr(value, "_fields") else items + if isinstance(value, list): + return [fn(item) for item in value] + if isinstance(value, dict): + return {key: fn(item) for key, item in value.items()} + return value + + def readonly_buffer_views(result: Any, snapshots: list[torch.Tensor]) -> Any: from ._tensors import _map_tensors @@ -866,14 +878,7 @@ def replace(value: Any) -> Any: value._head_key ] return captures[id(value)] - if isinstance(value, tuple): - items = tuple(replace(item) for item in value) - return type(value)(*items) if hasattr(value, "_fields") else items - if isinstance(value, list): - return [replace(item) for item in value] - if isinstance(value, dict): - return {key: replace(item) for key, item in value.items()} - return value + return _map_tensor_arguments(replace, value) return func(*replace(args), **replace(kwargs)) @@ -998,14 +1003,7 @@ def replace(value: Any) -> Any: ) originals[id(copies[id(value)])] = value return copies[id(value)] - if isinstance(value, tuple): - items = tuple(replace(item) for item in value) - return type(value)(*items) if hasattr(value, "_fields") else items - if isinstance(value, list): - return [replace(item) for item in value] - if isinstance(value, dict): - return {key: replace(item) for key, item in value.items()} - return value + return _map_tensor_arguments(replace, value) result = func(*replace(args), **replace(kwargs)) copied = [ diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index f46db981d..2342af0e7 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -7748,6 +7748,7 @@ def _tracked_tensor_function( kwargs: dict[str, object], ) -> object: from ._heads import ( + _map_tensor_arguments, _stage_local_buffers, head_call_arguments, mutates_tensor, @@ -7841,14 +7842,7 @@ def replace(value: object) -> object: result = value.as_subclass(torch.Tensor) replacements[id(value)] = result return result - if isinstance(value, tuple): - values = tuple(replace(item) for item in value) - return type(value)(*values) if hasattr(value, "_fields") else values - if isinstance(value, list): - return [replace(item) for item in value] - if isinstance(value, dict): - return {key: replace(item) for key, item in value.items()} - return value + return _map_tensor_arguments(replace, value) result = func( *cast(tuple[object, ...], replace(args)), From bcaa2f01433b14dda23f1f3a7e248ac5cb6c335d Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 20 Sep 2026 01:23:44 +0000 Subject: [PATCH 014/150] Consolidate trainer-rank test process scaffolding --- tests/unit/test_trainer_rank_check.py | 13 +-- tests/unit/test_trainer_rank_commands.py | 46 ++------ ...rainer_rank_memory_recovery_distributed.py | 12 +- .../unit/test_trainer_rank_planning_status.py | 13 +-- tests/unit/test_trainer_rank_validation.py | 107 ++++-------------- tests/unit/trainer_rank_test_support.py | 2 +- 6 files changed, 34 insertions(+), 159 deletions(-) diff --git a/tests/unit/test_trainer_rank_check.py b/tests/unit/test_trainer_rank_check.py index b6dceb68f..ccb3e348a 100644 --- a/tests/unit/test_trainer_rank_check.py +++ b/tests/unit/test_trainer_rank_check.py @@ -1,6 +1,5 @@ from __future__ import annotations -from datetime import timedelta import importlib from pathlib import Path import sys @@ -9,6 +8,7 @@ import torch import torch.distributed as dist import torch.multiprocessing as mp +from trainer_rank_test_support import gloo_group sys.path.insert(0, str(Path(__file__).parents[2] / "dev")) _compare_outputs = importlib.import_module("trainer_rank_check")._compare_outputs @@ -87,14 +87,7 @@ def _all_ranks_checked_worker( world_size: int, init_method: str, ) -> None: - dist.init_process_group( - "gloo", - init_method=init_method, - rank=rank, - world_size=world_size, - timeout=timedelta(seconds=30), - ) - try: + with gloo_group(rank, init_method, world_size=world_size): def check() -> None: if rank == 1: @@ -103,5 +96,3 @@ def check() -> None: with pytest.raises(AssertionError, match="injected rank-local failure"): all_ranks_checked("injected", check) dist.barrier() - finally: - dist.destroy_process_group() diff --git a/tests/unit/test_trainer_rank_commands.py b/tests/unit/test_trainer_rank_commands.py index 4dbc50d45..c9c19a44a 100644 --- a/tests/unit/test_trainer_rank_commands.py +++ b/tests/unit/test_trainer_rank_commands.py @@ -7,13 +7,14 @@ import gc import sys import threading -from types import ModuleType, SimpleNamespace +from types import SimpleNamespace from typing import Any import pytest import torch import torch.distributed as dist import torch.multiprocessing as mp +from trainer_rank_test_support import gloo_group, megatron_topology from art.trainer_rank import ( ForwardInput, @@ -87,14 +88,7 @@ def _input(value): def _suspended_abort_worker(physical, rendezvous): from test_trainer_rank_custom_tensors import _trainer - dist.init_process_group( - "gloo", - init_method=f"file://{rendezvous}", - rank=physical, - world_size=2, - timeout=timedelta(seconds=8), - ) - try: + with gloo_group(physical, f"file://{rendezvous}", timeout=8): native, _ = _trainer("student") native._checkpoint_process_group = dist.new_group( backend="gloo", timeout=timedelta(seconds=8) @@ -165,8 +159,6 @@ async def abort(): for mode in ("async", "stream", "cancel", "shutdown"): asyncio.run(run(mode)) - finally: - dist.destroy_process_group() def test_suspended_callbacks_leave_physical_checkpoint_abort_responsive(tmp_path): @@ -283,29 +275,12 @@ def backward(ctx, gradient): def _distributed_worker(physical, rendezvous, output): - dist.init_process_group( - "gloo", - init_method=f"file://{rendezvous}", - rank=physical, - world_size=4, - timeout=timedelta(seconds=30), - ) - try: - groups = [dist.new_group([0, 1]), dist.new_group([2, 3])] + with ( + gloo_group(physical, f"file://{rendezvous}", world_size=4), + megatron_topology(physical, dp_size=2, tp_size=2) as ps, + ): dp_groups = [dist.new_group([0, 2]), dist.new_group([1, 3])] dp, tp = divmod(physical, 2) - ps = SimpleNamespace( - get_tensor_model_parallel_rank=lambda: tp, - get_context_parallel_rank=lambda: 0, - get_data_parallel_rank=lambda: dp, - get_data_parallel_world_size=lambda: 2, - get_tensor_and_context_parallel_group=lambda **kwargs: groups[dp], - ) - megatron = ModuleType("megatron") - core = ModuleType("megatron.core") - setattr(core, "parallel_state", ps) - setattr(megatron, "core", core) - sys.modules.update({"megatron": megatron, "megatron.core": core}) rank: Any = _Rank(dp, 2) counts = [0, 0] @@ -534,13 +509,6 @@ def forward(self, value): dist.all_gather_object(gathered, (counts, result, per_dp, rank.steps)) if physical == 0: torch.save(gathered, output) - except BaseException: - import traceback - - traceback.print_exc() - raise - finally: - dist.destroy_process_group() def test_gloo_dp2_tp2_participation_and_gradients(tmp_path): diff --git a/tests/unit/test_trainer_rank_memory_recovery_distributed.py b/tests/unit/test_trainer_rank_memory_recovery_distributed.py index 1ce028c86..0f29c79f6 100644 --- a/tests/unit/test_trainer_rank_memory_recovery_distributed.py +++ b/tests/unit/test_trainer_rank_memory_recovery_distributed.py @@ -1,6 +1,5 @@ """Real collectives around injected transfer failure; no CUDA memory claims.""" -from datetime import timedelta import json from pathlib import Path from types import SimpleNamespace @@ -114,14 +113,7 @@ def _head_failure_worker(index: int, directory: str, failure: str) -> None: from art.trainer_rank._heads import LiveHead, export_head from art.trainer_rank._tensors import CotangentCollector - dist.init_process_group( - "gloo", - init_method=f"file://{directory}/rendezvous", - rank=index, - world_size=2, - timeout=timedelta(seconds=20), - ) - try: + with gloo_group(index, f"file://{directory}/rendezvous", timeout=20): trainer, api = _trainer("student") parameter = api.parameter("head", lambda: torch.ones(4), checkpoint="student") parameter.grad = torch.ones_like(parameter) @@ -166,8 +158,6 @@ def copy(tensor, *args, **kwargs): torch.testing.assert_close(parameter.grad, torch.ones_like(parameter)) assert trainer._version_state()._transaction is None dist.barrier() - finally: - dist.destroy_process_group() @pytest.mark.parametrize("failure", ["copy", "stage"]) diff --git a/tests/unit/test_trainer_rank_planning_status.py b/tests/unit/test_trainer_rank_planning_status.py index 7cb996e50..1a779bcbe 100644 --- a/tests/unit/test_trainer_rank_planning_status.py +++ b/tests/unit/test_trainer_rank_planning_status.py @@ -2,7 +2,6 @@ from __future__ import annotations -from datetime import timedelta from pathlib import Path import subprocess import sys @@ -77,6 +76,7 @@ def test_planning_failures_and_empty_ranks_use_aligned_status(tmp_path: Path) -> def _worker(index: int, directory: Path) -> None: import torch import torch.distributed as dist + from trainer_rank_test_support import gloo_group from art.trainer_rank import ForwardInput, TrainerRank, _impl from art.trainer_rank._prefix_tree_planner import ( @@ -85,14 +85,7 @@ def _worker(index: int, directory: Path) -> None: ) torch.set_num_threads(1) - dist.init_process_group( - "gloo", - rank=index, - world_size=2, - init_method=f"file://{directory / 'gloo'}", - timeout=timedelta(seconds=10), - ) - try: + with gloo_group(index, f"file://{directory / 'gloo'}", timeout=10): for mode in ( "estimate", "materialize", @@ -208,8 +201,6 @@ def retained_tokens(plan): finally: patches.undo() dist.barrier() - finally: - dist.destroy_process_group() if __name__ == "__main__": diff --git a/tests/unit/test_trainer_rank_validation.py b/tests/unit/test_trainer_rank_validation.py index f5fb2a140..555a02211 100644 --- a/tests/unit/test_trainer_rank_validation.py +++ b/tests/unit/test_trainer_rank_validation.py @@ -4,7 +4,6 @@ from collections.abc import Iterable from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass, fields, replace -from datetime import timedelta import gc from importlib.util import find_spec import inspect @@ -20,7 +19,7 @@ import pytest import torch import torch.distributed as dist -import torch.multiprocessing as mp +from trainer_rank_test_support import gloo_group, spawn_and_join from art.megatron.prefix_tree_packing import prefix_tree_pack from art.trainer_rank import ( @@ -2201,24 +2200,13 @@ def fail_snapshot(path: Path, ignore_errors: bool = False, **_: object) -> None: def _checkpoint_load_failure_worker( rank: int, world_size: int, init_method: str, phase: str ) -> None: - dist.init_process_group( - "gloo", - rank=rank, - world_size=world_size, - init_method=init_method, - timeout=timedelta(seconds=15), - ) from art.trainer_rank import _checkpoint as checkpoint_module from art.trainer_rank import _lora_export as lora_export_module - originals = ( - checkpoint_module._load_adapter, - checkpoint_module._optimizer_state, - checkpoint_module._commit_slot, - checkpoint_module._slot_snapshot, - checkpoint_module._restore_slots, - ) - try: + with ( + gloo_group(rank, init_method, world_size=world_size, timeout=15), + pytest.MonkeyPatch.context() as monkeypatch, + ): trainer = TrainerRank.__new__(TrainerRank) trainer.runtime = SimpleNamespace( model=[], @@ -2235,8 +2223,8 @@ def _checkpoint_load_failure_worker( trainer._validate_checkpoint_consistency = lambda *_args: () # type: ignore[method-assign] trainer._validate_loaded_checkpoint_config = lambda *_args: None # type: ignore[method-assign] trainer._restore_canonical_optimizer = lambda *_args: cast(Any, object()) # type: ignore[method-assign] - setattr(checkpoint_module, "_slot_snapshot", lambda *_args: ()) - setattr(checkpoint_module, "_restore_slots", lambda *_args: None) + monkeypatch.setattr(checkpoint_module, "_slot_snapshot", lambda *_args: ()) + monkeypatch.setattr(checkpoint_module, "_restore_slots", lambda *_args: None) if phase == "export": if rank == 1: trainer._checkpoint_slots["student"] = _CheckpointSlot( @@ -2291,7 +2279,7 @@ def _checkpoint_load_failure_worker( "digest", ) - setattr( + monkeypatch.setattr( checkpoint_module, "_load_adapter", ( @@ -2302,7 +2290,7 @@ def _checkpoint_load_failure_worker( ) ), ) - setattr( + monkeypatch.setattr( checkpoint_module, "_optimizer_state", ( @@ -2315,7 +2303,7 @@ def _checkpoint_load_failure_worker( ) ), ) - setattr( + monkeypatch.setattr( checkpoint_module, "_commit_slot", ( @@ -2336,40 +2324,18 @@ def _checkpoint_load_failure_worker( completed = torch.tensor(1) dist.all_reduce(completed) assert completed.item() == world_size - finally: - for name, value in zip( - ( - "_load_adapter", - "_optimizer_state", - "_commit_slot", - "_slot_snapshot", - "_restore_slots", - ), - originals, - strict=True, - ): - setattr(checkpoint_module, name, value) - dist.destroy_process_group() @pytest.mark.parametrize("phase", ("read", "optimizer", "commit", "export")) def test_checkpoint_load_failure_is_collective_and_transactional( tmp_path: Path, phase: str ) -> None: - context = mp.spawn( + spawn_and_join( _checkpoint_load_failure_worker, args=(2, f"file://{tmp_path / f'load-{phase}'}", phase), - nprocs=2, - join=False, + timeout=90, + failure=f"collective checkpoint {phase} failure test hung", ) - deadline = time.monotonic() + 90 - while time.monotonic() < deadline: - if context.join(timeout=1): - return - else: - for process in context.processes: - process.terminate() - pytest.fail(f"collective checkpoint {phase} failure test hung") @pytest.mark.skipif(find_spec("megatron") is None, reason="requires Megatron") @@ -3239,15 +3205,8 @@ def test_optim_step_rejects_invalid_live_graph_policy() -> None: def _live_graph_error_worker(rank: int, world_size: int, init_method: str) -> None: - dist.init_process_group( - "gloo", - rank=rank, - world_size=world_size, - init_method=init_method, - timeout=timedelta(seconds=30), - ) retained: torch.Tensor | None = None - try: + with gloo_group(rank, init_method, world_size=world_size): trainer = TrainerRank(_runtime()) cast(Any, trainer)._slot_ref = _slot_ref param = torch.nn.Parameter(torch.tensor([2.0])) @@ -3272,24 +3231,15 @@ def _live_graph_error_worker(rank: int, world_size: int, init_method: str) -> No dist.all_reduce(completed) assert completed.item() == world_size assert retained is None or retained.grad_fn is not None - finally: - dist.destroy_process_group() def test_optim_step_live_graph_error_is_collective(tmp_path: Path) -> None: - context = mp.spawn( + spawn_and_join( _live_graph_error_worker, args=(2, f"file://{tmp_path / 'live-graph'}"), - nprocs=2, - join=False, + timeout=45, + failure="collective live-graph policy test hung", ) - deadline = time.monotonic() + 45 - while time.monotonic() < deadline: - if context.join(timeout=1): - return - for process in context.processes: - process.terminate() - pytest.fail("collective live-graph policy test hung") def test_forward_preserves_nested_shape_for_inactive_requests() -> None: @@ -3458,14 +3408,7 @@ def test_skipped_forward_wave_cannot_resume_another_iterator( def _forward_yield_modes_worker(rank: int, world_size: int, init_method: str) -> None: - dist.init_process_group( - "gloo", - rank=rank, - world_size=world_size, - init_method=init_method, - timeout=timedelta(seconds=30), - ) - try: + with gloo_group(rank, init_method, world_size=world_size): with pytest.MonkeyPatch.context() as monkeypatch: try: from megatron.core import parallel_state @@ -3541,27 +3484,19 @@ def execute(plan, **_kwargs): completed = torch.tensor(1) trainer.reduce(completed) assert completed.item() == world_size - finally: - dist.destroy_process_group() @pytest.mark.parametrize("world_size", [2, 4]) def test_forward_batches_yield_modes_collectively( tmp_path: Path, world_size: int ) -> None: - context = mp.spawn( + spawn_and_join( _forward_yield_modes_worker, args=(world_size, f"file://{tmp_path / 'forward-yields'}"), nprocs=world_size, - join=False, + timeout=60, + failure="forward yield modes test hung", ) - deadline = time.monotonic() + 60 - while time.monotonic() < deadline: - if context.join(timeout=1): - return - for process in context.processes: - process.terminate() - pytest.fail("forward yield modes test hung") def test_forward_batches_syncs_fit_decision_across_dp( diff --git a/tests/unit/trainer_rank_test_support.py b/tests/unit/trainer_rank_test_support.py index 97911ba17..d4e4191a1 100644 --- a/tests/unit/trainer_rank_test_support.py +++ b/tests/unit/trainer_rank_test_support.py @@ -54,7 +54,7 @@ def megatron_topology(physical, *, dp_size, tp_size): ) setattr(megatron, "core", core) with patch.dict(sys.modules, {"megatron": megatron, "megatron.core": core}): - yield + yield getattr(core, "parallel_state") def spawn_and_join(worker, args, *, timeout, failure, nprocs=2): From 7365d9a62ec328f0d21f4646ff54c0d4b6751b4a Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 20 Sep 2026 01:32:09 +0000 Subject: [PATCH 015/150] Inherit duplicate TrainerRank public methods --- src/art/trainer_rank/__init__.py | 337 +------------------------------ src/art/trainer_rank/_impl.py | 62 +++++- 2 files changed, 61 insertions(+), 338 deletions(-) diff --git a/src/art/trainer_rank/__init__.py b/src/art/trainer_rank/__init__.py index 49b573815..eb6c2f89a 100644 --- a/src/art/trainer_rank/__init__.py +++ b/src/art/trainer_rank/__init__.py @@ -80,19 +80,12 @@ class TrainerRank(_impl.TrainerRank): profile is keyed by topology and calibrates itself online. """ - def __init__( - self, runtime: TrainingRuntime, *, options: ForwardOptions | None = None - ) -> None: - super().__init__(runtime, options=options) - @property def hidden_size(self) -> int: """Width of the returned hidden states.""" return self._hidden_size - def zero_grad(self) -> None: - super().zero_grad() - + # Keep ModuleHandle available to runtime type-hint resolution. def module( self, name: str, @@ -103,334 +96,6 @@ def module( """Retrieve a live module whose calls capture immutable checkpoint weights.""" return super().module(name, factory, checkpoint=checkpoint) - def parameter( - self, - name: str, - factory: Callable[[], torch.Tensor | torch.nn.Parameter], - *, - checkpoint: AdapterSelection = Unset, - ) -> torch.nn.Parameter: - """Register or retrieve a checkpoint-owned trainable tensor.""" - return super().parameter(name, factory, checkpoint=checkpoint) - - def buffer( - self, - name: str, - factory: Callable[[], torch.Tensor], - *, - checkpoint: AdapterSelection = Unset, - ) -> torch.Tensor: - """Register or retrieve a checkpoint-owned persistent buffer.""" - return super().buffer(name, factory, checkpoint=checkpoint) - - def prefetch_checkpoints( - self, - *checkpoints: str | MaterializedCheckpoint, - ) -> asyncio.Task[None]: - return super().prefetch_checkpoints(*checkpoints) - - def load_checkpoint(self, checkpoint: str | MaterializedCheckpoint | None) -> None: - super().load_checkpoint(checkpoint) - - def snapshot_checkpoint(self, source: str, destination: str) -> bool: - """Clone a loaded checkpoint into a forward-only resident snapshot.""" - return super().snapshot_checkpoint(source, destination) - - def push_checkpoint( - self, checkpoint: str | MaterializedCheckpoint | None - ) -> PushedCheckpoint: - return super().push_checkpoint(checkpoint) - - def pop_checkpoint(self) -> None: - super().pop_checkpoint() - - def save_checkpoint( - self, - output_dir: str, - checkpoint_path: str | Literal["active"] = "active", - ) -> None: - super().save_checkpoint(output_dir, checkpoint_path) - - def prepare_checkpoint_save( - self, - output_dir: str, - checkpoint_path: str | Literal["active"] = "active", - ) -> None: - super().prepare_checkpoint_save(output_dir, checkpoint_path) - - def finish_checkpoint_save(self, output_dir: str) -> None: - super().finish_checkpoint_save(output_dir) - - def abort_checkpoint_save(self, output_dir: str) -> None: - super().abort_checkpoint_save(output_dir) - - def export_lora( - self, - output_dir: str, - checkpoint_path: str | Literal["active"] = "active", - ) -> int: - return super().export_lora(output_dir, checkpoint_path) - - @overload - def forward_batches( - self, - inputs: Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]], - *, - options: ForwardOptions | None = None, - checkpoint: AdapterSelection = Unset, - no_grad: bool | None = None, - yield_empty: bool = False, - ) -> Iterator[ - MicroBatch[ - ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT], - ForwardOutput[LogprobsT, TopKT, LogitsT, HiddenStatesT], - ] - ]: ... - - @overload - def forward_batches( - self, - inputs: Iterable[ - Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]] - ], - *, - options: ForwardOptions | None = None, - checkpoint: AdapterSelection = Unset, - no_grad: bool | None = None, - yield_empty: bool = False, - ) -> Iterator[ - MicroBatch[ - Sequence[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]], - Sequence[ForwardOutput[LogprobsT, TopKT, LogitsT, HiddenStatesT]], - ] - ]: ... - - @overload - def forward_batches( - self, - inputs: Iterable[ - Iterable[Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]]] - ], - *, - options: ForwardOptions | None = None, - checkpoint: AdapterSelection = Unset, - no_grad: bool | None = None, - yield_empty: bool = False, - ) -> Iterator[ - MicroBatch[ - Sequence[Sequence[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]]], - Sequence[Sequence[ForwardOutput[LogprobsT, TopKT, LogitsT, HiddenStatesT]]], - ] - ]: ... - - @overload - def forward_batches( - self, - inputs: Iterable[ - Iterable[ - Iterable[ - Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]] - ] - ] - ], - *, - options: ForwardOptions | None = None, - checkpoint: AdapterSelection = Unset, - no_grad: bool | None = None, - yield_empty: bool = False, - ) -> Iterator[ - MicroBatch[ - Sequence[ - Sequence[ - Sequence[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]] - ] - ], - Sequence[ - Sequence[ - Sequence[ForwardOutput[LogprobsT, TopKT, LogitsT, HiddenStatesT]] - ] - ], - ] - ]: ... - - def forward_batches( - self, - inputs: Iterable[ForwardInputs], - *, - options: ForwardOptions | None = None, - checkpoint: AdapterSelection = Unset, - no_grad: bool | None = None, - yield_empty: bool = False, - ) -> Iterator[MicroBatch[ForwardInputs, ForwardOutputs]]: - """Forward replicated inputs in adaptive data-parallel microbatches. - - Per-input checkpoints and `no_grad` values override the method defaults. - `no_grad=None` inherits the ambient PyTorch grad mode; `True` disables - grads and `False` enables them. - Input and target tensors may be on a different device from the trainer; - ART moves its packed model inputs and labels internally without mutating - the caller-owned `ForwardInput` objects. - - Per-position outputs contain the full flattened input sequence in source - order, including with context parallelism. Logical callbacks execute - once per DP rank; use `backward(loss)` to route cotangents to internal - TP/CP participants. Direct physical callers must invoke matching - forwards and backwards on their TP/CP peers. `reduce` combines only - distinct data-parallel batches. - - Empty local microbatches are skipped unless `yield_empty=True`. Every - rank must use the same setting. When a wave skips ranks, TrainerRank - collective methods raise if called from its loop body; fully populated - waves permit them. Use `yield_empty=True` for per-wave collectives, - including reductions on ranks with no outputs. Exhaust or close a retained - iterator before making collective calls after an early exit. Guards apply - on the iterator's thread; raw torch.distributed calls are not guarded. - Collective calls must still match across ranks. - """ - forward = cast( - Callable[..., Iterator[MicroBatch[ForwardInputs, ForwardOutputs]]], - super().forward_batches, - ) - return forward( - inputs, - options=options, - checkpoint=checkpoint, - no_grad=no_grad, - yield_empty=yield_empty, - ) - - @overload - def forward( - self, - inputs: ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT], - *, - options: ForwardOptions | None = None, - checkpoint: AdapterSelection = Unset, - no_grad: bool | None = None, - ) -> ForwardOutput[LogprobsT, TopKT, LogitsT, HiddenStatesT]: ... - - @overload - def forward( - self, - inputs: Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]], - *, - options: ForwardOptions | None = None, - checkpoint: AdapterSelection = Unset, - no_grad: bool | None = None, - ) -> Sequence[ForwardOutput[LogprobsT, TopKT, LogitsT, HiddenStatesT]]: ... - - @overload - def forward( - self, - inputs: Iterable[ - Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]] - ], - *, - options: ForwardOptions | None = None, - checkpoint: AdapterSelection = Unset, - no_grad: bool | None = None, - ) -> Sequence[ - Sequence[ForwardOutput[LogprobsT, TopKT, LogitsT, HiddenStatesT]] - ]: ... - - @overload - def forward( - self, - inputs: Iterable[ - Iterable[Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]]] - ], - *, - options: ForwardOptions | None = None, - checkpoint: AdapterSelection = Unset, - no_grad: bool | None = None, - ) -> Sequence[ - Sequence[Sequence[ForwardOutput[LogprobsT, TopKT, LogitsT, HiddenStatesT]]] - ]: ... - - @overload - def forward( - self, - inputs: Iterable[ - Iterable[ - Iterable[ - Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]] - ] - ] - ], - *, - options: ForwardOptions | None = None, - checkpoint: AdapterSelection = Unset, - no_grad: bool | None = None, - ) -> Sequence[ - Sequence[ - Sequence[Sequence[ForwardOutput[LogprobsT, TopKT, LogitsT, HiddenStatesT]]] - ] - ]: ... - - def forward( - self, - inputs: ForwardInputs, - *, - options: ForwardOptions | None = None, - checkpoint: AdapterSelection = Unset, - no_grad: bool | None = None, - ) -> ForwardOutputs: - """Forward inputs already local to this data-parallel rank. - - Outputs contain full sequences in source order on every TP/CP rank, - with the same loss and reduction contract as `forward_batches`. - - Per-input checkpoints and `no_grad` values override the method defaults. - `no_grad=None` inherits the ambient PyTorch grad mode; `True` disables - grads and `False` enables them. - Input and target tensors may be on a different device from the trainer; - ART moves its packed model inputs and labels internally without mutating - the caller-owned `ForwardInput` objects. - """ - forward = cast( - Callable[..., ForwardOutputs], - super().forward, - ) - return forward(inputs, options=options, checkpoint=checkpoint, no_grad=no_grad) - - def reduce( - self, - tensor: torch.Tensor, - *, - op: dist.ReduceOp.RedOpType = dist.ReduceOp.SUM, - ) -> None: - """Reduce in place over data-parallel batches, excluding TP/CP replicas.""" - super().reduce(tensor, op=op) - - 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]: - """Step checkpoint slots that have accumulated gradients. - - A mapping assigns independent optimizer parameters to each checkpoint; - ``scale_grads`` may likewise map checkpoints to gradient scales. Mapping - keys select the checkpoints when ``checkpoints`` is omitted, and all - explicitly supplied checkpoint sets must match. Each checkpoint's gradient - norm is clipped independently. If any selected norm is nonfinite, no - selected checkpoint is updated. - - Retained forwards use immutable checkpoint versions and may be consumed - after this step within their captured `max_gradient_staleness` policy. - Pass `on_live_graphs="error"` to additionally refuse updates while a - selected checkpoint still has a live forward graph on any rank. - """ - return super().optim_step( - params=params, - scale_grads=scale_grads, - checkpoints=checkpoints, - on_live_graphs=on_live_graphs, - ) - from ._commands import ( RankCallbackResult, diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 2342af0e7..21a7ff19b 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1801,7 +1801,10 @@ def parameter( *, checkpoint: AdapterSelection = Unset, ) -> torch.nn.Parameter: - """Return a replicated checkpoint-owned trainable parameter.""" + """Register or retrieve a checkpoint-owned trainable tensor. + + The tensor is replicated across TrainerRank processes. + """ value = self._custom_object(name, "parameter", factory, checkpoint=checkpoint) return cast(torch.nn.Parameter, value) @@ -1812,7 +1815,10 @@ def buffer( *, checkpoint: AdapterSelection = Unset, ) -> torch.Tensor: - """Return a replicated checkpoint-owned persistent tensor.""" + """Register or retrieve a checkpoint-owned persistent buffer. + + The tensor is replicated across TrainerRank processes. + """ value = self._custom_object(name, "buffer", factory, checkpoint=checkpoint) return cast(torch.Tensor, value) @@ -2551,6 +2557,31 @@ def forward_batches( no_grad: bool | None = None, yield_empty: bool = False, ) -> Iterator[MicroBatch[ForwardInputs, ForwardOutputs]]: + """Forward replicated inputs in adaptive data-parallel microbatches. + + Per-input checkpoints and `no_grad` values override the method defaults. + `no_grad=None` inherits the ambient PyTorch grad mode; `True` disables + grads and `False` enables them. + Input and target tensors may be on a different device from the trainer; + ART moves its packed model inputs and labels internally without mutating + the caller-owned `ForwardInput` objects. + + Per-position outputs contain the full flattened input sequence in source + order, including with context parallelism. Logical callbacks execute + once per DP rank; use `backward(loss)` to route cotangents to internal + TP/CP participants. Direct physical callers must invoke matching + forwards and backwards on their TP/CP peers. `reduce` combines only + distinct data-parallel batches. + + Empty local microbatches are skipped unless `yield_empty=True`. Every + rank must use the same setting. When a wave skips ranks, TrainerRank + collective methods raise if called from its loop body; fully populated + waves permit them. Use `yield_empty=True` for per-wave collectives, + including reductions on ranks with no outputs. Exhaust or close a retained + iterator before making collective calls after an early exit. Guards apply + on the iterator's thread; raw torch.distributed calls are not guarded. + Collective calls must still match across ranks. + """ if not isinstance(yield_empty, bool): raise TypeError("yield_empty must be a bool") enabled = torch.is_grad_enabled() if no_grad is None else not no_grad @@ -2811,6 +2842,18 @@ def forward( checkpoint: AdapterSelection = Unset, no_grad: bool | None = None, ) -> ForwardOutputs: + """Forward inputs already local to this data-parallel rank. + + Outputs contain full sequences in source order on every TP/CP rank, + with the same loss and reduction contract as `forward_batches`. + + Per-input checkpoints and `no_grad` values override the method defaults. + `no_grad=None` inherits the ambient PyTorch grad mode; `True` disables + grads and `False` enables them. + Input and target tensors may be on a different device from the trainer; + ART moves its packed model inputs and labels internally without mutating + the caller-owned `ForwardInput` objects. + """ self._guard_forward_collective("forward") enabled = torch.is_grad_enabled() if no_grad is None else not no_grad with torch.set_grad_enabled(enabled): @@ -3438,6 +3481,7 @@ def reduce( *, op: dist.ReduceOp.RedOpType = dist.ReduceOp.SUM, ) -> None: + """Reduce in place over data-parallel batches, excluding TP/CP replicas.""" self._guard_forward_collective("reduce") from megatron.core import parallel_state as ps @@ -3486,6 +3530,20 @@ def optim_step( checkpoints: Sequence[str] | None = None, on_live_graphs: Literal["allow", "error"] = "allow", ) -> dict[str, float]: + """Step checkpoint slots that have accumulated gradients. + + A mapping assigns independent optimizer parameters to each checkpoint; + ``scale_grads`` may likewise map checkpoints to gradient scales. Mapping + keys select the checkpoints when ``checkpoints`` is omitted, and all + explicitly supplied checkpoint sets must match. Each checkpoint's gradient + norm is clipped independently. If any selected norm is nonfinite, no + selected checkpoint is updated. + + Retained forwards use immutable checkpoint versions and may be consumed + after this step within their captured `max_gradient_staleness` policy. + Pass `on_live_graphs="error"` to additionally refuse updates while a + selected checkpoint still has a live forward graph on any rank. + """ self._guard_forward_collective("optim_step") if on_live_graphs not in ("allow", "error"): raise ValueError( From 8f0f35e6195244407e0770fd29960277e1ebd820 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 20 Sep 2026 01:52:37 +0000 Subject: [PATCH 016/150] Remove obsolete facade type variables and imports --- src/art/trainer_rank/__init__.py | 12 ++---------- 1 file changed, 2 insertions(+), 10 deletions(-) diff --git a/src/art/trainer_rank/__init__.py b/src/art/trainer_rank/__init__.py index eb6c2f89a..faf148367 100644 --- a/src/art/trainer_rank/__init__.py +++ b/src/art/trainer_rank/__init__.py @@ -1,8 +1,7 @@ from __future__ import annotations -import asyncio -from collections.abc import Callable, Iterable, Iterator, Mapping, Sequence -from typing import TYPE_CHECKING, Literal, TypeVar, cast, overload +from collections.abc import Callable +from typing import TypeVar import torch import torch.distributed as dist @@ -29,10 +28,6 @@ MicroBatch = _impl.MicroBatch MicroBatchStats = _impl.MicroBatchStats TopK = _impl.TopK -LogprobsT = TypeVar("LogprobsT", bound=torch.Tensor | None, covariant=True) -TopKT = TypeVar("TopKT", bound=TopK | None, covariant=True) -LogitsT = TypeVar("LogitsT", bound=torch.Tensor | None, covariant=True) -HiddenStatesT = TypeVar("HiddenStatesT", bound=torch.Tensor | None, covariant=True) TrainerRankMemoryError = _impl.TrainerRankMemoryError TrainerRankPartialExecutionError = _impl.TrainerRankPartialExecutionError TrainerRankRuntimeSupportError = _impl.TrainerRankRuntimeSupportError @@ -41,9 +36,6 @@ MaterializedCheckpoint = _impl.MaterializedCheckpoint PushedCheckpoint = _impl.PushedCheckpoint -if TYPE_CHECKING: - from art.megatron.train import TrainingRuntime - ModuleT = TypeVar("ModuleT", bound=torch.nn.Module) From 3fc523ea58d1e6374b43ef6b22ebf3e007e593cc Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 20 Sep 2026 01:54:47 +0000 Subject: [PATCH 017/150] Reuse trainer test process group fixtures --- tests/unit/test_trainer_command_transport.py | 14 +- .../unit/test_trainer_rank_custom_tensors.py | 24 +-- .../unit/test_trainer_rank_head_recompute.py | 17 +- tests/unit/test_trainer_rank_live_heads.py | 40 +--- ...trainer_rank_recovery_slots_distributed.py | 172 +++++++++--------- tests/unit/trainer_rank_test_support.py | 16 +- 6 files changed, 109 insertions(+), 174 deletions(-) diff --git a/tests/unit/test_trainer_command_transport.py b/tests/unit/test_trainer_command_transport.py index 56ed7f1be..4c2266955 100644 --- a/tests/unit/test_trainer_command_transport.py +++ b/tests/unit/test_trainer_command_transport.py @@ -3,7 +3,6 @@ from __future__ import annotations from dataclasses import dataclass -from datetime import timedelta import gc import sys from types import ModuleType, SimpleNamespace @@ -12,8 +11,8 @@ import pytest import torch -import torch.distributed as dist import torch.multiprocessing as mp +from trainer_rank_test_support import gloo_group from art.trainer_rank import ForwardInput, ForwardOutput, TrainerRank from art.trainer_rank._commands import _Command, _encode_command, _Executor @@ -130,14 +129,7 @@ def _transport_worker(physical: int, rendezvous: str, cuda: bool) -> None: device = torch.device(f"cuda:{1 - physical}" if cuda else "cpu") if cuda: torch.cuda.set_device(device) - dist.init_process_group( - "gloo", - init_method=f"file://{rendezvous}", - rank=physical, - world_size=2, - timeout=timedelta(seconds=30), - ) - try: + with gloo_group(physical, f"file://{rendezvous}"): ps = SimpleNamespace( get_tensor_model_parallel_rank=lambda: physical, get_context_parallel_rank=lambda: 0, @@ -225,8 +217,6 @@ def forward(inputs: ForwardInput) -> ForwardOutput: ) if cuda: assert torch.cuda.memory_allocated(physical) == foreign_before - finally: - dist.destroy_process_group() @pytest.mark.parametrize("cuda", [False, True], ids=["cpu", "reversed-cuda-indices"]) diff --git a/tests/unit/test_trainer_rank_custom_tensors.py b/tests/unit/test_trainer_rank_custom_tensors.py index b46811dc7..a2ec3c6d6 100644 --- a/tests/unit/test_trainer_rank_custom_tensors.py +++ b/tests/unit/test_trainer_rank_custom_tensors.py @@ -4,7 +4,6 @@ from collections.abc import Callable import copy from dataclasses import dataclass -from datetime import timedelta from importlib.util import find_spec import io import json @@ -17,6 +16,7 @@ import torch import torch.distributed as dist import torch.multiprocessing as mp +from trainer_rank_test_support import gloo_group from art.trainer_rank import ( AdamParams, @@ -181,14 +181,7 @@ def _distributed_custom_registration_worker( init_method: str, mode: str, ) -> None: - dist.init_process_group( - "gloo", - init_method=init_method, - rank=rank, - world_size=world_size, - timeout=timedelta(seconds=30), - ) - try: + with gloo_group(rank, init_method, world_size=world_size): trainer, api = _trainer("student") slot = trainer._checkpoint_slots["student"] if mode == "trainability": @@ -232,8 +225,6 @@ def fail(*_args: object, **_kwargs: object) -> None: assert "head" not in slot.custom assert not slot.params dist.barrier() - finally: - dist.destroy_process_group() def _distributed_custom_grad_flags_worker( @@ -241,14 +232,7 @@ def _distributed_custom_grad_flags_worker( world_size: int, init_method: str, ) -> None: - dist.init_process_group( - "gloo", - init_method=init_method, - rank=rank, - world_size=world_size, - timeout=timedelta(seconds=30), - ) - try: + with gloo_group(rank, init_method, world_size=world_size): trainer, api = _trainer("student") used = api.parameter("used", lambda: torch.tensor(1.0), checkpoint="student") api.parameter("unused", lambda: torch.tensor(2.0), checkpoint="student") @@ -258,8 +242,6 @@ def _distributed_custom_grad_flags_worker( assert trainer._dynamic_param_step_flags( trainer._checkpoint_slots["student"].params ) == (True, False) - finally: - dist.destroy_process_group() def test_module_accepts_class_and_lambda_factories_and_is_idempotent() -> None: diff --git a/tests/unit/test_trainer_rank_head_recompute.py b/tests/unit/test_trainer_rank_head_recompute.py index 45d7b1083..54ccf0d4f 100644 --- a/tests/unit/test_trainer_rank_head_recompute.py +++ b/tests/unit/test_trainer_rank_head_recompute.py @@ -1,7 +1,6 @@ from __future__ import annotations from contextlib import nullcontext -from datetime import timedelta from types import SimpleNamespace from unittest.mock import patch @@ -9,6 +8,7 @@ import torch import torch.distributed as dist import torch.multiprocessing as mp +from trainer_rank_test_support import process_group from art.megatron.prefix_tree_packing import prefix_tree_pack from art.trainer_rank import ForwardInput, TrainerRank, _impl @@ -174,14 +174,13 @@ def _context_parallel_worker(rank, cp_size, dp_size, init_method, backend): device = torch.device("cpu" if backend == "gloo" else f"cuda:{rank}") if device.type == "cuda": torch.cuda.set_device(device) - dist.init_process_group( - backend, - init_method=init_method, - rank=rank, + with process_group( + rank, + init_method, world_size=cp_size * dp_size, - timeout=timedelta(seconds=90), - ) - try: + timeout=90, + backend=backend, + ): cp_groups = [ dist.new_group(list(range(dp * cp_size, (dp + 1) * cp_size))) for dp in range(dp_size) @@ -215,8 +214,6 @@ def _context_parallel_worker(rank, cp_size, dp_size, init_method, backend): _check_context_parallel_case( cp_rank, dp_rank, cp_size, dp_size, cp_group, device, mode ) - finally: - dist.destroy_process_group() def _check_context_parallel_case( diff --git a/tests/unit/test_trainer_rank_live_heads.py b/tests/unit/test_trainer_rank_live_heads.py index 712a220bf..07cf1a421 100644 --- a/tests/unit/test_trainer_rank_live_heads.py +++ b/tests/unit/test_trainer_rank_live_heads.py @@ -2,16 +2,15 @@ import asyncio from copy import deepcopy -from dataclasses import replace import pytest from test_trainer_rank_custom_tensors import _trainer import torch from torch.utils.checkpoint import checkpoint +from trainer_rank_test_support import gloo_group from art.trainer_rank import AdamParams, ModuleHandle, run_rank_callback from art.trainer_rank._heads import ( - HeadBufferUpdate, HeadRegistration, LiveHead, execute_head_operation, @@ -288,20 +287,9 @@ def test_constructor_staleness_applies_to_heads_before_mutating_gradients(): def _buffer_authority_worker(process_rank, init_method): - from datetime import timedelta - - import torch.distributed as dist - from art.trainer_rank._heads import synchronize_head_buffers - dist.init_process_group( - "gloo", - rank=process_rank, - world_size=2, - init_method=init_method, - timeout=timedelta(seconds=30), - ) - try: + with gloo_group(process_rank, init_method): trainer, rank = _trainer("student") head = rank.module("bn", lambda: torch.nn.BatchNorm1d(2), checkpoint="student") for _ in range(process_rank + 1): @@ -309,8 +297,6 @@ def _buffer_authority_worker(process_rank, init_method): synchronize_head_buffers(trainer) torch.testing.assert_close(head.running_mean, torch.full((2,), 0.1)) assert head.num_batches_tracked.item() == 1 - finally: - dist.destroy_process_group() def test_distributed_persistent_buffers_use_dp_zero_authority(tmp_path): @@ -897,20 +883,11 @@ def test_inplace_operation_snapshots_readonly_checkpoint_parameter(): def _live_buffer_authority_worker(process_rank, init_method, asymmetric=False): - from datetime import timedelta - import torch.distributed as dist from art.trainer_rank._heads import synchronize_head_buffers - dist.init_process_group( - "gloo", - rank=process_rank, - world_size=2, - init_method=init_method, - timeout=timedelta(seconds=30), - ) - try: + with gloo_group(process_rank, init_method): trainer, rank = _trainer("student") native = rank.module( "head", lambda: torch.nn.BatchNorm1d(2), checkpoint="student" @@ -952,8 +929,6 @@ def _live_buffer_authority_worker(process_rank, init_method, asymmetric=False): trainer, "head_publish", () if update is None else (update,) ) assert native.num_batches_tracked.item() == process_rank - finally: - dist.destroy_process_group() def test_distributed_live_buffer_refresh_accepts_dp_zero_authority(tmp_path): @@ -1305,19 +1280,12 @@ def test_cuda_live_buffer_views_and_functional_publication(client): @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") def test_cuda_buffer_sync_stages_cpu_authority_before_comparison(tmp_path): - import torch.distributed as dist - from art.trainer_rank._heads import synchronize_head_buffers - dist.init_process_group( - "gloo", rank=0, world_size=1, init_method=f"file://{tmp_path / 'cuda_sync'}" - ) - try: + with gloo_group(0, f"file://{tmp_path / 'cuda_sync'}", world_size=1, timeout=None): trainer, rank = _trainer("student") trainer.device = torch.device("cuda", 0) buffer = rank.buffer("mean", lambda: torch.ones(2), checkpoint="student") synchronize_head_buffers(trainer) torch.testing.assert_close(buffer, torch.ones(2, device="cuda")) assert export_head(trainer, "student", "mean").buffer_revision == 0 - finally: - dist.destroy_process_group() diff --git a/tests/unit/test_trainer_rank_recovery_slots_distributed.py b/tests/unit/test_trainer_rank_recovery_slots_distributed.py index 6192a1dea..b2ddefdac 100644 --- a/tests/unit/test_trainer_rank_recovery_slots_distributed.py +++ b/tests/unit/test_trainer_rank_recovery_slots_distributed.py @@ -62,100 +62,94 @@ def worker(index, mode, directory): import torch import torch.distributed as dist + from trainer_rank_test_support import gloo_group from art.trainer_rank import TrainerRank, _impl torch.set_num_threads(1) - dist.init_process_group( - "gloo", - rank=index, - world_size=2, - init_method=f"file://{directory}/rendezvous", - timeout=timedelta(seconds=3), - ) - rank = TrainerRank.__new__(TrainerRank) - rank.device = torch.device("cpu") - rank._checkpoint_mutation_lock = threading.RLock() - rank._checkpoint_prefetch_lock = threading.Lock() - rank._checkpoint_slots = {} - rank._checkpoint_group_lock = threading.Lock() - # Native _ensure_checkpoint_slots uses these all-rank groups unchanged. - rank._checkpoint_process_group = dist.new_group( - backend="gloo", timeout=timedelta(seconds=3) - ) - rank._checkpoint_finalize_process_group = dist.new_group( - backend="gloo", timeout=timedelta(seconds=3) - ) - groups = [ - dist.new_group([r], backend="gloo", timeout=timedelta(seconds=3)) - for r in range(2) - ] - rank._forward_memory_group = lambda: groups[index] - plan = SimpleNamespace( - packed_tokens=1, - logical_tokens=1, - active_logical_tokens=1, - grad_segment_count=0, - output_bytes=0, - signature=_impl._MemorySignature( - topology=(2, 1, 1, 1), - planner_coefficients=(0, None), - slot_group_count=1, - request_mix=("hidden_states",), - grad_enabled=False, - grad_modes=(False,), - ), - ) - rank._plan_flat_forward = lambda *args, **kwargs: plan - rank._estimate_required_memory_bytes_from_values = lambda **kwargs: 80 - rank._snapshot_planning_telemetry = lambda *args: None - reads = [] - ensures = [] - - def available(): - reads.append(1) - deficient = mode == "both" or (mode == "asymmetric" and index == 0) - return 10 if deficient and len(reads) <= 2 else 200 - - rank._available_memory_bytes = available - native = rank._ensure_checkpoint_slots - - def ensure(values): - ensures.append(1) - return native(values) - - rank._ensure_checkpoint_slots = ensure - request = SimpleNamespace( - target_tokens=None, - logits=False, - top_k=None, - hidden_states=True, - checkpoint=None, - ) - error = barrier_error = None - try: - rank._plan_admissible_forward([request], checkpoint=None, context="forward") - except BaseException as exc: - error = {"type": type(exc).__name__, "message": str(exc)} - try: - dist.barrier() - except BaseException as exc: - barrier_error = {"type": type(exc).__name__, "message": str(exc)} - (directory / f"rank-{index}.json").write_text( - json.dumps( - { - "source": _impl.__file__, - "rank": index, - "mode": mode, - "ensures": len(ensures), - "samples": len(reads), - "error": error, - "barrier_error": barrier_error, - }, - indent=2, + with gloo_group(index, f"file://{directory}/rendezvous", timeout=3): + rank = TrainerRank.__new__(TrainerRank) + rank.device = torch.device("cpu") + rank._checkpoint_mutation_lock = threading.RLock() + rank._checkpoint_prefetch_lock = threading.Lock() + rank._checkpoint_slots = {} + rank._checkpoint_group_lock = threading.Lock() + # Native _ensure_checkpoint_slots uses these all-rank groups unchanged. + rank._checkpoint_process_group = dist.new_group( + backend="gloo", timeout=timedelta(seconds=3) + ) + rank._checkpoint_finalize_process_group = dist.new_group( + backend="gloo", timeout=timedelta(seconds=3) + ) + groups = [ + dist.new_group([r], backend="gloo", timeout=timedelta(seconds=3)) + for r in range(2) + ] + rank._forward_memory_group = lambda: groups[index] + plan = SimpleNamespace( + packed_tokens=1, + logical_tokens=1, + active_logical_tokens=1, + grad_segment_count=0, + output_bytes=0, + signature=_impl._MemorySignature( + topology=(2, 1, 1, 1), + planner_coefficients=(0, None), + slot_group_count=1, + request_mix=("hidden_states",), + grad_enabled=False, + grad_modes=(False,), + ), + ) + rank._plan_flat_forward = lambda *args, **kwargs: plan + rank._estimate_required_memory_bytes_from_values = lambda **kwargs: 80 + rank._snapshot_planning_telemetry = lambda *args: None + reads = [] + ensures = [] + + def available(): + reads.append(1) + deficient = mode == "both" or (mode == "asymmetric" and index == 0) + return 10 if deficient and len(reads) <= 2 else 200 + + rank._available_memory_bytes = available + native = rank._ensure_checkpoint_slots + + def ensure(values): + ensures.append(1) + return native(values) + + rank._ensure_checkpoint_slots = ensure + request = SimpleNamespace( + target_tokens=None, + logits=False, + top_k=None, + hidden_states=True, + checkpoint=None, + ) + error = barrier_error = None + try: + rank._plan_admissible_forward([request], checkpoint=None, context="forward") + except BaseException as exc: + error = {"type": type(exc).__name__, "message": str(exc)} + try: + dist.barrier() + except BaseException as exc: + barrier_error = {"type": type(exc).__name__, "message": str(exc)} + (directory / f"rank-{index}.json").write_text( + json.dumps( + { + "source": _impl.__file__, + "rank": index, + "mode": mode, + "ensures": len(ensures), + "samples": len(reads), + "error": error, + "barrier_error": barrier_error, + }, + indent=2, + ) ) - ) - dist.destroy_process_group() if __name__ == "__main__": diff --git a/tests/unit/trainer_rank_test_support.py b/tests/unit/trainer_rank_test_support.py index d4e4191a1..da17ac031 100644 --- a/tests/unit/trainer_rank_test_support.py +++ b/tests/unit/trainer_rank_test_support.py @@ -1,11 +1,10 @@ -"""Real Gloo groups and lightweight topology for trainer-rank contract tests.""" +"""Real process groups and lightweight topology for trainer-rank contract tests.""" from contextlib import contextmanager from datetime import timedelta import sys import time from types import ModuleType, SimpleNamespace -from unittest.mock import patch import pytest import torch.distributed as dist @@ -13,13 +12,13 @@ @contextmanager -def gloo_group(rank, rendezvous, *, world_size=2, timeout=30): +def process_group(rank, rendezvous, *, world_size=2, timeout=30, backend="gloo"): dist.init_process_group( - "gloo", + backend, init_method=rendezvous, rank=rank, world_size=world_size, - timeout=timedelta(seconds=timeout), + timeout=None if timeout is None else timedelta(seconds=timeout), ) try: yield @@ -27,6 +26,9 @@ def gloo_group(rank, rendezvous, *, world_size=2, timeout=30): dist.destroy_process_group() +gloo_group = process_group + + @contextmanager def megatron_topology(physical, *, dp_size, tp_size): """Install just the callback topology, with each real TP group created in order.""" @@ -53,7 +55,9 @@ def megatron_topology(physical, *, dp_size, tp_size): ), ) setattr(megatron, "core", core) - with patch.dict(sys.modules, {"megatron": megatron, "megatron.core": core}): + with pytest.MonkeyPatch.context() as modules: + modules.setitem(sys.modules, "megatron", megatron) + modules.setitem(sys.modules, "megatron.core", core) yield getattr(core, "parallel_state") From 1bcf119680140a5e60aab1a1e6c4c2c2eb748bf4 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 20 Sep 2026 02:00:58 +0000 Subject: [PATCH 018/150] Reuse live parameter test fixtures --- .../unit/test_trainer_live_parameter_roots.py | 40 ++++++++----------- .../unit/test_trainer_rank_parameter_hooks.py | 33 +++++---------- 2 files changed, 28 insertions(+), 45 deletions(-) diff --git a/tests/unit/test_trainer_live_parameter_roots.py b/tests/unit/test_trainer_live_parameter_roots.py index 4fb0f7ade..37f7436ba 100644 --- a/tests/unit/test_trainer_live_parameter_roots.py +++ b/tests/unit/test_trainer_live_parameter_roots.py @@ -9,10 +9,23 @@ from test_trainer_rank_custom_tensors import _trainer import torch +from art.trainer_rank import TrainerRank from art.trainer_rank._heads import LiveHead, export_head, head_gradient_targets from art.trainer_rank._tensors import CotangentCollector +def _live_parameter( + factory=lambda: torch.tensor(2.0), +) -> tuple[TrainerRank, torch.nn.Parameter, CotangentCollector, LiveHead]: + trainer, rank = _trainer("student") + parameter = rank.parameter("weight", factory, checkpoint="student") + collector = CotangentCollector() + live = LiveHead( + export_head(trainer, "student", "weight"), parameter.detach(), collector + ) + return trainer, parameter, collector, live + + class _FailBackward(torch.autograd.Function): @staticmethod def forward(ctx, value): @@ -46,13 +59,8 @@ def test_native_direct_root_does_not_publish_when_another_root_fails(existing): @pytest.mark.parametrize("dtype", [torch.float32, torch.float64, torch.complex64]) @pytest.mark.parametrize("surface", ["native", "client"]) def test_direct_roots_preserve_aliases_explicit_gradients_and_dtype(dtype, surface): - trainer, rank = _trainer("student") - parameter = rank.parameter( - "weight", lambda: torch.tensor([2.0, 3.0], dtype=dtype), checkpoint="student" - ) - collector = CotangentCollector() - live = LiveHead( - export_head(trainer, "student", "weight"), parameter.detach(), collector + trainer, parameter, collector, live = _live_parameter( + lambda: torch.tensor([2.0, 3.0], dtype=dtype) ) root: Any = parameter if surface == "native" else live.value gradients = ( @@ -83,14 +91,7 @@ def test_direct_roots_preserve_aliases_explicit_gradients_and_dtype(dtype, surfa def test_direct_live_root_and_old_arithmetic_keep_their_own_versions(): - trainer, rank = _trainer("student") - parameter = rank.parameter( - "weight", lambda: torch.tensor(2.0), checkpoint="student" - ) - collector = CotangentCollector() - live = LiveHead( - export_head(trainer, "student", "weight"), parameter.detach(), collector - ) + trainer, parameter, collector, live = _live_parameter() root: Any = live.value old = root.square() parameter.data.fill_(5) @@ -109,14 +110,7 @@ def test_direct_live_root_and_old_arithmetic_keep_their_own_versions(): def test_client_direct_root_failure_discards_cotangents_and_revalidates_handle(): - trainer, rank = _trainer("student") - parameter = rank.parameter( - "weight", lambda: torch.tensor(2.0), checkpoint="student" - ) - collector = CotangentCollector() - live = LiveHead( - export_head(trainer, "student", "weight"), parameter.detach(), collector - ) + trainer, parameter, collector, live = _live_parameter() root: Any = live.value bad = _FailBackward.apply(torch.tensor(1.0, requires_grad=True)) with pytest.raises(RuntimeError, match="local backward failed"): diff --git a/tests/unit/test_trainer_rank_parameter_hooks.py b/tests/unit/test_trainer_rank_parameter_hooks.py index bd49e19b4..ae9628db1 100644 --- a/tests/unit/test_trainer_rank_parameter_hooks.py +++ b/tests/unit/test_trainer_rank_parameter_hooks.py @@ -12,11 +12,16 @@ from art.trainer_rank._tensors import CotangentCollector -def setup(client): +def setup(client, values=None): trainer, rank = _trainer("student") - native = rank.parameter("p", lambda: torch.tensor(2.0), checkpoint="student") + factory = lambda: torch.tensor(2.0) if values is None else values.clone() + native = rank.parameter("p", factory, checkpoint="student") collector = CotangentCollector() - live = LiveHead(export_head(trainer, "student", "p"), torch.tensor(2.0), collector) + live = LiveHead( + export_head(trainer, "student", "p"), + factory() if values is None else values, + collector, + ) parameter = live.value if client else native assert isinstance(parameter, torch.Tensor) @@ -144,13 +149,8 @@ def test_head_hook_registries_follow_graph_lifetime(): def test_sparse_hook_gradients_preserve_layout_and_mix_with_dense( client, sparse_first, mixed ): - trainer, rank = _trainer("student") values = torch.arange(1.0, 7).reshape(3, 2) - native = rank.parameter("p", lambda: values.clone(), checkpoint="student") - collector = CotangentCollector() - live = LiveHead(export_head(trainer, "student", "p"), values, collector) - parameter = live.value if client else native - assert isinstance(parameter, torch.Tensor) + _, native, parameter, _, _, backward = setup(client, values) seen = [] parameter.register_hook(lambda gradient: seen.append(gradient.layout) or gradient) indices = torch.tensor([0, 2, 0]) @@ -165,25 +165,14 @@ def test_sparse_hook_gradients_preserve_layout_and_mix_with_dense( ) native.grad = torch.ones_like(values) - def backward(): - if client: - packets = collector.backward(loss) - with trainer._gradient_transaction(): - for packet in packets: - trainer._commit_versioned_gradients( - head_gradient_targets(trainer, packet) - ) - else: - trainer.backward(loss) - if not mixed: with pytest.raises(ValueError, match="layout"): - backward() + backward(loss) torch.testing.assert_close(native.grad, torch.ones_like(values)) if client: assert seen == [torch.sparse_coo] else: - backward() + backward(loss) assert seen == [torch.strided] expected = 1 + 2 * values + torch.tensor([[2.0, 2.0], [0.0, 0.0], [1.0, 1.0]]) torch.testing.assert_close(native.grad, expected) From 0429b719c343af5bc06c0e0d47b8b8140ef5c891 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 20 Sep 2026 02:00:30 +0000 Subject: [PATCH 019/150] Reuse custom tensor wrapping for module members --- src/art/trainer_rank/_impl.py | 40 ++++++++++++----------------------- 1 file changed, 13 insertions(+), 27 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 21a7ff19b..2ee137867 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -7968,35 +7968,21 @@ def _track_custom_object( return _CustomObject(custom.kind, value, custom.generation) module = cast(torch.nn.Module, custom.value) - parameters: dict[int, torch.nn.Parameter] = {} - buffers: dict[int, torch.Tensor] = {} + replacements: dict[tuple[str, int], torch.Tensor] = {} for child in module.modules(): - for key, source in child._parameters.items(): - if source is None: - continue - value = parameters.get(id(source)) - if value is None: - with torch.no_grad(): - value = _TrackedParameter( - source.detach().clone(), tracker, source.requires_grad + for kind in ("parameter", "buffer"): + for key, source in getattr(child, f"_{kind}s").items(): + if source is None: + continue + identity = (kind, id(source)) + if identity not in replacements: + replacements[identity] = cast( + torch.Tensor, + _track_custom_object( + _CustomObject(kind, source, custom.generation), tracker + ).value, ) - value.__dict__.update( - (attribute, item) - for attribute, item in source.__dict__.items() - if attribute != "_art_tracker" - ) - value._art_tracker = tracker - parameters[id(source)] = value - child._parameters[key] = value - for key, source in child._buffers.items(): - if source is None: - continue - value = buffers.get(id(source)) - if value is None: - with torch.no_grad(): - value = _TrackedTensor(source.detach().clone(), tracker) - buffers[id(source)] = value - child._buffers[key] = value + getattr(child, f"_{kind}s")[key] = replacements[identity] from ._heads import native_module_handle trainer = tracker.validate() From 1d9d6ae43aa9bfe1a129ce09dfd0b0b595001ab5 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 20 Sep 2026 02:09:08 +0000 Subject: [PATCH 020/150] Share policy choice checks and preserve public hint bindings --- src/art/trainer_rank/__init__.py | 8 ++++++-- src/art/trainer_rank/_options.py | 18 ++++++------------ 2 files changed, 12 insertions(+), 14 deletions(-) diff --git a/src/art/trainer_rank/__init__.py b/src/art/trainer_rank/__init__.py index faf148367..825025223 100644 --- a/src/art/trainer_rank/__init__.py +++ b/src/art/trainer_rank/__init__.py @@ -1,7 +1,7 @@ from __future__ import annotations -from collections.abc import Callable -from typing import TypeVar +from collections.abc import Callable, Sequence +from typing import Literal, TypeVar import torch import torch.distributed as dist @@ -28,6 +28,10 @@ MicroBatch = _impl.MicroBatch MicroBatchStats = _impl.MicroBatchStats TopK = _impl.TopK +LogprobsT = TypeVar("LogprobsT", bound=torch.Tensor | None, covariant=True) +TopKT = TypeVar("TopKT", bound=TopK | None, covariant=True) +LogitsT = TypeVar("LogitsT", bound=torch.Tensor | None, covariant=True) +HiddenStatesT = TypeVar("HiddenStatesT", bound=torch.Tensor | None, covariant=True) TrainerRankMemoryError = _impl.TrainerRankMemoryError TrainerRankPartialExecutionError = _impl.TrainerRankPartialExecutionError TrainerRankRuntimeSupportError = _impl.TrainerRankRuntimeSupportError diff --git a/src/art/trainer_rank/_options.py b/src/art/trainer_rank/_options.py index 9535ae1d9..201e5a47f 100644 --- a/src/art/trainer_rank/_options.py +++ b/src/art/trainer_rank/_options.py @@ -119,19 +119,13 @@ def _validate_options(options: ForwardOptions | ResolvedForwardOptions) -> None: value = getattr(options, name) if value is not Unset and type(value) is not bool: raise ValueError(f"{name} must be a bool") - if options.backward_state is not Unset and options.backward_state not in ( - "auto", - "gpu", - "cpu", - "replay", + for name, choices in ( + ("backward_state", ("auto", "gpu", "cpu", "replay")), + ("output_device", ("auto", "model", "cpu")), ): - raise ValueError(f"unknown backward_state: {options.backward_state!r}") - if options.output_device is not Unset and options.output_device not in ( - "auto", - "model", - "cpu", - ): - raise ValueError(f"unknown output_device: {options.output_device!r}") + value = getattr(options, name) + if value is not Unset and value not in choices: + raise ValueError(f"unknown {name}: {value!r}") corrections = options.stale_gradient_corrections if corrections is not Unset: if any( From 101fd1fec462e01a129807e03a308c03658d6c1f Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 20 Sep 2026 02:27:22 +0000 Subject: [PATCH 021/150] Share argument container traversal with replay restore --- src/art/trainer_rank/_graphs.py | 11 +++-------- src/art/trainer_rank/_heads.py | 19 ++----------------- src/art/trainer_rank/_impl.py | 2 +- src/art/trainer_rank/_tensors.py | 12 ++++++++++++ 4 files changed, 18 insertions(+), 26 deletions(-) diff --git a/src/art/trainer_rank/_graphs.py b/src/art/trainer_rank/_graphs.py index 3294f413b..7bc08fd12 100644 --- a/src/art/trainer_rank/_graphs.py +++ b/src/art/trainer_rank/_graphs.py @@ -18,6 +18,8 @@ from art._tensor_residency import observe_resident_tensors +from ._tensors import _map_tensor_arguments + ForwardHandle = str type Retention = Literal["gpu", "cpu", "replay"] @@ -69,14 +71,7 @@ def _restore_inputs(value: Any) -> Any: if f.init }, ) - if isinstance(value, tuple): - items = tuple(_restore_inputs(item) for item in value) - return type(value)(*items) if hasattr(value, "_fields") else items - if isinstance(value, list): - return [_restore_inputs(item) for item in value] - if isinstance(value, dict): - return {key: _restore_inputs(item) for key, item in value.items()} - return value + return _map_tensor_arguments(_restore_inputs, value) def _tensors(value: Any) -> Iterable[torch.Tensor]: diff --git a/src/art/trainer_rank/_heads.py b/src/art/trainer_rank/_heads.py index b02cc82fb..37abb1824 100644 --- a/src/art/trainer_rank/_heads.py +++ b/src/art/trainer_rank/_heads.py @@ -11,6 +11,8 @@ import torch +from ._tensors import _map_tensor_arguments, _map_tensors + if TYPE_CHECKING: from ._impl import TrainerRank, _CustomObject, _CustomTensorTracker @@ -28,8 +30,6 @@ def head_call_arguments(args: tuple[Any, ...], kwargs: dict[str, Any]) -> Any: call = _head_call.get() if call is None: return None - from ._tensors import _map_tensors - changed = False def replace(value: torch.Tensor) -> torch.Tensor: @@ -163,21 +163,7 @@ def tensor_mutation_targets( } -def _map_tensor_arguments(fn: Callable[[Any], Any], value: Any) -> Any: - """Map one argument container, preserving namedtuples and opaque values.""" - if isinstance(value, tuple): - items = tuple(fn(item) for item in value) - return type(value)(*items) if hasattr(value, "_fields") else items - if isinstance(value, list): - return [fn(item) for item in value] - if isinstance(value, dict): - return {key: fn(item) for key, item in value.items()} - return value - - def readonly_buffer_views(result: Any, snapshots: list[torch.Tensor]) -> Any: - from ._tensors import _map_tensors - def wrap(value: torch.Tensor) -> torch.Tensor: with torch._C.DisableTorchFunctionSubclass(): if any( @@ -204,7 +190,6 @@ def __reduce_ex__(self, proto: SupportsIndex) -> Any: @classmethod def __torch_function__(cls, func, types, args=(), kwargs=None): from ._impl import _walk_objects - from ._tensors import _map_tensors kwargs = kwargs or {} if getattr(func, "__name__", "") in { diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 2ee137867..23691b532 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -7806,7 +7806,6 @@ def _tracked_tensor_function( kwargs: dict[str, object], ) -> object: from ._heads import ( - _map_tensor_arguments, _stage_local_buffers, head_call_arguments, mutates_tensor, @@ -7814,6 +7813,7 @@ def _tracked_tensor_function( tensor_metadata_function, tensor_mutation_targets, ) + from ._tensors import _map_tensor_arguments if (captured := head_call_arguments(args, kwargs)) is not None: return func(*captured[0], **captured[1]) diff --git a/src/art/trainer_rank/_tensors.py b/src/art/trainer_rank/_tensors.py index cffe6512e..aec5cea3c 100644 --- a/src/art/trainer_rank/_tensors.py +++ b/src/art/trainer_rank/_tensors.py @@ -108,6 +108,18 @@ def _map_tensors(fn: Callable[[torch.Tensor], torch.Tensor], tree: Any) -> Any: return unflatten_tensors(spec, tuple(map(fn, tensors))) +def _map_tensor_arguments(fn: Callable[[Any], Any], value: Any) -> Any: + """Map one argument container, preserving namedtuples and opaque values.""" + if isinstance(value, tuple): + items = tuple(fn(item) for item in value) + return type(value)(*items) if hasattr(value, "_fields") else items + if isinstance(value, list): + return [fn(item) for item in value] + if isinstance(value, dict): + return {key: fn(item) for key, item in value.items()} + return value + + def _plain(tensor: torch.Tensor) -> torch.Tensor: if isinstance(tensor, ManagedTensor): with torch.enable_grad(), torch._C.DisableTorchFunctionSubclass(): From 6e6155310593b0ec7663bff6f1a614863584be6c Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 20 Sep 2026 02:31:14 +0000 Subject: [PATCH 022/150] Use real trainer typing and share CPU gradient stubs --- .../unit/test_trainer_rank_custom_tensors.py | 100 +++++------------- tests/unit/test_trainer_rank_live_heads.py | 13 +-- 2 files changed, 26 insertions(+), 87 deletions(-) diff --git a/tests/unit/test_trainer_rank_custom_tensors.py b/tests/unit/test_trainer_rank_custom_tensors.py index a2ec3c6d6..4e3866fb8 100644 --- a/tests/unit/test_trainer_rank_custom_tensors.py +++ b/tests/unit/test_trainer_rank_custom_tensors.py @@ -10,7 +10,7 @@ from pathlib import Path import shutil from types import SimpleNamespace -from typing import Any, Protocol, TypeVar, cast +from typing import Any, cast import pytest import torch @@ -24,7 +24,6 @@ ModuleHandle, TrainerRank, TrainerRankSlotStateError, - Unset, ) from art.trainer_rank._checkpoint import ( PreparedCustomPayload, @@ -35,34 +34,6 @@ ) from art.trainer_rank._impl import _CheckpointSlot -ModuleT = TypeVar("ModuleT", bound=torch.nn.Module) - - -class _CustomTensorAPI(Protocol): - def module( - self, - name: str, - factory: Callable[[], ModuleT], - *, - checkpoint: str | object = Unset, - ) -> ModuleHandle: ... - - def parameter( - self, - name: str, - factory: Callable[[], torch.Tensor | torch.nn.Parameter], - *, - checkpoint: str | object = Unset, - ) -> torch.nn.Parameter: ... - - def buffer( - self, - name: str, - factory: Callable[[], torch.Tensor], - *, - checkpoint: str | object = Unset, - ) -> torch.Tensor: ... - class _ClassFactoryHead(torch.nn.Module): def __init__(self) -> None: @@ -160,11 +131,24 @@ def _config() -> dict[str, object]: } -def _trainer(*names: str) -> tuple[TrainerRank, _CustomTensorAPI]: +def _trainer(*names: str) -> tuple[TrainerRank, TrainerRank]: trainer = TrainerRank(_runtime()) for name in names: trainer._checkpoint_slots[name] = _CheckpointSlot(config=cast(Any, _config())) - return trainer, cast(_CustomTensorAPI, trainer) + return trainer, trainer + + +def _use_local_gradients(trainer: TrainerRank, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr( + trainer, + "_reduce_dynamic_grads", + lambda params, **_kwargs: tuple( + torch.zeros_like(param, dtype=torch.float32) + if param.grad is None + else param.grad.float() + for param in params + ), + ) def test_custom_head_uses_model_hidden_size_before_forward() -> None: @@ -633,16 +617,7 @@ def test_selected_optimizer_step_updates_only_its_custom_checkpoint( trainer, rank = _trainer("A", "B") a = rank.parameter("gain", lambda: torch.tensor(1.0), checkpoint="A") b = rank.parameter("gain", lambda: torch.tensor(2.0), checkpoint="B") - monkeypatch.setattr( - trainer, - "_reduce_dynamic_grads", - lambda params, **_kwargs: tuple( - torch.zeros_like(param, dtype=torch.float32) - if param.grad is None - else param.grad.float() - for param in params - ), - ) + _use_local_gradients(trainer, monkeypatch) with trainer._gradient_transaction(): (a * 3).backward() before_b = b.detach().clone() @@ -660,16 +635,7 @@ def test_optimizer_skips_unused_custom_parameters( trainer, rank = _trainer("student") used = rank.parameter("used", lambda: torch.tensor(1.0), checkpoint="student") unused = rank.parameter("unused", lambda: torch.tensor(3.0), checkpoint="student") - monkeypatch.setattr( - trainer, - "_reduce_dynamic_grads", - lambda params, **_kwargs: tuple( - torch.zeros_like(param, dtype=torch.float32) - if param.grad is None - else param.grad.float() - for param in params - ), - ) + _use_local_gradients(trainer, monkeypatch) with trainer._gradient_transaction(): (used * 2).backward() @@ -943,16 +909,7 @@ def test_custom_tensor_names_cannot_corrupt_lora_optimizer_metadata( ) -> None: trainer, rank = _real_lora_trainer() head = rank.module("layer", _CollisionHead, checkpoint="student") - monkeypatch.setattr( - trainer, - "_reduce_dynamic_grads", - lambda params, **_kwargs: tuple( - torch.zeros_like(param, dtype=torch.float32) - if param.grad is None - else param.grad.float() - for param in params - ), - ) + _use_local_gradients(trainer, monkeypatch) with trainer._gradient_transaction(): head.q_proj.lora_A.weight.sum().backward() trainer.optim_step( @@ -1242,7 +1199,7 @@ def factory() -> _BufferLayoutHead: api.module("head", factory, checkpoint="student") -def _real_lora_trainer() -> tuple[TrainerRank, _CustomTensorAPI]: +def _real_lora_trainer() -> tuple[TrainerRank, TrainerRank]: trainer, api = _empty_real_lora_trainer() adapter = { "layer.q_proj.lora_A.weight": torch.tensor([[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]]), @@ -1256,16 +1213,16 @@ def _real_lora_trainer() -> tuple[TrainerRank, _CustomTensorAPI]: return trainer, api -def _empty_real_lora_trainer() -> tuple[TrainerRank, _CustomTensorAPI]: +def _empty_real_lora_trainer() -> tuple[TrainerRank, TrainerRank]: from art.megatron.lora import LoRA lora = LoRA("layer.q_proj", 3, 4, 2, 2, torch.float32, torch.device("cpu")) trainer = TrainerRank(_runtime(lora)) - return trainer, cast(_CustomTensorAPI, trainer) + return trainer, trainer def _register_custom_tensors( - rank: _CustomTensorAPI, + rank: TrainerRank, ) -> tuple[ModuleHandle, torch.nn.Parameter, torch.Tensor]: head = rank.module("value_head", lambda: _ValueHead(3), checkpoint="student") temperature = rank.parameter( @@ -1285,16 +1242,7 @@ def _step_custom_tensors( *, scale: float = 1.0, ) -> None: - monkeypatch.setattr( - trainer, - "_reduce_dynamic_grads", - lambda params, **_kwargs: tuple( - torch.zeros_like(param, dtype=torch.float32) - if param.grad is None - else param.grad.float() - for param in params - ), - ) + _use_local_gradients(trainer, monkeypatch) hidden = torch.tensor([[0.25, -0.5, 1.0]]) with trainer._gradient_transaction(): (head(hidden).sum() + temperature * scale).backward() diff --git a/tests/unit/test_trainer_rank_live_heads.py b/tests/unit/test_trainer_rank_live_heads.py index 07cf1a421..fdf6b8e54 100644 --- a/tests/unit/test_trainer_rank_live_heads.py +++ b/tests/unit/test_trainer_rank_live_heads.py @@ -4,7 +4,7 @@ from copy import deepcopy import pytest -from test_trainer_rank_custom_tensors import _trainer +from test_trainer_rank_custom_tensors import _trainer, _use_local_gradients import torch from torch.utils.checkpoint import checkpoint from trainer_rank_test_support import gloo_group @@ -66,16 +66,7 @@ def _tensor(live: LiveHead) -> torch.Tensor: def _step(trainer, monkeypatch): - monkeypatch.setattr( - trainer, - "_reduce_dynamic_grads", - lambda params, **kwargs: tuple( - torch.zeros_like(p, dtype=torch.float32) - if p.grad is None - else p.grad.float() - for p in params - ), - ) + _use_local_gradients(trainer, monkeypatch) trainer.optim_step( params=AdamParams(learning_rate=0.1, weight_decay=0.0), checkpoints=["student"] ) From 1a0bc409af3f29bea3fddad2eef8f8bd4e43c3a9 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 20 Sep 2026 02:47:56 +0000 Subject: [PATCH 023/150] Share native live-head test registration setup --- tests/unit/test_trainer_rank_live_heads.py | 99 ++++++++-------------- 1 file changed, 37 insertions(+), 62 deletions(-) diff --git a/tests/unit/test_trainer_rank_live_heads.py b/tests/unit/test_trainer_rank_live_heads.py index fdf6b8e54..e74655cd0 100644 --- a/tests/unit/test_trainer_rank_live_heads.py +++ b/tests/unit/test_trainer_rank_live_heads.py @@ -45,6 +45,12 @@ def compute(x): ) +def _native_head(kind="module", name="head", factory=TiedHead): + trainer, rank = _trainer("student") + native = getattr(rank, kind)(name, factory, checkpoint="student") + return trainer, native + + def _live_head( trainer, name, source, collector=None, *, checkpoint="student" ) -> LiveHead: @@ -76,8 +82,7 @@ def _step(trainer, monkeypatch): def test_native_old_head_graph_uses_original_tied_weights_after_step( monkeypatch, checkpointed ): - trainer, rank = _trainer("student") - head = rank.module("head", lambda: TiedHead(checkpointed), checkpoint="student") + trainer, head = _native_head(factory=lambda: TiedHead(checkpointed)) old_input = torch.tensor(3.0, requires_grad=True) old_loss = head(old_input) with trainer._gradient_transaction(): @@ -96,9 +101,8 @@ def test_native_old_head_graph_uses_original_tied_weights_after_step( def test_batchnorm_buffers_publish_once_and_old_graph_keeps_original_buffers(): - trainer, rank = _trainer("student") - head = rank.module( - "bn", lambda: torch.nn.BatchNorm1d(2, dtype=torch.float64), checkpoint="student" + trainer, head = _native_head( + "module", "bn", lambda: torch.nn.BatchNorm1d(2, dtype=torch.float64) ) initial = deepcopy(head).eval() head.eval() @@ -124,8 +128,7 @@ def forward(self, input): super().forward(input) raise RuntimeError("failed after buffer mutation") - trainer, rank = _trainer("student") - head = rank.module("bn", lambda: Failing(2), checkpoint="student") + trainer, head = _native_head("module", "bn", lambda: Failing(2)) with pytest.raises(RuntimeError, match="failed after"): head(torch.ones(4, 2)) assert head.num_batches_tracked.item() == 0 @@ -133,8 +136,7 @@ def forward(self, input): def test_client_tied_head_old_backward_after_native_refresh(monkeypatch): - trainer, rank = _trainer("student") - native = rank.module("head", TiedHead, checkpoint="student") + trainer, native = _native_head() collector = CotangentCollector() live = _live_head(trainer, "head", TiedHead(), collector) old_input = torch.tensor(3.0, requires_grad=True) @@ -153,8 +155,7 @@ def test_client_tied_head_old_backward_after_native_refresh(monkeypatch): def test_live_parameter_reuses_handle_and_earlier_capture_retains_version(): - trainer, rank = _trainer("student") - parameter = rank.parameter("gain", lambda: torch.tensor(2.0), checkpoint="student") + trainer, parameter = _native_head("parameter", "gain", lambda: torch.tensor(2.0)) collector = CotangentCollector() live = _live_head(trainer, "gain", torch.tensor(2.0), collector) handle = _tensor(live) @@ -170,8 +171,7 @@ def test_live_parameter_reuses_handle_and_earlier_capture_retains_version(): def test_client_buffer_publication_conflict_is_atomic(): - trainer, rank = _trainer("student") - native = rank.module("bn", lambda: torch.nn.BatchNorm1d(2), checkpoint="student") + trainer, native = _native_head("module", "bn", lambda: torch.nn.BatchNorm1d(2)) live = _live_head(trainer, "bn", torch.nn.BatchNorm1d(2)) _module(live)(torch.ones(4, 2)) update = live.take_publication() @@ -220,8 +220,7 @@ def test_registration_materializes_factory_state_and_persistent_buffers(): @pytest.mark.parametrize("reentrant", (False, True)) def test_external_checkpoint_requires_explicit_snapshot(reentrant): - trainer, rank = _trainer("student") - head = rank.module("head", TiedHead, checkpoint="student") + trainer, head = _native_head() x = torch.tensor(3.0, requires_grad=True) old = checkpoint(head, x, use_reentrant=reentrant) head.left.data.fill_(4) @@ -233,8 +232,7 @@ def test_external_checkpoint_requires_explicit_snapshot(reentrant): @pytest.mark.parametrize("reentrant", (False, True)) def test_explicit_snapshot_supports_external_checkpoint(reentrant): - trainer, rank = _trainer("student") - head = rank.module("head", TiedHead, checkpoint="student") + trainer, head = _native_head() captured = head.snapshot() x = torch.tensor(3.0, requires_grad=True) old = checkpoint(captured, x, use_reentrant=reentrant) @@ -340,8 +338,7 @@ def test_cuda_checkpoint_head_across_completed_optimizer_update(monkeypatch, cli def test_client_buffer_reads_keep_old_graph_and_replacement_invalidates_handle(): - trainer, rank = _trainer("student") - native = rank.buffer("scale", lambda: torch.tensor(2.0), checkpoint="student") + trainer, native = _native_head("buffer", "scale", lambda: torch.tensor(2.0)) collector = CotangentCollector() live = _live_head(trainer, "scale", torch.tensor(2.0), collector) x = torch.tensor(3.0, requires_grad=True) @@ -361,8 +358,7 @@ def test_client_buffer_reads_keep_old_graph_and_replacement_invalidates_handle() def test_client_module_explicit_dtype_move_retains_ties_and_live_parameters(): - trainer, rank = _trainer("student") - native = rank.module("head", TiedHead, checkpoint="student") + trainer, native = _native_head() live = _live_head(trainer, "head", TiedHead()) head = _module(live).to("cpu").to(dtype=torch.float64) assert head.left.dtype == torch.float64 @@ -382,8 +378,7 @@ def test_client_module_explicit_dtype_move_retains_ties_and_live_parameters(): def test_client_explicit_snapshot_supports_nonreentrant_checkpoint(): - trainer, rank = _trainer("student") - native = rank.module("head", TiedHead, checkpoint="student") + trainer, native = _native_head() collector = CotangentCollector() live = _live_head(trainer, "head", TiedHead(), collector) captured = _module(live).snapshot() @@ -400,8 +395,7 @@ def test_client_explicit_snapshot_supports_nonreentrant_checkpoint(): @pytest.mark.parametrize("client", (False, True)) def test_buffer_item_and_bitwise_mutations_publish_without_losing_handle(client): - trainer, rank = _trainer("student") - native = rank.buffer("mask", lambda: torch.tensor([1, 2]), checkpoint="student") + trainer, native = _native_head("buffer", "mask", lambda: torch.tensor([1, 2])) live = _live_head(trainer, "mask", torch.tensor([1, 2])) if client else None value = native if live is None else _tensor(live) value[0] = 4 @@ -464,8 +458,7 @@ def test_reusing_module_factory_value_does_not_share_checkpoint_storage(): def test_registration_under_no_grad_preserves_authoritative_trainability(): - trainer, rank = _trainer("student") - native = rank.module("head", TiedHead, checkpoint="student") + trainer, native = _native_head() collector = CotangentCollector() with torch.no_grad(): live = _live_head(trainer, "head", TiedHead(), collector) @@ -479,8 +472,7 @@ def test_registration_under_no_grad_preserves_authoritative_trainability(): @pytest.mark.parametrize("external", (False, True)) def test_client_reentrant_checkpoint_rejects_nested_remote_bridges(external): - trainer, rank = _trainer("student") - native = rank.module("head", TiedHead, checkpoint="student") + trainer, native = _native_head() collector = CotangentCollector() live = _live_head(trainer, "head", TiedHead(None if external else True), collector) x = torch.tensor(3.0, requires_grad=True) @@ -563,9 +555,8 @@ def test_client_reused_module_factory_does_not_share_checkpoint_handles(): def test_managed_model_operand_captures_live_parameter_and_keeps_old_version(operation): from art.trainer_rank._tensors import detach_tree - trainer, rank = _trainer("student") - native = rank.parameter( - "weight", lambda: torch.tensor([2.0, 4.0]), checkpoint="student" + trainer, native = _native_head( + "parameter", "weight", lambda: torch.tensor([2.0, 4.0]) ) collector = CotangentCollector() live = _live_head(trainer, "weight", torch.zeros(2), collector) @@ -604,8 +595,7 @@ def forward(self, value): self.count = self.count + 1 return value + self.count - trainer, rank = _trainer("student") - native = rank.module("head", Counter, checkpoint="student") + trainer, native = _native_head(factory=Counter) live = _live_head(trainer, "head", Counter()) if client else None head = native if live is None else _module(live) assert head(torch.tensor(0)).item() == 1 @@ -632,8 +622,7 @@ def forward(self, value): self.b.resize_(2) return value - trainer, rank = _trainer("student") - native = rank.module("head", Resize, checkpoint="student") + trainer, native = _native_head(factory=Resize) live = _live_head(trainer, "head", Resize()) if client else None head = native if live is None else _module(live) before = export_head(trainer, "student", "head").buffer_revision @@ -648,8 +637,7 @@ def forward(self, value): @pytest.mark.parametrize("client", (False, True)) def test_failing_handle_forward_hook_does_not_publish_buffers(client): - trainer, rank = _trainer("student") - native = rank.module("head", lambda: torch.nn.BatchNorm1d(2), checkpoint="student") + trainer, native = _native_head(factory=lambda: torch.nn.BatchNorm1d(2)) live = _live_head(trainer, "head", torch.nn.BatchNorm1d(2)) if client else None head = native if live is None else _module(live) @@ -694,8 +682,7 @@ def forward(self, value): raise RuntimeError("forward failed") return value + self.count - trainer, rank = _trainer("student") - native = rank.module("head", Counter, checkpoint="student") + trainer, native = _native_head(factory=Counter) live = _live_head(trainer, "head", Counter()) if client else None head = native if live is None else _module(live) initial = 7 if pending_before else 0 @@ -741,8 +728,7 @@ def mutate(module, amount, phase): @pytest.mark.parametrize("client", (False, True)) def test_handle_hook_parameter_gradients_keep_original_version(client): - trainer, rank = _trainer("student") - native = rank.module("head", TiedHead, checkpoint="student") + trainer, native = _native_head() collector = CotangentCollector() live = _live_head(trainer, "head", TiedHead(), collector) if client else None head = native if live is None else _module(live) @@ -767,8 +753,7 @@ def test_handle_hook_parameter_gradients_keep_original_version(client): @pytest.mark.parametrize("client", (False, True)) def test_recursive_handle_hook_does_not_publish_before_outer_failure(client): - trainer, rank = _trainer("student") - native = rank.module("head", TiedHead, checkpoint="student") + trainer, native = _native_head() live = _live_head(trainer, "head", TiedHead()) if client else None head = native if live is None else _module(live) @@ -851,10 +836,7 @@ def inner_hook(module, args): def test_inplace_operation_snapshots_readonly_checkpoint_parameter(): - trainer, rank = _trainer("student") - parameter = rank.parameter( - "weight", lambda: torch.tensor(2.0), checkpoint="student" - ) + trainer, parameter = _native_head("parameter", "weight", lambda: torch.tensor(2.0)) input_value = torch.tensor(3.0, requires_grad=True) original = (input_value * 1).mul_(parameter) parameter.data.fill_(7) @@ -1065,8 +1047,7 @@ def recover(view): def test_live_buffer_view_mutation_rejects_without_silent_write( client, grad_enabled, view ): - trainer, rank = _trainer("student") - native = rank.buffer("stats", lambda: torch.zeros(4), checkpoint="student") + trainer, native = _native_head("buffer", "stats", lambda: torch.zeros(4)) live = _live_head(trainer, "stats", torch.zeros(4)) if client else None buffer = native if live is None else _tensor(live) revision = export_head(trainer, "student", "stats").buffer_revision @@ -1089,8 +1070,7 @@ def test_live_buffer_view_mutation_rejects_without_silent_write( @pytest.mark.parametrize("client", (False, True)) def test_functional_batchnorm_no_grad_publishes_buffer_changes(client): - trainer, rank = _trainer("student") - native = rank.buffer("mean", lambda: torch.zeros(2), checkpoint="student") + trainer, native = _native_head("buffer", "mean", lambda: torch.zeros(2)) live = _live_head(trainer, "mean", torch.zeros(2)) if client else None mean = native if live is None else _tensor(live) before = export_head(trainer, "student", "mean").buffer_revision @@ -1108,8 +1088,7 @@ def test_functional_batchnorm_no_grad_publishes_buffer_changes(client): @pytest.mark.parametrize("client", (False, True)) def test_live_parameter_metadata_does_not_capture_weights(monkeypatch, client): - trainer, rank = _trainer("student") - native = rank.parameter("weight", lambda: torch.zeros(3, 4), checkpoint="student") + trainer, native = _native_head("parameter", "weight", lambda: torch.zeros(3, 4)) live = _live_head(trainer, "weight", torch.zeros(3, 4)) if client else None value = native if live is None else _tensor(live) @@ -1142,8 +1121,7 @@ def test_parameter_tensor_properties_capture_immutable_versions( if property_name == "imag": pytest.skip("imag requires a complex tensor") initial = initial.real.clone() - trainer, rank = _trainer("student") - native = rank.parameter("weight", lambda: initial.clone(), checkpoint="student") + trainer, native = _native_head("parameter", "weight", lambda: initial.clone()) collector = CotangentCollector() live = _live_head(trainer, "weight", initial, collector) if client else None value = native if live is None else _tensor(live) @@ -1181,8 +1159,7 @@ def test_parameter_tensor_properties_capture_immutable_versions( @pytest.mark.parametrize("property_name", ("T", "mT", "H", "mH", "real", "imag")) def test_buffer_tensor_properties_reject_unpublished_mutation(client, property_name): initial = torch.ones(2, 2, dtype=torch.complex64) - trainer, rank = _trainer("student") - native = rank.buffer("stats", lambda: initial.clone(), checkpoint="student") + trainer, native = _native_head("buffer", "stats", lambda: initial.clone()) live = _live_head(trainer, "stats", initial) if client else None value = native if live is None else _tensor(live) with pytest.raises(RuntimeError, match="Views of live checkpoint buffers"): @@ -1207,8 +1184,7 @@ def test_distributed_buffer_registration_mismatch_fails_on_every_rank( @pytest.mark.parametrize("client", (False, True)) def test_module_buffer_view_mutation_rejects_without_publishing(client): - trainer, rank = _trainer("student") - native = rank.module("head", lambda: torch.nn.BatchNorm1d(2), checkpoint="student") + trainer, native = _native_head(factory=lambda: torch.nn.BatchNorm1d(2)) live = _live_head(trainer, "head", torch.nn.BatchNorm1d(2)) if client else None head = native if live is None else _module(live) with pytest.raises( @@ -1222,8 +1198,7 @@ def test_module_buffer_view_mutation_rejects_without_publishing(client): @pytest.mark.parametrize("client", (False, True)) def test_stateful_function_on_buffer_snapshot_view_rejects(client): - trainer, rank = _trainer("student") - native = rank.buffer("mean", lambda: torch.zeros(2), checkpoint="student") + trainer, native = _native_head("buffer", "mean", lambda: torch.zeros(2)) live = _live_head(trainer, "mean", torch.zeros(2)) if client else None mean = native if live is None else _tensor(live) with ( From d99dde4539e5e7fa7b80445a8c23c7114941e425 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 20 Sep 2026 03:01:38 +0000 Subject: [PATCH 024/150] Share strict gradient reducer installation in validation tests --- tests/unit/test_trainer_rank_validation.py | 82 ++++++---------------- 1 file changed, 22 insertions(+), 60 deletions(-) diff --git a/tests/unit/test_trainer_rank_validation.py b/tests/unit/test_trainer_rank_validation.py index 555a02211..6c7f78eaf 100644 --- a/tests/unit/test_trainer_rank_validation.py +++ b/tests/unit/test_trainer_rank_validation.py @@ -212,6 +212,16 @@ def _output_shape(outputs: object) -> object: return [_output_shape(item) for item in outputs] +def _use_strict_local_gradients( + trainer: TrainerRank, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr( + trainer, + "_reduce_dynamic_grads", + lambda params, **_kwargs: tuple(item.grad.float() for item in params), + ) + + def _trainer_with_checkpoint( monkeypatch: pytest.MonkeyPatch, value: torch.Tensor, @@ -219,11 +229,7 @@ def _trainer_with_checkpoint( trainer = TrainerRank(_runtime()) param = torch.nn.Parameter(value.clone()) trainer._checkpoint_slots.setdefault("student", _CheckpointSlot()).params = (param,) - monkeypatch.setattr( - trainer, - "_reduce_dynamic_grads", - lambda params, **_kwargs: tuple(item.grad.float() for item in params), - ) + _use_strict_local_gradients(trainer, monkeypatch) return trainer, param @@ -1650,11 +1656,7 @@ def test_weights_only_load_replaces_stale_optimizer_and_recreates_it_lazily( assert stale is not slot.optimizer replacement.grad = torch.ones_like(replacement) - monkeypatch.setattr( - trainer, - "_reduce_dynamic_grads", - lambda params, **_kwargs: tuple(item.grad.float() for item in params), - ) + _use_strict_local_gradients(trainer, monkeypatch) result = trainer.optim_step( checkpoints=["student"], params=AdamParams(learning_rate=1e-3, weight_decay=0), @@ -2386,11 +2388,7 @@ def make_trainer() -> TrainerRank: "student", loaded, set(adapter) ) trainer._checkpoint_slots["student"] = _CheckpointSlot(params, config) - monkeypatch.setattr( - trainer, - "_reduce_dynamic_grads", - lambda params, **_kwargs: tuple(item.grad.float() for item in params), - ) + _use_strict_local_gradients(trainer, monkeypatch) return trainer original = make_trainer() @@ -2411,11 +2409,7 @@ def make_trainer() -> TrainerRank: restored_lora = LoRA("layer.q_proj", 3, 4, 2, 2, torch.float32, torch.device("cpu")) restored = TrainerRank(_runtime(restored_lora)) - monkeypatch.setattr( - restored, - "_reduce_dynamic_grads", - lambda params, **_kwargs: tuple(item.grad.float() for item in params), - ) + _use_strict_local_gradients(restored, monkeypatch) checkpoint_module.load_checkpoint(restored, prepared, "student") assert ( restored._checkpoint_slots["student"].optimizer is not None @@ -2528,11 +2522,7 @@ def test_optim_step_rejects_explicit_slot_subset_with_missing_grads( trainer._checkpoint_slots.setdefault("missing", _CheckpointSlot()).params = ( missing, ) - monkeypatch.setattr( - trainer, - "_reduce_dynamic_grads", - lambda params, **_kwargs: tuple(param.grad.float() for param in params), - ) + _use_strict_local_gradients(trainer, monkeypatch) with pytest.raises(TrainerRankSlotStateError, match="missing"): trainer.optim_step( @@ -2552,11 +2542,7 @@ def test_optim_step_implicitly_steps_only_slots_with_grads( trainer._checkpoint_slots.setdefault("untouched", _CheckpointSlot()).params = ( untouched, ) - monkeypatch.setattr( - trainer, - "_reduce_dynamic_grads", - lambda params, **_kwargs: tuple(param.grad.float() for param in params), - ) + _use_strict_local_gradients(trainer, monkeypatch) before_ready = ready.detach().clone() before_untouched = untouched.detach().clone() @@ -2645,11 +2631,7 @@ def zero_grad(self, *, set_to_none: bool = False) -> None: master_params=(master,), optimizer=RecordingOptimizer(name, master) ) - monkeypatch.setattr( - trainer, - "_reduce_dynamic_grads", - lambda params, **_kwargs: tuple(param.grad.float() for param in params), - ) + _use_strict_local_gradients(trainer, monkeypatch) monkeypatch.setattr( trainer, "_dynamic_optimizer", lambda name, _params: dynamics[name] ) @@ -2691,11 +2673,7 @@ def test_optim_step_checks_all_checkpoint_grads_before_stepping( param = torch.nn.Parameter(torch.ones(1)) param.grad = torch.full_like(param, grad) trainer._checkpoint_slots[name] = _CheckpointSlot(params=(param,)) - monkeypatch.setattr( - trainer, - "_reduce_dynamic_grads", - lambda params, **_kwargs: tuple(param.grad.float() for param in params), - ) + _use_strict_local_gradients(trainer, monkeypatch) monkeypatch.setattr( trainer, "_dynamic_optimizer", @@ -2754,11 +2732,7 @@ def test_optim_step_allows_either_configuration_to_be_mapped( param = torch.nn.Parameter(torch.ones(1)) param.grad = torch.ones_like(param) trainer._checkpoint_slots["student"] = _CheckpointSlot(params=(param,)) - monkeypatch.setattr( - trainer, - "_reduce_dynamic_grads", - lambda params, **_kwargs: tuple(param.grad.float() for param in params), - ) + _use_strict_local_gradients(trainer, monkeypatch) adam = AdamParams(learning_rate=1e-3, weight_decay=0.0) trainer.optim_step( @@ -2778,11 +2752,7 @@ def test_optim_step_prepares_all_optimizers_before_first_update( param = torch.nn.Parameter(torch.ones(1)) param.grad = torch.ones_like(param) trainer._checkpoint_slots[name] = _CheckpointSlot(params=(param,)) - monkeypatch.setattr( - trainer, - "_reduce_dynamic_grads", - lambda params, **_kwargs: tuple(param.grad.float() for param in params), - ) + _use_strict_local_gradients(trainer, monkeypatch) class RecordingOptimizer: def step(self) -> None: @@ -2850,11 +2820,7 @@ def test_optim_step_implicitly_ignores_resident_forward_snapshot( list( trainer.forward_batches([_target_request(1)], checkpoint="saved", no_grad=True) ) - monkeypatch.setattr( - trainer, - "_reduce_dynamic_grads", - lambda params, **_kwargs: tuple(param.grad.float() for param in params), - ) + _use_strict_local_gradients(trainer, monkeypatch) before_student = student.detach().clone() before_snapshot = snapshot.detach().clone() @@ -2900,11 +2866,7 @@ def zero_padding_grads(_model: object) -> None: ) trainer = TrainerRank(runtime) trainer._checkpoint_slots.setdefault("student", _CheckpointSlot()).params = (param,) - monkeypatch.setattr( - trainer, - "_reduce_dynamic_grads", - lambda params, **_kwargs: tuple(item.grad.float() for item in params), - ) + _use_strict_local_gradients(trainer, monkeypatch) trainer.optim_step( params=AdamParams( From ee5415eebbd5ed3a615f689d7d630874cdbac530 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 20 Sep 2026 03:05:42 +0000 Subject: [PATCH 025/150] Share native and live head tensor inventories --- src/art/trainer_rank/_heads.py | 53 ++++++++++++++-------------------- 1 file changed, 21 insertions(+), 32 deletions(-) diff --git a/src/art/trainer_rank/_heads.py b/src/art/trainer_rank/_heads.py index 37abb1824..e65978edf 100644 --- a/src/art/trainer_rank/_heads.py +++ b/src/art/trainer_rank/_heads.py @@ -545,25 +545,22 @@ def _custom_tracker(custom: _CustomObject) -> _CustomTensorTracker: return cast(Any, custom.value)._art_tracker -def _custom_parameters(custom: _CustomObject) -> dict[str, torch.nn.Parameter]: - if custom.kind == "module": - return dict(cast(torch.nn.Module, custom.value).named_parameters()) - return ( - {"": cast(torch.nn.Parameter, custom.value)} - if custom.kind == "parameter" - else {} - ) - - -def _custom_buffers(custom: _CustomObject) -> dict[str, torch.Tensor]: - if custom.kind == "module": - return dict(cast(torch.nn.Module, custom.value).named_buffers()) - return {"": cast(torch.Tensor, custom.value)} if custom.kind == "buffer" else {} +def _head_tensors( + kind: HeadKind, source: torch.nn.Module | torch.Tensor, *, parameters: bool +) -> dict[str, torch.Tensor]: + if kind == "module": + module = cast(torch.nn.Module, source) + return dict(module.named_parameters() if parameters else module.named_buffers()) + target_kind = "parameter" if parameters else "buffer" + return {"": cast(torch.Tensor, source)} if kind == target_kind else {} def export_head(trainer: TrainerRank, checkpoint: str, name: str) -> HeadState: custom = trainer._checkpoint_slots[checkpoint].custom[name] - parameters, buffers = _custom_parameters(custom), _custom_buffers(custom) + parameters, buffers = ( + _head_tensors(custom.kind, custom.value, parameters=True), + _head_tensors(custom.kind, custom.value, parameters=False), + ) return HeadState( trainer._capture_checkpoint_version(checkpoint), name, @@ -721,7 +718,7 @@ def _stage_buffer_publications(trainer: TrainerRank, updates: Any) -> list[Any]: raise RuntimeError( f"Custom module {update.name!r} buffers changed before publication" ) - targets = _custom_buffers(custom) + targets = _head_tensors(custom.kind, custom.value, parameters=False) if update.buffers.keys() != targets.keys(): raise ValueError("Custom module buffer keys changed before publication") values = [] @@ -752,7 +749,7 @@ def head_gradient_targets( version, max_gradient_staleness=metadata["max_gradient_staleness"] ) custom = trainer._checkpoint_slots[version.checkpoint].custom[metadata["name"]] - parameters = _custom_parameters(custom) + parameters = _head_tensors(custom.kind, custom.value, parameters=True) if len(metadata["keys"]) != len(packet.gradients): raise ValueError("Custom head cotangent count does not match parameters") for key, gradient in zip(metadata["keys"], packet.gradients, strict=True): @@ -772,7 +769,7 @@ def head_gradient_targets( ( version, metadata["max_gradient_staleness"], - parameters[key], + cast(torch.nn.Parameter, parameters[key]), gradient.to(device=parameters[key].device, dtype=parameters[key].dtype) if materialize else gradient, @@ -1086,20 +1083,10 @@ def invalidate(self, reason: str) -> None: self.invalid_reason = reason def parameters(self) -> dict[str, torch.Tensor]: - if self.state.kind == "module": - return dict(cast(torch.nn.Module, self.source).named_parameters()) - return ( - {"": cast(torch.Tensor, self.source)} - if self.state.kind == "parameter" - else {} - ) + return _head_tensors(self.state.kind, self.source, parameters=True) def buffers(self) -> dict[str, torch.Tensor]: - if self.state.kind == "module": - return dict(cast(torch.nn.Module, self.source).named_buffers()) - return ( - {"": cast(torch.Tensor, self.source)} if self.state.kind == "buffer" else {} - ) + return _head_tensors(self.state.kind, self.source, parameters=False) def capture(self, keys: tuple[str, ...] | None = None) -> dict[str, torch.Tensor]: import json @@ -1235,7 +1222,7 @@ def synchronize_head_buffers(trainer: TrainerRank, checkpoints: Any = None) -> N if custom.kind == "parameter": continue if custom.kind == "buffer": - buffers = _custom_buffers(custom) + buffers = _head_tensors(custom.kind, custom.value, parameters=False) else: persistent = { id(buffer) @@ -1246,7 +1233,9 @@ def synchronize_head_buffers(trainer: TrainerRank, checkpoints: Any = None) -> N } buffers = { key: value - for key, value in _custom_buffers(custom).items() + for key, value in _head_tensors( + custom.kind, custom.value, parameters=False + ).items() if id(value) in persistent } targets[(checkpoint, name)] = (_custom_tracker(custom), buffers) From fb3be23d0b2ca0884cbdc271c8aa1e92caa0999b Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 20 Sep 2026 03:17:06 +0000 Subject: [PATCH 026/150] test: share trainer snapshot setup in version tests --- tests/unit/test_trainer_rank_versions.py | 51 +++++++----------------- 1 file changed, 15 insertions(+), 36 deletions(-) diff --git a/tests/unit/test_trainer_rank_versions.py b/tests/unit/test_trainer_rank_versions.py index 1caba182f..19cf9d2b2 100644 --- a/tests/unit/test_trainer_rank_versions.py +++ b/tests/unit/test_trainer_rank_versions.py @@ -24,6 +24,12 @@ def _trainer() -> tuple[TrainerRank, torch.nn.Parameter]: return trainer, parameter +def _snapshot_trainer() -> tuple[TrainerRank, torch.nn.Parameter, torch.nn.Parameter]: + trainer, current = _trainer() + version = trainer._capture_checkpoint_version("student") + return trainer, current, trainer._snapshot_parameter(current, version) + + @pytest.mark.parametrize("reentrant", (False, True)) def test_snapshot_recompute_routes_original_gradient_after_update( reentrant: bool, @@ -44,10 +50,7 @@ def test_snapshot_recompute_routes_original_gradient_after_update( def test_coupled_versions_accumulate_and_repeated_backward_routes_once() -> None: - trainer, current = _trainer() - old = trainer._snapshot_parameter( - current, trainer._capture_checkpoint_version("student") - ) + trainer, current, old = _snapshot_trainer() with torch.no_grad(): current.fill_(3) trainer._checkpoint_slots["student"].revision += 1 @@ -81,10 +84,7 @@ def test_stale_backward_preserves_existing_current_gradient() -> None: def test_failed_backward_discards_staged_gradients_and_releases_batch() -> None: - trainer, current = _trainer() - old = trainer._snapshot_parameter( - current, trainer._capture_checkpoint_version("student") - ) + trainer, current, old = _snapshot_trainer() def fail(_gradient: torch.Tensor) -> None: raise RuntimeError("local autograd failed") @@ -123,10 +123,7 @@ def test_snapshot_backward_requires_atomic_scope_for_outer_failure( def test_reentrant_backward_transaction_rolls_back_completed_nested_task() -> None: - trainer, current = _trainer() - old = trainer._snapshot_parameter( - current, trainer._capture_checkpoint_version("student") - ) + trainer, current, old = _snapshot_trainer() x = torch.tensor(3.0, dtype=torch.float64, requires_grad=True) loss = checkpoint(lambda x: x * old.square(), x, use_reentrant=True) with pytest.raises(RuntimeError, match="later backward failed"): @@ -196,10 +193,7 @@ def test_accumulated_origin_is_checked_before_optimizer_mutation() -> None: def test_snapshot_lifetime_follows_graph_references() -> None: - trainer, current = _trainer() - snapshot = trainer._snapshot_parameter( - current, trainer._capture_checkpoint_version("student") - ) + trainer, current, snapshot = _snapshot_trainer() reference = weakref.ref(snapshot) loss = snapshot.square() del snapshot @@ -225,10 +219,7 @@ def test_replay_capture_does_not_reset_origin_age() -> None: def test_collective_preflight_failure_does_not_commit_local_gradients() -> None: - trainer, current = _trainer() - old = trainer._snapshot_parameter( - current, trainer._capture_checkpoint_version("student") - ) + trainer, current, old = _snapshot_trainer() def other_rank_failed(validate) -> None: validate() @@ -258,10 +249,7 @@ def test_gradient_dtype_is_validated_before_any_accumulator_is_published() -> No def test_nested_transaction_still_participates_in_collective_preflight() -> None: - trainer, current = _trainer() - old = trainer._snapshot_parameter( - current, trainer._capture_checkpoint_version("student") - ) + trainer, current, old = _snapshot_trainer() calls = [] def collective(validate) -> None: @@ -279,10 +267,7 @@ def collective(validate) -> None: def test_caught_nested_backward_failure_invalidates_whole_transaction() -> None: - trainer, current = _trainer() - old = trainer._snapshot_parameter( - current, trainer._capture_checkpoint_version("student") - ) + trainer, current, old = _snapshot_trainer() with pytest.raises(RuntimeError, match="nested gradient transaction failed"): with trainer._gradient_transaction(): with pytest.raises(RuntimeError, match="failed"): @@ -345,10 +330,7 @@ def test_divergent_optimizer_provenance_rejects_collectively_before_mutation( def test_transaction_coalesces_many_children_but_keeps_each_origin() -> None: - trainer, current = _trainer() - old = trainer._snapshot_parameter( - current, trainer._capture_checkpoint_version("student") - ) + trainer, current, old = _snapshot_trainer() trainer._checkpoint_slots["student"].revision = 1 new = trainer._snapshot_parameter( current, trainer._capture_checkpoint_version("student") @@ -373,10 +355,7 @@ def test_transaction_coalesces_many_children_but_keeps_each_origin() -> None: def test_retained_failed_transaction_traceback_releases_staging() -> None: - trainer, current = _trainer() - old = trainer._snapshot_parameter( - current, trainer._capture_checkpoint_version("student") - ) + trainer, current, old = _snapshot_trainer() saved_error = None reference = None try: From f197898ba208d547266ac8ed0cbb3125aa410146 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 20 Sep 2026 03:57:37 +0000 Subject: [PATCH 027/150] test: share CP2 attention layout setup --- .../cp_attn/test_cpu_offload_residency.py | 43 +-------------- .../cp_attn/test_retained_backward.py | 41 +------------- tests/support/cp_attention.py | 54 +++++++++++++++++++ 3 files changed, 58 insertions(+), 80 deletions(-) create mode 100644 tests/support/cp_attention.py diff --git a/tests/integration/megatron/cp_attn/test_cpu_offload_residency.py b/tests/integration/megatron/cp_attn/test_cpu_offload_residency.py index dcda63aad..c80657e20 100644 --- a/tests/integration/megatron/cp_attn/test_cpu_offload_residency.py +++ b/tests/integration/megatron/cp_attn/test_cpu_offload_residency.py @@ -5,7 +5,6 @@ import gc import json import os -from typing import cast from unittest.mock import patch import pytest @@ -16,21 +15,14 @@ pytest.importorskip("megatron.core") from art.megatron.context_parallel import executor # noqa: E402 -from art.megatron.context_parallel.runtime import ( - prepare_megatron_context_parallel_state, -) # noqa: E402 -from art.megatron.context_parallel.types import ( # noqa: E402 - ContextParallelConfig, - ParallelTopology, -) from art.megatron.flex_attn import compiled # noqa: E402 from art.megatron.runtime.compile_cache import configure_reusable_backward # noqa: E402 -from art.preprocessing.pack import PackedTensors # noqa: E402 from art.trainer_rank._graphs import GraphCache # noqa: E402 from art.trainer_rank._memory_policy import ( # noqa: E402 ForwardMemoryCost, placement_cost, ) +from tests.support.cp_attention import prepare_cp2_attention # noqa: E402 @pytest.mark.skipif( @@ -72,38 +64,7 @@ def _worker(rank, rendezvous): def _check(rank, device): length, heads, dim = 512, 2, 64 - micro = cast( - PackedTensors, - { - "tokens": torch.arange(length)[None], - "group_ids": torch.ones((1, length), dtype=torch.long), - "parent_ids": torch.ones((1, length), dtype=torch.long), - "input_pos": torch.arange(length)[None], - }, - ) - state, plan, _, _ = prepare_megatron_context_parallel_state( - micro=micro, - topology=ParallelTopology(cp=2), - config=ContextParallelConfig( - planner_chunk_size=128, planner_owned_token_ms=1.0 - ), - cp_group=dist.group.WORLD, - cp_rank=rank, - target_device=device, - ) - indices = torch.tensor( - [ - index - for start, end, _local_start in plan.token_layout_index.ownership_ranges_by_rank[ - rank - ] - for index in range(start, end) - ], - device=device, - ) - assert indices.numel() > 0 - assert indices.numel() == sum(plan.local_valid_lengths) - executor.prepare_context_parallel_execution_state(state=state, device=device) + micro, state, plan, indices = prepare_cp2_attention(rank, device, length) torch.manual_seed(841) full = tuple(torch.randn(length, 1, heads, dim).to(device) for _ in range(3)) local = tuple(value.index_select(0, indices) for value in full) diff --git a/tests/integration/megatron/cp_attn/test_retained_backward.py b/tests/integration/megatron/cp_attn/test_retained_backward.py index 0ece0cc27..870d1adc5 100644 --- a/tests/integration/megatron/cp_attn/test_retained_backward.py +++ b/tests/integration/megatron/cp_attn/test_retained_backward.py @@ -3,7 +3,6 @@ from datetime import timedelta import os from pathlib import Path -from typing import cast from unittest.mock import patch import weakref @@ -15,17 +14,10 @@ pytest.importorskip("megatron.core") from art.megatron.context_parallel import executor # noqa: E402 -from art.megatron.context_parallel.runtime import ( # noqa: E402 - prepare_megatron_context_parallel_state, -) -from art.megatron.context_parallel.types import ( # noqa: E402 - ContextParallelConfig, - ParallelTopology, -) from art.megatron.flex_attn import compiled # noqa: E402 from art.megatron.runtime.compile_cache import configure_reusable_backward # noqa: E402 -from art.preprocessing.pack import PackedTensors # noqa: E402 from art.trainer_rank._graphs import GraphCache # noqa: E402 +from tests.support.cp_attention import prepare_cp2_attention # noqa: E402 def test_cp_retained_failure_releases_original_records(monkeypatch): @@ -113,36 +105,7 @@ def _check_repeated_backward( rank: int, device: torch.device, backend: str, dim: int ) -> None: length, heads = 512, 2 - micro = cast( - PackedTensors, - { - "tokens": torch.arange(length)[None], - "group_ids": torch.ones((1, length), dtype=torch.long), - "parent_ids": torch.ones((1, length), dtype=torch.long), - "input_pos": torch.arange(length)[None], - }, - ) - state, plan, _, _ = prepare_megatron_context_parallel_state( - micro=micro, - topology=ParallelTopology(cp=2), - config=ContextParallelConfig( - planner_chunk_size=128, planner_owned_token_ms=1.0 - ), - cp_group=dist.group.WORLD, - cp_rank=rank, - target_device=device, - ) - indices = torch.tensor( - [ - index - for start, end, _ in plan.token_layout_index.ownership_ranges_by_rank[rank] - for index in range(start, end) - ], - device=device, - ) - assert indices.numel() > 0 - assert indices.numel() == sum(plan.local_valid_lengths) - executor.prepare_context_parallel_execution_state(state=state, device=device) + micro, state, plan, indices = prepare_cp2_attention(rank, device, length) torch.manual_seed(841) dtype = torch.bfloat16 if backend == "FLASH" else torch.float32 full = tuple( diff --git a/tests/support/cp_attention.py b/tests/support/cp_attention.py new file mode 100644 index 000000000..aa07831c2 --- /dev/null +++ b/tests/support/cp_attention.py @@ -0,0 +1,54 @@ +"""Shared CP2 layout setup for attention and graph-residency tests.""" + +from typing import cast + +import torch +import torch.distributed as dist + +from art.megatron.context_parallel import executor +from art.megatron.context_parallel.runtime import ( + prepare_megatron_context_parallel_state, +) +from art.megatron.context_parallel.types import ( + ArtContextParallelState, + ContextParallelConfig, + ParallelTopology, + RankRuntimePlan, +) +from art.preprocessing.pack import PackedTensors + + +def prepare_cp2_attention( + rank: int, device: torch.device, length: int +) -> tuple[PackedTensors, ArtContextParallelState, RankRuntimePlan, torch.Tensor]: + micro = cast( + PackedTensors, + { + "tokens": torch.arange(length)[None], + "group_ids": torch.ones((1, length), dtype=torch.long), + "parent_ids": torch.ones((1, length), dtype=torch.long), + "input_pos": torch.arange(length)[None], + }, + ) + state, plan, _, _ = prepare_megatron_context_parallel_state( + micro=micro, + topology=ParallelTopology(cp=2), + config=ContextParallelConfig( + planner_chunk_size=128, planner_owned_token_ms=1.0 + ), + cp_group=dist.group.WORLD, + cp_rank=rank, + target_device=device, + ) + indices = torch.tensor( + [ + index + for start, end, _ in plan.token_layout_index.ownership_ranges_by_rank[rank] + for index in range(start, end) + ], + device=device, + ) + assert indices.numel() > 0 + assert indices.numel() == sum(plan.local_valid_lengths) + executor.prepare_context_parallel_execution_state(state=state, device=device) + return micro, state, plan, indices From 36b127f65d4c2e87e68610ed89b136d32c15ea0f Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 20 Sep 2026 03:56:43 +0000 Subject: [PATCH 028/150] test: finish sharing initial version snapshots --- tests/unit/test_trainer_rank_versions.py | 8 ++------ 1 file changed, 2 insertions(+), 6 deletions(-) diff --git a/tests/unit/test_trainer_rank_versions.py b/tests/unit/test_trainer_rank_versions.py index 19cf9d2b2..1a89fa519 100644 --- a/tests/unit/test_trainer_rank_versions.py +++ b/tests/unit/test_trainer_rank_versions.py @@ -34,9 +34,7 @@ def _snapshot_trainer() -> tuple[TrainerRank, torch.nn.Parameter, torch.nn.Param def test_snapshot_recompute_routes_original_gradient_after_update( reentrant: bool, ) -> None: - trainer, current = _trainer() - version = trainer._capture_checkpoint_version("student") - old = trainer._snapshot_parameter(current, version) + trainer, current, old = _snapshot_trainer() x = torch.tensor(3.0, dtype=torch.float64, requires_grad=True) loss = checkpoint(lambda x: x * old.square(), x, use_reentrant=reentrant) with torch.no_grad(): @@ -150,9 +148,7 @@ def test_cotangent_batch_preflights_all_versions_before_any_mutation() -> None: def test_replacement_invalidates_origin_and_old_target() -> None: - trainer, current = _trainer() - version = trainer._capture_checkpoint_version("student") - snapshot = trainer._snapshot_parameter(current, version) + trainer, current, snapshot = _snapshot_trainer() new = torch.nn.Parameter(torch.tensor(5.0, dtype=torch.float64)) trainer._checkpoint_slots["student"] = _CheckpointSlot(params=(new,), generation=1) with pytest.raises(TrainerRankSlotStateError, match="replaced"): From de6a8d952b2b92060e12d7e4ab3a9504b0f2997b Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 20 Sep 2026 04:16:16 +0000 Subject: [PATCH 029/150] test: share inline trainer operation capture --- tests/unit/test_trainer_operations.py | 93 ++++++++------------------- 1 file changed, 28 insertions(+), 65 deletions(-) diff --git a/tests/unit/test_trainer_operations.py b/tests/unit/test_trainer_operations.py index a6ca3d5f5..32ad649c9 100644 --- a/tests/unit/test_trainer_operations.py +++ b/tests/unit/test_trainer_operations.py @@ -1,4 +1,5 @@ import asyncio +from collections.abc import Coroutine import gc from types import SimpleNamespace from typing import Any, cast @@ -8,12 +9,20 @@ import torch from art.trainer_rank._operations import ( + OperationId, OperationResultReleasedError, TrainerOperation, execute_operation, ) +def _execute( + rank: Any, identity: OperationId, kind: str, payload: object +) -> Coroutine[Any, Any, Any]: + """Capture inline requests before returning the execution coroutine.""" + return execute_operation(rank, TrainerOperation.capture(identity, kind, payload)) + + def test_update_identity_replays_outcome_without_applying_again(): async def run(): calls = [] @@ -28,15 +37,8 @@ async def run(): assert await execute_operation(rank, operation) == {"step": 1} assert len(calls) == 1 with pytest.raises(ValueError, match="different arguments"): - await execute_operation( - rank, - TrainerOperation.capture( - ("client", 1), "optim_step", {"params": {"lr": 2}} - ), - ) - await execute_operation( - rank, TrainerOperation.capture(("client", 1), "acknowledge", ((), ())) - ) + await _execute(rank, ("client", 1), "optim_step", {"params": {"lr": 2}}) + await _execute(rank, ("client", 1), "acknowledge", ((), ())) with pytest.raises(OperationResultReleasedError): await execute_operation(rank, operation) assert len(calls) == 1 @@ -65,9 +67,7 @@ def backward_packets(**kwargs): assert type(error.value) is type(stale) assert str(error.value) == str(stale) assert len(calls) == 1 - await execute_operation( - rank, TrainerOperation.capture(operation.id, "acknowledge", ((), ())) - ) + await _execute(rank, operation.id, "acknowledge", ((), ())) with pytest.raises(OperationResultReleasedError): await execute_operation(rank, operation) assert len(calls) == 1 @@ -158,15 +158,8 @@ async def run(): ) # ID 1 is still pending. Odd IDs after it were cancelled before admission. for sequence in range(2, 2002, 2): - await execute_operation( - rank, TrainerOperation.capture(("client", sequence), "optim_step", {}) - ) - await execute_operation( - rank, - TrainerOperation.capture( - ("client", sequence), "acknowledge", ((1,), ()) - ), - ) + await _execute(rank, ("client", sequence), "optim_step", {}) + await _execute(rank, ("client", sequence), "acknowledge", ((1,), ())) ledger = rank._rank._operation_outcomes assert not ledger.outcomes assert len(ledger.acknowledged) == 1 @@ -174,16 +167,9 @@ async def run(): assert state.through == 2000 and state.pending == {1} for sequence in (2, 3, 1999, 2000): with pytest.raises(OperationResultReleasedError): - await execute_operation( - rank, - TrainerOperation.capture(("client", sequence), "optim_step", {}), - ) - await execute_operation( - rank, TrainerOperation.capture(("client", 1), "optim_step", {}) - ) - await execute_operation( - rank, TrainerOperation.capture(("client", 2000), "acknowledge", ((), ())) - ) + await _execute(rank, ("client", sequence), "optim_step", {}) + await _execute(rank, ("client", 1), "optim_step", {}) + await _execute(rank, ("client", 2000), "acknowledge", ((), ())) assert not state.pending and not ledger.outcomes assert len(calls) == 1001 @@ -202,20 +188,12 @@ async def run(): (6, (1, 4, 6)), (5, (1, 2, 3)), ): - await execute_operation( - rank, - TrainerOperation.capture( - ("client", through), "acknowledge", (pending, ()) - ), - ) + await _execute(rank, ("client", through), "acknowledge", (pending, ())) state = rank._rank._operation_outcomes.acknowledged["client"] assert state.through == 6 and state.pending == {1, 6} for sequence in (2, 3, 4, 5): with pytest.raises(OperationResultReleasedError): - await execute_operation( - rank, - TrainerOperation.capture(("client", sequence), "optim_step", {}), - ) + await _execute(rank, ("client", sequence), "optim_step", {}) for identity in (("client", 1), ("client", 6), ("client", 7), ("other", 3)): operation = TrainerOperation.capture(identity, "optim_step", {}) await execute_operation(rank, operation) @@ -240,9 +218,7 @@ async def optim_step(): operation = TrainerOperation.capture(("client", 1), "optim_step", {}) original = asyncio.create_task(execute_operation(rank, operation)) await entered.wait() - await execute_operation( - rank, TrainerOperation.capture(operation.id, "acknowledge", ((), ())) - ) + await _execute(rank, operation.id, "acknowledge", ((), ())) with pytest.raises(OperationResultReleasedError): await execute_operation(rank, operation) release.set() @@ -275,13 +251,9 @@ async def run(): first = await execute_operation(rank, next_wave) assert await execute_operation(rank, next_wave) is first # Other results can be acknowledged while this wave's reply is lost. - await execute_operation( - rank, TrainerOperation.capture(("client", 3), "acknowledge", ((2, 3), ())) - ) + await _execute(rank, ("client", 3), "acknowledge", ((2, 3), ())) # A lost pull is abandoned independently of closing its iterator. - await execute_operation( - rank, TrainerOperation.capture(("client", 3), "acknowledge", ((3,), (2,))) - ) + await _execute(rank, ("client", 3), "acknowledge", ((3,), (2,))) close = TrainerOperation.capture( ("client", 3), "batches_close", @@ -295,9 +267,7 @@ async def run(): ("release", ("packet",)), ("close", "iterator"), ] - await execute_operation( - rank, TrainerOperation.capture(close.id, "acknowledge", ((), ())) - ) + await _execute(rank, close.id, "acknowledge", ((), ())) assert not rank._rank._operation_outcomes.outcomes with pytest.raises(OperationResultReleasedError): await execute_operation(rank, next_wave) @@ -447,12 +417,7 @@ def fail(*args, **kwargs): hook.remove() if retain_graph: setattr(view, "_submit_backward", lambda *args, **kwargs: None) - await execute_operation( - view, - TrainerOperation.capture( - ("client", 2), "backward", {"packets": packets} - ), - ) + await _execute(view, ("client", 2), "backward", {"packets": packets}) gc.collect() assert packet.handle not in state.exports and references[0]() is None @@ -473,13 +438,11 @@ async def run(): cast(Any, SimpleNamespace(rank=SimpleNamespace(), state=state)) ) with pytest.raises((ValueError, KeyError)): - await execute_operation( + await _execute( view, - TrainerOperation.capture( - ("client", 1), - "backward", - {"packets": (CotangentPacket(handle, gradients),)}, - ), + ("client", 1), + "backward", + {"packets": (CotangentPacket(handle, gradients),)}, ) assert "unrelated" in state.exports assert ("known" in state.exports) is (handle == "missing") From 30ef5154c36e875887e6b7ac8ba3378fe21b027f Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 20 Sep 2026 04:33:13 +0000 Subject: [PATCH 030/150] test: run trainer operation cases directly as async tests --- tests/unit/test_trainer_operations.py | 709 ++++++++++++-------------- 1 file changed, 330 insertions(+), 379 deletions(-) diff --git a/tests/unit/test_trainer_operations.py b/tests/unit/test_trainer_operations.py index 32ad649c9..d6b70146e 100644 --- a/tests/unit/test_trainer_operations.py +++ b/tests/unit/test_trainer_operations.py @@ -23,260 +23,223 @@ def _execute( return execute_operation(rank, TrainerOperation.capture(identity, kind, payload)) -def test_update_identity_replays_outcome_without_applying_again(): - async def run(): - calls = [] - rank = SimpleNamespace( - _rank=SimpleNamespace(), - optim_step=lambda **kwargs: calls.append(kwargs) or {"step": len(calls)}, - ) - operation = TrainerOperation.capture( - ("client", 1), "optim_step", {"params": {"lr": 1}} - ) - assert await execute_operation(rank, operation) == {"step": 1} - assert await execute_operation(rank, operation) == {"step": 1} - assert len(calls) == 1 - with pytest.raises(ValueError, match="different arguments"): - await _execute(rank, ("client", 1), "optim_step", {"params": {"lr": 2}}) - await _execute(rank, ("client", 1), "acknowledge", ((), ())) - with pytest.raises(OperationResultReleasedError): - await execute_operation(rank, operation) - assert len(calls) == 1 - - asyncio.run(run()) - - -def test_failed_gradient_identity_preserves_original_error(): +async def test_update_identity_replays_outcome_without_applying_again(): + calls = [] + rank = SimpleNamespace( + _rank=SimpleNamespace(), + optim_step=lambda **kwargs: calls.append(kwargs) or {"step": len(calls)}, + ) + operation = TrainerOperation.capture( + ("client", 1), "optim_step", {"params": {"lr": 1}} + ) + assert await execute_operation(rank, operation) == {"step": 1} + assert await execute_operation(rank, operation) == {"step": 1} + assert len(calls) == 1 + with pytest.raises(ValueError, match="different arguments"): + await _execute(rank, ("client", 1), "optim_step", {"params": {"lr": 2}}) + await _execute(rank, ("client", 1), "acknowledge", ((), ())) + with pytest.raises(OperationResultReleasedError): + await execute_operation(rank, operation) + assert len(calls) == 1 + + +async def test_failed_gradient_identity_preserves_original_error(): from art.trainer_rank import TrainerRankSlotStateError - async def run(): - stale = TrainerRankSlotStateError("original forward is stale") - calls = [] + stale = TrainerRankSlotStateError("original forward is stale") + calls = [] - def backward_packets(**kwargs): - calls.append(kwargs) - raise stale + def backward_packets(**kwargs): + calls.append(kwargs) + raise stale - rank = SimpleNamespace( - _rank=SimpleNamespace(), backward_packets=backward_packets - ) - operation = TrainerOperation.capture(("client", 1), "backward", {"packets": ()}) - for _ in range(2): - with pytest.raises(TrainerRankSlotStateError) as error: - await execute_operation(rank, operation) - assert type(error.value) is type(stale) - assert str(error.value) == str(stale) - assert len(calls) == 1 - await _execute(rank, operation.id, "acknowledge", ((), ())) - with pytest.raises(OperationResultReleasedError): + rank = SimpleNamespace(_rank=SimpleNamespace(), backward_packets=backward_packets) + operation = TrainerOperation.capture(("client", 1), "backward", {"packets": ()}) + for _ in range(2): + with pytest.raises(TrainerRankSlotStateError) as error: await execute_operation(rank, operation) - assert len(calls) == 1 - - asyncio.run(run()) - - -def test_concurrent_retry_waiter_cancellation_does_not_cancel_update(): - async def run(): - entered, release = asyncio.Event(), asyncio.Event() - calls = 0 - - async def optim_step(): - nonlocal calls - calls += 1 - entered.set() - await release.wait() - return calls - - rank = SimpleNamespace(_rank=SimpleNamespace(), optim_step=optim_step) - operation = TrainerOperation.capture(("client", 1), "optim_step", {}) - original = asyncio.create_task(execute_operation(rank, operation)) - await entered.wait() - retry = asyncio.create_task(execute_operation(rank, operation)) - await asyncio.sleep(0) - retry.cancel() - with pytest.raises(asyncio.CancelledError): - await retry - release.set() - assert await original == 1 - assert await execute_operation(rank, operation) == 1 - - asyncio.run(run()) - - -def test_operation_captures_tensor_arguments_at_submission(): - async def run(): - source = torch.tensor([2.0]) - operation = TrainerOperation.capture( - ("client", 1), "forward", {"inputs": source} - ) - source.add_(10) - rank = SimpleNamespace( - _rank=SimpleNamespace(), - forward=lambda inputs: inputs * 3, - export_forward=lambda output: output, - ) - assert torch.equal( - await execute_operation(rank, operation), torch.tensor([6.0]) - ) - - asyncio.run(run()) + assert type(error.value) is type(stale) + assert str(error.value) == str(stale) + assert len(calls) == 1 + await _execute(rank, operation.id, "acknowledge", ((), ())) + with pytest.raises(OperationResultReleasedError): + await execute_operation(rank, operation) + assert len(calls) == 1 + + +async def test_concurrent_retry_waiter_cancellation_does_not_cancel_update(): + entered, release = asyncio.Event(), asyncio.Event() + calls = 0 + + async def optim_step(): + nonlocal calls + calls += 1 + entered.set() + await release.wait() + return calls + + rank = SimpleNamespace(_rank=SimpleNamespace(), optim_step=optim_step) + operation = TrainerOperation.capture(("client", 1), "optim_step", {}) + original = asyncio.create_task(execute_operation(rank, operation)) + await entered.wait() + retry = asyncio.create_task(execute_operation(rank, operation)) + await asyncio.sleep(0) + retry.cancel() + with pytest.raises(asyncio.CancelledError): + await retry + release.set() + assert await original == 1 + assert await execute_operation(rank, operation) == 1 + + +async def test_operation_captures_tensor_arguments_at_submission(): + source = torch.tensor([2.0]) + operation = TrainerOperation.capture(("client", 1), "forward", {"inputs": source}) + source.add_(10) + rank = SimpleNamespace( + _rank=SimpleNamespace(), + forward=lambda inputs: inputs * 3, + export_forward=lambda output: output, + ) + assert torch.equal(await execute_operation(rank, operation), torch.tensor([6.0])) @pytest.mark.parametrize("cuda_tag", [False, True]) -def test_operation_codec_remaps_storages_and_preserves_aliases(monkeypatch, cuda_tag): +async def test_operation_codec_remaps_storages_and_preserves_aliases( + monkeypatch, cuda_tag +): from test_trainer_command_transport import _check_payload, _payload - async def run(): - source = _payload("cpu") - with monkeypatch.context() as capture: - if cuda_tag: - capture.setattr( - torch.serialization, - "_package_registry", - [ - (0, lambda storage: "cuda:7", lambda storage, location: None), - *torch.serialization._package_registry, - ], - ) - operation = TrainerOperation.capture( - ("client", 1), "optim_step", {"value": source} + source = _payload("cpu") + with monkeypatch.context() as capture: + if cuda_tag: + capture.setattr( + torch.serialization, + "_package_registry", + [ + (0, lambda storage: "cuda:7", lambda storage, location: None), + *torch.serialization._package_registry, + ], ) - monkeypatch.setattr(torch.cuda, "is_available", lambda: False) - rank = SimpleNamespace(_rank=SimpleNamespace(), optim_step=lambda value: value) - result = await execute_operation(rank, operation) - _check_payload(result) - assert result.base.untyped_storage() is not source.base.untyped_storage() - - asyncio.run(run()) - - -def test_acknowledgement_bounds_history_including_unadmitted_holes(): - async def run(): - calls = [] - rank = SimpleNamespace( - _rank=SimpleNamespace(), optim_step=lambda: calls.append(1) - ) - # ID 1 is still pending. Odd IDs after it were cancelled before admission. - for sequence in range(2, 2002, 2): - await _execute(rank, ("client", sequence), "optim_step", {}) - await _execute(rank, ("client", sequence), "acknowledge", ((1,), ())) - ledger = rank._rank._operation_outcomes - assert not ledger.outcomes - assert len(ledger.acknowledged) == 1 - state = ledger.acknowledged["client"] - assert state.through == 2000 and state.pending == {1} - for sequence in (2, 3, 1999, 2000): - with pytest.raises(OperationResultReleasedError): - await _execute(rank, ("client", sequence), "optim_step", {}) - await _execute(rank, ("client", 1), "optim_step", {}) - await _execute(rank, ("client", 2000), "acknowledge", ((), ())) - assert not state.pending and not ledger.outcomes - assert len(calls) == 1001 - - asyncio.run(run()) - - -def test_out_of_order_acknowledgements_never_resurrect_ids_or_retire_other_sessions(): - async def run(): - calls = [] - rank = SimpleNamespace( - _rank=SimpleNamespace(), optim_step=lambda: calls.append(1) + operation = TrainerOperation.capture( + ("client", 1), "optim_step", {"value": source} ) - for through, pending in ( - (5, (1, 3, 4)), - (3, (1,)), - (6, (1, 4, 6)), - (5, (1, 2, 3)), - ): - await _execute(rank, ("client", through), "acknowledge", (pending, ())) - state = rank._rank._operation_outcomes.acknowledged["client"] - assert state.through == 6 and state.pending == {1, 6} - for sequence in (2, 3, 4, 5): - with pytest.raises(OperationResultReleasedError): - await _execute(rank, ("client", sequence), "optim_step", {}) - for identity in (("client", 1), ("client", 6), ("client", 7), ("other", 3)): - operation = TrainerOperation.capture(identity, "optim_step", {}) - await execute_operation(rank, operation) - await execute_operation(rank, operation) - assert len(calls) == 4 - - asyncio.run(run()) - - -def test_retiring_running_update_fences_retries_until_completion_is_dropped(): - async def run(): - entered, release = asyncio.Event(), asyncio.Event() - calls = [] - - async def optim_step(): - calls.append(1) - entered.set() - await release.wait() - raise ValueError("failed after mutation") - - rank = SimpleNamespace(_rank=SimpleNamespace(), optim_step=optim_step) - operation = TrainerOperation.capture(("client", 1), "optim_step", {}) - original = asyncio.create_task(execute_operation(rank, operation)) - await entered.wait() - await _execute(rank, operation.id, "acknowledge", ((), ())) + monkeypatch.setattr(torch.cuda, "is_available", lambda: False) + rank = SimpleNamespace(_rank=SimpleNamespace(), optim_step=lambda value: value) + result = await execute_operation(rank, operation) + _check_payload(result) + assert result.base.untyped_storage() is not source.base.untyped_storage() + + +async def test_acknowledgement_bounds_history_including_unadmitted_holes(): + calls = [] + rank = SimpleNamespace(_rank=SimpleNamespace(), optim_step=lambda: calls.append(1)) + # ID 1 is still pending. Odd IDs after it were cancelled before admission. + for sequence in range(2, 2002, 2): + await _execute(rank, ("client", sequence), "optim_step", {}) + await _execute(rank, ("client", sequence), "acknowledge", ((1,), ())) + ledger = rank._rank._operation_outcomes + assert not ledger.outcomes + assert len(ledger.acknowledged) == 1 + state = ledger.acknowledged["client"] + assert state.through == 2000 and state.pending == {1} + for sequence in (2, 3, 1999, 2000): with pytest.raises(OperationResultReleasedError): - await execute_operation(rank, operation) - release.set() - with pytest.raises(ValueError, match="failed after mutation"): - await original - assert len(calls) == 1 and not rank._rank._operation_outcomes.outcomes - - asyncio.run(run()) - - -def test_batch_pulls_and_close_are_identified_without_advancing_twice(): - async def run(): - events = [] - rank = SimpleNamespace( - _rank=SimpleNamespace(), - open_forward_batches=lambda **kwargs: events.append("open") or "iterator", - next_forward_batch=lambda **kwargs: events.append("next") or "batch", - export_forward=lambda batch: SimpleNamespace(handle="packet", batch=batch), - close_forward_batches=lambda handle: events.append(("close", handle)), - release_forward=lambda handles: events.append(("release", tuple(handles))), - ) - opening = TrainerOperation.capture( - ("client", 1), "batches_open", {"inputs": []} - ) - assert await execute_operation(rank, opening) == "iterator" - assert await execute_operation(rank, opening) == "iterator" - next_wave = TrainerOperation.capture( - ("client", 2), "batches_next", {"handle": "iterator"} - ) - first = await execute_operation(rank, next_wave) - assert await execute_operation(rank, next_wave) is first - # Other results can be acknowledged while this wave's reply is lost. - await _execute(rank, ("client", 3), "acknowledge", ((2, 3), ())) - # A lost pull is abandoned independently of closing its iterator. - await _execute(rank, ("client", 3), "acknowledge", ((3,), (2,))) - close = TrainerOperation.capture( - ("client", 3), - "batches_close", - {"handle": "iterator"}, - ) - await execute_operation(rank, close) - await execute_operation(rank, close) - assert events == [ - "open", - "next", - ("release", ("packet",)), - ("close", "iterator"), - ] - await _execute(rank, close.id, "acknowledge", ((), ())) - assert not rank._rank._operation_outcomes.outcomes + await _execute(rank, ("client", sequence), "optim_step", {}) + await _execute(rank, ("client", 1), "optim_step", {}) + await _execute(rank, ("client", 2000), "acknowledge", ((), ())) + assert not state.pending and not ledger.outcomes + assert len(calls) == 1001 + + +async def test_out_of_order_acknowledgements_never_resurrect_ids_or_retire_other_sessions(): + calls = [] + rank = SimpleNamespace(_rank=SimpleNamespace(), optim_step=lambda: calls.append(1)) + for through, pending in ( + (5, (1, 3, 4)), + (3, (1,)), + (6, (1, 4, 6)), + (5, (1, 2, 3)), + ): + await _execute(rank, ("client", through), "acknowledge", (pending, ())) + state = rank._rank._operation_outcomes.acknowledged["client"] + assert state.through == 6 and state.pending == {1, 6} + for sequence in (2, 3, 4, 5): with pytest.raises(OperationResultReleasedError): - await execute_operation(rank, next_wave) - - asyncio.run(run()) + await _execute(rank, ("client", sequence), "optim_step", {}) + for identity in (("client", 1), ("client", 6), ("client", 7), ("other", 3)): + operation = TrainerOperation.capture(identity, "optim_step", {}) + await execute_operation(rank, operation) + await execute_operation(rank, operation) + assert len(calls) == 4 + + +async def test_retiring_running_update_fences_retries_until_completion_is_dropped(): + entered, release = asyncio.Event(), asyncio.Event() + calls = [] + + async def optim_step(): + calls.append(1) + entered.set() + await release.wait() + raise ValueError("failed after mutation") + + rank = SimpleNamespace(_rank=SimpleNamespace(), optim_step=optim_step) + operation = TrainerOperation.capture(("client", 1), "optim_step", {}) + original = asyncio.create_task(execute_operation(rank, operation)) + await entered.wait() + await _execute(rank, operation.id, "acknowledge", ((), ())) + with pytest.raises(OperationResultReleasedError): + await execute_operation(rank, operation) + release.set() + with pytest.raises(ValueError, match="failed after mutation"): + await original + assert len(calls) == 1 and not rank._rank._operation_outcomes.outcomes + + +async def test_batch_pulls_and_close_are_identified_without_advancing_twice(): + events = [] + rank = SimpleNamespace( + _rank=SimpleNamespace(), + open_forward_batches=lambda **kwargs: events.append("open") or "iterator", + next_forward_batch=lambda **kwargs: events.append("next") or "batch", + export_forward=lambda batch: SimpleNamespace(handle="packet", batch=batch), + close_forward_batches=lambda handle: events.append(("close", handle)), + release_forward=lambda handles: events.append(("release", tuple(handles))), + ) + opening = TrainerOperation.capture(("client", 1), "batches_open", {"inputs": []}) + assert await execute_operation(rank, opening) == "iterator" + assert await execute_operation(rank, opening) == "iterator" + next_wave = TrainerOperation.capture( + ("client", 2), "batches_next", {"handle": "iterator"} + ) + first = await execute_operation(rank, next_wave) + assert await execute_operation(rank, next_wave) is first + # Other results can be acknowledged while this wave's reply is lost. + await _execute(rank, ("client", 3), "acknowledge", ((2, 3), ())) + # A lost pull is abandoned independently of closing its iterator. + await _execute(rank, ("client", 3), "acknowledge", ((3,), (2,))) + close = TrainerOperation.capture( + ("client", 3), + "batches_close", + {"handle": "iterator"}, + ) + await execute_operation(rank, close) + await execute_operation(rank, close) + assert events == [ + "open", + "next", + ("release", ("packet",)), + ("close", "iterator"), + ] + await _execute(rank, close.id, "acknowledge", ((), ())) + assert not rank._rank._operation_outcomes.outcomes + with pytest.raises(OperationResultReleasedError): + await execute_operation(rank, next_wave) @pytest.mark.parametrize("unserializable", [False, True]) -def test_failed_outcome_releases_traceback_activations_and_replays_error( +async def test_failed_outcome_releases_traceback_activations_and_replays_error( unserializable, ): from art.trainer_rank import TrainerRankMemoryError @@ -285,166 +248,154 @@ class UnserializableError(RuntimeError): def __reduce__(self): raise TypeError("cannot serialize this error") - async def run(): - references = [] - - def forward(): - activation = torch.ones(64, requires_grad=True).square() - references.append(weakref.ref(activation)) - if unserializable: - raise UnserializableError("failed forward") - raise TrainerRankMemoryError( - "failed forward", predicted_peak_bytes=123, usable_limit_bytes=100 - ) + references = [] - rank = SimpleNamespace( - _rank=SimpleNamespace(), - forward=forward, - export_forward=lambda output: output, - ) - operation = TrainerOperation.capture(("client", 1), "forward", {}) - for _ in range(3): - try: - await execute_operation(rank, operation) - except RuntimeError as error: - assert "failed forward" in str(error) - if not unserializable: - assert isinstance(error, TrainerRankMemoryError) - assert error.predicted_peak_bytes == 123 - assert error.usable_limit_bytes == 100 - else: - pytest.fail("failed operation unexpectedly succeeded") - gc.collect() - assert references[0]() is None - assert len(references) == 1 - assert ( - rank._rank._operation_outcomes.outcomes[operation.id].completion.exception() - is None + def forward(): + activation = torch.ones(64, requires_grad=True).square() + references.append(weakref.ref(activation)) + if unserializable: + raise UnserializableError("failed forward") + raise TrainerRankMemoryError( + "failed forward", predicted_peak_bytes=123, usable_limit_bytes=100 ) - asyncio.run(run()) - - -def test_concurrent_failed_retry_has_independent_error_without_retained_traceback(): - async def run(): - entered, release = asyncio.Event(), asyncio.Event() - references = [] - - async def optim_step(): - activation = torch.ones(64, requires_grad=True).square() - references.append(weakref.ref(activation)) - entered.set() - await release.wait() - raise ValueError("update rejected") - - rank = SimpleNamespace(_rank=SimpleNamespace(), optim_step=optim_step) - operation = TrainerOperation.capture(("client", 1), "optim_step", {}) - original = asyncio.create_task(execute_operation(rank, operation)) - await entered.wait() - retry = asyncio.create_task(execute_operation(rank, operation)) - await asyncio.sleep(0) - release.set() - results = await asyncio.gather(original, retry, return_exceptions=True) - assert all(isinstance(error, ValueError) for error in results) - assert results[0] is not results[1] - assert len(references) == 1 - del original, retry, results - await asyncio.sleep(0) + rank = SimpleNamespace( + _rank=SimpleNamespace(), + forward=forward, + export_forward=lambda output: output, + ) + operation = TrainerOperation.capture(("client", 1), "forward", {}) + for _ in range(3): + try: + await execute_operation(rank, operation) + except RuntimeError as error: + assert "failed forward" in str(error) + if not unserializable: + assert isinstance(error, TrainerRankMemoryError) + assert error.predicted_peak_bytes == 123 + assert error.usable_limit_bytes == 100 + else: + pytest.fail("failed operation unexpectedly succeeded") gc.collect() assert references[0]() is None - - asyncio.run(run()) + assert len(references) == 1 + assert ( + rank._rank._operation_outcomes.outcomes[operation.id].completion.exception() + is None + ) + + +async def test_concurrent_failed_retry_has_independent_error_without_retained_traceback(): + entered, release = asyncio.Event(), asyncio.Event() + references = [] + + async def optim_step(): + activation = torch.ones(64, requires_grad=True).square() + references.append(weakref.ref(activation)) + entered.set() + await release.wait() + raise ValueError("update rejected") + + rank = SimpleNamespace(_rank=SimpleNamespace(), optim_step=optim_step) + operation = TrainerOperation.capture(("client", 1), "optim_step", {}) + original = asyncio.create_task(execute_operation(rank, operation)) + await entered.wait() + retry = asyncio.create_task(execute_operation(rank, operation)) + await asyncio.sleep(0) + release.set() + results = await asyncio.gather(original, retry, return_exceptions=True) + assert all(isinstance(error, ValueError) for error in results) + assert results[0] is not results[1] + assert len(references) == 1 + del original, retry, results + await asyncio.sleep(0) + gc.collect() + assert references[0]() is None @pytest.mark.parametrize("retain_graph", [False, True]) @pytest.mark.parametrize("failure_stage", ["collect", "remote"]) -def test_failed_exported_backward_replays_once_and_releases_only_consumed_graphs( +async def test_failed_exported_backward_replays_once_and_releases_only_consumed_graphs( retain_graph, failure_stage, ): from art.trainer_rank import TrainerRankZero from art.trainer_rank._tensors import CotangentCollector, detach_tree - async def run(): - state = SimpleNamespace(collector=CotangentCollector(), sequence=0, exports={}) - owner = SimpleNamespace() - view = TrainerRankZero( - cast(Any, SimpleNamespace(rank=owner, state=state, dp_rank=0)) - ) - references, attempts = [], [] + state = SimpleNamespace(collector=CotangentCollector(), sequence=0, exports={}) + owner = SimpleNamespace() + view = TrainerRankZero( + cast(Any, SimpleNamespace(rank=owner, state=state, dp_rank=0)) + ) + references, attempts = [], [] - def export(handle): - graph = state.collector.attach( - detach_tree(handle, torch.tensor(2.0, requires_grad=True)) - ) - references.append(weakref.ref(graph)) - return view.export_forward(graph) - - packet = export("used") - unrelated = export("unrelated") - client = CotangentCollector() - loss = client.attach(packet).square() - packets = client.backward(loss, retain_graph=retain_graph) - - def fail(*args, **kwargs): - attempts.append("attempt") - raise ValueError("remote gradient rejected") - - hook = None - if failure_stage == "collect": - hook = references[0]().grad_fn.register_hook(fail) - setattr(view, "_submit_backward", lambda *args, **kwargs: None) - else: - setattr(view, "_submit_backward", fail) - operation = TrainerOperation.capture( - ("client", 1), - "backward", - {"packets": packets, "retain_graph": retain_graph}, + def export(handle): + graph = state.collector.attach( + detach_tree(handle, torch.tensor(2.0, requires_grad=True)) ) - for _ in range(2): - try: - await execute_operation(view, operation) - except ValueError as error: - assert str(error) == "remote gradient rejected" - else: - pytest.fail("failed backward unexpectedly succeeded") - gc.collect() - assert attempts == ["attempt"] - assert (packet.handle in state.exports) is retain_graph - assert (references[0]() is not None) is retain_graph - assert unrelated.handle in state.exports and references[1]() is not None - if hook is not None: - hook.remove() - if retain_graph: - setattr(view, "_submit_backward", lambda *args, **kwargs: None) - await _execute(view, ("client", 2), "backward", {"packets": packets}) - gc.collect() - assert packet.handle not in state.exports and references[0]() is None - - asyncio.run(run()) - - -def test_malformed_nonretained_backward_preserves_unrelated_exports(): + references.append(weakref.ref(graph)) + return view.export_forward(graph) + + packet = export("used") + unrelated = export("unrelated") + client = CotangentCollector() + loss = client.attach(packet).square() + packets = client.backward(loss, retain_graph=retain_graph) + + def fail(*args, **kwargs): + attempts.append("attempt") + raise ValueError("remote gradient rejected") + + hook = None + if failure_stage == "collect": + hook = references[0]().grad_fn.register_hook(fail) + setattr(view, "_submit_backward", lambda *args, **kwargs: None) + else: + setattr(view, "_submit_backward", fail) + operation = TrainerOperation.capture( + ("client", 1), + "backward", + {"packets": packets, "retain_graph": retain_graph}, + ) + for _ in range(2): + try: + await execute_operation(view, operation) + except ValueError as error: + assert str(error) == "remote gradient rejected" + else: + pytest.fail("failed backward unexpectedly succeeded") + gc.collect() + assert attempts == ["attempt"] + assert (packet.handle in state.exports) is retain_graph + assert (references[0]() is not None) is retain_graph + assert unrelated.handle in state.exports and references[1]() is not None + if hook is not None: + hook.remove() + if retain_graph: + setattr(view, "_submit_backward", lambda *args, **kwargs: None) + await _execute(view, ("client", 2), "backward", {"packets": packets}) + gc.collect() + assert packet.handle not in state.exports and references[0]() is None + + +async def test_malformed_nonretained_backward_preserves_unrelated_exports(): from art.trainer_rank import TrainerRankZero from art.trainer_rank._tensors import CotangentCollector, CotangentPacket - async def run(): - for handle, gradients in (("known", ()), ("missing", (torch.ones(1),))): - state = SimpleNamespace( - collector=CotangentCollector(), - exports={"known": (torch.ones(1),), "unrelated": (torch.ones(1),)}, - ) - view = TrainerRankZero( - cast(Any, SimpleNamespace(rank=SimpleNamespace(), state=state)) + for handle, gradients in (("known", ()), ("missing", (torch.ones(1),))): + state = SimpleNamespace( + collector=CotangentCollector(), + exports={"known": (torch.ones(1),), "unrelated": (torch.ones(1),)}, + ) + view = TrainerRankZero( + cast(Any, SimpleNamespace(rank=SimpleNamespace(), state=state)) + ) + with pytest.raises((ValueError, KeyError)): + await _execute( + view, + ("client", 1), + "backward", + {"packets": (CotangentPacket(handle, gradients),)}, ) - with pytest.raises((ValueError, KeyError)): - await _execute( - view, - ("client", 1), - "backward", - {"packets": (CotangentPacket(handle, gradients),)}, - ) - assert "unrelated" in state.exports - assert ("known" in state.exports) is (handle == "missing") - - asyncio.run(run()) + assert "unrelated" in state.exports + assert ("known" in state.exports) is (handle == "missing") From c2fb6f850252237ee4657919e6fc9b1e6b2408d4 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 20 Sep 2026 04:38:13 +0000 Subject: [PATCH 031/150] test: remove redundant trainer async wrappers --- tests/unit/test_trainer_driver_transport.py | 360 +++++++++--------- .../test_trainer_rank_release_completion.py | 186 +++++---- 2 files changed, 257 insertions(+), 289 deletions(-) diff --git a/tests/unit/test_trainer_driver_transport.py b/tests/unit/test_trainer_driver_transport.py index 97292eb16..8c85b6b13 100644 --- a/tests/unit/test_trainer_driver_transport.py +++ b/tests/unit/test_trainer_driver_transport.py @@ -54,185 +54,174 @@ def _operation(view, kind, payload, identity=None): @pytest.mark.parametrize("device", ["cpu", "cuda"]) @pytest.mark.parametrize("batches", [False, True]) -def test_driver_cpu_exports_preserve_old_gradients_and_worker_policy(device, batches): +async def test_driver_cpu_exports_preserve_old_gradients_and_worker_policy( + device, batches +): if device == "cuda" and not torch.cuda.is_available(): pytest.skip("requires CUDA") - async def run(): - rank: Any = _TransportRank(device) - view = _view(_Executor(rank, "zero")) - client = CotangentCollector() - request = replace( - _input(3), - input_tokens=torch.tensor([3], device=device), - options=ForwardOptions(output_device="model"), - ) - # Physical worker succeeds, but no aggregate GPU copy can be admitted. - rank._available_memory_bytes = lambda: 0 + rank: Any = _TransportRank(device) + view = _view(_Executor(rank, "zero")) + client = CotangentCollector() + request = replace( + _input(3), + input_tokens=torch.tensor([3], device=device), + options=ForwardOptions(output_device="model"), + ) + # Physical worker succeeds, but no aggregate GPU copy can be admitted. + rank._available_memory_bytes = lambda: 0 + if device == "cuda": + with pytest.raises(MemoryError, match="Gathered model-device outputs"): + view.forward(request) + seen = [] + attach = view._attach + + def check_cpu(packet): + assert all(packet.cpu) + assert all(tensor.device.type == "cpu" for tensor in packet.packet.tensors) + seen.append(packet.packet.handle) + before = torch.cuda.memory_allocated() if device == "cuda" else 0 + result = attach(packet) if device == "cuda": - with pytest.raises(MemoryError, match="Gathered model-device outputs"): - view.forward(request) - seen = [] - attach = view._attach - - def check_cpu(packet): - assert all(packet.cpu) - assert all(tensor.device.type == "cpu" for tensor in packet.packet.tensors) - seen.append(packet.packet.handle) - before = torch.cuda.memory_allocated() if device == "cuda" else 0 - result = attach(packet) - if device == "cuda": - assert torch.cuda.memory_allocated() == before - return result - - # A shallow operation view preserves this bound observer; its delegate - # inspects the actual placement before collector.attach can copy anything. - setattr(view, "_attach", check_cpu) - - async def forward(identity): - if batches: - handle = await _operation( - view, "batches_open", {"inputs": [request]}, identity + ":open" - ) - packet = await _operation( - view, "batches_next", {"handle": handle}, identity - ) - await _operation( - view, "batches_close", {"handle": handle}, identity + ":close" - ) - return client.attach(packet, managed=True).outputs[0].hidden_states - packet = await _operation(view, "forward", {"inputs": request}, identity) - return client.attach(packet, managed=True).hidden_states - - old = await forward("old") - with torch.no_grad(): - rank.weight.add_(1) - fresh = await forward("fresh") - assert view._transport_handles is None - assert old.device.type == fresh.device.type == "cpu" - assert len(seen) == 2 - assert all(policy.output_device == "model" for policy in rank.policies) - state = rank._rank_command_state - assert len(state.exports) == 2 - assert all( - tensor.device.type == "cpu" - for tensors in state.exports.values() - for tensor in tensors - ) - assert all( - tensor.device == rank.device - for tensors in state.graphs.values() - for tensor in tensors - ) - head = torch.nn.Parameter(torch.tensor(5.0, device=device)) - loss = (old * fresh * head).sum() - for retained in (True, False): - packets = client.backward(loss, retain_graph=retained) + assert torch.cuda.memory_allocated() == before + return result + + # A shallow operation view preserves this bound observer; its delegate + # inspects the actual placement before collector.attach can copy anything. + setattr(view, "_attach", check_cpu) + + async def forward(identity): + if batches: + handle = await _operation( + view, "batches_open", {"inputs": [request]}, identity + ":open" + ) + packet = await _operation( + view, "batches_next", {"handle": handle}, identity + ) await _operation( - view, - "backward", - {"packets": packets, "retain_graph": retained}, - "backward:" + str(retained), + view, "batches_close", {"handle": handle}, identity + ":close" ) - assert len(state.exports) == (2 if retained else 0) - assert rank.weight.grad is not None and head.grad is not None - # 3*w_old^2 * 3*w_new^2 * head, differentiated into the same parameter. - torch.testing.assert_close( - rank.weight.grad, torch.tensor(5400.0, device=device) + return client.attach(packet, managed=True).outputs[0].hidden_states + packet = await _operation(view, "forward", {"inputs": request}, identity) + return client.attach(packet, managed=True).hidden_states + + old = await forward("old") + with torch.no_grad(): + rank.weight.add_(1) + fresh = await forward("fresh") + assert view._transport_handles is None + assert old.device.type == fresh.device.type == "cpu" + assert len(seen) == 2 + assert all(policy.output_device == "model" for policy in rank.policies) + state = rank._rank_command_state + assert len(state.exports) == 2 + assert all( + tensor.device.type == "cpu" + for tensors in state.exports.values() + for tensor in tensors + ) + assert all( + tensor.device == rank.device + for tensors in state.graphs.values() + for tensor in tensors + ) + head = torch.nn.Parameter(torch.tensor(5.0, device=device)) + loss = (old * fresh * head).sum() + for retained in (True, False): + packets = client.backward(loss, retain_graph=retained) + await _operation( + view, + "backward", + {"packets": packets, "retain_graph": retained}, + "backward:" + str(retained), ) - torch.testing.assert_close(head.grad, torch.tensor(648.0, device=device)) - assert not state.graphs - setattr(view, "_attach", attach) - rank._available_memory_bytes = lambda: 1 << 60 - native = view.forward(request) - assert native.hidden_states.device == rank.device - view.backward(native.hidden_states.sum()) - assert not state.graphs - - asyncio.run(run()) + assert len(state.exports) == (2 if retained else 0) + assert rank.weight.grad is not None and head.grad is not None + # 3*w_old^2 * 3*w_new^2 * head, differentiated into the same parameter. + torch.testing.assert_close(rank.weight.grad, torch.tensor(5400.0, device=device)) + torch.testing.assert_close(head.grad, torch.tensor(648.0, device=device)) + assert not state.graphs + setattr(view, "_attach", attach) + rank._available_memory_bytes = lambda: 1 << 60 + native = view.forward(request) + assert native.hidden_states.device == rank.device + view.backward(native.hidden_states.sum()) + assert not state.graphs @pytest.mark.parametrize("policy", ["model", "cpu", "auto"]) -def test_transport_preserves_worker_admission_and_native_view(policy): - async def run(): - rank: Any = _TransportRank() - view = _view(_Executor(rank, "zero")) - rank.worker_reject = True - request = replace(_input(3), options=ForwardOptions(output_device=policy)) - with pytest.raises(MemoryError, match="worker admission rejected"): - await _operation(view, "forward", {"inputs": request}) - assert rank.policies[0].output_device == policy - assert view._transport_handles is None - assert not rank._rank_command_state.exports - assert not rank._rank_command_state.graphs - - asyncio.run(run()) +async def test_transport_preserves_worker_admission_and_native_view(policy): + rank: Any = _TransportRank() + view = _view(_Executor(rank, "zero")) + rank.worker_reject = True + request = replace(_input(3), options=ForwardOptions(output_device=policy)) + with pytest.raises(MemoryError, match="worker admission rejected"): + await _operation(view, "forward", {"inputs": request}) + assert rank.policies[0].output_device == policy + assert view._transport_handles is None + assert not rank._rank_command_state.exports + assert not rank._rank_command_state.graphs @pytest.mark.parametrize("failure_kind", ["budget", "allocation"]) -def test_export_failure_releases_only_failed_operation_and_counts_live_exports( +async def test_export_failure_releases_only_failed_operation_and_counts_live_exports( monkeypatch, failure_kind ): - async def run(): - rank: Any = _TransportRank() - view = _view(_Executor(rank, "zero")) - state = rank._rank_command_state - # Additional headroom excludes the CPU payloads held by earlier exports, - # mirroring fresh host/cgroup accounting used by the real rank. - total = 64 - observed = [] - - def available(): - used = sum( - t.untyped_storage().nbytes() - for values in state.exports.values() - for t in values - ) - observed.append(used) - return total - used - - rank._available_cpu_memory_bytes = available - request = _input(3) - good = await _operation(view, "forward", {"inputs": request}, "good") - assert ( - sum(t.numel() * t.element_size() for t in state.exports[good.handle]) == 4 + rank: Any = _TransportRank() + view = _view(_Executor(rank, "zero")) + state = rank._rank_command_state + # Additional headroom excludes the CPU payloads held by earlier exports, + # mirroring fresh host/cgroup accounting used by the real rank. + total = 64 + observed = [] + + def available(): + used = sum( + t.untyped_storage().nbytes() + for values in state.exports.values() + for t in values ) - original_graphs = set(state.graphs) - # The worker snapshot fits. A budget drop at export must fail before - # its clone and clean only the newly created physical graphs. - calls = 0 - - def constrained(): - nonlocal calls - calls += 1 - return available() if calls == 1 else 0 - - if failure_kind == "budget": - rank._available_cpu_memory_bytes = constrained - else: - from art.trainer_rank import _tensors - - detach = _tensors.detach_tree - - def reject_clone(handle, *args, **kwargs): - if handle.startswith("client:"): - raise MemoryError("export snapshot allocation failed") - return detach(handle, *args, **kwargs) - - monkeypatch.setattr(_tensors, "detach_tree", reject_clone) - with pytest.raises(MemoryError, match="export snapshot") as failure: - await _operation(view, "forward", {"inputs": request}, "failure") - assert failure.value.__traceback__ is not None - assert set(state.exports) == {good.handle} - assert set(state.graphs) == original_graphs - assert 4 in observed - rank._available_cpu_memory_bytes = available - await _operation(view, "release", {"handles": [good.handle]}, "release") - assert not state.exports - assert not state.graphs - assert available() == total - - asyncio.run(run()) + observed.append(used) + return total - used + + rank._available_cpu_memory_bytes = available + request = _input(3) + good = await _operation(view, "forward", {"inputs": request}, "good") + assert sum(t.numel() * t.element_size() for t in state.exports[good.handle]) == 4 + original_graphs = set(state.graphs) + # The worker snapshot fits. A budget drop at export must fail before + # its clone and clean only the newly created physical graphs. + calls = 0 + + def constrained(): + nonlocal calls + calls += 1 + return available() if calls == 1 else 0 + + if failure_kind == "budget": + rank._available_cpu_memory_bytes = constrained + else: + from art.trainer_rank import _tensors + + detach = _tensors.detach_tree + + def reject_clone(handle, *args, **kwargs): + if handle.startswith("client:"): + raise MemoryError("export snapshot allocation failed") + return detach(handle, *args, **kwargs) + + monkeypatch.setattr(_tensors, "detach_tree", reject_clone) + with pytest.raises(MemoryError, match="export snapshot") as failure: + await _operation(view, "forward", {"inputs": request}, "failure") + assert failure.value.__traceback__ is not None + assert set(state.exports) == {good.handle} + assert set(state.graphs) == original_graphs + assert 4 in observed + rank._available_cpu_memory_bytes = available + await _operation(view, "release", {"handles": [good.handle]}, "release") + assert not state.exports + assert not state.graphs + assert available() == total def test_transport_release_drops_cpu_payload_without_collecting_cycles(): @@ -256,28 +245,23 @@ async def run(): @pytest.mark.parametrize("kind", ["forward", "batches_open"]) -def test_abandoned_reply_releases_native_graphs_and_iterators(kind): - async def run(): - rank = cast(Any, _TransportRank()) - view = _view(_Executor(rank, "zero")) - state = rank._rank_command_state - operation = TrainerOperation.capture( - ("client", 1), - kind, - {"inputs": _input(3) if kind == "forward" else [_input(3)]}, - ) - await execute_operation(view, operation) - if kind == "forward": - assert state.graphs and state.exports - else: - assert state.iterators and state.batch_inputs - acknowledgement = TrainerOperation.capture( - operation.id, "acknowledge", ((), (1,)) - ) - await execute_operation(view, acknowledgement) - await execute_operation(view, acknowledgement) - assert not state.graphs and not state.exports - assert not state.iterators and not state.batch_inputs - assert not rank._operation_outcomes.outcomes - - asyncio.run(run()) +async def test_abandoned_reply_releases_native_graphs_and_iterators(kind): + rank = cast(Any, _TransportRank()) + view = _view(_Executor(rank, "zero")) + state = rank._rank_command_state + operation = TrainerOperation.capture( + ("client", 1), + kind, + {"inputs": _input(3) if kind == "forward" else [_input(3)]}, + ) + await execute_operation(view, operation) + if kind == "forward": + assert state.graphs and state.exports + else: + assert state.iterators and state.batch_inputs + acknowledgement = TrainerOperation.capture(operation.id, "acknowledge", ((), (1,))) + await execute_operation(view, acknowledgement) + await execute_operation(view, acknowledgement) + assert not state.graphs and not state.exports + assert not state.iterators and not state.batch_inputs + assert not rank._operation_outcomes.outcomes diff --git a/tests/unit/test_trainer_rank_release_completion.py b/tests/unit/test_trainer_rank_release_completion.py index 234dbe02b..f0d6356bb 100644 --- a/tests/unit/test_trainer_rank_release_completion.py +++ b/tests/unit/test_trainer_rank_release_completion.py @@ -10,7 +10,7 @@ @pytest.mark.parametrize("peer_buffers", [False, True]) -def test_completed_release_is_finalized_before_queued_done_callback( +async def test_completed_release_is_finalized_before_queued_done_callback( monkeypatch, peer_buffers ): synchronized = [] @@ -18,106 +18,90 @@ def test_completed_release_is_finalized_before_queued_done_callback( "art.trainer_rank._heads.synchronize_head_buffers", synchronized.append ) - async def run(): - rank = cast(Any, _Rank()) - executor = _Executor(rank, "zero") - state = executor.state - state.graphs["zero:old:dp:0"] = (rank.weight,) - state.released.add("zero:old:dp:0") - completed = asyncio.get_running_loop().create_future() - release = state.pending_release = _Release( - completed, - [(tuple(state.released), False), ((), peer_buffers)], - executor._finish_release, - ) - completed.set_result(None) - completed.add_done_callback(lambda _: executor._finish_release(release)) - await executor._join_release() - assert not state.graphs and not state.released - assert state.pending_release is None - # The already queued callback cannot alter the next release's ownership. - state.graphs["zero:new:dp:0"] = (rank.weight,) - await asyncio.sleep(0) - assert tuple(state.graphs) == ("zero:new:dp:0",) - assert synchronized == ([rank] if peer_buffers else []) - - asyncio.run(run()) + rank = cast(Any, _Rank()) + executor = _Executor(rank, "zero") + state = executor.state + state.graphs["zero:old:dp:0"] = (rank.weight,) + state.released.add("zero:old:dp:0") + completed = asyncio.get_running_loop().create_future() + release = state.pending_release = _Release( + completed, + [(tuple(state.released), False), ((), peer_buffers)], + executor._finish_release, + ) + completed.set_result(None) + completed.add_done_callback(lambda _: executor._finish_release(release)) + await executor._join_release() + assert not state.graphs and not state.released + assert state.pending_release is None + # The already queued callback cannot alter the next release's ownership. + state.graphs["zero:new:dp:0"] = (rank.weight,) + await asyncio.sleep(0) + assert tuple(state.graphs) == ("zero:new:dp:0",) + assert synchronized == ([rank] if peer_buffers else []) @pytest.mark.parametrize("cancelled", [False, True]) -def test_background_cleanup_failure_is_reported_and_blocks_next_entry(cancelled): - async def run(): - rank = cast(Any, _Rank()) - executor = _Executor(rank, "zero") - loop = asyncio.get_running_loop() - reports = [] - loop.set_exception_handler(lambda _loop, context: reports.append(context)) - completed = loop.create_future() - release = executor.state.pending_release = _Release( - completed, [((), False)], executor._finish_release - ) - completed.add_done_callback(lambda _: executor._finish_release(release)) - if cancelled: - completed.cancel() - else: - completed.set_exception(RuntimeError("injected release transport error")) - with pytest.raises( - RuntimeError, match="Callback release reconciliation failed" - ): - await asyncio.wait_for(executor._join_release(), 1) - assert len(reports) == 1 - with pytest.raises( - RuntimeError, match="Callback release reconciliation failed" - ): - await executor.reconcile_releases() - assert executor.state.pending_release is None - - asyncio.run(run()) - - -def test_success_in_unrelated_exception_handler_still_awaits_cleanup(monkeypatch): - async def run(): - rank = cast(Any, _Rank()) - executor = _Executor(rank, "zero") - started, finish = asyncio.Event(), asyncio.Event() - - async def reconcile(**kwargs): - started.set() - await finish.wait() - - monkeypatch.setattr(executor, "reconcile_releases", reconcile) - - async def callback(): - try: - raise ValueError("unrelated handled exception") - except ValueError: - async with executor.release_on_exit(): - pass - - pending = asyncio.create_task(callback()) - await started.wait() - assert not pending.done() - finish.set() - await pending - - asyncio.run(run()) - - -def test_checkpoint_fence_defers_cancellation_without_cancelling_release(): - async def run(): - rank = cast(Any, _Rank()) - executor = _Executor(rank, "zero") - completed = asyncio.get_running_loop().create_future() - executor.state.pending_release = _Release( - completed, [((), False)], executor._finish_release - ) - pending = asyncio.create_task(join_rank_callback_release(rank)) - await asyncio.sleep(0) - pending.cancel() - await asyncio.sleep(0) - assert not pending.done() and not completed.cancelled() - completed.set_result(None) - assert isinstance(await pending, asyncio.CancelledError) - assert executor.state.pending_release is None - - asyncio.run(run()) +async def test_background_cleanup_failure_is_reported_and_blocks_next_entry(cancelled): + rank = cast(Any, _Rank()) + executor = _Executor(rank, "zero") + loop = asyncio.get_running_loop() + reports = [] + loop.set_exception_handler(lambda _loop, context: reports.append(context)) + completed = loop.create_future() + release = executor.state.pending_release = _Release( + completed, [((), False)], executor._finish_release + ) + completed.add_done_callback(lambda _: executor._finish_release(release)) + if cancelled: + completed.cancel() + else: + completed.set_exception(RuntimeError("injected release transport error")) + with pytest.raises(RuntimeError, match="Callback release reconciliation failed"): + await asyncio.wait_for(executor._join_release(), 1) + assert len(reports) == 1 + with pytest.raises(RuntimeError, match="Callback release reconciliation failed"): + await executor.reconcile_releases() + assert executor.state.pending_release is None + + +async def test_success_in_unrelated_exception_handler_still_awaits_cleanup(monkeypatch): + rank = cast(Any, _Rank()) + executor = _Executor(rank, "zero") + started, finish = asyncio.Event(), asyncio.Event() + + async def reconcile(**kwargs): + started.set() + await finish.wait() + + monkeypatch.setattr(executor, "reconcile_releases", reconcile) + + async def callback(): + try: + raise ValueError("unrelated handled exception") + except ValueError: + async with executor.release_on_exit(): + pass + + pending = asyncio.create_task(callback()) + await started.wait() + assert not pending.done() + finish.set() + await pending + + +async def test_checkpoint_fence_defers_cancellation_without_cancelling_release(): + rank = cast(Any, _Rank()) + executor = _Executor(rank, "zero") + completed = asyncio.get_running_loop().create_future() + executor.state.pending_release = _Release( + completed, [((), False)], executor._finish_release + ) + pending = asyncio.create_task(join_rank_callback_release(rank)) + await asyncio.sleep(0) + pending.cancel() + await asyncio.sleep(0) + assert not pending.done() and not completed.cancelled() + completed.set_result(None) + assert isinstance(await pending, asyncio.CancelledError) + assert executor.state.pending_release is None From 797f4639f950339f305f208f18aaa36c9e7cc0e5 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 21 Sep 2026 06:44:03 +0000 Subject: [PATCH 032/150] test: share trainer operation imports --- tests/unit/test_trainer_driver_transport.py | 2 +- tests/unit/test_trainer_operations.py | 16 ++++++---------- 2 files changed, 7 insertions(+), 11 deletions(-) diff --git a/tests/unit/test_trainer_driver_transport.py b/tests/unit/test_trainer_driver_transport.py index 8c85b6b13..9537c6c84 100644 --- a/tests/unit/test_trainer_driver_transport.py +++ b/tests/unit/test_trainer_driver_transport.py @@ -16,7 +16,7 @@ from art.trainer_rank._commands import _Executor, _view from art.trainer_rank._operations import TrainerOperation, execute_operation from art.trainer_rank._options import resolve_forward_options -from art.trainer_rank._tensors import CotangentCollector, flatten_tensors +from art.trainer_rank._tensors import CotangentCollector class _TransportRank(_Rank): diff --git a/tests/unit/test_trainer_operations.py b/tests/unit/test_trainer_operations.py index d6b70146e..edea9d709 100644 --- a/tests/unit/test_trainer_operations.py +++ b/tests/unit/test_trainer_operations.py @@ -8,12 +8,18 @@ import pytest import torch +from art.trainer_rank import ( + TrainerRankMemoryError, + TrainerRankSlotStateError, + TrainerRankZero, +) from art.trainer_rank._operations import ( OperationId, OperationResultReleasedError, TrainerOperation, execute_operation, ) +from art.trainer_rank._tensors import CotangentCollector, CotangentPacket, detach_tree def _execute( @@ -44,8 +50,6 @@ async def test_update_identity_replays_outcome_without_applying_again(): async def test_failed_gradient_identity_preserves_original_error(): - from art.trainer_rank import TrainerRankSlotStateError - stale = TrainerRankSlotStateError("original forward is stale") calls = [] @@ -242,8 +246,6 @@ async def test_batch_pulls_and_close_are_identified_without_advancing_twice(): async def test_failed_outcome_releases_traceback_activations_and_replays_error( unserializable, ): - from art.trainer_rank import TrainerRankMemoryError - class UnserializableError(RuntimeError): def __reduce__(self): raise TypeError("cannot serialize this error") @@ -319,9 +321,6 @@ async def test_failed_exported_backward_replays_once_and_releases_only_consumed_ retain_graph, failure_stage, ): - from art.trainer_rank import TrainerRankZero - from art.trainer_rank._tensors import CotangentCollector, detach_tree - state = SimpleNamespace(collector=CotangentCollector(), sequence=0, exports={}) owner = SimpleNamespace() view = TrainerRankZero( @@ -379,9 +378,6 @@ def fail(*args, **kwargs): async def test_malformed_nonretained_backward_preserves_unrelated_exports(): - from art.trainer_rank import TrainerRankZero - from art.trainer_rank._tensors import CotangentCollector, CotangentPacket - for handle, gradients in (("known", ()), ("missing", (torch.ones(1),))): state = SimpleNamespace( collector=CotangentCollector(), From ba8e72028a21aa52cb1ce8c42fe17c85f0ca0949 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 21 Sep 2026 07:14:22 +0000 Subject: [PATCH 033/150] Consolidate initialized trainer test imports --- tests/unit/test_trainer_rank_commands.py | 21 +++++---------------- tests/unit/test_trainer_rank_graphs.py | 19 ++++++------------- tests/unit/test_trainer_rank_head_memory.py | 8 +------- tests/unit/test_trainer_rank_live_heads.py | 12 +++--------- 4 files changed, 15 insertions(+), 45 deletions(-) diff --git a/tests/unit/test_trainer_rank_commands.py b/tests/unit/test_trainer_rank_commands.py index c9c19a44a..5ed9f04c6 100644 --- a/tests/unit/test_trainer_rank_commands.py +++ b/tests/unit/test_trainer_rank_commands.py @@ -18,6 +18,7 @@ from art.trainer_rank import ( ForwardInput, + ForwardOptions, ForwardOutput, MicroBatch, MicroBatchStats, @@ -26,6 +27,10 @@ run_rank_callback, run_rank_callback_stream, ) +from art.trainer_rank._commands import _Executor, _OutputPacket, _view +from art.trainer_rank._heads import LiveHead +from art.trainer_rank._impl import _rebuild_forward_tree +from art.trainer_rank._tensors import CotangentCollector, detach_tree class _Rank: @@ -51,8 +56,6 @@ def forward(self, tree, **kwargs): with torch.set_grad_enabled(enabled): value = tree.input_tokens.float() * self.weight return ForwardOutput(None, None, None, value) - from art.trainer_rank._impl import _rebuild_forward_tree - return _rebuild_forward_tree( tree, [self.forward(child, **kwargs) for child in tree] ) @@ -194,8 +197,6 @@ def callback(view): def test_client_packet_registry_survives_callbacks(): rank: Any = _Rank() - from art.trainer_rank._tensors import CotangentCollector - packet = asyncio.run( run_rank_callback( rank, @@ -549,8 +550,6 @@ def test_released_client_graph_releases_physical_bridge(): def test_forward_batches_captures_policy_before_iteration(): - from art.trainer_rank import ForwardOptions - rank = object.__new__(TrainerRank) rank._forward_options = ForwardOptions(max_gradient_staleness=1, allow_replay=False) rank._skipped_forward_waves = {} @@ -591,9 +590,6 @@ def test_tuple_root_and_nested_tuple_shape(mode): def test_logical_head_factory_runs_once_and_head_only_client_backward(): from test_trainer_rank_custom_tensors import _trainer - from art.trainer_rank._heads import LiveHead - from art.trainer_rank._tensors import CotangentCollector - trainer, _ = _trainer("student") calls = [] @@ -661,9 +657,6 @@ def callback(view): def test_persistent_iterator_binds_policy_and_checkpoint_and_pulls_one_wave(): - from art.trainer_rank import ForwardOptions - from art.trainer_rank._tensors import CotangentCollector - class Rank(_Rank): _capture_forward_options = TrainerRank._capture_forward_options @@ -748,10 +741,6 @@ def run(callback): def test_nested_aggregate_outputs_admit_before_any_model_copy(): from dataclasses import replace - from art.trainer_rank import ForwardOptions - from art.trainer_rank._commands import _Executor, _OutputPacket, _view - from art.trainer_rank._tensors import detach_tree - rank: Any = _Rank() rank._available_memory_bytes = lambda: 1024 * 1024 view = _view(_Executor(rank, "zero")) diff --git a/tests/unit/test_trainer_rank_graphs.py b/tests/unit/test_trainer_rank_graphs.py index 44bf1e721..01741817f 100644 --- a/tests/unit/test_trainer_rank_graphs.py +++ b/tests/unit/test_trainer_rank_graphs.py @@ -7,7 +7,13 @@ import torch from torch.utils.checkpoint import checkpoint +from art.trainer_rank import ForwardOutput, TopK, TrainerRank from art.trainer_rank._graphs import GraphCache +from art.trainer_rank._impl import _CheckpointSlot +from art.trainer_rank._options import ( + ImportanceSamplingGradientCorrection, + ResolvedForwardOptions, +) def test_quantized_te_retained_backward_rejects_before_destructive_call(): @@ -275,12 +281,7 @@ def test_backward_uses_creation_order_not_wire_handle_order(): def _corrected_cache(*, retention="gpu", policy="when_available", stale=True): from contextlib import contextmanager - from art.trainer_rank import ForwardOutput from art.trainer_rank._corrections import capture_forward_corrections - from art.trainer_rank._options import ( - ImportanceSamplingGradientCorrection, - ResolvedForwardOptions, - ) cache = GraphCache() original = torch.nn.Parameter(torch.tensor(-1.0)) @@ -353,12 +354,7 @@ def test_newer_replay_opportunistically_corrects_current_jacobian(): def test_current_replay_rejects_changed_selected_token_events(corrections): from contextlib import nullcontext - from art.trainer_rank import ForwardOutput, TopK from art.trainer_rank._corrections import capture_forward_corrections - from art.trainer_rank._options import ( - ImportanceSamplingGradientCorrection, - ResolvedForwardOptions, - ) cache = GraphCache() parameter = torch.nn.Parameter(torch.tensor([1.0, 2.0])) @@ -491,9 +487,6 @@ def test_replay_restores_original_autocast_context(): def test_replay_failure_discards_transaction_and_releases_participating_records(): - from art.trainer_rank import TrainerRank - from art.trainer_rank._impl import _CheckpointSlot - trainer = TrainerRank.__new__(TrainerRank) parameter = torch.nn.Parameter(torch.tensor(2.0)) trainer._checkpoint_slots = {"student": _CheckpointSlot(params=(parameter,))} diff --git a/tests/unit/test_trainer_rank_head_memory.py b/tests/unit/test_trainer_rank_head_memory.py index ec8eeac72..3dce0d426 100644 --- a/tests/unit/test_trainer_rank_head_memory.py +++ b/tests/unit/test_trainer_rank_head_memory.py @@ -10,7 +10,7 @@ from art.trainer_rank import ForwardOptions, TrainerRankMemoryError, _impl from art.trainer_rank._commands import _Executor from art.trainer_rank._heads import LiveHead, export_head -from art.trainer_rank._tensors import CotangentCollector +from art.trainer_rank._tensors import CotangentCollector, detach_tree class _Head(torch.nn.Module): @@ -67,8 +67,6 @@ def test_registered_head_forward_admission_and_native_backward( torch.ones(4, requires_grad=True), retention="replay", ) - from art.trainer_rank._tensors import detach_tree - value = rank._forward_cotangent_collector().attach( detach_tree(handle, tensors) )[0] @@ -113,8 +111,6 @@ def watch(parameters): assert cache.handles() == (handle,) monkeypatch.setattr(rank, "_available_memory_bytes", lambda: 292) head = rank.module("head", _Head, checkpoint="student") - from art.trainer_rank._tensors import detach_tree - value = rank._forward_cotangent_collector().attach(detach_tree(handle, outputs))[0] rank.backward(head(value).sum()) torch.testing.assert_close( @@ -191,8 +187,6 @@ def register(): monkeypatch.setattr(rank, "_available_memory_bytes", lambda: 100) register() assert rank._checkpoint_slots["student"].params == () - from art.trainer_rank._tensors import detach_tree - value = rank._forward_cotangent_collector().attach(detach_tree(handle, outputs))[0] rank.backward(value.sum()) assert not cache.handles() diff --git a/tests/unit/test_trainer_rank_live_heads.py b/tests/unit/test_trainer_rank_live_heads.py index e74655cd0..513e0fe64 100644 --- a/tests/unit/test_trainer_rank_live_heads.py +++ b/tests/unit/test_trainer_rank_live_heads.py @@ -16,8 +16,10 @@ execute_head_operation, export_head, head_gradient_targets, + logical_register_head, ) -from art.trainer_rank._tensors import CotangentCollector +from art.trainer_rank._options import ForwardOptions +from art.trainer_rank._tensors import CotangentCollector, detach_tree class TiedHead(torch.nn.Module): @@ -262,8 +264,6 @@ def factory(): def test_constructor_staleness_applies_to_heads_before_mutating_gradients(): - from art.trainer_rank._options import ForwardOptions - trainer, rank = _trainer("student") setattr(trainer, "_forward_options", ForwardOptions(max_gradient_staleness=0)) head = rank.module("head", TiedHead, checkpoint="student") @@ -429,8 +429,6 @@ def test_client_parameter_mutations_fail_before_changing_owned_values(mutation): def test_head_export_preserves_strict_constructor_policy_for_client(): - from art.trainer_rank._options import ForwardOptions - trainer, rank = _trainer("student") setattr(trainer, "_forward_options", ForwardOptions(max_gradient_staleness=0)) rank.parameter("gain", lambda: torch.tensor(2.0), checkpoint="student") @@ -553,8 +551,6 @@ def test_client_reused_module_factory_does_not_share_checkpoint_handles(): @pytest.mark.parametrize("operation", ("model_first", "parameter_first", "linear")) def test_managed_model_operand_captures_live_parameter_and_keeps_old_version(operation): - from art.trainer_rank._tensors import detach_tree - trainer, native = _native_head( "parameter", "weight", lambda: torch.tensor([2.0, 4.0]) ) @@ -937,8 +933,6 @@ def test_inplace_operation_snapshots_readonly_client_tensor(kind): def test_logical_callback_reentrant_head_rejects_before_gradient_publication(): from types import SimpleNamespace - from art.trainer_rank._heads import logical_register_head - trainer, _ = _trainer("student") collector = CotangentCollector() From 38b53decf3359753e98205c766421a4ed8412a20 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 21 Sep 2026 07:46:58 +0000 Subject: [PATCH 034/150] Share standard-library imports in trainer tests --- tests/unit/test_trainer_rank_commands.py | 3 +-- tests/unit/test_trainer_rank_graphs.py | 14 ++------------ tests/unit/test_trainer_rank_live_heads.py | 3 +-- 3 files changed, 4 insertions(+), 16 deletions(-) diff --git a/tests/unit/test_trainer_rank_commands.py b/tests/unit/test_trainer_rank_commands.py index 5ed9f04c6..95328aa92 100644 --- a/tests/unit/test_trainer_rank_commands.py +++ b/tests/unit/test_trainer_rank_commands.py @@ -3,6 +3,7 @@ from __future__ import annotations import asyncio +from dataclasses import replace from datetime import timedelta import gc import sys @@ -739,8 +740,6 @@ def run(callback): def test_nested_aggregate_outputs_admit_before_any_model_copy(): - from dataclasses import replace - rank: Any = _Rank() rank._available_memory_bytes = lambda: 1024 * 1024 view = _view(_Executor(rank, "zero")) diff --git a/tests/unit/test_trainer_rank_graphs.py b/tests/unit/test_trainer_rank_graphs.py index 01741817f..7a7658021 100644 --- a/tests/unit/test_trainer_rank_graphs.py +++ b/tests/unit/test_trainer_rank_graphs.py @@ -1,7 +1,9 @@ from __future__ import annotations +from contextlib import contextmanager, nullcontext import gc from types import SimpleNamespace +import weakref import pytest import torch @@ -279,8 +281,6 @@ def test_backward_uses_creation_order_not_wire_handle_order(): def _corrected_cache(*, retention="gpu", policy="when_available", stale=True): - from contextlib import contextmanager - from art.trainer_rank._corrections import capture_forward_corrections cache = GraphCache() @@ -352,8 +352,6 @@ def test_newer_replay_opportunistically_corrects_current_jacobian(): @pytest.mark.parametrize("corrections", [False, True]) def test_current_replay_rejects_changed_selected_token_events(corrections): - from contextlib import nullcontext - from art.trainer_rank._corrections import capture_forward_corrections cache = GraphCache() @@ -386,11 +384,6 @@ def execute(_): @pytest.mark.parametrize("checkpointing", [False, True]) def test_abandoned_release_frees_physical_record_without_autograd(checkpointing): - import gc - import weakref - - from torch.utils.checkpoint import checkpoint - cache = GraphCache() parameter = torch.nn.Parameter(torch.randn(4, 4)) @@ -417,9 +410,6 @@ def compute(value): def test_backward_releases_unused_differentiable_output_branches(): - import gc - import weakref - cache = GraphCache() parameter = torch.nn.Parameter(torch.randn(4, 4)) handle, outputs = cache.run( diff --git a/tests/unit/test_trainer_rank_live_heads.py b/tests/unit/test_trainer_rank_live_heads.py index 513e0fe64..193b990cb 100644 --- a/tests/unit/test_trainer_rank_live_heads.py +++ b/tests/unit/test_trainer_rank_live_heads.py @@ -2,6 +2,7 @@ import asyncio from copy import deepcopy +from types import SimpleNamespace import pytest from test_trainer_rank_custom_tensors import _trainer, _use_local_gradients @@ -931,8 +932,6 @@ def test_inplace_operation_snapshots_readonly_client_tensor(kind): def test_logical_callback_reentrant_head_rejects_before_gradient_publication(): - from types import SimpleNamespace - trainer, _ = _trainer("student") collector = CotangentCollector() From 5b7e1c6e3dc5fcb47a4d30dc7a2697bcb62bc8b7 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 21 Sep 2026 07:44:31 +0000 Subject: [PATCH 035/150] Consolidate remaining trainer test imports --- tests/unit/test_trainer_driver_transport.py | 4 +-- tests/unit/test_trainer_rank_tensors.py | 31 ++++++--------------- 2 files changed, 10 insertions(+), 25 deletions(-) diff --git a/tests/unit/test_trainer_driver_transport.py b/tests/unit/test_trainer_driver_transport.py index 9537c6c84..bcd039230 100644 --- a/tests/unit/test_trainer_driver_transport.py +++ b/tests/unit/test_trainer_driver_transport.py @@ -12,7 +12,7 @@ from test_trainer_rank_commands import _input, _Rank import torch -from art.trainer_rank import ForwardInput, ForwardOptions, ForwardOutput +from art.trainer_rank import ForwardInput, ForwardOptions, ForwardOutput, _tensors from art.trainer_rank._commands import _Executor, _view from art.trainer_rank._operations import TrainerOperation, execute_operation from art.trainer_rank._options import resolve_forward_options @@ -201,8 +201,6 @@ def constrained(): if failure_kind == "budget": rank._available_cpu_memory_bytes = constrained else: - from art.trainer_rank import _tensors - detach = _tensors.detach_tree def reject_clone(handle, *args, **kwargs): diff --git a/tests/unit/test_trainer_rank_tensors.py b/tests/unit/test_trainer_rank_tensors.py index ebc754e37..780bfc86a 100644 --- a/tests/unit/test_trainer_rank_tensors.py +++ b/tests/unit/test_trainer_rank_tensors.py @@ -1,8 +1,16 @@ from __future__ import annotations -from collections import OrderedDict, namedtuple +from collections import OrderedDict, defaultdict, namedtuple +from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass, field +import gc import pickle +import subprocess +import sys +from threading import Event +from types import SimpleNamespace +from typing import cast +import weakref import pytest import torch @@ -306,9 +314,6 @@ def test_cpu_output_bridge_gpu_head_and_remote_model_cotangents(): @pytest.mark.parametrize("managed", [False, True]) def test_release_follows_dependent_loss_not_temporary_output(managed): - import gc - import weakref - collector = CotangentCollector() released = [] value = collector.attach( @@ -332,8 +337,6 @@ def test_release_follows_dependent_loss_not_temporary_output(managed): def test_release_follows_retained_graph_and_all_output_branches(): - import gc - collector = CotangentCollector() released = [] outputs = collector.attach( @@ -360,8 +363,6 @@ def test_release_follows_retained_graph_and_all_output_branches(): def test_release_dropped_and_nondifferentiable_packets(): - import gc - collector = CotangentCollector() released = [] output = collector.attach( @@ -458,9 +459,6 @@ def reject(gradient): @pytest.mark.parametrize("managed", [False, True]) def test_unrelated_failed_backward_cannot_enter_an_active_collection(managed): - from concurrent.futures import ThreadPoolExecutor - from threading import Event - collector = CotangentCollector() a = collector.attach( detach_tree("a", torch.tensor(2.0, requires_grad=True)), managed=managed @@ -693,9 +691,6 @@ def test_managed_ambiguous_cuda_devices_require_explicit_transfer(): "container", ["set", "object", "tensor_key", "default_factory"] ) def test_output_packets_reject_opaque_tensor_bearing_metadata(container): - from collections import defaultdict - from types import SimpleNamespace - source = torch.tensor(3.0, requires_grad=True) tree = ( {source} @@ -711,9 +706,6 @@ def test_output_packets_reject_opaque_tensor_bearing_metadata(container): def test_graph_release_callback_does_not_run_at_interpreter_shutdown(): - import subprocess - import sys - result = subprocess.run( [ sys.executable, @@ -879,9 +871,6 @@ def __torch_function__(cls, func, types, args=(), kwargs=None): @pytest.mark.parametrize("fail", [False, True]) def test_flatten_releases_tensor_references_without_cyclic_gc(fail): - import gc - import weakref - @dataclass class Broken: value: int = 0 @@ -974,8 +963,6 @@ def test_duplicate_handle_sparse_cotangents_accumulate_in_any_order(device, layo def test_transformers_dataclass_mapping_packet_preserves_fields_entries_and_aliases(): - from typing import cast - from transformers.modeling_outputs import BaseModelOutput source = cast(torch.FloatTensor, torch.tensor([2.0, 3.0], requires_grad=True)) From 5ef848904428414891100b8dddaeff429daf83d3 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 21 Sep 2026 10:08:02 +0000 Subject: [PATCH 036/150] Share trainer head mutation-name constants --- src/art/trainer_rank/_heads.py | 118 +++++++++++++++------------------ 1 file changed, 53 insertions(+), 65 deletions(-) diff --git a/src/art/trainer_rank/_heads.py b/src/art/trainer_rank/_heads.py index e65978edf..a8040c575 100644 --- a/src/art/trainer_rank/_heads.py +++ b/src/art/trainer_rank/_heads.py @@ -106,6 +106,57 @@ def replace(value: torch.Tensor) -> torch.Tensor: "volatile", } ) +_TENSOR_MUTATION_DUNDERS = frozenset( + { + "__setitem__", + "__set__", + "__iadd__", + "__isub__", + "__imul__", + "__itruediv__", + "__ifloordiv__", + "__imod__", + "__ipow__", + "__imatmul__", + "__iand__", + "__ior__", + "__ixor__", + "__ilshift__", + "__irshift__", + } +) +_CLIENT_BUFFER_MUTATIONS = (_TENSOR_MUTATION_DUNDERS - {"__set__", "__imatmul__"}) | { + "copy_", + "fill_", + "zero_", + "add_", + "sub_", + "mul_", + "div_", + "true_divide_", + "floor_divide_", + "remainder_", + "fmod_", + "pow_", + "lerp_", + "bitwise_and_", + "bitwise_or_", + "bitwise_xor_", + "bitwise_left_shift_", + "bitwise_right_shift_", + "masked_fill_", + "masked_scatter_", + "scatter_", + "scatter_add_", + "index_copy_", + "index_add_", + "index_fill_", + "index_put_", + "put_", + "clamp_", + "clamp_min_", + "clamp_max_", +} def tensor_metadata_function(func: Callable[..., Any]) -> bool: @@ -127,24 +178,7 @@ def mutates_tensor(func: Callable[..., Any], kwargs: Mapping[str, Any]) -> bool: name = getattr(func, "__name__", "") return ( (name.endswith("_") and not name.endswith("__")) - or name - in { - "__setitem__", - "__set__", - "__iadd__", - "__isub__", - "__imul__", - "__itruediv__", - "__ifloordiv__", - "__imod__", - "__ipow__", - "__imatmul__", - "__iand__", - "__ior__", - "__ixor__", - "__ilshift__", - "__irshift__", - } + or name in _TENSOR_MUTATION_DUNDERS or kwargs.get("out") is not None or kwargs.get("inplace") is True ) @@ -914,53 +948,7 @@ def __torch_function__( for value in _walk_objects((args, kwargs)) ) if mutating and ( - kwargs.get("out") is not None - or name - not in { - "__setitem__", - "__iadd__", - "__isub__", - "__imul__", - "__itruediv__", - "__ifloordiv__", - "__imod__", - "__ipow__", - "__iand__", - "__ior__", - "__ixor__", - "__ilshift__", - "__irshift__", - "copy_", - "fill_", - "zero_", - "add_", - "sub_", - "mul_", - "div_", - "true_divide_", - "floor_divide_", - "remainder_", - "fmod_", - "pow_", - "lerp_", - "bitwise_and_", - "bitwise_or_", - "bitwise_xor_", - "bitwise_left_shift_", - "bitwise_right_shift_", - "masked_fill_", - "masked_scatter_", - "scatter_", - "scatter_add_", - "index_copy_", - "index_add_", - "index_fill_", - "index_put_", - "put_", - "clamp_", - "clamp_min_", - "clamp_max_", - } + kwargs.get("out") is not None or name not in _CLIENT_BUFFER_MUTATIONS ): raise RuntimeError( f"Unsupported checkpoint buffer mutation {name}; use buffer.copy_() with unchanged shape and dtype" From 1a13ee60f6cd5ae3526006f1fbbde13e42862558 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 21 Sep 2026 10:20:40 +0000 Subject: [PATCH 037/150] Coordinate authority buffer snapshot failures across ranks --- src/art/trainer_rank/_heads.py | 24 ++++++---- tests/unit/test_trainer_rank_live_heads.py | 53 +++++++++++++++++++++- 2 files changed, 66 insertions(+), 11 deletions(-) diff --git a/src/art/trainer_rank/_heads.py b/src/art/trainer_rank/_heads.py index a8040c575..c1d7d10de 100644 --- a/src/art/trainer_rank/_heads.py +++ b/src/art/trainer_rank/_heads.py @@ -1234,16 +1234,20 @@ def synchronize_head_buffers(trainer: TrainerRank, checkpoints: Any = None) -> N raise trainer._slot_state_error( "Custom buffer registrations differ across ranks" ) - payload = ( - { - key: ( - tracker.buffer_revision, - {name: _plain(value).cpu() for name, value in buffers.items()}, - ) - for key, (tracker, buffers) in targets.items() - } - if dist.get_rank(group) == 0 - else None + payload = _checkpoint._phase( + lambda: ( + { + key: ( + tracker.buffer_revision, + {name: _plain(value).cpu() for name, value in buffers.items()}, + ) + for key, (tracker, buffers) in targets.items() + } + if dist.get_rank(group) == 0 + else None + ), + "snapshot synchronized buffers", + group, ) authoritative = _checkpoint._gather(payload, group)[0] assert authoritative is not None diff --git a/tests/unit/test_trainer_rank_live_heads.py b/tests/unit/test_trainer_rank_live_heads.py index 193b990cb..ee30c63b9 100644 --- a/tests/unit/test_trainer_rank_live_heads.py +++ b/tests/unit/test_trainer_rank_live_heads.py @@ -8,7 +8,7 @@ from test_trainer_rank_custom_tensors import _trainer, _use_local_gradients import torch from torch.utils.checkpoint import checkpoint -from trainer_rank_test_support import gloo_group +from trainer_rank_test_support import gloo_group, spawn_and_join from art.trainer_rank import AdamParams, ModuleHandle, run_rank_callback from art.trainer_rank._heads import ( @@ -298,6 +298,57 @@ def test_distributed_persistent_buffers_use_dp_zero_authority(tmp_path): ) +def _buffer_snapshot_failure_worker(process_rank, init_method): + import torch.distributed as dist + + from art.trainer_rank import _heads + + with gloo_group(process_rank, init_method, timeout=15): + trainer, rank = _trainer("student") + trainer._checkpoint_process_group = dist.group.WORLD + trainer._checkpoint_finalize_process_group = dist.group.WORLD + buffer = rank.buffer("counter", lambda: torch.tensor(1.0), checkpoint="student") + if process_rank == 1: + buffer.add_(1) + failure = MemoryError("authority buffer snapshot failed") + + def fail_snapshot(value): + raise failure + + with pytest.MonkeyPatch.context() as patch: + if process_rank == 0: + patch.setattr(_heads, "_plain", fail_snapshot) + with pytest.raises( + (MemoryError, RuntimeError), match="authority buffer snapshot failed" + ) as caught: + _heads.synchronize_head_buffers(trainer) + if process_rank == 0: + assert caught.value is failure + else: + assert isinstance(caught.value, RuntimeError) + assert "snapshot synchronized buffers" in str(caught.value) + + assert buffer.item() == process_rank + 1 + assert ( + export_head(trainer, "student", "counter").buffer_revision == process_rank + ) + completed = torch.tensor(1) + dist.all_reduce(completed, group=trainer._checkpoint_group()) + assert completed.item() == 2 + _heads.synchronize_head_buffers(trainer) + assert buffer.item() == 1 + assert export_head(trainer, "student", "counter").buffer_revision == 2 + + +def test_distributed_authority_buffer_snapshot_failure_keeps_group_usable(tmp_path): + spawn_and_join( + _buffer_snapshot_failure_worker, + args=(f"file://{tmp_path / 'snapshot_failure'}",), + timeout=60, + failure="Authority buffer snapshot failure did not exit on every rank", + ) + + def test_remote_reregistration_rejects_changed_ties(): trainer, rank = _trainer("student") rank.module("head", TiedHead, checkpoint="student") From b73cd1bb467b31ef2dff7dcd65c25b742b6b6093 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 21 Sep 2026 10:34:37 +0000 Subject: [PATCH 038/150] Release undelivered trainer output graphs on attachment failure --- src/art/trainer_rank/_commands.py | 58 ++++- .../unit/test_trainer_rank_output_delivery.py | 242 ++++++++++++++++++ 2 files changed, 289 insertions(+), 11 deletions(-) create mode 100644 tests/unit/test_trainer_rank_output_delivery.py diff --git a/src/art/trainer_rank/_commands.py b/src/art/trainer_rank/_commands.py index 1afb121a0..bc0d4f4ca 100644 --- a/src/art/trainer_rank/_commands.py +++ b/src/art/trainer_rank/_commands.py @@ -4,7 +4,7 @@ import asyncio from collections.abc import AsyncGenerator, Callable, Generator, Iterator, Sequence -from contextlib import asynccontextmanager, nullcontext +from contextlib import asynccontextmanager, contextmanager, nullcontext from dataclasses import dataclass, field, replace from functools import partial import inspect @@ -764,16 +764,37 @@ def forward(self, inputs: _impl.ForwardInputs, **kwargs: Any) -> Any: materialized = self._rank._capture_forward_options( materialized, kwargs.get("options") ) - (packet,) = self._invoke("forward", materialized, **kwargs) - return self._attach(self._place_outputs([(packet, materialized)])[0]) + packets = self._invoke("forward", materialized, **kwargs) + with self._release_on_error([packet.packet.handle for packet in packets]): + (packet,) = packets + return self._attach(self._place_outputs([(packet, materialized)])[0]) single = isinstance(materialized, _impl.ForwardInput) roots = [materialized] if single else materialized outputs: list[Any] = [] - for batch in self.forward_batches(roots, **kwargs): - outputs.extend(batch.outputs) - return ( - outputs[0] if single else _impl._rebuild_forward_tree(materialized, outputs) - ) + handles: list[str] = [] + with self._release_on_error(handles): + items = self._prepare_batches(roots, kwargs) + for batch in self._iterate_batches(items, kwargs, handles): + outputs.extend(batch.outputs) + return ( + outputs[0] + if single + else _impl._rebuild_forward_tree(materialized, outputs) + ) + + @contextmanager + def _release_on_error(self, handles: Sequence[str]) -> Iterator[None]: + try: + yield + except BaseException: + if handles: + try: + # Followers must release their graphs too; head publication + # is unrelated to reclaiming outputs never delivered. + self._executor.invoke("release", tuple(handles)) + except BaseException: + self._executor.state.released.update(handles) + raise def forward_batches(self, inputs: Any, **kwargs: Any) -> Iterator[_impl.MicroBatch]: items = self._prepare_batches(inputs, kwargs) @@ -819,12 +840,17 @@ def close_forward_batches(self, handle: str) -> None: self._invoke("batches_close", handle) def _iterate_batches( - self, items: Any, kwargs: dict[str, Any] + self, + items: Any, + kwargs: dict[str, Any], + handles: list[str] | None = None, ) -> Iterator[_impl.MicroBatch]: identifier = self._invoke("batches", items, **kwargs) try: while True: - batch = self._combine_wave(self._invoke("next", identifier), items) + batch = self._combine_wave( + self._invoke("next", identifier), items, handles + ) if batch is None: return yield batch @@ -834,7 +860,17 @@ def _iterate_batches( if not self._executor.stopped: self._invoke("close", identifier) - def _combine_wave(self, wave: Any, items: Any) -> _impl.MicroBatch | None: + def _combine_wave( + self, wave: Any, items: Any, accumulated: list[str] | None = None + ) -> _impl.MicroBatch | None: + handles = [packet.packet.handle for _, packet in wave if packet is not None] + with self._release_on_error(handles): + batch = self._assemble_wave(wave, items) + if accumulated is not None: + accumulated.extend(handles) + return batch + + def _assemble_wave(self, wave: Any, items: Any) -> _impl.MicroBatch | None: if all(batch is None for batch, _ in wave): return None if any(batch is None for batch, _ in wave): diff --git a/tests/unit/test_trainer_rank_output_delivery.py b/tests/unit/test_trainer_rank_output_delivery.py new file mode 100644 index 000000000..98cd6aa22 --- /dev/null +++ b/tests/unit/test_trainer_rank_output_delivery.py @@ -0,0 +1,242 @@ +"""Undelivered logical outputs must not retain physical forward graphs.""" + +from dataclasses import replace +from functools import partial +import weakref + +import pytest +from test_trainer_rank_commands import _input, _Rank +import torch + +from art.trainer_rank import ( + ForwardInput, + ForwardOutput, + MicroBatch, + MicroBatchStats, + _tensors, +) +from art.trainer_rank._commands import _Executor, _view +from art.trainer_rank._graphs import GraphCache +from art.trainer_rank._tensors import CotangentCollector, detach_tree + + +class _CachedRank(_Rank): + def __init__(self, *args): + super().__init__(*args) + self.cache = GraphCache() + self.collector = CotangentCollector() + self.records = [] + + def forward(self, tree, **kwargs): + if not isinstance(tree, ForwardInput): + return super().forward(tree, **kwargs) + handle, (value,) = self.cache.run( + lambda tokens: (tokens.float() * self.weight,), tree.input_tokens + ) + self.records.append(weakref.ref(self.cache._records[handle])) + return self.collector.attach( + detach_tree(handle, ForwardOutput(None, None, None, value)), + on_release=partial(self.cache.release, handle), + ) + + def _forward_graph_cache(self): + return self.cache + + def _forward_cotangent_collector(self): + return self.collector + + def _forward_memory_group(self): + return None + + +def _fail_delivery(monkeypatch, executor, kind, *, packet_number=1): + """Fail inside the real copy or collector after physical graph registration.""" + error = MemoryError("injected logical output delivery failure") + packet = executor._packet + targets, partial_outputs = set(), [] + count = 0 + + def capture(*args): + nonlocal count + output = packet(*args) + count += 1 + if count == packet_number: + targets.update(id(tensor) for tensor in output.packet.tensors) + return replace(output, cpu=(False,) * len(output.cpu), managed=kind == "attach") + + monkeypatch.setattr(executor, "_packet", capture) + if kind == "copy": + to = torch.Tensor.to + + def fail_copy(tensor, *args, **kwargs): + if id(tensor) in targets: + raise error + return to(tensor, *args, **kwargs) + + monkeypatch.setattr(torch.Tensor, "to", fail_copy) + else: + managed = _tensors.managed_tensor + calls = 0 + + def fail_attach(tensor): + nonlocal calls + calls += 1 + if calls == packet_number: + # The bridge and its finalizer already exist. Retain its proxy + # even after the error, so collection cannot repair the leak. + partial_outputs.append(tensor) + raise error + return managed(tensor) + + monkeypatch.setattr(_tensors, "managed_tensor", fail_attach) + return error, partial_outputs + + +@pytest.mark.parametrize("mode", ["rank", "zero"]) +@pytest.mark.parametrize("kind", ["copy", "attach"]) +def test_failed_delivery_releases_registered_graph_and_native_cache( + monkeypatch, mode, kind +): + rank = _CachedRank() + executor = _Executor(rank, mode) + view = _view(executor) + with monkeypatch.context() as patch: + error, partial_outputs = _fail_delivery(patch, executor, kind) + with pytest.raises(MemoryError) as failure: + view.forward(_input(3)) + assert failure.value is error and error.__traceback__ is not None + assert bool(partial_outputs) is (kind == "attach") + assert not executor.state.graphs + assert not rank.cache.handles() + assert all(record() is None for record in rank.records) + assert not executor.iterators + view.backward(view.forward(_input(7)).hidden_states.sum()) + assert rank.weight.grad.item() == 7 + assert not executor.state.graphs and not rank.cache.handles() + + +@pytest.mark.parametrize("delivery", ["forward", "iterator", "persistent"]) +def test_later_wave_failure_keeps_only_previously_delivered_outputs( + monkeypatch, delivery +): + rank = _CachedRank() + executor = _Executor(rank, "zero") + view = _view(executor) + requests = [_input(3), _input(5)] + delivered = None + with monkeypatch.context() as patch: + error, _ = _fail_delivery(patch, executor, "copy", packet_number=2) + if delivery == "iterator": + iterator = view.forward_batches(requests) + delivered = next(iterator).outputs[0] + advance = lambda: next(iterator) + elif delivery == "persistent": + handle = view.open_forward_batches(requests) + delivered = view.next_forward_batch(handle).outputs[0] + advance = lambda: view.next_forward_batch(handle) + else: + advance = lambda: view.forward(requests) + previous = set(executor.state.graphs) + with pytest.raises(MemoryError) as failure: + advance() + assert failure.value is error and error.__traceback__ is not None + assert set(executor.state.graphs) == previous + assert len(rank.cache.handles()) == int(delivered is not None) + assert not executor.iterators and not executor.state.iterators + assert not executor.state.batch_inputs + assert rank.closed == 1 + if delivered is not None: + assert delivered.hidden_states.item() == 6 + view.backward(delivered.hidden_states.sum()) + assert rank.weight.grad.item() == 3 + view.zero_grad() + view.backward(view.forward(_input(7)).hidden_states.sum()) + assert rank.weight.grad.item() == 7 + assert not executor.state.graphs and not rank.cache.handles() + + +@pytest.mark.parametrize("kind", ["copy", "attach", "assembly"]) +def test_later_packet_failure_releases_every_physical_owner(monkeypatch, kind): + ranks = [_CachedRank(dp, 2) for dp in range(2)] + executors = [_Executor(rank, "zero") for rank in ranks] + view = _view(executors[0]) + invoke = executors[0].invoke + releases = [] + + def dispatch(operation, *args, **kwargs): + result = invoke(operation, *args, **kwargs) + if operation == "release": + releases.append(args[0]) + executors[1].invoke(operation, *args, **kwargs) + return result + + monkeypatch.setattr(executors[0], "invoke", dispatch) + requests = [_input(3), _input(5)] + attached = [] + attach = executors[0].state.collector.attach + + def remember(*args, **kwargs): + result = attach(*args, **kwargs) + attached.append(result) + return result + + monkeypatch.setattr(executors[0].state.collector, "attach", remember) + with monkeypatch.context() as patch: + if kind != "assembly": + error, partial_outputs = _fail_delivery(patch, executors[1], kind) + wave = [] + for index, executor in enumerate(executors): + packet = executor._packet([ranks[index].forward(requests[index])], 1) + batch = MicroBatch( + [], [], [index], MicroBatchStats(0, 1, 2, 1, 0, 0, 0, 0, 0, False) + ) + wave.append((batch, packet)) + if kind == "assembly": + wave[0] = (replace(wave[0][0], stats=None), wave[0][1]) + with pytest.raises((MemoryError, TypeError)) as failure: + view._combine_wave(wave, requests) + if kind != "assembly": + assert failure.value is error + assert bool(partial_outputs) is (kind == "attach") + assert failure.value.__traceback__ is not None + assert attached # A prior packet's proxy remains strongly referenced. + assert releases and set(releases[-1]) == {"zero:1:dp:0", "zero:1:dp:1"} + assert all(not executor.state.graphs for executor in executors) + assert all(not rank.cache.handles() for rank in ranks) + assert all(record() is None for rank in ranks for record in rank.records) + for rank, executor in zip(ranks, executors, strict=True): + local = _view(_Executor(rank, "rank")) + local.backward(local.forward(_input(7)).hidden_states.sum()) + assert rank.weight.grad.item() == 7 + + +def test_failed_release_preserves_delivery_error_and_retries_without_head_flush( + monkeypatch, +): + rank = _CachedRank() + executor = _Executor(rank, "rank") + view = _view(executor) + invoke = executor.invoke + failed = False + flushes = [] + + def release_once(operation, *args, **kwargs): + nonlocal failed + if operation == "release" and not failed: + failed = True + raise RuntimeError("injected release failure") + return invoke(operation, *args, **kwargs) + + monkeypatch.setattr(executor, "invoke", release_once) + monkeypatch.setattr(view, "_flush_heads", lambda: flushes.append(True)) + with monkeypatch.context() as patch: + error, _ = _fail_delivery(patch, executor, "copy") + with pytest.raises(MemoryError) as failure: + view.forward(_input(3)) + assert failure.value is error and error.__traceback__ is not None + assert failed and len(flushes) == 1 + assert executor.state.released == set(executor.state.graphs) + assert executor.state.graphs and rank.cache.handles() + assert view.optim_step() == {"steps": 1} + assert not executor.state.released and not executor.state.graphs + assert not rank.cache.handles() From a75373b1629eec9263bf1873e74996bb0dca1cfb Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 21 Sep 2026 11:06:09 +0000 Subject: [PATCH 039/150] Preserve delivery errors while closing trainer iterators --- src/art/trainer_rank/_commands.py | 25 ++++- .../unit/test_trainer_rank_output_delivery.py | 103 ++++++++++++++++-- 2 files changed, 112 insertions(+), 16 deletions(-) diff --git a/src/art/trainer_rank/_commands.py b/src/art/trainer_rank/_commands.py index bc0d4f4ca..17ec154d0 100644 --- a/src/art/trainer_rank/_commands.py +++ b/src/art/trainer_rank/_commands.py @@ -830,7 +830,7 @@ def next_forward_batch(self, handle: str) -> _impl.MicroBatch | None: try: batch = self._combine_wave(self._invoke("batches_next", handle), items) except BaseException: - self.close_forward_batches(handle) + self._close_failed_iterator("batches_close", handle) raise if batch is None: self.close_forward_batches(handle) @@ -839,6 +839,13 @@ def next_forward_batch(self, handle: str) -> _impl.MicroBatch | None: def close_forward_batches(self, handle: str) -> None: self._invoke("batches_close", handle) + def _close_failed_iterator(self, operation: str, handle: int | str) -> None: + try: + # Pending releases and head publication must not prevent closure. + self._executor.invoke(operation, handle) + except BaseException: + pass # Preserve the delivery error and any queued graph release. + def _iterate_batches( self, items: Any, @@ -846,11 +853,16 @@ def _iterate_batches( handles: list[str] | None = None, ) -> Iterator[_impl.MicroBatch]: identifier = self._invoke("batches", items, **kwargs) + delivery_failed = False try: while True: - batch = self._combine_wave( - self._invoke("next", identifier), items, handles - ) + try: + batch = self._combine_wave( + self._invoke("next", identifier), items, handles + ) + except BaseException: + delivery_failed = True + raise if batch is None: return yield batch @@ -858,7 +870,10 @@ def _iterate_batches( # This iterator belongs to its creating callback, even if a retained # traceback delays its finalizer until a later callback is serving. if not self._executor.stopped: - self._invoke("close", identifier) + if delivery_failed: + self._close_failed_iterator("close", identifier) + else: + self._invoke("close", identifier) def _combine_wave( self, wave: Any, items: Any, accumulated: list[str] | None = None diff --git a/tests/unit/test_trainer_rank_output_delivery.py b/tests/unit/test_trainer_rank_output_delivery.py index 98cd6aa22..801455504 100644 --- a/tests/unit/test_trainer_rank_output_delivery.py +++ b/tests/unit/test_trainer_rank_output_delivery.py @@ -210,33 +210,114 @@ def remember(*args, **kwargs): assert rank.weight.grad.item() == 7 +@pytest.mark.parametrize( + "delivery,close_error", + [ + ("rank", False), + ("iterator", False), + ("persistent", False), + ("iterator", True), + ("persistent", True), + ], +) def test_failed_release_preserves_delivery_error_and_retries_without_head_flush( - monkeypatch, + monkeypatch, delivery, close_error ): rank = _CachedRank() - executor = _Executor(rank, "rank") + executor = _Executor(rank, "rank" if delivery == "rank" else "zero") view = _view(executor) invoke = executor.invoke - failed = False + release_calls = 0 flushes = [] - def release_once(operation, *args, **kwargs): - nonlocal failed - if operation == "release" and not failed: - failed = True - raise RuntimeError("injected release failure") + def fail_release(operation, *args, **kwargs): + nonlocal release_calls + if operation == "release": + release_calls += 1 + if release_calls <= (1 if delivery == "rank" else 2): + raise RuntimeError("injected release failure") return invoke(operation, *args, **kwargs) - monkeypatch.setattr(executor, "invoke", release_once) + monkeypatch.setattr(executor, "invoke", fail_release) monkeypatch.setattr(view, "_flush_heads", lambda: flushes.append(True)) with monkeypatch.context() as patch: + if close_error: + batches = rank.forward_batches + + def fail_close(*args, **kwargs): + try: + yield from batches(*args, **kwargs) + finally: + raise RuntimeError("injected iterator close failure") + + patch.setattr(rank, "forward_batches", fail_close) error, _ = _fail_delivery(patch, executor, "copy") + if delivery == "iterator": + iterator = view.forward_batches([_input(3)]) + advance = lambda: next(iterator) + elif delivery == "persistent": + handle = view.open_forward_batches([_input(3)]) + advance = lambda: view.next_forward_batch(handle) + else: + advance = lambda: view.forward(_input(3)) with pytest.raises(MemoryError) as failure: - view.forward(_input(3)) + advance() assert failure.value is error and error.__traceback__ is not None - assert failed and len(flushes) == 1 + assert release_calls == 1 + assert len(flushes) == (1 if delivery == "rank" else 2) + assert not executor.iterators and not executor.state.iterators + assert not executor.state.batch_inputs + assert rank.closed == int(delivery != "rank") + assert executor.state.released == set(executor.state.graphs) + assert executor.state.graphs and rank.cache.handles() + if delivery != "rank": + with pytest.raises(RuntimeError, match="injected release failure"): + view.optim_step() + assert release_calls == 2 and rank.steps == 0 and len(flushes) == 2 assert executor.state.released == set(executor.state.graphs) assert executor.state.graphs and rank.cache.handles() assert view.optim_step() == {"steps": 1} assert not executor.state.released and not executor.state.graphs assert not rank.cache.handles() + assert all(record() is None for record in rank.records) + view.backward(view.forward(_input(7)).hidden_states.sum()) + assert rank.weight.grad.item() == 7 + assert not executor.state.graphs and not rank.cache.handles() + + +@pytest.mark.parametrize("ending", ["exhaust", "close", "throw", "close_error"]) +def test_non_delivery_iterator_closure_keeps_head_publication(monkeypatch, ending): + rank = _CachedRank() + executor = _Executor(rank, "zero") + view = _view(executor) + flushes = [] + fail = False + + def flush(): + flushes.append(True) + if fail: + raise RuntimeError("injected close head publication failure") + + monkeypatch.setattr(view, "_flush_heads", flush) + iterator = view.forward_batches([_input(3)]) + output = next(iterator).outputs[0] + assert len(flushes) == 2 + if ending == "exhaust": + with pytest.raises(StopIteration): + next(iterator) + elif ending == "throw": + with pytest.raises(ValueError, match="consumer failure"): + iterator.throw(ValueError("consumer failure")) + elif ending == "close_error": + fail = True + with pytest.raises(RuntimeError, match="close head publication failure"): + iterator.close() + assert rank.closed == 0 and executor.iterators + else: + iterator.close() + assert len(flushes) == (4 if ending == "exhaust" else 3) + executor.stop() + assert rank.closed == 1 and not executor.iterators + _view(_Executor(rank, "zero")).backward(output.hidden_states.sum()) + assert rank.weight.grad.item() == 3 + assert not executor.state.graphs and not rank.cache.handles() From 022eeaa6bf1b6b1fae0dbfef7bddd11f0ac5b5ff Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 21 Sep 2026 14:32:59 +0000 Subject: [PATCH 040/150] Fix trainer output delivery test typing --- tests/unit/test_trainer_rank_output_delivery.py | 17 +++++++++++------ 1 file changed, 11 insertions(+), 6 deletions(-) diff --git a/tests/unit/test_trainer_rank_output_delivery.py b/tests/unit/test_trainer_rank_output_delivery.py index 801455504..67a561672 100644 --- a/tests/unit/test_trainer_rank_output_delivery.py +++ b/tests/unit/test_trainer_rank_output_delivery.py @@ -1,7 +1,9 @@ """Undelivered logical outputs must not retain physical forward graphs.""" +from collections.abc import Generator from dataclasses import replace from functools import partial +from typing import Any import weakref import pytest @@ -97,7 +99,7 @@ def fail_attach(tensor): def test_failed_delivery_releases_registered_graph_and_native_cache( monkeypatch, mode, kind ): - rank = _CachedRank() + rank: Any = _CachedRank() executor = _Executor(rank, mode) view = _view(executor) with monkeypatch.context() as patch: @@ -119,7 +121,7 @@ def test_failed_delivery_releases_registered_graph_and_native_cache( def test_later_wave_failure_keeps_only_previously_delivered_outputs( monkeypatch, delivery ): - rank = _CachedRank() + rank: Any = _CachedRank() executor = _Executor(rank, "zero") view = _view(executor) requests = [_input(3), _input(5)] @@ -132,7 +134,9 @@ def test_later_wave_failure_keeps_only_previously_delivered_outputs( advance = lambda: next(iterator) elif delivery == "persistent": handle = view.open_forward_batches(requests) - delivered = view.next_forward_batch(handle).outputs[0] + batch = view.next_forward_batch(handle) + assert batch is not None + delivered = batch.outputs[0] advance = lambda: view.next_forward_batch(handle) else: advance = lambda: view.forward(requests) @@ -157,7 +161,7 @@ def test_later_wave_failure_keeps_only_previously_delivered_outputs( @pytest.mark.parametrize("kind", ["copy", "attach", "assembly"]) def test_later_packet_failure_releases_every_physical_owner(monkeypatch, kind): - ranks = [_CachedRank(dp, 2) for dp in range(2)] + ranks: list[Any] = [_CachedRank(dp, 2) for dp in range(2)] executors = [_Executor(rank, "zero") for rank in ranks] view = _view(executors[0]) invoke = executors[0].invoke @@ -223,7 +227,7 @@ def remember(*args, **kwargs): def test_failed_release_preserves_delivery_error_and_retries_without_head_flush( monkeypatch, delivery, close_error ): - rank = _CachedRank() + rank: Any = _CachedRank() executor = _Executor(rank, "rank" if delivery == "rank" else "zero") view = _view(executor) invoke = executor.invoke @@ -287,7 +291,7 @@ def fail_close(*args, **kwargs): @pytest.mark.parametrize("ending", ["exhaust", "close", "throw", "close_error"]) def test_non_delivery_iterator_closure_keeps_head_publication(monkeypatch, ending): - rank = _CachedRank() + rank: Any = _CachedRank() executor = _Executor(rank, "zero") view = _view(executor) flushes = [] @@ -301,6 +305,7 @@ def flush(): monkeypatch.setattr(view, "_flush_heads", flush) iterator = view.forward_batches([_input(3)]) output = next(iterator).outputs[0] + assert isinstance(iterator, Generator) assert len(flushes) == 2 if ending == "exhaust": with pytest.raises(StopIteration): From 595f12cebe17299c9962e69a9844a947053bcd47 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 21 Sep 2026 14:37:57 +0000 Subject: [PATCH 041/150] Keep Tinker client imports isolated from backend initialization --- src/art/tinker/__init__.py | 26 +++++- tests/unit/test_tinker_import_boundary.py | 102 ++++++++++++++++++++++ 2 files changed, 126 insertions(+), 2 deletions(-) create mode 100644 tests/unit/test_tinker_import_boundary.py diff --git a/src/art/tinker/__init__.py b/src/art/tinker/__init__.py index b706cd4f5..1949b24e1 100644 --- a/src/art/tinker/__init__.py +++ b/src/art/tinker/__init__.py @@ -1,5 +1,27 @@ -from .backend import TinkerBackend +from typing import TYPE_CHECKING, Any + from .renderers import get_renderer_name -from .server import OpenAICompatibleTinkerServer + +if TYPE_CHECKING: + from .backend import TinkerBackend + from .server import OpenAICompatibleTinkerServer __all__ = ["TinkerBackend", "get_renderer_name", "OpenAICompatibleTinkerServer"] + + +def __getattr__(name: str) -> Any: + if name == "TinkerBackend": + from .backend import TinkerBackend + + globals()[name] = TinkerBackend + return TinkerBackend + if name == "OpenAICompatibleTinkerServer": + from .server import OpenAICompatibleTinkerServer + + globals()[name] = OpenAICompatibleTinkerServer + return OpenAICompatibleTinkerServer + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") + + +def __dir__() -> list[str]: + return sorted(set(globals()) | set(__all__)) diff --git a/tests/unit/test_tinker_import_boundary.py b/tests/unit/test_tinker_import_boundary.py new file mode 100644 index 000000000..d810036b6 --- /dev/null +++ b/tests/unit/test_tinker_import_boundary.py @@ -0,0 +1,102 @@ +import os +from pathlib import Path +import subprocess +import sys +import textwrap + +import pytest + + +def _run(script: str) -> None: + root = Path(__file__).resolve().parents[2] + result = subprocess.run( + [sys.executable, "-c", textwrap.dedent(script)], + cwd=root, + env={ + **os.environ, + "PYTHONPATH": os.pathsep.join( + (str(root / "src"), os.getenv("PYTHONPATH", "")) + ), + "PYTHON_DOTENV_DISABLED": "1", + }, + capture_output=True, + text=True, + timeout=60, + ) + assert result.returncode == 0, result.stdout + result.stderr + + +def test_client_import_preserves_native_asyncio() -> None: + _run( + """ + import _asyncio + import asyncio + import importlib + import multiprocessing + import sys + + def snapshot(): + policy = asyncio.get_event_loop_policy() + loop = asyncio.new_event_loop() + try: + assert not getattr(policy, "_nest_patched", False) + assert not getattr(loop, "_nest_patched", False) + return ( + asyncio.Task, asyncio.tasks.Task, + asyncio.Future, asyncio.futures.Future, + asyncio.run, asyncio.get_event_loop, + asyncio.events.get_event_loop, type(policy).get_event_loop, + type(loop).run_until_complete, type(loop).run_forever, + type(loop)._run_once, + multiprocessing.get_start_method(allow_none=True), + ) + finally: + loop.close() + + assert asyncio.Task is _asyncio.Task + assert asyncio.Future is _asyncio.Future + native = snapshot() + for name in ("art", "art.tinker", "art.tinker.client"): + importlib.import_module(name) + assert snapshot() == native, name + assert not any(m == "mp_actors" or m.startswith("mp_actors.") + for m in sys.modules), name + assert {"art.tinker.backend", "art.tinker.server", "art.local.backend", + "nest_asyncio"}.isdisjoint(sys.modules), name + """ + ) + + +@pytest.mark.parametrize("name", ["TinkerBackend", "OpenAICompatibleTinkerServer"]) +def test_public_exports_keep_original_classes(name: str) -> None: + _run( + f""" + import importlib + import art.tinker as package + + public = ["TinkerBackend", "get_renderer_name", "OpenAICompatibleTinkerServer"] + assert package.__all__ == public + assert set(public) <= set(dir(package)) + try: + package.unknown_export + except AttributeError as error: + assert str(error) == "module 'art.tinker' has no attribute 'unknown_export'" + else: + raise AssertionError("unknown export did not raise AttributeError") + + first = getattr(package, {name!r}) + namespace = {{}} + exec("from art.tinker import *", namespace) + from art.tinker import TinkerBackend, OpenAICompatibleTinkerServer + for export, module in (("TinkerBackend", "backend"), + ("OpenAICompatibleTinkerServer", "server"), + ("get_renderer_name", "renderers")): + original = getattr(importlib.import_module("art.tinker." + module), export) + assert getattr(package, export) is original + assert vars(package)[export] is original + assert namespace[export] is original + assert first is namespace[{name!r}] + assert TinkerBackend is namespace["TinkerBackend"] + assert OpenAICompatibleTinkerServer is namespace["OpenAICompatibleTinkerServer"] + """ + ) From 4d9c63111aa961a84676130061124bee76158c2c Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 21 Sep 2026 16:35:24 +0000 Subject: [PATCH 042/150] Release logical forward batches before advancing --- src/art/trainer_rank/_commands.py | 1 + tests/unit/test_trainer_rank_commands.py | 29 ++++++++++++++++++++++++ 2 files changed, 30 insertions(+) diff --git a/src/art/trainer_rank/_commands.py b/src/art/trainer_rank/_commands.py index 17ec154d0..9682ef26b 100644 --- a/src/art/trainer_rank/_commands.py +++ b/src/art/trainer_rank/_commands.py @@ -866,6 +866,7 @@ def _iterate_batches( if batch is None: return yield batch + del batch finally: # This iterator belongs to its creating callback, even if a retained # traceback delays its finalizer until a later callback is serving. diff --git a/tests/unit/test_trainer_rank_commands.py b/tests/unit/test_trainer_rank_commands.py index 95328aa92..fe57cac5d 100644 --- a/tests/unit/test_trainer_rank_commands.py +++ b/tests/unit/test_trainer_rank_commands.py @@ -10,6 +10,7 @@ import threading from types import SimpleNamespace from typing import Any +import weakref import pytest import torch @@ -657,6 +658,34 @@ def callback(view): asyncio.run(run_rank_callback(rank, callback, mode="zero")) +@pytest.mark.parametrize("mode", ["rank", "zero"]) +def test_logical_iterator_releases_previous_batch_before_next_forward(mode): + previous = [] + + class Rank(_Rank): + def forward(self, tree, **kwargs): + assert all(reference() is None for reference in previous) + return super().forward(tree, **kwargs) + + rank: Any = Rank() + + def callback(view): + for batch in view.forward_batches([_input(3), _input(5)]): + previous[:] = [ + weakref.ref(batch), + weakref.ref(batch.outputs[0]), + weakref.ref(batch.outputs[0].hidden_states), + ] + view.backward(batch.outputs[0].hidden_states.sum()) + assert not rank._rank_command_state.graphs + del batch + + asyncio.run(run_rank_callback(rank, callback, mode=mode)) + assert rank.weight.grad.item() == 8 + assert rank.closed == 1 + assert all(reference() is None for reference in previous) + + def test_persistent_iterator_binds_policy_and_checkpoint_and_pulls_one_wave(): class Rank(_Rank): _capture_forward_options = TrainerRank._capture_forward_options From 3b6f05e95b91b548c7d9c93cc879ba2384fe2ed3 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 16:28:15 +0000 Subject: [PATCH 043/150] Preserve caller RNG while advancing and replaying trainer graphs --- scripts/ci/trainer-rank-gpu-tests.sh | 1 + src/art/trainer_rank/_graphs.py | 36 +- src/art/trainer_rank/_impl.py | 25 +- src/art/trainer_rank/_rng.py | 127 +++++ .../megatron/lora/test_dynamic_lora_slots.py | 3 + .../unit/test_trainer_batch_input_capture.py | 2 + tests/unit/test_trainer_rank_commands.py | 3 + tests/unit/test_trainer_rank_rng.py | 515 ++++++++++++++++++ 8 files changed, 674 insertions(+), 38 deletions(-) create mode 100644 src/art/trainer_rank/_rng.py create mode 100644 tests/unit/test_trainer_rank_rng.py diff --git a/scripts/ci/trainer-rank-gpu-tests.sh b/scripts/ci/trainer-rank-gpu-tests.sh index d6729a6b4..939ffcd93 100755 --- a/scripts/ci/trainer-rank-gpu-tests.sh +++ b/scripts/ci/trainer-rank-gpu-tests.sh @@ -11,6 +11,7 @@ test -x "${runtime_python}" "${runtime_python}" -m pytest --tb=short \ tests/unit/test_trainer_rank_head_recompute.py \ + tests/unit/test_trainer_rank_rng.py \ tests/unit/test_trainer_rank_custom_tensors.py \ tests/unit/test_trainer_rank_tensors.py \ tests/unit/test_trainer_rank_graphs_cuda.py \ diff --git a/src/art/trainer_rank/_graphs.py b/src/art/trainer_rank/_graphs.py index 7bc08fd12..e6aa39acf 100644 --- a/src/art/trainer_rank/_graphs.py +++ b/src/art/trainer_rank/_graphs.py @@ -6,7 +6,6 @@ from contextlib import AbstractContextManager, ExitStack, contextmanager, nullcontext from copy import deepcopy from dataclasses import dataclass, field, fields, is_dataclass, replace -import random from time import perf_counter from typing import Any, Literal from uuid import uuid4 @@ -18,6 +17,7 @@ from art._tensor_residency import observe_resident_tensors +from ._rng import RNGState as _RNGState from ._tensors import _map_tensor_arguments ForwardHandle = str @@ -99,40 +99,6 @@ def _storage_sizes(tensors: Iterable[torch.Tensor]) -> tuple[int, int]: ) -@dataclass -class _RNGState: - cpu: torch.Tensor - cuda: dict[int, torch.Tensor] - python: tuple[Any, ...] - tracker: Any = None - - @classmethod - def capture(cls, devices: Sequence[int], tracker: Any) -> _RNGState: - return cls( - torch.get_rng_state(), - {device: torch.cuda.get_rng_state(device) for device in devices}, - random.getstate(), - None if tracker is None else _snapshot(tracker.get_states()), - ) - - def restore(self, tracker: Any) -> None: - torch.set_rng_state(self.cpu) - for device, state in self.cuda.items(): - torch.cuda.set_rng_state(state, device) - random.setstate(self.python) - if tracker is not None: - tracker.set_states(_snapshot(self.tracker)) - - @contextmanager - def replay(self, tracker: Any): - ambient = self.capture(tuple(self.cuda), tracker) - try: - self.restore(tracker) - yield - finally: - ambient.restore(tracker) - - @dataclass class _TransferStats: """Completed saved-storage copies, cumulative across released graphs.""" diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 5ec1244ec..5dc252ce8 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -86,6 +86,7 @@ prefix_tree_layout_candidates, select_prefix_tree_layout, ) +from art.trainer_rank._rng import TrainerRNG, caller_group from art.trainer_rank._telemetry import phase as _telemetry_phase from art.trainer_rank._versions import ( CheckpointVersion, @@ -1928,6 +1929,7 @@ def __init__( resolve_forward_options(options) self.runtime: TrainingRuntime = runtime self.device: torch.device = next(runtime.model[0].parameters()).device + self._rng = TrainerRNG(self.device) self._param_dtype_size = _dtype_size(next(runtime.model[0].parameters()).dtype) try: metadata_model = _language_model(runtime.model[0]) @@ -2366,6 +2368,7 @@ def _custom_object( raise TrainerRankSlotStateError( "Custom checkpoint object registration differs across ranks" ) + self._rng.synchronize(caller_group()) slot = self._checkpoint_slots[checkpoint_name] existing = slot.custom.get(name) registered = None if existing is None else existing.kind @@ -3075,6 +3078,16 @@ def forward_batches( forwards and backwards on their TP/CP peers. `reduce` combines only distinct data-parallel batches. + Model PyTorch randomness advances separately from caller randomness, + seeded when the physical TrainerRank is constructed. Direct physical + callers continue their TP/CP leader's default CPU and trainer-device CUDA + streams before each yield/forward return and custom-object factory. This + keeps matching random masks and custom-head dropout consistent without + synchronizing DP workers. Python/NumPy RNGs, explicit generators, other + devices, concurrent RNG use and rank-dependent control flow are outside + this contract. Checkpoint saves do not persist RNG state; activation + checkpointing must preserve RNG for correct recomputation. + Empty local microbatches are skipped unless `yield_empty=True`. Every rank must use the same setting. When a wave skips ranks, TrainerRank collective methods raise if called from its loop body; fully populated @@ -3116,13 +3129,14 @@ def _yield_forward_batches( try: while True: self._guard_forward_collective("forward_batches") - with torch.set_grad_enabled(enabled): + with torch.set_grad_enabled(enabled), self._rng.model(): try: batch = next(batches) except StopIteration: return if not yield_empty and not batch.outputs: continue + self._rng.synchronize(caller_group()) if ( not yield_empty and batch.stats.global_count < self._dp_rank_and_size()[1] @@ -3426,9 +3440,12 @@ def forward( backward.harvest() enabled = torch.is_grad_enabled() if no_grad is None else not no_grad with torch.set_grad_enabled(enabled): - self._reset_planning_telemetry() + # Caller iterators may draw their own inputs; only ART's internal + # execution belongs to the private model stream. materialized = self._capture_forward_options(inputs, options) requests = list(_flatten(materialized)) + with torch.set_grad_enabled(enabled), self._rng.model(): + self._reset_planning_telemetry() plan, check = self._plan_admissible_forward( requests, checkpoint=checkpoint, context="forward" ) @@ -3437,7 +3454,9 @@ def forward( ) if backward is not None: backward.attach(tracked_outputs) - return _unflatten(materialized, iter(tracked_outputs)) + outputs = _unflatten(materialized, iter(tracked_outputs)) + self._rng.synchronize(caller_group()) + return outputs def _execute_admitted_plan( self, plan: _AnyForwardPlan, *, check: _MemoryCheck, context: str diff --git a/src/art/trainer_rank/_rng.py b/src/art/trainer_rank/_rng.py new file mode 100644 index 000000000..15aa946f3 --- /dev/null +++ b/src/art/trainer_rank/_rng.py @@ -0,0 +1,127 @@ +"""Model/caller RNG ownership and replay of recorded forward randomness.""" + +from collections.abc import Iterator, Sequence +from contextlib import contextmanager +from copy import deepcopy +from dataclasses import dataclass +import hashlib +import random +from typing import Any + +import torch +import torch.distributed as dist + + +@dataclass +class RNGState: + cpu: torch.Tensor + cuda: dict[int, torch.Tensor] + python: tuple[Any, ...] | None = None + tracker: Any = None + + @classmethod + def capture( + cls, + devices: Sequence[int], + tracker: Any = None, + *, + torch_only: bool = False, + ) -> "RNGState": + return cls( + torch.get_rng_state(), + {device: torch.cuda.get_rng_state(device) for device in devices}, + None if torch_only else random.getstate(), + None if tracker is None else deepcopy(tracker.get_states()), + ) + + def restore(self, tracker: Any = None) -> None: + torch.set_rng_state(self.cpu) + for device, state in self.cuda.items(): + torch.cuda.set_rng_state(state, device) + if self.python is not None: + random.setstate(self.python) + if tracker is not None: + tracker.set_states(deepcopy(self.tracker)) + + @contextmanager + def replay(self, tracker: Any = None) -> Iterator[None]: + ambient = self.capture( + tuple(self.cuda), tracker, torch_only=self.python is None + ) + try: + self.restore(tracker) + yield + finally: + ambient.restore(tracker) + + +class TrainerRNG: + def __init__(self, device: torch.device) -> None: + self.device = device + self.devices = ( + (torch.cuda.current_device() if device.index is None else device.index,) + if device.type == "cuda" + else () + ) + initial = RNGState.capture(self.devices, torch_only=True) + + def derive(state: torch.Tensor, target: torch.device | str) -> torch.Tensor: + digest = hashlib.sha256(b"art.trainer_rank.model" + state.numpy().tobytes()) + seed = int.from_bytes(digest.digest()[:8], "little") + return torch.Generator(device=target).manual_seed(seed).get_state() + + # Capture before logical leaders run user code or initialize custom heads; + # those draws must not change the first model stream relative to peers. + self._model = RNGState( + derive(initial.cpu, "cpu"), + {index: derive(state, device) for index, state in initial.cuda.items()}, + ) + self._depth = 0 + + @contextmanager + def model(self) -> Iterator[None]: + """Advance the private torch stream, restoring caller state even on error. + + Never span a public iterator yield. Python and Megatron's separate RNG + tracker are untouched; activation checkpointing must preserve its RNG. + """ + if self._depth: + yield + return + with self._model.replay(): + self._depth += 1 + try: + yield + finally: + self._depth -= 1 + self._model = RNGState.capture(self.devices, torch_only=True) + + def synchronize(self, group: dist.ProcessGroup | None) -> None: + """Continue the TP×CP leader's caller stream, never synchronizing DP. + + None means no model-parallel group, not WORLD. Explicit caller reseeding + or restoration remains authoritative, and equal DP seeds remain equal. + """ + if group is None or dist.get_world_size(group) == 1: + return + state = RNGState.capture(self.devices, torch_only=True) + states = [state.cpu, *state.cuda.values()] + payload = torch.cat(states).to( + self.device if dist.get_backend(group) == "nccl" else "cpu" + ) + dist.broadcast(payload, src=dist.get_global_rank(group, 0), group=group) + received = payload.cpu().split([value.numel() for value in states]) + RNGState( + received[0], dict(zip(state.cuda, received[1:], strict=True)) + ).restore() + + +def caller_group() -> dist.ProcessGroup | None: + if not (dist.is_available() and dist.is_initialized()): + return None + try: + from megatron.core import parallel_state as ps + + return ps.get_tensor_and_context_parallel_group(check_initialized=False) + except (AssertionError, ImportError, RuntimeError, ValueError): + return None diff --git a/tests/integration/megatron/lora/test_dynamic_lora_slots.py b/tests/integration/megatron/lora/test_dynamic_lora_slots.py index 0782d35e0..d902b6d2a 100644 --- a/tests/integration/megatron/lora/test_dynamic_lora_slots.py +++ b/tests/integration/megatron/lora/test_dynamic_lora_slots.py @@ -595,6 +595,8 @@ def _optimizer_state(trainer: TrainerRank, name: str) -> LocalOptimizerState: def _trainer_for(lora: LoRA, device: torch.device) -> TrainerRank: + from art.trainer_rank._rng import TrainerRNG + trainer = TrainerRank.__new__(TrainerRank) trainer.runtime = SimpleNamespace( model=[lora], @@ -602,6 +604,7 @@ def _trainer_for(lora: LoRA, device: torch.device) -> TrainerRank: model_support_handler=_IdentityModelSupportHandler(), ) trainer.device = device + trainer._rng = TrainerRNG(device) trainer._slot_stack = [] trainer._default_slot_ref = None trainer._skipped_forward_waves = {} diff --git a/tests/unit/test_trainer_batch_input_capture.py b/tests/unit/test_trainer_batch_input_capture.py index e1353eb47..349b44af4 100644 --- a/tests/unit/test_trainer_batch_input_capture.py +++ b/tests/unit/test_trainer_batch_input_capture.py @@ -18,6 +18,7 @@ Unset, run_rank_callback, ) +from art.trainer_rank._rng import TrainerRNG class _CapturingRank(_Rank): @@ -33,6 +34,7 @@ def test_batches_snapshot_tokens_targets_and_structure_before_first_pull( rank: Any if surface == "native": rank = object.__new__(TrainerRank) + rank._rng = TrainerRNG(torch.device("cpu")) rank._skipped_forward_waves = {} def batches(inputs, **kwargs): diff --git a/tests/unit/test_trainer_rank_commands.py b/tests/unit/test_trainer_rank_commands.py index fe57cac5d..c701d91ac 100644 --- a/tests/unit/test_trainer_rank_commands.py +++ b/tests/unit/test_trainer_rank_commands.py @@ -552,7 +552,10 @@ def test_released_client_graph_releases_physical_bridge(): def test_forward_batches_captures_policy_before_iteration(): + from art.trainer_rank._rng import TrainerRNG + rank = object.__new__(TrainerRank) + rank._rng = TrainerRNG(torch.device("cpu")) rank._forward_options = ForwardOptions(max_gradient_staleness=1, allow_replay=False) rank._skipped_forward_waves = {} request = _input(3) diff --git a/tests/unit/test_trainer_rank_rng.py b/tests/unit/test_trainer_rank_rng.py new file mode 100644 index 000000000..b51efaca6 --- /dev/null +++ b/tests/unit/test_trainer_rank_rng.py @@ -0,0 +1,515 @@ +from __future__ import annotations + +from contextlib import nullcontext +from datetime import timedelta +from types import SimpleNamespace +from typing import Any, cast + +import pytest +import torch +import torch.distributed as dist +import torch.multiprocessing as mp +from torch.utils.checkpoint import checkpoint + +from art.trainer_rank import AdamParams, ForwardInput, ForwardOutput, TrainerRank +from art.trainer_rank._impl import _CheckpointSlot, _GatherContextParallelRows +from art.trainer_rank._rng import RNGState, TrainerRNG, caller_group + + +def test_model_seed_does_not_depend_on_logical_leader_draws(): + torch.manual_seed(517) + leader = TrainerRNG(torch.device("cpu")) + peer = TrainerRNG(torch.device("cpu")) + # Only the logical leader runs the caller's head factory and random masks. + torch.nn.Linear(7, 11) + torch.rand(37) + with leader.model(): + expected = torch.rand(23) + torch.manual_seed(919) + with peer.model(): + actual = torch.rand(23) + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + + +@pytest.mark.parametrize("retention", ("gpu", "cpu", "replay")) +def test_graph_replay_preserves_caller_and_private_model_progress( + monkeypatch, retention +): + from art.trainer_rank._graphs import GraphCache + + trainer = _trainer() + cache = GraphCache() + weight = torch.nn.Parameter(torch.tensor(2.0)) + states, handles = [], [] + + def execute(_): + output = torch.nn.functional.dropout(torch.ones(29), 0.3) * weight + states.append(torch.get_rng_state()) + return (output,) + + def forward(): + handle, (output,) = cache.run(execute, None, retention=retention) + handles.append(handle) + return [ForwardOutput(output, None, None, None)] + + _stub_forward(monkeypatch, trainer, forward) + torch.manual_seed(791) + caller = torch.get_rng_state() + output = trainer.forward([ForwardInput(input_tokens=torch.arange(29))])[0] + assert output.target_logprobs is not None + assert torch.equal(torch.get_rng_state(), caller) + torch.rand(17) + caller = torch.get_rng_state() + cache.backward(handles[0], (torch.ones_like(output.target_logprobs),)) + assert torch.equal(torch.get_rng_state(), caller) + assert weight.grad is not None + torch.testing.assert_close(weight.grad, (output.target_logprobs / 2).sum()) + generator = torch.Generator().set_state(states[0]) + with trainer._rng.model(): + torch.testing.assert_close(torch.rand(19), torch.rand(19, generator=generator)) + + +def _trainer(device="cpu"): + model = torch.nn.Linear(3, 4, bias=False, device=device) + runtime = SimpleNamespace( + model=[model], + optimizer=None, + provider=SimpleNamespace(hidden_size=4, num_layers=1), + model_support_handler=SimpleNamespace( + build_gdn_execution_spec=False, + zero_internal_padding_grads=lambda _: None, + ), + ) + return TrainerRank(cast(Any, runtime)) + + +def _stub_forward(monkeypatch, trainer, execute): + monkeypatch.setattr( + trainer, "_plan_admissible_forward", lambda *a, **k: (None, None) + ) + monkeypatch.setattr(trainer, "_execute_admitted_plan", lambda *a, **k: execute()) + + +def _state(device): + return RNGState.capture( + (device.index,) if device.type == "cuda" else (), torch_only=True + ) + + +def _assert_state_equal(left, right): + assert torch.equal(left.cpu, right.cpu) + assert left.cuda.keys() == right.cuda.keys() + for device in left.cuda: + assert torch.equal(left.cuda[device], right.cuda[device]) + + +def test_model_stream_advances_without_advancing_caller(monkeypatch): + trainer = _trainer() + torch.manual_seed(811) + caller = torch.get_rng_state() + expected = torch.rand(3, 12) + torch.set_rng_state(caller) + observed = [] + internal_states = [] + + def execute(): + # Nested internal work shares the model stream instead of restarting it. + with trainer._rng.model(): + observed.append(torch.rand(12)) + internal_states.append(torch.get_rng_state()) + return [] + + _stub_forward(monkeypatch, trainer, execute) + trainer.forward([], no_grad=True) + assert torch.equal(torch.get_rng_state(), caller) + # Caller randomness advances independently between model forwards. + assert torch.equal(torch.rand(12), expected[0]) + trainer.forward([]) + assert not torch.equal(observed[0], expected[0]) + generator = torch.Generator().set_state(internal_states[0]) + torch.testing.assert_close(observed[1], torch.rand(12, generator=generator)) + assert torch.equal(torch.rand(12), expected[1]) + + +def test_forward_failure_restores_caller_and_advances_model(monkeypatch): + trainer = _trainer() + torch.manual_seed(981) + caller = torch.get_rng_state() + internal_states = [] + + def fail(): + torch.rand(7) + internal_states.append(torch.get_rng_state()) + raise ValueError("model failed") + + _stub_forward(monkeypatch, trainer, fail) + with pytest.raises(ValueError, match="model failed"): + trainer.forward([]) + assert torch.equal(torch.get_rng_state(), caller) + generator = torch.Generator().set_state(internal_states[0]) + with trainer._rng.model(): + assert torch.equal(torch.rand(7), torch.rand(7, generator=generator)) + + +@pytest.mark.parametrize("yield_empty", (False, True)) +def test_microbatch_yields_and_close_do_not_hold_rng_context(monkeypatch, yield_empty): + trainer = _trainer() + torch.manual_seed(177) + caller = torch.get_rng_state() + draws = torch.rand(5, 8) + torch.set_rng_state(caller) + observed = [] + internal_states = [] + + def batches(*args, **kwargs): + for index in range(3): + observed.append(torch.rand(8)) + internal_states.append(torch.get_rng_state()) + yield SimpleNamespace( + outputs=[] if index == 0 else [index], + stats=SimpleNamespace(global_count=1), + ) + + monkeypatch.setattr(trainer, "_forward_batches", batches) + iterator = trainer.forward_batches([], yield_empty=yield_empty) + next(iterator) + assert trainer._rng._depth == 0 + assert torch.equal(torch.get_rng_state(), caller) + assert torch.equal(torch.rand(8), draws[0]) + next(iterator) + assert trainer._rng._depth == 0 + iterator.close() + assert torch.equal(torch.rand(8), draws[1]) + generator = torch.Generator().set_state(internal_states[0]) + for draw in observed[1:]: + torch.testing.assert_close(draw, torch.rand(8, generator=generator)) + + +def test_uninitialized_model_parallel_group_does_not_mean_world(monkeypatch): + monkeypatch.setattr(dist, "is_initialized", lambda: False) + assert caller_group() is None + monkeypatch.setattr( + dist, "broadcast", lambda *a, **k: pytest.fail("unexpected WORLD broadcast") + ) + TrainerRNG(torch.device("cpu")).synchronize(None) + + +def test_caller_reseed_and_restore_between_forwards(monkeypatch): + trainer = _trainer() + _stub_forward(monkeypatch, trainer, lambda: (torch.rand(9), [])[1]) + trainer.forward([]) + torch.manual_seed(191) + saved = torch.get_rng_state() + expected = torch.rand(9) + torch.set_rng_state(saved) + trainer.forward([]) + assert torch.equal(torch.rand(9), expected) + torch.set_rng_state(saved) + trainer.forward([]) + assert torch.equal(torch.rand(9), expected) + + +@pytest.mark.parametrize("microbatches", (False, True)) +def test_input_iterators_use_caller_rng(monkeypatch, microbatches): + trainer = _trainer() + torch.manual_seed(617) + state = torch.get_rng_state() + expected = torch.rand(2, 5) + torch.set_rng_state(state) + request = ForwardInput(input_tokens=torch.arange(8), hidden_states=True) + + def inputs(): + assert torch.equal(torch.rand(5), expected[0]) + yield request + + def execute(): + torch.rand(11) + return [ForwardOutput(None, None, None, torch.ones(8, 4))] + + _stub_forward(monkeypatch, trainer, execute) + if microbatches: + + def batches(items, **kwargs): + assert len(items) == 1 + torch.testing.assert_close(items[0].input_tokens, request.input_tokens) + yield SimpleNamespace( + outputs=execute(), stats=SimpleNamespace(global_count=1) + ) + + monkeypatch.setattr(trainer, "_forward_batches", batches) + list(trainer.forward_batches(inputs())) + else: + trainer.forward(inputs()) + assert torch.equal(torch.rand(5), expected[1]) + + +@pytest.mark.parametrize("dp_size", (1, 2)) +def test_replicated_caller_randomness_cpu(dp_size, tmp_path): + pytest.importorskip("megatron.core") + mp.spawn( + _distributed_worker, + args=(dp_size, "cp", "gloo", f"file://{tmp_path / 'rng'}"), + nprocs=2 * dp_size, + join=True, + ) + + +@pytest.mark.parametrize("parallelism", ("tp", "cp")) +def test_replicated_caller_randomness_cuda(parallelism, tmp_path): + if not torch.cuda.is_available() or torch.cuda.device_count() < 2: + pytest.skip("requires two CUDA devices") + pytest.importorskip("megatron.core") + mp.spawn( + _distributed_worker, + args=(1, parallelism, "nccl", f"file://{tmp_path / 'rng'}"), + nprocs=2, + join=True, + ) + + +def _distributed_worker(rank, dp_size, parallelism, backend, init_method): + from megatron.core import parallel_state as ps + + device = torch.device("cpu" if backend == "gloo" else f"cuda:{rank}") + if device.type == "cuda": + torch.cuda.set_device(device) + dist.init_process_group( + backend, + init_method=init_method, + rank=rank, + world_size=2 * dp_size, + timeout=timedelta(seconds=90), + ) + try: + replica_groups = [dist.new_group([2 * dp, 2 * dp + 1]) for dp in range(dp_size)] + dp_groups = [ + dist.new_group(list(range(replica, 2 * dp_size, 2))) for replica in range(2) + ] + dp_rank, replica_rank = divmod(rank, 2) + replica_group, dp_group = replica_groups[dp_rank], dp_groups[replica_rank] + with pytest.MonkeyPatch.context() as patch: + patch.setattr( + ps, "get_tensor_and_context_parallel_group", lambda **_: replica_group + ) + patch.setattr( + ps, + "get_tensor_model_parallel_world_size", + lambda: 2 if parallelism == "tp" else 1, + ) + patch.setattr( + ps, + "get_context_parallel_world_size", + lambda: 2 if parallelism == "cp" else 1, + ) + patch.setattr( + ps, + "get_tensor_model_parallel_group", + lambda **_: replica_group if parallelism == "tp" else None, + ) + patch.setattr( + ps, + "get_data_parallel_group", + lambda **_: dp_group if parallelism == "tp" else dist.group.WORLD, + ) + patch.setattr(ps, "get_data_parallel_rank", lambda: dp_rank) + patch.setattr(ps, "get_data_parallel_world_size", lambda: dp_size) + _gradient_oracle( + patch, + device, + rank, + dp_rank, + dp_size, + replica_rank, + replica_group, + dp_group, + parallelism, + ) + finally: + dist.destroy_process_group() + + +def _gradient_oracle( + patch, + device, + rank, + dp_rank, + dp_size, + replica_rank, + replica_group, + dp_group, + parallelism, +): + trainer = _trainer(device) + decoder = trainer.runtime.model[0].weight + with torch.no_grad(): + decoder.copy_(torch.arange(12, device=device).reshape(4, 3) / 30) + if parallelism == "tp": + decoder.grad_sync_op = "sum" + trainer._checkpoint_slots["student"] = _CheckpointSlot( + config={ + "base_model_name_or_path": "test", + "r": 1, + "lora_alpha": 1, + "target_modules": [], + }, + params=(decoder,), + ) + # Registration repairs initially different CPU/CUDA states within each DP + # worker, while the registered weights remain common to all DP workers. + torch.manual_seed(711 + 100 * dp_rank + replica_rank) + head = trainer.module( + "head", + lambda: torch.nn.Sequential(torch.nn.Dropout(0.3), torch.nn.Linear(4, 2)), + checkpoint="student", + ) + params = trainer._checkpoint_slots["student"].params + reference = tuple(torch.nn.Parameter(param.detach().clone()) for param in params) + optimizer = torch.optim.AdamW(reference, lr=0.01, weight_decay=0.0) + features = ( + torch.arange(24, device=device, dtype=torch.float32).reshape(8, 3) / 20 + + dp_rank / 10 + ) + rows = torch.arange(replica_rank * 4, (replica_rank + 1) * 4, device=device) + request = ForwardInput(input_tokens=torch.arange(8), hidden_states=True) + masks = [] + recompute_masks = [] + cpu_draws = [] + previous_caller_mask = None + tracker = None + if device.type == "cuda" and parallelism == "cp": + from megatron.core.tensor_parallel.random import get_cuda_rng_tracker + + tracker = get_cuda_rng_tracker() + tracker.add("art-test-model", 3199 + rank) + + def execute(): + # Internal consumption deliberately differs across physical ranks. It + # must not leak into caller masks or custom-head dropout. + torch.rand(rank + 1) + torch.rand(rank + 3, device=device) + records = [] + local_rows = rows + if parallelism == "cp" and not masks: + # One CP peer owns no tokens but still consumes the caller's full + # output and participates in backward and RNG synchronization. + local_rows = torch.arange(8 if replica_rank == 0 else 0, device=device) + + def model(weight): + with ( + tracker.fork("art-test-model") if tracker is not None else nullcontext() + ): + mask = torch.nn.functional.dropout( + torch.ones_like(features[local_rows]), 0.2 + ) + records.append(mask.detach().clone()) + return (features[local_rows] * mask) @ weight.T + + if tracker is not None: + from megatron.core.tensor_parallel import checkpoint as megatron_checkpoint + + hidden = megatron_checkpoint(model, False, decoder) + else: + hidden = checkpoint(model, decoder, use_reentrant=False) + recompute_masks.append(records) + full_mask = torch.zeros_like(features) + full_mask[local_rows] = records[0] + dist.all_reduce(full_mask, group=replica_group) + masks.append(full_mask) + if parallelism == "tp": + hidden = trainer._gather_sequence_parallel_hidden(hidden[:, None]) + else: + hidden = _GatherContextParallelRows.apply( + hidden, local_rows, len(features), replica_group + ) + return [ForwardOutput(None, None, None, hidden)] + + _stub_forward(patch, trainer, execute) + + def batches(*args, **kwargs): + yield SimpleNamespace( + outputs=execute(), stats=SimpleNamespace(global_count=dp_size) + ) + + patch.setattr(trainer, "_forward_batches", batches) + for step in range(2): + losses = [] + expected_losses = [] + for micro in range(2): + if step == micro == 0: + # Disagree again after registration: forward return must repair + # the caller even though the model uses a separate stream. + torch.manual_seed(1231 + 100 * dp_rank + replica_rank) + before = _state(device) + if step == 0: + output = trainer.forward([request])[0] + else: + iterator = trainer.forward_batches([request]) + output = next(iterator).outputs[0] + assert trainer._rng._depth == 0 + iterator.close() + if replica_rank == 0: + _assert_state_equal(_state(device), before) + caller = _state(device) + cpu_draw = torch.rand(16) + cpu_draws.append(cpu_draw) + mask = torch.rand(len(features), device=device) > 0.35 + if previous_caller_mask is not None: + assert not torch.equal(cpu_draw, previous_caller_mask) + previous_caller_mask = cpu_draw + loss = head(output.hidden_states[mask]).square().sum() + losses.append(loss) + # The unsplit reference replays the exact caller stream, including + # the CPU mask draw and dropout; model dropout is taken from the + # actual shards so this checks caller/model gradient consistency. + with torch.random.fork_rng( + devices=[device.index] if device.type == "cuda" else [] + ): + caller.restore() + torch.testing.assert_close(torch.rand(16), cpu_draw) + expected_mask = torch.rand(len(features), device=device) > 0.35 + assert torch.equal(expected_mask, mask) + hidden = (features * masks[-1]) @ reference[0].T + dropped = torch.nn.functional.dropout(hidden[expected_mask], 0.3) + expected_loss = ( + torch.nn.functional.linear(dropped, reference[1], reference[2]) + .square() + .sum() + ) + torch.testing.assert_close(loss, expected_loss) + expected_losses.append(expected_loss) + # The collective comparison is independent of the reference replay. + copies = [torch.empty_like(cpu_draw, device=device) for _ in range(2)] + dist.all_gather(copies, cpu_draw.to(device), group=replica_group) + assert torch.equal(copies[0], copies[1]) + before_backward = _state(device) + tracker_before = tracker.get_states() if tracker is not None else {} + trainer.backward(torch.stack(losses).sum()) + _assert_state_equal(_state(device), before_backward) + if tracker is not None: + for key, state in tracker_before.items(): + assert torch.equal(tracker.get_states()[key], state) + torch.stack(expected_losses).sum().backward() + reduced = trainer._reduce_dynamic_grads(params, scale_grads=1 / dp_size) + for actual, expected in zip(reduced, reference, strict=True): + assert expected.grad is not None + dist.all_reduce(expected.grad, group=dp_group) + expected.grad.div_(dp_size) + torch.testing.assert_close(actual, expected.grad, rtol=2e-5, atol=2e-5) + optimizer.step() + optimizer.zero_grad() + metrics = trainer.optim_step( + params=AdamParams(learning_rate=0.01, weight_decay=0.0, grad_clip_norm=0), + scale_grads=1 / dp_size, + checkpoints=["student"], + ) + assert metrics["update_successful"] == 1 + for actual, expected in zip(params, reference, strict=True): + torch.testing.assert_close(actual, expected, rtol=2e-5, atol=2e-5) + for records in recompute_masks: + assert len(records) == 2 + assert torch.equal(records[0], records[1]) + assert not torch.equal(masks[0], masks[1]) + if dp_size > 1: + copies = [torch.empty_like(cpu_draws[0]) for _ in range(dp_size)] + dist.all_gather(copies, cpu_draws[0], group=dp_group) + assert not torch.equal(copies[0], copies[1]) From d1273f1f58c4c7255a70b8856c18238b7309eda0 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 16:55:14 +0000 Subject: [PATCH 044/150] fix(trainer-rank): meter physical cached graph backward --- .github/workflows/prek.yml | 2 + src/art/trainer_rank/_impl.py | 9 +- .../test_trainer_rank_graph_backward_work.py | 265 ++++++++++++++++++ 3 files changed, 272 insertions(+), 4 deletions(-) create mode 100644 tests/unit/test_trainer_rank_graph_backward_work.py diff --git a/.github/workflows/prek.yml b/.github/workflows/prek.yml index 4ffccbe71..6da96cebf 100644 --- a/.github/workflows/prek.yml +++ b/.github/workflows/prek.yml @@ -224,6 +224,7 @@ jobs: tests/unit/test_prefix_tree_grad_parity.py \ tests/unit/test_prefix_tree_packing.py \ tests/unit/test_trainer_rank_handoff_budget.py \ + tests/unit/test_trainer_rank_graph_backward_work.py \ tests/unit/test_qwen35_adapter_config.py \ tests/unit/test_trainer_rank_physical_reserve.py \ tests/unit/test_trainer_rank_validation.py \ @@ -272,6 +273,7 @@ jobs: --ignore=tests/unit/test_prefix_tree_grad_parity.py \ --ignore=tests/unit/test_prefix_tree_packing.py \ --ignore=tests/unit/test_trainer_rank_handoff_budget.py \ + --ignore=tests/unit/test_trainer_rank_graph_backward_work.py \ --ignore=tests/unit/test_qwen35_adapter_config.py \ --ignore=tests/unit/test_trainer_rank_physical_reserve.py \ --ignore=tests/unit/test_trainer_rank_validation.py \ diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 5dc252ce8..53c5f5280 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -3253,8 +3253,6 @@ def _forward_batches( # Do not retain our completed graph through a new handoff traceback. del tracked_outputs, flat_outputs, outputs raise - if backward is not None: - backward.attach(tracked_outputs) stop = start + candidate.stats_global_count if stop < len(items): self._last_global_micro_batch_size = max( @@ -3452,8 +3450,6 @@ def forward( tracked_outputs = self._execute_admitted_plan( plan, check=check, context="forward" ) - if backward is not None: - backward.attach(tracked_outputs) outputs = _unflatten(materialized, iter(tracked_outputs)) self._rng.synchronize(caller_group()) return outputs @@ -7309,6 +7305,11 @@ def execute(captured: _ForwardGroupPlan) -> tuple[torch.Tensor, ...]: if spec is not None and captured_spec != spec: raise RuntimeError("Forward replay changed its output tree") spec = captured_spec + # Observe physical backward, including replay, before the cache + # replaces these outputs with detached caller cotangent proxies. + backward = self._backward_work() + if backward is not None: + backward.attach(outputs) return tensors finally: if hybrid is not None: diff --git a/tests/unit/test_trainer_rank_graph_backward_work.py b/tests/unit/test_trainer_rank_graph_backward_work.py new file mode 100644 index 000000000..e5f5d9b89 --- /dev/null +++ b/tests/unit/test_trainer_rank_graph_backward_work.py @@ -0,0 +1,265 @@ +"""Real CPU graph/bridge engines with simulated CUDA observer devices and tails. + +Only physical model tensors expose a CUDA device to the observer. Caller CPU +proxies retain their actual device, so attaching to them disables accounting. +This exercises engine boundaries, not native CUDA timing or memory behavior. +""" + +import asyncio +from dataclasses import replace +from types import SimpleNamespace +import weakref + +import pytest +from test_trainer_rank_backward_work import CUDA, Clock +from test_trainer_rank_validation import _runtime +import torch + +from art.megatron.context_parallel.types import ParallelTopology +from art.trainer_rank import ( + ForwardInput, + ForwardOptions, + ForwardOutput, + TrainerRank, + _backward_work, + run_rank_callback, +) +from art.trainer_rank._tensors import CotangentCollector + +MODEL_NS = 1_000_000 +CALLER_NS = 1_000_000_000 + + +class _PhysicalTensor: + device = torch.device("cuda:0") + + def __init__(self, tensor, released): + self.ref = weakref.ref(tensor, released) + + @property + def requires_grad(self): + tensor = self.ref() + assert tensor is not None + return tensor.requires_grad + + def register_hook(self, callback): + tensor = self.ref() + assert tensor is not None + return tensor.register_hook(callback) + + +@pytest.fixture +def rig(monkeypatch): + model = torch.nn.Linear(1, 1, bias=False) + rank = TrainerRank(_runtime(model)) + weight = model.weight + with torch.no_grad(): + weight.fill_(2) + clock, cuda = Clock(), CUDA() + # Preserve the real engine task IDs and completion callbacks. Only device + # labels, event readiness, and elapsed time are controlled by this fixture. + monkeypatch.setattr(_backward_work, "time", clock) + monkeypatch.setattr( + _backward_work, + "torch", + SimpleNamespace( + cuda=cuda, _C=torch._C, compiler=torch.compiler, autograd=torch.autograd + ), + ) + work = _backward_work.BackwardWork( + rank._recovery_state().lock, _PhysicalTensor.device + ) + monkeypatch.setattr(rank, "_backward_work", lambda: work) + physical, references = {}, [] + state = SimpleNamespace(before_backward=lambda: None, forwards=0) + attach = work.attach + + def observe(outputs): + attach( + [ + replace( + output, + hidden_states=physical.get( + id(output.hidden_states), output.hidden_states + ), + ) + for output in outputs + ] + ) + + monkeypatch.setattr(work, "attach", observe) + + class Model(torch.autograd.Function): + @staticmethod + def forward(ctx, parameter, tokens): + ctx.save_for_backward(tokens) + state.forwards += 1 + clock.value += CALLER_NS + return parameter.sum() * tokens + + @staticmethod + def backward(ctx, *gradients): + (gradient,) = gradients + state.before_backward() + clock.value += MODEL_NS + (tokens,) = ctx.saved_tensors + return (gradient * tokens).sum().reshape_as(weight), None + + def forward(items, prepared): + outputs = [] + for item in items: + value = Model.apply(weight, item.input_ids.float()) + key = id(value) + physical[key] = _PhysicalTensor( + value, lambda _, key=key: physical.pop(key, None) + ) + references.append(weakref.ref(value)) + outputs.append(ForwardOutput(None, None, None, value)) + return outputs + + monkeypatch.setattr(rank, "_topology", lambda: ParallelTopology()) + monkeypatch.setattr(rank, "_dp_rank_and_size", lambda: (0, 1)) + monkeypatch.setattr(rank, "_configure_hybridep", lambda *args, **kwargs: None) + monkeypatch.setattr(rank, "_prepare_packed_forward", lambda packed: None) + monkeypatch.setattr(rank, "_forward_packed", forward) + yield SimpleNamespace( + rank=rank, + weight=weight, + clock=clock, + cuda=cuda, + work=work, + state=state, + references=references, + ) + work.close() + + +def _input(retention="gpu", output_device="cpu"): + return ForwardInput( + input_tokens=torch.tensor([1, 2]), + hidden_states=True, + options=ForwardOptions(backward_state=retention, output_device=output_device), + ) + + +@pytest.mark.parametrize("api", ["forward", "forward_batches"]) +@pytest.mark.parametrize("retention", ["gpu", "cpu", "replay", "evict"]) +def test_physical_backward_excludes_proxy_idle_and_replay_time(rig, api, retention): + rank, work = rig.rank, rig.work + request = _input("gpu" if retention == "evict" else retention) + batches = rank.forward_batches([request]) if api == "forward_batches" else None + output = next(batches).outputs[0] if batches else rank.forward(request) + value = output.hidden_states + assert value is not None and value.device.type == "cpu" + cache = rank._forward_graph_cache() + if retention == "evict": + cache.evict(cache.handles()[0]) + if retention in {"replay", "evict"}: + assert rig.references[0]() is None + assert work.work_ns == 0 and not work.rows and not work.disabled + + def caller(_): + rig.clock.value += CALLER_NS + + value.register_hook(caller) + rig.clock.value += CALLER_NS # Caller idle time is outside either engine. + rank.backward(value.sum(), retain_graph=True) + assert rig.state.forwards == (2 if retention in {"replay", "evict"} else 1) + assert len(work.rows) == len(rig.cuda.events) == 1 + (row,) = work.rows.values() + assert row.ended is not None and not row.blocked + elapsed = row.ended - row.started + assert MODEL_NS <= elapsed < MODEL_NS + 100 + assert rig.cuda.events[0].stream == ("caller", 0) + work.harvest() + assert work.work_ns == 0 # Engine completion alone cannot credit GPU work. + rig.clock.value += CALLER_NS + rig.cuda.events[0].ready = True + work.harvest() + assert work.work_ns == elapsed and not work.rows + rank.backward(value.sum()) + rig.cuda.events[-1].ready = True + work.harvest() + assert 2 * MODEL_NS <= work.work_ns < 2 * MODEL_NS + 200 + torch.testing.assert_close(rig.weight.grad, torch.tensor([[6.0]])) + assert not cache.handles() + assert all(reference() is None for reference in rig.references) + assert not work.disabled + if batches is not None: + assert next(batches, None) is None + + +@pytest.mark.parametrize("mode", ["rank", "zero"]) +@pytest.mark.parametrize("remote", [False, True]) +def test_logical_and_remote_backwards_credit_one_physical_engine(rig, mode, remote): + def run(callback): + return asyncio.run(run_rank_callback(rig.rank, callback, mode=mode)).value + + if remote: + packet = run(lambda view: view.export_forward(view.forward(_input("replay")))) + collector = CotangentCollector() + output = collector.attach(packet) + packets = collector.backward(output.hidden_states.sum()) + assert not rig.work.rows and not rig.cuda.events + run(lambda view: view.backward_packets(packets)) + else: + + def train(view): + output = view.forward(_input("replay")) + assert not rig.work.rows and not rig.cuda.events + view.backward(output.hidden_states.sum()) + + run(train) + assert len(rig.work.rows) == len(rig.cuda.events) == 1 + rig.cuda.events[0].ready = True + rig.work.harvest() + assert MODEL_NS <= rig.work.work_ns < MODEL_NS + 100 + torch.testing.assert_close(rig.weight.grad, torch.tensor([[3.0]])) + assert not rig.rank._forward_graph_cache().handles() + + +def test_failed_physical_engine_prevents_later_credit(rig): + primary = RuntimeError("physical backward failed") + + def fail(): + raise primary + + output = rig.rank.forward(_input()) + rig.state.before_backward = fail + with pytest.raises(RuntimeError) as caught: + rig.rank.backward(output.hidden_states.sum()) + assert caught.value is primary + assert len(rig.work.rows) == 1 and not rig.cuda.events + assert next(iter(rig.work.rows.values())).ended is None + assert rig.weight.grad is None + assert not rig.rank._forward_graph_cache().handles() + rig.state.before_backward = lambda: None + output = rig.rank.forward(_input("replay")) + rig.rank.backward(output.hidden_states.sum()) + assert len(rig.cuda.events) == 1 + rig.cuda.events[0].ready = True + rig.work.harvest() + assert rig.work.work_ns == 0 and not rig.work.disabled + + +@pytest.mark.parametrize("nested", ["forward", "backward"]) +def test_nested_physical_work_is_excluded(rig, nested): + output = rig.rank.forward(_input()) + other = rig.rank.forward(_input()) if nested == "backward" else None + + def reenter(): + rig.state.before_backward = lambda: None + if other is None: + rig.rank.forward(_input(), no_grad=True) + else: + with torch.enable_grad(): + rig.rank.backward(other.hidden_states.sum()) + + rig.state.before_backward = reenter + rig.rank.backward(output.hidden_states.sum()) + assert len(rig.work.rows) == (2 if nested == "backward" else 1) + assert all(row.blocked for row in rig.work.rows.values()) + for event in rig.cuda.events: + event.ready = True + rig.work.harvest() + assert rig.work.work_ns == 0 and not rig.work.rows and not rig.work.disabled From 0004fc051f860ca11070697a9ee4f7bb47f8d570 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 16:54:14 +0000 Subject: [PATCH 045/150] Keep RNG and command failure collectives aligned Fingerprint graph placement policy sources in planner replay reports. --- src/art/trainer_rank/_impl.py | 25 +++++---- src/art/trainer_rank/_planner_misses.py | 2 + .../unit/test_trainer_rank_planner_reports.py | 13 +++-- tests/unit/test_trainer_rank_rng.py | 56 ++++++++++++++++++- 4 files changed, 79 insertions(+), 17 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 53c5f5280..235a3cb26 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -3442,16 +3442,21 @@ def forward( # execution belongs to the private model stream. materialized = self._capture_forward_options(inputs, options) requests = list(_flatten(materialized)) - with torch.set_grad_enabled(enabled), self._rng.model(): - self._reset_planning_telemetry() - plan, check = self._plan_admissible_forward( - requests, checkpoint=checkpoint, context="forward" - ) - tracked_outputs = self._execute_admitted_plan( - plan, check=check, context="forward" - ) - outputs = _unflatten(materialized, iter(tracked_outputs)) - self._rng.synchronize(caller_group()) + try: + with torch.set_grad_enabled(enabled), self._rng.model(): + self._reset_planning_telemetry() + plan, check = self._plan_admissible_forward( + requests, checkpoint=checkpoint, context="forward" + ) + tracked_outputs = self._execute_admitted_plan( + plan, check=check, context="forward" + ) + outputs = _unflatten(materialized, iter(tracked_outputs)) + finally: + # Failed peers must leave this frontier before the command layer's + # error exchange, just as successful peers do. Caller RNG is restored + # by model() before this collective, including on execution failure. + self._rng.synchronize(caller_group()) return outputs def _execute_admitted_plan( diff --git a/src/art/trainer_rank/_planner_misses.py b/src/art/trainer_rank/_planner_misses.py index 09c0462b3..d10fca02d 100644 --- a/src/art/trainer_rank/_planner_misses.py +++ b/src/art/trainer_rank/_planner_misses.py @@ -50,6 +50,8 @@ "_prefix_tree_performance_search.py", "_planner_misses.py", "_gdn_memory.py", + "_memory_policy.py", + "_options.py", "_planner_evidence.py", "_planner_retention.py", ) diff --git a/tests/unit/test_trainer_rank_planner_reports.py b/tests/unit/test_trainer_rank_planner_reports.py index cd865e3f5..59b2e7083 100644 --- a/tests/unit/test_trainer_rank_planner_reports.py +++ b/tests/unit/test_trainer_rank_planner_reports.py @@ -351,11 +351,14 @@ def test_replay_reruns_real_memory_estimator_and_prefix_layout(tmp_path): unrecorded["replay"]["memory_replay"]["rank"]["one_layer_recompute"] = None with pytest.raises(ValueError, match="recompute mode is not recorded"): reports.replay(unrecorded) - drifted = reports.validate_report(path.read_bytes()) - drifted["replay"]["source_files"]["_impl.py"]["sha256"] = "0" * 64 - with pytest.raises(ValueError, match="source differs"): - reports.replay(drifted) - assert reports.replay(drifted, allow_source_drift=True)["source_matches"] is False + for name in ("_impl.py", "_memory_policy.py", "_options.py"): + drifted = reports.validate_report(path.read_bytes()) + drifted["replay"]["source_files"][name]["sha256"] = "0" * 64 + with pytest.raises(ValueError, match="source differs"): + reports.replay(drifted) + assert ( + reports.replay(drifted, allow_source_drift=True)["source_matches"] is False + ) assert "_gdn_memory.py" in reports._source_files() # Frozen stages are independent inputs, not a recorded total substituted # for the estimator. Changing a fixed stage changes actual recomputation. diff --git a/tests/unit/test_trainer_rank_rng.py b/tests/unit/test_trainer_rank_rng.py index b51efaca6..ae3471d98 100644 --- a/tests/unit/test_trainer_rank_rng.py +++ b/tests/unit/test_trainer_rank_rng.py @@ -1,5 +1,6 @@ from __future__ import annotations +import asyncio from contextlib import nullcontext from datetime import timedelta from types import SimpleNamespace @@ -10,8 +11,16 @@ import torch.distributed as dist import torch.multiprocessing as mp from torch.utils.checkpoint import checkpoint - -from art.trainer_rank import AdamParams, ForwardInput, ForwardOutput, TrainerRank +from trainer_rank_test_support import gloo_group, megatron_topology, spawn_and_join + +from art.trainer_rank import ( + AdamParams, + ForwardInput, + ForwardOutput, + TrainerRank, + run_rank_callback, +) +from art.trainer_rank._commands import join_rank_callback_release from art.trainer_rank._impl import _CheckpointSlot, _GatherContextParallelRows from art.trainer_rank._rng import RNGState, TrainerRNG, caller_group @@ -151,6 +160,49 @@ def fail(): assert torch.equal(torch.rand(7), torch.rand(7, generator=generator)) +def test_failed_forward_keeps_command_collectives_aligned(tmp_path): + spawn_and_join( + _failed_forward_worker, + (f"file://{tmp_path / 'failed-forward'}",), + timeout=90, + failure="Forward failure stranded a peer before command error exchange", + ) + + +def _failed_forward_worker(physical, rendezvous): + with gloo_group(physical, rendezvous, timeout=10): + trainer = _trainer() + with ( + megatron_topology(physical, dp_size=1, tp_size=2), + pytest.MonkeyPatch.context() as patch, + ): + calls = 0 + + def execute(): + nonlocal calls + calls += 1 + torch.rand(7) + if calls == 1 and physical == 1: + raise ValueError("injected local model failure") + return [] + + _stub_forward(patch, trainer, execute) + + def callback(view): + with pytest.raises(RuntimeError, match="injected local model failure"): + view.forward([]) + assert view.forward([]) == [] + return "recovered" + + async def run(): + result = await run_rank_callback(trainer, callback) + await join_rank_callback_release(trainer) + assert result.value == ("recovered" if physical == 0 else None) + + asyncio.run(run()) + assert calls == 2 + + @pytest.mark.parametrize("yield_empty", (False, True)) def test_microbatch_yields_and_close_do_not_hold_rng_context(monkeypatch, yield_empty): trainer = _trainer() From b7d53fcb1271b92489f4729b8747e50e0db84e9e Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 17:14:55 +0000 Subject: [PATCH 046/150] fix(trainer-rank): preserve forward errors during RNG sync --- src/art/trainer_rank/_impl.py | 20 +++++++++-- tests/unit/test_trainer_rank_rng.py | 51 +++++++++++++++++++++++++++++ 2 files changed, 68 insertions(+), 3 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 235a3cb26..d48545b26 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -3442,6 +3442,7 @@ def forward( # execution belongs to the private model stream. materialized = self._capture_forward_options(inputs, options) requests = list(_flatten(materialized)) + error: BaseException | None = None try: with torch.set_grad_enabled(enabled), self._rng.model(): self._reset_planning_telemetry() @@ -3452,11 +3453,21 @@ def forward( plan, check=check, context="forward" ) outputs = _unflatten(materialized, iter(tracked_outputs)) + except BaseException as exc: + error = exc + raise finally: # Failed peers must leave this frontier before the command layer's # error exchange, just as successful peers do. Caller RNG is restored # by model() before this collective, including on execution failure. - self._rng.synchronize(caller_group()) + try: + self._rng.synchronize(caller_group()) + except BaseException as sync_error: + if error is None: + raise + self._memory_error_with_reduction_note( + error, sync_error, operation="RNG synchronization" + ) return outputs def _execute_admitted_plan( @@ -8614,7 +8625,10 @@ def _recovery_reduce( @staticmethod def _memory_error_with_reduction_note( - error: BaseException, exchange_error: BaseException | None + error: BaseException, + exchange_error: BaseException | None, + *, + operation: str = "memory reduction", ) -> BaseException: # Raise outside the exchange handler to preserve the local error's chain. # A secondary poisoned-communicator failure is diagnostic, not the primary. @@ -8622,7 +8636,7 @@ def _memory_error_with_reduction_note( try: BaseException.add_note( error, - "Secondary memory reduction failure:\n" + f"Secondary {operation} failure:\n" + "".join(traceback.format_exception(exchange_error)), ) except BaseException: diff --git a/tests/unit/test_trainer_rank_rng.py b/tests/unit/test_trainer_rank_rng.py index ae3471d98..2554e393f 100644 --- a/tests/unit/test_trainer_rank_rng.py +++ b/tests/unit/test_trainer_rank_rng.py @@ -160,6 +160,57 @@ def fail(): assert torch.equal(torch.rand(7), torch.rand(7, generator=generator)) +@pytest.mark.parametrize( + "primary_type", (None, ValueError, asyncio.CancelledError, KeyboardInterrupt) +) +@pytest.mark.parametrize("sync_type", (None, OSError, asyncio.CancelledError)) +def test_forward_sync_preserves_primary_failure(monkeypatch, primary_type, sync_type): + trainer = _trainer() + caller = torch.get_rng_state() + primary = None if primary_type is None else primary_type("model failure") + secondary = None if sync_type is None else sync_type("sync failure") + cause, context = LookupError("existing cause"), RuntimeError("existing context") + expected = primary if primary is not None else secondary + if expected is not None: + expected.__cause__, expected.__context__ = cause, context + expected.__suppress_context__ = False + expected.add_note("existing note") + synchronized = [] + + def execute(): + torch.rand(7) + if primary is not None: + raise primary + return [] + + def synchronize(group): + synchronized.append(group) + assert torch.equal(torch.get_rng_state(), caller) + if secondary is not None: + raise secondary + + _stub_forward(monkeypatch, trainer, execute) + monkeypatch.setattr(trainer._rng, "synchronize", synchronize) + if expected is None: + assert trainer.forward([]) == [] + else: + with pytest.raises(type(expected)) as caught: + trainer.forward([]) + assert caught.value is expected + assert expected.__cause__ is cause and expected.__context__ is context + assert not expected.__suppress_context__ + assert expected.__notes__[0] == "existing note" + assert len(expected.__notes__) == ( + 2 if primary is not None and secondary is not None else 1 + ) + if primary is not None and secondary is not None: + note = expected.__notes__[1] + assert "Secondary RNG synchronization failure:" in note + assert f"{type(secondary).__name__}: sync failure" in note + assert synchronized == [None] + assert torch.equal(torch.get_rng_state(), caller) + + def test_failed_forward_keeps_command_collectives_aligned(tmp_path): spawn_and_join( _failed_forward_worker, From ff1cb785f03b363476fd4172932838863f927645 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 17:30:38 +0000 Subject: [PATCH 047/150] fix(trainer-rank): drain commands after physical cancellation --- src/art/trainer_rank/_commands.py | 18 ++--- src/art/trainer_rank/_impl.py | 26 +++---- tests/unit/test_trainer_rank_rng.py | 102 +++++++++++++++++++--------- 3 files changed, 94 insertions(+), 52 deletions(-) diff --git a/src/art/trainer_rank/_commands.py b/src/art/trainer_rank/_commands.py index 9682ef26b..233bf792e 100644 --- a/src/art/trainer_rank/_commands.py +++ b/src/art/trainer_rank/_commands.py @@ -29,7 +29,7 @@ def _coordinate_call(call: Callable[[], T], *, group: dist.ProcessGroup | None) result, error = None, None try: result = call() - except Exception as exc: + except BaseException as exc: error = exc failures = [None if error is None else f"{type(error).__name__}: {error}"] if dist.is_initialized(): @@ -324,7 +324,7 @@ def invoke(self, operation: str, *args: Any, **kwargs: Any) -> Any: return self._execute(command) async def serve(self) -> None: - cancelled = None + deferred: BaseException | None = None try: while True: objects: list[Any] = [None] @@ -348,19 +348,21 @@ async def serve(self) -> None: except asyncio.CancelledError as error: # An abandoned receive could consume the next callback's # command. Drain this session through its leader stop. - cancelled = error + deferred = error if deferred is None else deferred command = self._decode(objects[0]) self.state.sequence = max(self.state.sequence, command.sequence) if command.operation == "stop": - if cancelled is not None: - raise cancelled + if deferred is not None: + raise deferred return try: self._execute(command) - except Exception: + except BaseException as error: # The leader receives the same coordinated error and chooses # whether to catch it, continue, or stop the callback. - pass + # Cancellation must not abandon the leader's stop command. + if not isinstance(error, Exception): + deferred = error if deferred is None else deferred finally: self.stopped = True self._close_iterators() @@ -392,7 +394,7 @@ def _execute(self, command: _Command) -> Any: try: with torch.set_grad_enabled(command.grad_enabled): result = self._dispatch(command) - except Exception as exc: + except BaseException as exc: error = exc errors = self._gather( None if error is None else f"{type(error).__name__}: {error}" diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index d48545b26..c74f8ec6d 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -3455,19 +3455,19 @@ def forward( outputs = _unflatten(materialized, iter(tracked_outputs)) except BaseException as exc: error = exc - raise - finally: - # Failed peers must leave this frontier before the command layer's - # error exchange, just as successful peers do. Caller RNG is restored - # by model() before this collective, including on execution failure. - try: - self._rng.synchronize(caller_group()) - except BaseException as sync_error: - if error is None: - raise - self._memory_error_with_reduction_note( - error, sync_error, operation="RNG synchronization" - ) + # Failed peers must leave this frontier before the command layer's + # error exchange, just as successful peers do. Caller RNG is restored + # by model() before this collective, including on execution failure. + try: + self._rng.synchronize(caller_group()) + except BaseException as sync_error: + if error is None: + raise + self._memory_error_with_reduction_note( + error, sync_error, operation="RNG synchronization" + ) + if error is not None: + raise error return outputs def _execute_admitted_plan( diff --git a/tests/unit/test_trainer_rank_rng.py b/tests/unit/test_trainer_rank_rng.py index 2554e393f..161144c74 100644 --- a/tests/unit/test_trainer_rank_rng.py +++ b/tests/unit/test_trainer_rank_rng.py @@ -20,7 +20,7 @@ TrainerRank, run_rank_callback, ) -from art.trainer_rank._commands import join_rank_callback_release +from art.trainer_rank._commands import _coordinate_call, join_rank_callback_release from art.trainer_rank._impl import _CheckpointSlot, _GatherContextParallelRows from art.trainer_rank._rng import RNGState, TrainerRNG, caller_group @@ -207,6 +207,7 @@ def synchronize(group): note = expected.__notes__[1] assert "Secondary RNG synchronization failure:" in note assert f"{type(secondary).__name__}: sync failure" in note + assert "model failure" not in note assert synchronized == [None] assert torch.equal(torch.get_rng_state(), caller) @@ -221,37 +222,76 @@ def test_failed_forward_keeps_command_collectives_aligned(tmp_path): def _failed_forward_worker(physical, rendezvous): - with gloo_group(physical, rendezvous, timeout=10): - trainer = _trainer() - with ( - megatron_topology(physical, dp_size=1, tp_size=2), - pytest.MonkeyPatch.context() as patch, + with ( + gloo_group(physical, rendezvous, timeout=10), + megatron_topology(physical, dp_size=1, tp_size=2), + ): + for primary_type, sync_failure, preflight in ( + (ValueError, False, False), + (asyncio.CancelledError, True, False), + (KeyboardInterrupt, True, False), + (asyncio.CancelledError, False, False), + (KeyboardInterrupt, False, False), + (ValueError, True, False), + (asyncio.CancelledError, True, True), + (KeyboardInterrupt, False, True), ): - calls = 0 - - def execute(): - nonlocal calls - calls += 1 - torch.rand(7) - if calls == 1 and physical == 1: - raise ValueError("injected local model failure") - return [] - - _stub_forward(patch, trainer, execute) - - def callback(view): - with pytest.raises(RuntimeError, match="injected local model failure"): - view.forward([]) - assert view.forward([]) == [] - return "recovered" - - async def run(): - result = await run_rank_callback(trainer, callback) - await join_rank_callback_release(trainer) - assert result.value == ("recovered" if physical == 0 else None) - - asyncio.run(run()) - assert calls == 2 + with pytest.MonkeyPatch.context() as patch: + _check_forward_failure( + patch, physical, primary_type, sync_failure, preflight + ) + + +def _check_forward_failure(patch, physical, primary_type, sync_failure, preflight): + trainer = _trainer() + calls = 0 + primary = primary_type("injected local model failure") + synchronize = trainer._rng.synchronize + + def execute(): + nonlocal calls + calls += 1 + torch.rand(7) + + def compute(): + if calls == 1 and physical == 1: + raise primary + return [] + + return _coordinate_call(compute, group=None) if preflight else compute() + + def sync(group): + synchronize(group) + if calls == 1 and sync_failure: + raise OSError("injected synchronization failure") + + _stub_forward(patch, trainer, execute) + patch.setattr(trainer._rng, "synchronize", sync) + + def callback(view): + expected = OSError if sync_failure and not preflight else RuntimeError + with pytest.raises(expected, match="injected"): + view.forward([]) + assert view.forward([]) == [] + return "recovered" + + async def run(): + deferred = physical == 1 and not isinstance(primary, Exception) + try: + result = await run_rank_callback(trainer, callback) + except BaseException as error: + assert deferred and error is primary + else: + assert not deferred + assert result.value == ("recovered" if physical == 0 else None) + await join_rank_callback_release(trainer) + # A follower must drain the first stop, not consume the next session. + result = await run_rank_callback(trainer, lambda view: view.forward([])) + await join_rank_callback_release(trainer) + assert result.value == ([] if physical == 0 else None) + + asyncio.run(run()) + assert calls == 3 @pytest.mark.parametrize("yield_empty", (False, True)) From 667f8be113518592a25c5381f78314043edb25a7 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 17:29:41 +0000 Subject: [PATCH 048/150] fix(trainer-rank): coordinate logical head factory failures --- src/art/trainer_rank/_heads.py | 42 ++++++++- tests/unit/test_trainer_rank_live_heads.py | 101 ++++++++++++++++++++- 2 files changed, 137 insertions(+), 6 deletions(-) diff --git a/src/art/trainer_rank/_heads.py b/src/art/trainer_rank/_heads.py index c1d7d10de..86d611c34 100644 --- a/src/art/trainer_rank/_heads.py +++ b/src/art/trainer_rank/_heads.py @@ -538,7 +538,8 @@ class HeadRegistration: checkpoint: Any name: str kind: HeadKind - value: torch.nn.Module | torch.Tensor + value: torch.nn.Module | torch.Tensor | None + factory_error: str | None = None @dataclass(frozen=True) @@ -643,7 +644,18 @@ def execute_head_operation( trainer, checkpoint, name ) if name in trainer._checkpoint_slots[checkpoint].custom else None if kind == "head_register": + from . import _checkpoint + registration: HeadRegistration = payload + # Every DP/TP/CP participant must settle leader-side construction before + # any rank enters checkpoint resolution or mutates its registered heads. + _checkpoint.raise_distributed( + None + if registration.factory_error is None + else RuntimeError(registration.factory_error), + f"construct custom object {registration.name!r}", + trainer._checkpoint_group(), + ) checkpoint = trainer._resolve_custom_checkpoint(registration.checkpoint) existing = trainer._checkpoint_slots[checkpoint].custom.get(registration.name) if existing is not None: @@ -1309,12 +1321,32 @@ def logical_register_head( if not current.invalid: return current.value if state is None: - value = factory() - state = view._invoke( - "head", "head_register", HeadRegistration(checkpoint, name, kind, value) - ).state + value, error, factory_error = None, None, None + try: + value = factory() + except BaseException as exc: + error = exc + factory_error = type(exc).__name__ + try: + factory_error += f": {exc}" + except BaseException: + pass + try: + state = view._invoke( + "head", + "head_register", + HeadRegistration(checkpoint, name, kind, value, factory_error), + ).state + except BaseException: + if error is None: + raise + # Preserve the factory's original type, identity and chain at its owner; + # physical command peers receive an ordinary coordinated failure. + if error is not None: + raise error else: value = deepcopy(view._rank._checkpoint_slots[checkpoint].custom[name].value) + assert value is not None value = ( move_module(value, view.device) if isinstance(value, torch.nn.Module) diff --git a/tests/unit/test_trainer_rank_live_heads.py b/tests/unit/test_trainer_rank_live_heads.py index ee30c63b9..fb8da46fd 100644 --- a/tests/unit/test_trainer_rank_live_heads.py +++ b/tests/unit/test_trainer_rank_live_heads.py @@ -2,15 +2,18 @@ import asyncio from copy import deepcopy +from datetime import timedelta from types import SimpleNamespace import pytest from test_trainer_rank_custom_tensors import _trainer, _use_local_gradients import torch +import torch.distributed as dist from torch.utils.checkpoint import checkpoint -from trainer_rank_test_support import gloo_group, spawn_and_join +from trainer_rank_test_support import gloo_group, megatron_topology, spawn_and_join from art.trainer_rank import AdamParams, ModuleHandle, run_rank_callback +from art.trainer_rank._commands import join_rank_callback_release from art.trainer_rank._heads import ( HeadRegistration, LiveHead, @@ -54,6 +57,102 @@ def _native_head(kind="module", name="head", factory=TiedHead): return trainer, native +def _factory_failure_worker(physical, rendezvous, mode, dp_size): + with gloo_group(physical, rendezvous, timeout=10): + trainer, _ = _trainer("student") + # Import native support before the topology facade. + trainer._slot_ref("student") + for attribute in ( + "_checkpoint_process_group", + "_checkpoint_finalize_process_group", + ): + setattr( + trainer, + attribute, + dist.new_group(backend="gloo", timeout=timedelta(seconds=5)), + ) + slot = trainer._checkpoint_slots["student"] + with megatron_topology(physical, dp_size=dp_size, tp_size=2 // dp_size): + + async def run(): + leader = physical == 0 or (mode == "rank" and dp_size == 2) + for error_type in (ValueError, asyncio.CancelledError): + primary = error_type("injected head factory failure") + cause, context = KeyError("cause"), LookupError("context") + primary.__cause__, primary.__context__ = cause, context + calls = [] + + def factory(): + calls.append(True) + if physical == 0: + raise primary + return TiedHead() + + def failed(view): + view.module("failed", factory, checkpoint="student") + + before = dict(slot.custom), slot.params, slot.optimizer + if leader: + expected = error_type if physical == 0 else RuntimeError + with pytest.raises( + expected, match="injected head factory failure" + ) as caught: + await run_rank_callback(trainer, failed, mode=mode) + if physical == 0: + assert caught.value is primary + assert ( + primary.__cause__ is cause + and primary.__context__ is context + ) + else: + await run_rank_callback(trainer, failed, mode=mode) + await join_rank_callback_release(trainer) + assert len(calls) == int(leader) + assert (slot.custom, slot.params, slot.optimizer) == before + completed = torch.tensor(1) + dist.all_reduce(completed, group=trainer._checkpoint_group()) + assert completed.item() == 2 + + calls = [] + + def factory(): + calls.append(True) + return TiedHead() + + def register(view): + head = view.module("head", factory, checkpoint="student") + assert view.module("head", factory, checkpoint="student") is head + assert head(torch.tensor(3.0)).item() == 15 + + await run_rank_callback(trainer, register, mode=mode) + assert len(calls) == int(leader) + assert "failed" not in slot.custom and "head" in slot.custom + + def local_lookup(view): + if physical == 0: + head = view.module("head", factory, checkpoint="student") + assert head(torch.tensor(3.0)).item() == 15 + + # DP1 enters callback cleanup while DP0 reopens the existing + # head. Lookup must remain local instead of entering WORLD. + await run_rank_callback(trainer, local_lookup, mode=mode) + assert len(calls) == int(leader) + + asyncio.run(run()) + + +@pytest.mark.parametrize("mode,dp_size", [("rank", 2), ("rank", 1), ("zero", 2)]) +def test_head_factory_failure_keeps_registration_and_callback_groups_usable( + tmp_path, mode, dp_size +): + spawn_and_join( + _factory_failure_worker, + (f"file://{tmp_path / 'factory-failure'}", mode, dp_size), + timeout=90, + failure="Head factory failure stranded a registration or callback peer", + ) + + def _live_head( trainer, name, source, collector=None, *, checkpoint="student" ) -> LiveHead: From 9d6304c61bd6c149959be78a605749c8d40bfb3d Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 18:14:09 +0000 Subject: [PATCH 049/150] Restore split-peak counter fixture recompute assumptions --- tests/unit/test_trainer_rank_split_peak.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/tests/unit/test_trainer_rank_split_peak.py b/tests/unit/test_trainer_rank_split_peak.py index 31bff76af..49f38099d 100644 --- a/tests/unit/test_trainer_rank_split_peak.py +++ b/tests/unit/test_trainer_rank_split_peak.py @@ -143,6 +143,9 @@ def test_profile_order_change_cannot_drop_completed_split_floor(monkeypatch): def _counter_split(monkeypatch): rank = _rank() + # Preserve this fixture's original recompute mode: one-layer packed pricing + # adds 6 KiB per token, dwarfing the synthetic 10,000-byte split budget. + rank._recompute_method = rank._recompute_num_layers = None # This executor injects allocator counters without creating cached graphs. monkeypatch.setattr(rank, "_graph_memory_policy_enabled", lambda: False) monkeypatch.setattr(rank, "_dp_rank_and_size", lambda: (0, 1)) From ea19d6223918bd652e165965f469f1c36cd5afe7 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 18:13:52 +0000 Subject: [PATCH 050/150] test(trainer-rank): complete command transport topology fixture --- tests/unit/test_trainer_command_transport.py | 12 +++++++++--- 1 file changed, 9 insertions(+), 3 deletions(-) diff --git a/tests/unit/test_trainer_command_transport.py b/tests/unit/test_trainer_command_transport.py index 4c2266955..e68c4b80a 100644 --- a/tests/unit/test_trainer_command_transport.py +++ b/tests/unit/test_trainer_command_transport.py @@ -11,8 +11,8 @@ import pytest import torch -import torch.multiprocessing as mp -from trainer_rank_test_support import gloo_group +import torch.distributed as dist +from trainer_rank_test_support import gloo_group, spawn_and_join from art.trainer_rank import ForwardInput, ForwardOutput, TrainerRank from art.trainer_rank._commands import _Command, _encode_command, _Executor @@ -133,6 +133,7 @@ def _transport_worker(physical: int, rendezvous: str, cuda: bool) -> None: ps = SimpleNamespace( get_tensor_model_parallel_rank=lambda: physical, get_context_parallel_rank=lambda: 0, + get_tensor_and_context_parallel_group=lambda **kwargs: dist.group.WORLD, ) core, megatron = ModuleType("megatron.core"), ModuleType("megatron") setattr(core, "parallel_state", ps) @@ -225,4 +226,9 @@ def test_commands_decode_on_cpu_and_native_handlers_place_locally( ) -> None: if cuda and torch.cuda.device_count() < 2: pytest.skip("requires two CUDA devices") - mp.spawn(_transport_worker, args=(str(tmp_path / "init"), cuda), nprocs=2) + spawn_and_join( + _transport_worker, + (str(tmp_path / "init"), cuda), + timeout=90, + failure="command transport did not complete", + ) From 09b64fba1c76441664a84f53e9d805ccfa39def8 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 18:30:23 +0000 Subject: [PATCH 051/150] Preserve split-floor oracles and share transport topology fixture --- tests/unit/test_trainer_command_transport.py | 19 +++++-------------- tests/unit/test_trainer_rank_split_peak.py | 17 +++++++++++++---- 2 files changed, 18 insertions(+), 18 deletions(-) diff --git a/tests/unit/test_trainer_command_transport.py b/tests/unit/test_trainer_command_transport.py index e68c4b80a..aa50a31d8 100644 --- a/tests/unit/test_trainer_command_transport.py +++ b/tests/unit/test_trainer_command_transport.py @@ -4,15 +4,12 @@ from dataclasses import dataclass import gc -import sys -from types import ModuleType, SimpleNamespace from typing import Any, Literal, cast import weakref import pytest import torch -import torch.distributed as dist -from trainer_rank_test_support import gloo_group, spawn_and_join +from trainer_rank_test_support import gloo_group, megatron_topology, spawn_and_join from art.trainer_rank import ForwardInput, ForwardOutput, TrainerRank from art.trainer_rank._commands import _Command, _encode_command, _Executor @@ -129,16 +126,10 @@ def _transport_worker(physical: int, rendezvous: str, cuda: bool) -> None: device = torch.device(f"cuda:{1 - physical}" if cuda else "cpu") if cuda: torch.cuda.set_device(device) - with gloo_group(physical, f"file://{rendezvous}"): - ps = SimpleNamespace( - get_tensor_model_parallel_rank=lambda: physical, - get_context_parallel_rank=lambda: 0, - get_tensor_and_context_parallel_group=lambda **kwargs: dist.group.WORLD, - ) - core, megatron = ModuleType("megatron.core"), ModuleType("megatron") - setattr(core, "parallel_state", ps) - setattr(megatron, "core", core) - sys.modules.update({"megatron": megatron, "megatron.core": core}) + with ( + gloo_group(physical, f"file://{rendezvous}"), + megatron_topology(physical, dp_size=1, tp_size=2), + ): runtime = _runtime(torch.nn.Linear(1, 1).to(device)) runtime.rank, runtime.world_size = physical, 2 rank: Any = TrainerRank(runtime) diff --git a/tests/unit/test_trainer_rank_split_peak.py b/tests/unit/test_trainer_rank_split_peak.py index 49f38099d..40a429ca2 100644 --- a/tests/unit/test_trainer_rank_split_peak.py +++ b/tests/unit/test_trainer_rank_split_peak.py @@ -7,12 +7,20 @@ from typing import Any import pytest -from test_trainer_rank_active_memory import _rank +from test_trainer_rank_active_memory import _rank as _active_rank import torch from art.trainer_rank import _impl as tr +def _rank(): + rank = _active_rank() + # Preserve this module's original recompute mode: packed pricing would + # overwhelm the synthetic budgets and bypass the split-floor oracles. + rank._recompute_method = rank._recompute_num_layers = None + return rank + + def _requests(count=2, length=100): return [ tr.ForwardInput( @@ -130,6 +138,10 @@ def test_profile_order_change_cannot_drop_completed_split_floor(monkeypatch): assert rank._plan_cost(b).ephemeral > rank._plan_cost(a).ephemeral before = dict(rank._memory_profiles) monkeypatch.setattr(rank, "_available_memory_bytes", lambda: 10_000) + assert ( + rank._split_required_memory([rank._plan_cost(p) for p in plan.subforwards]) + < 10_000 + ) accepted, check = rank._admit_split_rung( ((0,), (1,)), requests, @@ -143,9 +155,6 @@ def test_profile_order_change_cannot_drop_completed_split_floor(monkeypatch): def _counter_split(monkeypatch): rank = _rank() - # Preserve this fixture's original recompute mode: one-layer packed pricing - # adds 6 KiB per token, dwarfing the synthetic 10,000-byte split budget. - rank._recompute_method = rank._recompute_num_layers = None # This executor injects allocator counters without creating cached graphs. monkeypatch.setattr(rank, "_graph_memory_policy_enabled", lambda: False) monkeypatch.setattr(rank, "_dp_rank_and_size", lambda: (0, 1)) From 229a7f5a03d1f4bbf731d5b1be1365551e5eb1cb Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 18:31:33 +0000 Subject: [PATCH 052/150] test(trainer-rank): kill and reap surviving test workers --- tests/unit/trainer_rank_test_support.py | 30 ++++++++++++++++++++----- 1 file changed, 25 insertions(+), 5 deletions(-) diff --git a/tests/unit/trainer_rank_test_support.py b/tests/unit/trainer_rank_test_support.py index da17ac031..6dbab5fef 100644 --- a/tests/unit/trainer_rank_test_support.py +++ b/tests/unit/trainer_rank_test_support.py @@ -64,15 +64,35 @@ def megatron_topology(physical, *, dp_size, tp_size): def spawn_and_join(worker, args, *, timeout, failure, nprocs=2): """Bound a collective test while preserving spawned-worker tracebacks.""" processes = mp.spawn(worker, args=args, nprocs=nprocs, join=False) + error = None try: deadline = time.monotonic() + timeout while time.monotonic() < deadline: if processes.join(timeout=1): return pytest.fail(failure) + except BaseException as exc: + error = exc + raise finally: - for process in processes.processes: - if process.is_alive(): - process.terminate() - for process in processes.processes: - process.join(timeout=5) + try: + for process in processes.processes: + if process.is_alive(): + process.terminate() + for process in processes.processes: + process.join(timeout=5) + if process.is_alive(): + process.kill() + process.join(timeout=5) + survivors = [p.pid for p in processes.processes if p.is_alive()] + if survivors: + pytest.fail(f"Spawned workers survived SIGKILL: {survivors}") + except BaseException as cleanup_error: + if error is None: + raise + try: + BaseException.add_note( + error, f"Worker cleanup failed: {cleanup_error!r}" + ) + except BaseException: + pass From b3b520111b9518b040f9cb55864bc5dba8239d59 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 18:55:35 +0000 Subject: [PATCH 053/150] fix(trainer-rank): exclude consumed graphs from HybridEP floor --- src/art/trainer_rank/_impl.py | 2 +- tests/unit/test_trainer_rank_moe_memory.py | 44 +++++++++++++++++++++- 2 files changed, 44 insertions(+), 2 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index d0ff96023..7d924d31b 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -4585,7 +4585,7 @@ def _checkpoint_memory_floor( # this is separate from already-held native buffer capacity. Do not # prune graph references or reset execution state while estimating. rows = max(rows for rows, _ in group_rows) - if any(ref() is not None for ref in self._pending_hybridep_graphs): + if any(_graph_marker_is_live(ref) for ref in self._pending_hybridep_graphs): rows = max(rows, self._hybridep_rows_high_water) workspace = max(workspace, -(-rows // 4) * 4 * self._hidden_size * 2) return retained, workspace diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index 3fbdcee89..a537ccd8a 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -8,7 +8,7 @@ import pytest import torch -from art.trainer_rank import ForwardInput, TrainerRank +from art.trainer_rank import ForwardInput, ForwardOutput, TrainerRank from art.trainer_rank._impl import ( _PACKED_PRICED_LOGICAL_ROW_BYTES, _MemoryProfile, @@ -780,6 +780,48 @@ def test_hybridep_high_water_needs_a_live_larger_graph( assert tuple(refs) == before and rank._pending_hybridep_graphs is refs +def test_hybridep_admission_ignores_consumed_graph_with_retained_sibling( + hybrid_checkpoint_rank, +): + rank = hybrid_checkpoint_rank + values = dict( + packed_tokens=2, + logical_tokens=2, + output_bytes=8, + signature=replace(_signature(), topology=(1, 1, 2, 1)), + group_rows=((2, True),), + ) + baseline = rank._subforward_cost(**values) + rank._available_memory_bytes = lambda: 600000000 + assert rank._memory_check_required(baseline.required).fits + rank._hybridep_rows_high_water = 218751 + rank._hybridep_graph_tracking = True + value = torch.tensor(2.0, requires_grad=True) + (output,) = rank._track_slot_graph_outputs( + None, [ForwardOutput(None, None, value.square(), value.pow(3))] + ) + refs = rank._pending_hybridep_graphs + (marker_ref,) = refs + assert marker_ref() is not None and not marker_ref().item() + live = rank._subforward_cost(**values) + assert live.checkpoint_workspace == 218752 * 2048 * 2 + assert not rank._memory_check_required(live.required).fits + + assert output.hidden_states is not None + output.hidden_states.backward() + assert output.logits is not None and output.logits.grad_fn is not None + assert marker_ref() is not None and marker_ref().item() + # The unused sibling retains the consumed marker. Price and admit before + # any execution helper prunes it or resets the communication high-water. + consumed = rank._subforward_cost(**values) + assert consumed == baseline + assert rank._memory_check_required(consumed.required).fits + assert rank._pending_hybridep_graphs is refs and refs == [marker_ref] + assert marker_ref() is not None and marker_ref().item() + assert rank._hybridep_rows_high_water == 218751 + assert rank._hybridep_graph_tracking and rank._hybridep_buffer_id is None + + @pytest.mark.parametrize( "mode", ["empty", "no_grad", "cp1", "ep1", "unsupported", "selective"] ) From 7e5b34a3342b85f5bc834fbca0d4e760af1014e8 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 19:11:10 +0000 Subject: [PATCH 054/150] test(trainer-rank): reuse head construction helpers --- .../unit/test_trainer_live_parameter_roots.py | 14 ++++-------- .../unit/test_trainer_rank_parameter_hooks.py | 22 ++++++------------- .../test_trainer_rank_parameter_no_grad.py | 9 ++++---- 3 files changed, 15 insertions(+), 30 deletions(-) diff --git a/tests/unit/test_trainer_live_parameter_roots.py b/tests/unit/test_trainer_live_parameter_roots.py index 37f7436ba..dcea359d5 100644 --- a/tests/unit/test_trainer_live_parameter_roots.py +++ b/tests/unit/test_trainer_live_parameter_roots.py @@ -6,7 +6,7 @@ from typing import Any import pytest -from test_trainer_rank_custom_tensors import _trainer +from test_trainer_rank_live_heads import _live_head, _native_head import torch from art.trainer_rank import TrainerRank @@ -17,12 +17,9 @@ def _live_parameter( factory=lambda: torch.tensor(2.0), ) -> tuple[TrainerRank, torch.nn.Parameter, CotangentCollector, LiveHead]: - trainer, rank = _trainer("student") - parameter = rank.parameter("weight", factory, checkpoint="student") + trainer, parameter = _native_head("parameter", "weight", factory) collector = CotangentCollector() - live = LiveHead( - export_head(trainer, "student", "weight"), parameter.detach(), collector - ) + live = _live_head(trainer, "weight", parameter.detach(), collector) return trainer, parameter, collector, live @@ -38,10 +35,7 @@ def backward(ctx, *gradients): @pytest.mark.parametrize("existing", [False, True]) def test_native_direct_root_does_not_publish_when_another_root_fails(existing): - trainer, rank = _trainer("student") - parameter = rank.parameter( - "weight", lambda: torch.tensor(2.0), checkpoint="student" - ) + trainer, parameter = _native_head("parameter", "weight", lambda: torch.tensor(2.0)) if existing: parameter.grad = torch.tensor(7.0) original = parameter.grad diff --git a/tests/unit/test_trainer_rank_parameter_hooks.py b/tests/unit/test_trainer_rank_parameter_hooks.py index ae9628db1..b7a4727dd 100644 --- a/tests/unit/test_trainer_rank_parameter_hooks.py +++ b/tests/unit/test_trainer_rank_parameter_hooks.py @@ -4,24 +4,19 @@ import json import pytest -from test_trainer_rank_custom_tensors import _trainer +from test_trainer_rank_live_heads import TiedHead, _live_head, _native_head import torch from art.trainer_rank import ModuleHandle -from art.trainer_rank._heads import LiveHead, export_head, head_gradient_targets +from art.trainer_rank._heads import export_head, head_gradient_targets from art.trainer_rank._tensors import CotangentCollector def setup(client, values=None): - trainer, rank = _trainer("student") factory = lambda: torch.tensor(2.0) if values is None else values.clone() - native = rank.parameter("p", factory, checkpoint="student") + trainer, native = _native_head("parameter", "p", factory) collector = CotangentCollector() - live = LiveHead( - export_head(trainer, "student", "p"), - factory() if values is None else values, - collector, - ) + live = _live_head(trainer, "p", factory() if values is None else values, collector) parameter = live.value if client else native assert isinstance(parameter, torch.Tensor) @@ -110,7 +105,7 @@ def test_hook_failure_leaves_all_authoritative_gradients_unchanged(client, bad_r trainer, native, parameter, _, collector, backward = setup(client) rank = trainer other = rank.parameter("q", lambda: torch.tensor(3.0), checkpoint="student") - qlive = LiveHead(export_head(trainer, "student", "q"), torch.tensor(3.0), collector) + qlive = _live_head(trainer, "q", torch.tensor(3.0), collector) q = qlive.value if client else other native.grad, other.grad = torch.tensor(5.0), torch.tensor(6.0) parameter.register_hook(lambda gradient: gradient * 0) @@ -180,12 +175,9 @@ def test_sparse_hook_gradients_preserve_layout_and_mix_with_dense( @pytest.mark.parametrize("client", (False, True)) def test_tied_module_parameter_hook_sums_all_calls(client): - from test_trainer_rank_live_heads import TiedHead - - trainer, rank = _trainer("student") - native = rank.module("head", TiedHead, checkpoint="student") + trainer, native = _native_head(factory=TiedHead) collector = CotangentCollector() - live = LiveHead(export_head(trainer, "student", "head"), TiedHead(), collector) + live = _live_head(trainer, "head", TiedHead(), collector) head = live.value if client else native assert isinstance(head, ModuleHandle) seen = [] diff --git a/tests/unit/test_trainer_rank_parameter_no_grad.py b/tests/unit/test_trainer_rank_parameter_no_grad.py index 73053308e..4065f4226 100644 --- a/tests/unit/test_trainer_rank_parameter_no_grad.py +++ b/tests/unit/test_trainer_rank_parameter_no_grad.py @@ -4,11 +4,11 @@ import pytest from test_trainer_rank_custom_tensors import _trainer -from test_trainer_rank_live_heads import _step +from test_trainer_rank_live_heads import _live_head, _native_head, _step import torch from art.trainer_rank._commands import run_rank_callback -from art.trainer_rank._heads import LiveHead, export_head +from art.trainer_rank._heads import export_head from art.trainer_rank._tensors import CotangentCollector @@ -29,11 +29,10 @@ def _read(parameter, operation): def test_no_grad_parameter_reads_preserve_saved_loss_after_optimizer_step( monkeypatch, surface, operation ): - trainer, rank = _trainer("student") initial = torch.tensor([[2.0, 3.0], [4.0, 5.0]]) - native = rank.parameter("p", lambda: initial.clone(), checkpoint="student") + trainer, native = _native_head("parameter", "p", lambda: initial.clone()) collector = CotangentCollector() - live = LiveHead(export_head(trainer, "student", "p"), initial, collector) + live = _live_head(trainer, "p", initial, collector) def capture(parameter): with torch.no_grad(): From b61edd34729c1642b50e8d2a0309d8afebe3cbc8 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 19:49:00 +0000 Subject: [PATCH 055/150] fix(trainer-rank): reject fractional custom optimizer counters --- src/art/trainer_rank/_impl.py | 9 +++-- .../unit/test_trainer_rank_custom_tensors.py | 35 +++++++++++++++++-- 2 files changed, 40 insertions(+), 4 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 7d924d31b..715c8c19c 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -10714,11 +10714,16 @@ def _validate_custom_optimizer_state( for name, tensor in tensors.items() if tuple(tensor.shape) != expected_shape or tensor.dtype != torch.float32 ] - if invalid or not math.isfinite(state.step) or state.step < 0: + if ( + invalid + or not math.isfinite(state.step) + or state.step < 0 + or state.step != int(state.step) + ): raise TrainerRankSlotStateError( f"Custom optimizer state for {checkpoint!r}/{key!r} is invalid; " f"expected FP32 tensors with shape {expected_shape} and a nonnegative " - f"finite step (invalid={invalid}, step={state.step})." + f"finite integer step (invalid={invalid}, step={state.step})." ) diff --git a/tests/unit/test_trainer_rank_custom_tensors.py b/tests/unit/test_trainer_rank_custom_tensors.py index 4e3866fb8..309bd6ab3 100644 --- a/tests/unit/test_trainer_rank_custom_tensors.py +++ b/tests/unit/test_trainer_rank_custom_tensors.py @@ -1055,10 +1055,11 @@ def test_prepared_forward_snapshot_restores_frozen_custom_tensors( @pytest.mark.skipif(find_spec("megatron") is None, reason="requires Megatron") +@pytest.mark.parametrize("step,valid", [(0.5, False), (1.0, True)]) def test_custom_tensors_and_optimizer_restore_lazily_and_survive_unmaterialized_save( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, step: float, valid: bool ) -> None: - from safetensors.torch import load_file + from safetensors.torch import load_file, save_file from art.trainer_rank import _checkpoint @@ -1071,6 +1072,15 @@ def test_custom_tensors_and_optimizer_restore_lazily_and_survive_unmaterialized_ saved = tmp_path / "saved" original.save_checkpoint(str(saved), "student") + relative = "optimizer/custom.safetensors" + payload = load_file(saved / relative) + payload["step/value_head.proj.bias"].fill_(step) + save_file(payload, saved / relative) + manifest = json.loads((saved / "checkpoint.json").read_text()) + manifest["files"][relative] = _file_digest(saved / relative) + manifest["digest"] = _manifest_digest(manifest) + (saved / "checkpoint.json").write_text(json.dumps(manifest)) + restored, restored_api = _empty_real_lora_trainer() async def prefetch() -> None: @@ -1098,6 +1108,27 @@ def head_factory() -> _ValueHead: calls += 1 return _ValueHead(3) + if not valid: + slot = restored._checkpoint_slots["student"] + params, optimizer = slot.params, slot.optimizer + assert optimizer is not None + before = copy.deepcopy((params, optimizer.optimizer.state_dict())) + with pytest.raises( + TrainerRankSlotStateError, + match="value_head.proj.bias.*nonnegative finite integer step.*step=0.5", + ): + restored_api.module("value_head", head_factory, checkpoint="student") + assert calls == 1 and not slot.custom + assert slot.params is params and slot.optimizer is optimizer + torch.testing.assert_close( + (params, optimizer.optimizer.state_dict()), before, atol=0, rtol=0 + ) + temperature = restored_api.parameter( + "temperature", lambda: torch.tensor(-99.0), checkpoint="student" + ) + torch.testing.assert_close(temperature, original_temperature, atol=0, rtol=0) + return + restored_head = restored_api.module( "value_head", head_factory, checkpoint="student" ) From 77ebaa6ee1723978c7972c28d4720354602475cb Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 20:18:10 +0000 Subject: [PATCH 056/150] ci(trainer-rank): allow time for public GPU checks --- scripts/ci/trainer-rank-gpu.sky.yaml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/scripts/ci/trainer-rank-gpu.sky.yaml b/scripts/ci/trainer-rank-gpu.sky.yaml index 7f399c332..a654ccf9a 100644 --- a/scripts/ci/trainer-rank-gpu.sky.yaml +++ b/scripts/ci/trainer-rank-gpu.sky.yaml @@ -12,7 +12,7 @@ setup: | --frozen --no-install-project --inexact run: | - timeout --signal=TERM --kill-after=30s 20m \ + timeout --signal=TERM --kill-after=30s 30m \ bash scripts/ci/trainer-rank-gpu-tests.sh config: From fa036b93bfb2317edbb52f22c14d8093b1d68fcc Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 20:19:29 +0000 Subject: [PATCH 057/150] fix(trainer-rank): preserve custom counter writer coverage Check the original saved optimizer counter and reseal only the invalid case. Preserve float-counter validation with is_integer(). --- src/art/trainer_rank/_impl.py | 7 +------ tests/unit/test_trainer_rank_custom_tensors.py | 14 ++++++++------ 2 files changed, 9 insertions(+), 12 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 715c8c19c..113df7666 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -10714,12 +10714,7 @@ def _validate_custom_optimizer_state( for name, tensor in tensors.items() if tuple(tensor.shape) != expected_shape or tensor.dtype != torch.float32 ] - if ( - invalid - or not math.isfinite(state.step) - or state.step < 0 - or state.step != int(state.step) - ): + if invalid or state.step < 0 or not state.step.is_integer(): raise TrainerRankSlotStateError( f"Custom optimizer state for {checkpoint!r}/{key!r} is invalid; " f"expected FP32 tensors with shape {expected_shape} and a nonnegative " diff --git a/tests/unit/test_trainer_rank_custom_tensors.py b/tests/unit/test_trainer_rank_custom_tensors.py index 309bd6ab3..ff638f06f 100644 --- a/tests/unit/test_trainer_rank_custom_tensors.py +++ b/tests/unit/test_trainer_rank_custom_tensors.py @@ -1074,12 +1074,14 @@ def test_custom_tensors_and_optimizer_restore_lazily_and_survive_unmaterialized_ relative = "optimizer/custom.safetensors" payload = load_file(saved / relative) - payload["step/value_head.proj.bias"].fill_(step) - save_file(payload, saved / relative) - manifest = json.loads((saved / "checkpoint.json").read_text()) - manifest["files"][relative] = _file_digest(saved / relative) - manifest["digest"] = _manifest_digest(manifest) - (saved / "checkpoint.json").write_text(json.dumps(manifest)) + assert payload["step/value_head.proj.bias"].item() == 1.0 + if not valid: + payload["step/value_head.proj.bias"].fill_(step) + save_file(payload, saved / relative) + manifest = json.loads((saved / "checkpoint.json").read_text()) + manifest["files"][relative] = _file_digest(saved / relative) + manifest["digest"] = _manifest_digest(manifest) + (saved / "checkpoint.json").write_text(json.dumps(manifest)) restored, restored_api = _empty_real_lora_trainer() From 229f22c892e8b8c4905f0d844d446c480a5c3ad6 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 21:06:52 +0000 Subject: [PATCH 058/150] test(trainer-rank): scope temporary command fault injections --- tests/unit/test_trainer_rank_commands.py | 23 ++++++----------------- 1 file changed, 6 insertions(+), 17 deletions(-) diff --git a/tests/unit/test_trainer_rank_commands.py b/tests/unit/test_trainer_rank_commands.py index c701d91ac..04620dd19 100644 --- a/tests/unit/test_trainer_rank_commands.py +++ b/tests/unit/test_trainer_rank_commands.py @@ -10,6 +10,7 @@ import threading from types import SimpleNamespace from typing import Any +from unittest.mock import patch import weakref import pytest @@ -361,8 +362,7 @@ def track_loads(payload): decoded_outputs.append(value[1].packet.tensors) return value - setattr(_commands.cloudpickle, "loads", track_loads) - try: + with patch.object(_commands.cloudpickle, "loads", track_loads): large = run( lambda view: view.forward( [ @@ -380,8 +380,6 @@ def track_loads(payload): no_grad=True, ) ) - finally: - setattr(_commands.cloudpickle, "loads", loads) assert bool(decoded_outputs) is (physical == 0) if physical == 0: assert ( @@ -420,11 +418,8 @@ def decode_refusal(view): with pytest.raises(ValueError, match="result decode failure"): view.forward([_input(11)]) - setattr(_commands.cloudpickle, "loads", fail_result_decode) - try: + with patch.object(_commands.cloudpickle, "loads", fail_result_decode): run(decode_refusal) - finally: - setattr(_commands.cloudpickle, "loads", loads) assert set(rank._rank_command_state.graphs) <= retained_before_failure assert not rank._rank_command_state.iterators run(lambda view: view.backward(_loss_tree(view.forward([_input(11)])))) @@ -441,11 +436,10 @@ def track_dumps(value): return dumps(value) rank._available_cpu_memory_bytes = lambda: 1024 if physical == 0 else 1 << 60 - setattr(_commands.cloudpickle, "dumps", track_dumps) try: - run(host_refusal) + with patch.object(_commands.cloudpickle, "dumps", track_dumps): + run(host_refusal) finally: - setattr(_commands.cloudpickle, "dumps", dumps) del rank._available_cpu_memory_bytes assert not serialized @@ -453,8 +447,6 @@ def track_dumps(value): retained_before_failure = set(rank._rank_command_state.graphs) def failing_forward(tree, **kwargs): - from dataclasses import replace - result = forward(tree, **kwargs) if isinstance(result, ForwardOutput): result = replace( @@ -467,11 +459,8 @@ def backward_refusal(view): with pytest.raises(RuntimeError, match="physical backward failure"): view.backward(_loss_tree(view.forward([_input(7)]))) - rank.forward = failing_forward - try: + with patch.object(rank, "forward", failing_forward): run(backward_refusal) - finally: - rank.forward = forward assert set(rank._rank_command_state.graphs) <= retained_before_failure run(lambda view: view.zero_grad()) del sys.modules["megatron"], sys.modules["megatron.core"] From fe20204949ce2c823ac4dc659d7cac7c655fa263 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 21:44:02 +0000 Subject: [PATCH 059/150] refactor(trainer-rank): share iterator dispatch and scope test patches --- src/art/trainer_rank/_commands.py | 28 ++++++++----------- tests/unit/test_trainer_rank_live_heads.py | 4 --- ...rainer_rank_memory_recovery_distributed.py | 6 ++-- 3 files changed, 13 insertions(+), 25 deletions(-) diff --git a/src/art/trainer_rank/_commands.py b/src/art/trainer_rank/_commands.py index 233bf792e..ad31b05f6 100644 --- a/src/art/trainer_rank/_commands.py +++ b/src/art/trainer_rank/_commands.py @@ -514,24 +514,18 @@ def _dispatch(self, command: _Command) -> Any: handle = f"{self.mode}:batches:{command.sequence}" self.state.iterators[handle] = self.rank.forward_batches(*args, **kwargs) return handle - if op in ("next", "batches_next"): - iterator = ( - self.iterators[args[0]] - if op == "next" - else self.state.iterators[args[0]] - ) - batch = next(iterator, None) - if batch is None: - return (None, None) - return replace(batch, inputs=[], outputs=[]), self._packet( - batch.outputs, command.sequence - ) - if op in ("close", "batches_close"): - iterator = ( - self.iterators.pop(args[0], None) - if op == "close" - else self.state.iterators.pop(args[0], None) + if op in ("next", "batches_next", "close", "batches_close"): + iterators = ( + self.iterators if op in ("next", "close") else self.state.iterators ) + if op in ("next", "batches_next"): + batch = next(iterators[args[0]], None) + if batch is None: + return (None, None) + return replace(batch, inputs=[], outputs=[]), self._packet( + batch.outputs, command.sequence + ) + iterator = iterators.pop(args[0], None) if op == "batches_close": self.state.batch_inputs.pop(args[0], None) if iterator is not None: diff --git a/tests/unit/test_trainer_rank_live_heads.py b/tests/unit/test_trainer_rank_live_heads.py index fb8da46fd..eb87e7d89 100644 --- a/tests/unit/test_trainer_rank_live_heads.py +++ b/tests/unit/test_trainer_rank_live_heads.py @@ -398,8 +398,6 @@ def test_distributed_persistent_buffers_use_dp_zero_authority(tmp_path): def _buffer_snapshot_failure_worker(process_rank, init_method): - import torch.distributed as dist - from art.trainer_rank import _heads with gloo_group(process_rank, init_method, timeout=15): @@ -1003,8 +1001,6 @@ def test_inplace_operation_snapshots_readonly_checkpoint_parameter(): def _live_buffer_authority_worker(process_rank, init_method, asymmetric=False): - import torch.distributed as dist - from art.trainer_rank._heads import synchronize_head_buffers with gloo_group(process_rank, init_method): diff --git a/tests/unit/test_trainer_rank_memory_recovery_distributed.py b/tests/unit/test_trainer_rank_memory_recovery_distributed.py index 0f29c79f6..33fef6e29 100644 --- a/tests/unit/test_trainer_rank_memory_recovery_distributed.py +++ b/tests/unit/test_trainer_rank_memory_recovery_distributed.py @@ -149,12 +149,10 @@ def copy(tensor, *args, **kwargs): return convert(tensor, *args, **kwargs) setattr(trainer, "_commit_versioned_gradients", stage) - setattr(torch.Tensor, "to", copy) - try: + with pytest.MonkeyPatch.context() as patch: + patch.setattr(torch.Tensor, "to", copy) with pytest.raises((MemoryError, RuntimeError), match="head .* allocation"): _Executor(trainer, "zero")._backward(packets, retain_graph=False) - finally: - setattr(torch.Tensor, "to", convert) torch.testing.assert_close(parameter.grad, torch.ones_like(parameter)) assert trainer._version_state()._transaction is None dist.barrier() From 6eea24b38af7beda53fa20fff393459deabb4e00 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 22:00:22 +0000 Subject: [PATCH 060/150] fix(trainer-rank): preserve output placement errors during cleanup --- src/art/trainer_rank/_commands.py | 12 +++--------- tests/unit/test_trainer_rank_output_delivery.py | 16 ++++++++++++---- tests/unit/test_trainer_rank_output_memory.py | 4 ---- .../unit/test_trainer_rank_output_memory_cuda.py | 3 --- 4 files changed, 15 insertions(+), 20 deletions(-) diff --git a/src/art/trainer_rank/_commands.py b/src/art/trainer_rank/_commands.py index ad31b05f6..12fdfa092 100644 --- a/src/art/trainer_rank/_commands.py +++ b/src/art/trainer_rank/_commands.py @@ -963,15 +963,9 @@ def visit(request: Any, result: Any) -> None: ) if hasattr(self._rank, "_pending_backward_memory"): available -= sum(self._rank._pending_backward_memory()) - try: - placements = iter( - choose_output_placements(costs, gpu_available_bytes=available) - ) - except BaseException: - self._invoke( - "release", tuple(output.packet.handle for output, _ in outputs) - ) - raise + placements = iter( + choose_output_placements(costs, gpu_available_bytes=available) + ) result = [] for output, _ in outputs: cpu = tuple(next(placements) == "cpu" for _ in output.cpu) diff --git a/tests/unit/test_trainer_rank_output_delivery.py b/tests/unit/test_trainer_rank_output_delivery.py index 67a561672..4b63d43bd 100644 --- a/tests/unit/test_trainer_rank_output_delivery.py +++ b/tests/unit/test_trainer_rank_output_delivery.py @@ -15,6 +15,7 @@ ForwardOutput, MicroBatch, MicroBatchStats, + _memory_policy, _tensors, ) from art.trainer_rank._commands import _Executor, _view @@ -52,7 +53,7 @@ def _forward_memory_group(self): def _fail_delivery(monkeypatch, executor, kind, *, packet_number=1): - """Fail inside the real copy or collector after physical graph registration.""" + """Fail placement, copy, or collection after physical graph registration.""" error = MemoryError("injected logical output delivery failure") packet = executor._packet targets, partial_outputs = set(), [] @@ -76,6 +77,12 @@ def fail_copy(tensor, *args, **kwargs): return to(tensor, *args, **kwargs) monkeypatch.setattr(torch.Tensor, "to", fail_copy) + elif kind == "placement": + + def fail_placement(*args, **kwargs): + raise error + + monkeypatch.setattr(_memory_policy, "choose_output_placements", fail_placement) else: managed = _tensors.managed_tensor calls = 0 @@ -95,7 +102,7 @@ def fail_attach(tensor): @pytest.mark.parametrize("mode", ["rank", "zero"]) -@pytest.mark.parametrize("kind", ["copy", "attach"]) +@pytest.mark.parametrize("kind", ["copy", "attach", "placement"]) def test_failed_delivery_releases_registered_graph_and_native_cache( monkeypatch, mode, kind ): @@ -214,6 +221,7 @@ def remember(*args, **kwargs): assert rank.weight.grad.item() == 7 +@pytest.mark.parametrize("kind", ["copy", "placement"]) @pytest.mark.parametrize( "delivery,close_error", [ @@ -225,7 +233,7 @@ def remember(*args, **kwargs): ], ) def test_failed_release_preserves_delivery_error_and_retries_without_head_flush( - monkeypatch, delivery, close_error + monkeypatch, delivery, close_error, kind ): rank: Any = _CachedRank() executor = _Executor(rank, "rank" if delivery == "rank" else "zero") @@ -255,7 +263,7 @@ def fail_close(*args, **kwargs): raise RuntimeError("injected iterator close failure") patch.setattr(rank, "forward_batches", fail_close) - error, _ = _fail_delivery(patch, executor, "copy") + error, _ = _fail_delivery(patch, executor, kind) if delivery == "iterator": iterator = view.forward_batches([_input(3)]) advance = lambda: next(iterator) diff --git a/tests/unit/test_trainer_rank_output_memory.py b/tests/unit/test_trainer_rank_output_memory.py index adb5a5d80..4422564cc 100644 --- a/tests/unit/test_trainer_rank_output_memory.py +++ b/tests/unit/test_trainer_rank_output_memory.py @@ -40,15 +40,11 @@ def test_logical_copy_preserves_pending_restore(rank, monkeypatch, policy): _pending(rank, monkeypatch, SimpleNamespace(restore_workspace_bytes=100)) monkeypatch.setattr(rank, "_available_memory_bytes", lambda: 120) view = _view(_Executor(rank, "zero")) - released = [] - monkeypatch.setattr(view, "_invoke", lambda *args: released.append(args)) if policy == "model": with pytest.raises(MemoryError, match="only 20 bytes"): view._place_outputs([_output(policy=policy)]) - assert released == [("release", ("new",))] else: assert view._place_outputs([_output(policy=policy)])[0].cpu == (True,) - assert released == [] def test_logical_copy_reserves_distinct_checkpoints_and_standalone_heads( diff --git a/tests/unit/test_trainer_rank_output_memory_cuda.py b/tests/unit/test_trainer_rank_output_memory_cuda.py index 9a18be164..61a83e684 100644 --- a/tests/unit/test_trainer_rank_output_memory_cuda.py +++ b/tests/unit/test_trainer_rank_output_memory_cuda.py @@ -38,8 +38,6 @@ def test_logical_outputs_leave_existing_replay_backward_admissible(monkeypatch): 0 ] view = _view(_Executor(trainer, "zero")) - released = [] - monkeypatch.setattr(view, "_invoke", lambda *args: released.append(args)) gc.collect() torch.cuda.synchronize() baseline = torch.cuda.memory_allocated() @@ -53,7 +51,6 @@ def test_logical_outputs_leave_existing_replay_backward_admissible(monkeypatch): large = 32 * 1024**2 with pytest.raises(MemoryError): view._place_outputs([_output(large, policy="model")]) - assert released == [("release", ("new",))] assert cache.handles() == (handle,) auto = view._attach(view._place_outputs([_output(large)])[0]) assert auto.hidden_states.device.type == "cpu" From 0c6ea44410d8d7defa177f5ba75e1f79eec825d9 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 22:13:22 +0000 Subject: [PATCH 061/150] fix(trainer-rank): preserve transport failures during cleanup --- src/art/trainer_rank/_operations.py | 16 ++-- tests/unit/test_trainer_driver_transport.py | 89 +++++++++++++++++++-- tests/unit/test_trainer_operations.py | 3 + 3 files changed, 91 insertions(+), 17 deletions(-) diff --git a/src/art/trainer_rank/_operations.py b/src/art/trainer_rank/_operations.py index 4ad76edbd..401790002 100644 --- a/src/art/trainer_rank/_operations.py +++ b/src/art/trainer_rank/_operations.py @@ -185,17 +185,13 @@ async def execute_operation(rank_zero: Any, operation: TrainerOperation) -> Any: # receives the original policy, and native callback views are intact. transport = copy(rank_zero) transport._transport_handles = [] - try: - tree = ( - transport.forward(**payload) - if operation.kind == "forward" - else transport.next_forward_batch(**payload) - ) + tree = ( + transport.forward(**payload) + if operation.kind == "forward" + else transport.next_forward_batch(**payload) + ) + with transport._release_on_error(transport._transport_handles): result = None if tree is None else transport.export_forward(tree) - except BaseException: - if transport._transport_handles: - transport._invoke("release", tuple(transport._transport_handles)) - raise elif operation.kind == "backward": result = rank_zero.backward_packets(**payload) elif operation.kind == "optim_step": diff --git a/tests/unit/test_trainer_driver_transport.py b/tests/unit/test_trainer_driver_transport.py index bcd039230..47f19f143 100644 --- a/tests/unit/test_trainer_driver_transport.py +++ b/tests/unit/test_trainer_driver_transport.py @@ -5,6 +5,7 @@ import asyncio from dataclasses import replace import gc +import sys from typing import Any, cast import weakref @@ -163,9 +164,12 @@ async def test_transport_preserves_worker_admission_and_native_view(policy): assert not rank._rank_command_state.graphs -@pytest.mark.parametrize("failure_kind", ["budget", "allocation"]) +@pytest.mark.parametrize( + "failure_kind", ["budget", "allocation", "release", "cancel", "attach"] +) +@pytest.mark.parametrize("batches", [False, True]) async def test_export_failure_releases_only_failed_operation_and_counts_live_exports( - monkeypatch, failure_kind + monkeypatch, failure_kind, batches ): rank: Any = _TransportRank() view = _view(_Executor(rank, "zero")) @@ -189,6 +193,32 @@ def available(): good = await _operation(view, "forward", {"inputs": request}, "good") assert sum(t.numel() * t.element_size() for t in state.exports[good.handle]) == 4 original_graphs = set(state.graphs) + handle = ( + await _operation(view, "batches_open", {"inputs": [request, _input(7)]}, "open") + if batches + else None + ) + operation = TrainerOperation.capture( + ("failure", 1), + "batches_next" if batches else "forward", + {"handle": handle} if batches else {"inputs": request}, + ) + error = (asyncio.CancelledError if failure_kind == "cancel" else MemoryError)( + "injected transport delivery failure" + ) + release_fails = failure_kind in ("release", "cancel", "attach") + invoke = view._executor.invoke + release_calls, flushes = [], [] + + def release(operation, *args, **kwargs): + if operation == "release": + release_calls.append(args[0]) + if release_fails and len(release_calls) <= 2: + raise RuntimeError("injected release failure") + return invoke(operation, *args, **kwargs) + + monkeypatch.setattr(view._executor, "invoke", release) + monkeypatch.setattr(view, "_flush_heads", lambda: flushes.append(sys.exc_info()[1])) # The worker snapshot fits. A budget drop at export must fail before # its clone and clean only the newly created physical graphs. calls = 0 @@ -198,28 +228,73 @@ def constrained(): calls += 1 return available() if calls == 1 else 0 + detach, managed = _tensors.detach_tree, _tensors.managed_tensor if failure_kind == "budget": rank._available_cpu_memory_bytes = constrained + elif failure_kind == "attach": + + def reject_attach(tensor): + raise error + + monkeypatch.setattr(_tensors, "managed_tensor", reject_attach) else: - detach = _tensors.detach_tree def reject_clone(handle, *args, **kwargs): if handle.startswith("client:"): - raise MemoryError("export snapshot allocation failed") + raise error return detach(handle, *args, **kwargs) monkeypatch.setattr(_tensors, "detach_tree", reject_clone) - with pytest.raises(MemoryError, match="export snapshot") as failure: - await _operation(view, "forward", {"inputs": request}, "failure") + with pytest.raises( + type(error), match="export snapshot|transport delivery" + ) as failure: + await execute_operation(view, operation) assert failure.value.__traceback__ is not None + if failure_kind != "budget": + assert failure.value is error + assert len(release_calls) == 1 and flushes and not any(flushes) + assert (set(state.graphs) != original_graphs) is release_fails + assert state.released == set(state.graphs) - original_graphs + with pytest.raises(type(error), match=str(failure.value)) as replay: + await execute_operation(view, operation) + assert replay.value is not failure.value and len(release_calls) == 1 assert set(state.exports) == {good.handle} - assert set(state.graphs) == original_graphs assert 4 in observed rank._available_cpu_memory_bytes = available + monkeypatch.setattr(_tensors, "detach_tree", detach) + monkeypatch.setattr(_tensors, "managed_tensor", managed) + if release_fails: + before = len(flushes) + with pytest.raises(RuntimeError, match="injected release failure"): + view.optim_step() + assert len(release_calls) == 2 and len(flushes) == before and rank.steps == 0 + assert state.released == set(state.graphs) - original_graphs + assert view.optim_step() == {"steps": 1} + assert not state.released and set(state.graphs) == original_graphs await _operation(view, "release", {"handles": [good.handle]}, "release") assert not state.exports assert not state.graphs assert available() == total + if batches and failure_kind == "attach": + assert not state.iterators and not state.batch_inputs + handle = await _operation( + view, "batches_open", {"inputs": [_input(7)]}, "reopen" + ) + packet = await _operation( + view, + "batches_next" if batches else "forward", + {"handle": handle} if batches else {"inputs": _input(7)}, + "recovered", + ) + client = CotangentCollector() + output = client.attach(packet) + value = (output.outputs[0] if batches else output).hidden_states + await _operation(view, "backward", {"packets": client.backward(value.sum())}) + assert rank.weight.grad.item() == 28 + if batches: + await _operation(view, "batches_close", {"handle": handle}, "close") + assert not state.graphs and not state.exports + assert not state.iterators and not state.batch_inputs def test_transport_release_drops_cpu_payload_without_collecting_cycles(): diff --git a/tests/unit/test_trainer_operations.py b/tests/unit/test_trainer_operations.py index edea9d709..6a5dfa2c3 100644 --- a/tests/unit/test_trainer_operations.py +++ b/tests/unit/test_trainer_operations.py @@ -1,5 +1,6 @@ import asyncio from collections.abc import Coroutine +from contextlib import nullcontext import gc from types import SimpleNamespace from typing import Any, cast @@ -104,6 +105,7 @@ async def test_operation_captures_tensor_arguments_at_submission(): _rank=SimpleNamespace(), forward=lambda inputs: inputs * 3, export_forward=lambda output: output, + _release_on_error=lambda handles: nullcontext(), ) assert torch.equal(await execute_operation(rank, operation), torch.tensor([6.0])) @@ -208,6 +210,7 @@ async def test_batch_pulls_and_close_are_identified_without_advancing_twice(): open_forward_batches=lambda **kwargs: events.append("open") or "iterator", next_forward_batch=lambda **kwargs: events.append("next") or "batch", export_forward=lambda batch: SimpleNamespace(handle="packet", batch=batch), + _release_on_error=lambda handles: nullcontext(), close_forward_batches=lambda handle: events.append(("close", handle)), release_forward=lambda handles: events.append(("release", tuple(handles))), ) From 59f970417815b0cf99a35ea85678b4d2d1875ba2 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 22:06:30 +0000 Subject: [PATCH 062/150] refactor(trainer-rank): inline gradient batch commit helper --- src/art/trainer_rank/_versions.py | 16 +++++----------- 1 file changed, 5 insertions(+), 11 deletions(-) diff --git a/src/art/trainer_rank/_versions.py b/src/art/trainer_rank/_versions.py index 4776e891b..e6391b39d 100644 --- a/src/art/trainer_rank/_versions.py +++ b/src/art/trainer_rank/_versions.py @@ -260,13 +260,17 @@ def accumulate(self, gradients: Sequence[VersionedGradient]) -> None: def commit(self, gradients: Sequence[VersionedGradient]) -> None: batch = _GradientBatch() + prepared = None try: self.validate_gradients(gradients) for entry in gradients: batch.add(*entry) - self._commit_batch(batch) + prepared = self._prepare_batch(batch) + self._publish(prepared) finally: batch.clear() + if prepared is not None: + prepared.clear() def _prepare_batch(self, batch: _GradientBatch) -> _PreparedGradients: from ._parameter_hooks import apply_parameter_hooks @@ -313,16 +317,6 @@ def _publish(self, prepared: _PreparedGradients) -> None: finally: parameter = gradient = previous = None - def _commit_batch(self, batch: _GradientBatch) -> None: - prepared = None - try: - prepared = self._prepare_batch(batch) - self._publish(prepared) - finally: - batch.clear() - if prepared is not None: - prepared.clear() - def validate_accumulated(self, names: Sequence[str]) -> None: for name in names: for version, maximum in self._origins.get(name, ()): From a5082b187c015470d281f682d4ef318c8d674ef4 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 22:19:03 +0000 Subject: [PATCH 063/150] fix(trainer-rank): distinguish callback stream exhaustion --- src/art/trainer_rank/_commands.py | 6 +++- .../test_trainer_rank_callback_lifecycle.py | 28 +++++++++++++++- tests/unit/test_trainer_rank_commands.py | 32 +++++++++++++++++-- 3 files changed, 62 insertions(+), 4 deletions(-) diff --git a/src/art/trainer_rank/_commands.py b/src/art/trainer_rank/_commands.py index 12fdfa092..5457f7864 100644 --- a/src/art/trainer_rank/_commands.py +++ b/src/art/trainer_rank/_commands.py @@ -1199,7 +1199,11 @@ async def run_rank_callback_stream( if inspect.isasyncgen(iterator) else iterator.send(sent) ) - except (StopIteration, StopAsyncIteration): + except ( + StopAsyncIteration + if inspect.isasyncgen(iterator) + else StopIteration + ): return sent = yield RankCallbackResult( 0 if mode == "zero" else executor.dp_rank, value diff --git a/tests/unit/test_trainer_rank_callback_lifecycle.py b/tests/unit/test_trainer_rank_callback_lifecycle.py index 211cceb35..0310e82f9 100644 --- a/tests/unit/test_trainer_rank_callback_lifecycle.py +++ b/tests/unit/test_trainer_rank_callback_lifecycle.py @@ -10,7 +10,7 @@ import torch.multiprocessing as mp from trainer_rank_test_support import gloo_group, megatron_topology -from art.trainer_rank import run_rank_callback +from art.trainer_rank import run_rank_callback, run_rank_callback_stream class _CheckpointRank(_Rank): @@ -157,6 +157,32 @@ def mismatched(view): asyncio.run(run_rank_callback(rank, following, mode=mode)) dist.barrier() + wrong_stop = StopAsyncIteration("user generator failure") + closed = [] + + def generate(view): + try: + view.zero_grad() + yield 17 + raise wrong_stop + finally: + closed.append(True) + + async def consume(): + async for result in run_rank_callback_stream(rank, generate, mode=mode): + assert result.value == (17 if physical == 0 else None) + + if physical == 0: + with pytest.raises(RuntimeError, match="generator raised") as failure: + asyncio.run(consume()) + assert failure.value.__cause__ is wrong_stop + else: + asyncio.run(consume()) + assert closed == ([True] if physical == 0 else []) + dist.barrier() + asyncio.run(run_rank_callback(rank, following, mode=mode)) + dist.barrier() + def test_delayed_callback_cleanup_and_checkpoint_scopes_leave_gloo_reusable(tmp_path): mp.spawn(_lifecycle_worker, args=(str(tmp_path / "lifecycle"),), nprocs=2) diff --git a/tests/unit/test_trainer_rank_commands.py b/tests/unit/test_trainer_rank_commands.py index 04620dd19..4c32e0744 100644 --- a/tests/unit/test_trainer_rank_commands.py +++ b/tests/unit/test_trainer_rank_commands.py @@ -229,22 +229,50 @@ def callback(view): assert asyncio.run(run_rank_callback(rank, callback)).value == {"steps": 1} -def test_stream_forwards_sends_and_closes(): +@pytest.mark.parametrize("asynchronous", [False, True]) +@pytest.mark.parametrize( + "ending", ["close", "return", StopIteration, StopAsyncIteration] +) +def test_stream_forwards_sends_and_closes(asynchronous, ending): rank: Any = _Rank() closed = [] + error = ending("user generator failure") if isinstance(ending, type) else None def callback(view): try: sent = yield view.forward(_input(3)).hidden_states.item() yield sent * 2 + if error is not None: + raise error + finally: + closed.append(True) + + async def async_callback(view): + try: + sent = yield view.forward(_input(3)).hidden_states.item() + yield sent * 2 + if error is not None: + raise error finally: closed.append(True) async def run(): - stream = run_rank_callback_stream(rank, callback, mode="zero") + stream = run_rank_callback_stream( + rank, async_callback if asynchronous else callback, mode="zero" + ) assert (await anext(stream)).value == 6 assert (await stream.asend(9)).value == 18 + if ending == "return": + with pytest.raises(StopAsyncIteration): + await anext(stream) + elif error is not None: + with pytest.raises(RuntimeError, match="generator raised") as failure: + await anext(stream) + assert failure.value.__cause__ is error await stream.aclose() + assert ( + await run_rank_callback(rank, lambda view: view.optim_step()) + ).value == {"steps": 1} asyncio.run(run()) assert closed == [True] From cac496c13fa7aa2e004b70f806529aa02bd7196b Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 22:41:55 +0000 Subject: [PATCH 064/150] refactor(trainer-rank): simplify checkpoint and stream internals --- src/art/trainer_rank/_checkpoint.py | 42 +++++++++-------------------- src/art/trainer_rank/_commands.py | 9 +++---- 2 files changed, 17 insertions(+), 34 deletions(-) diff --git a/src/art/trainer_rank/_checkpoint.py b/src/art/trainer_rank/_checkpoint.py index 9ed42095a..f3117c94c 100644 --- a/src/art/trainer_rank/_checkpoint.py +++ b/src/art/trainer_rank/_checkpoint.py @@ -1509,12 +1509,6 @@ def _load_adapter( return {key: handle.get_tensor(key) for key in keys if key in available} -def _localized( - module: LoRA, tensor: torch.Tensor, parameter: torch.nn.Parameter -) -> torch.Tensor: - return module._localized_weight(tensor, into=parameter).contiguous() - - def _slot_snapshot(trainer: TrainerRank) -> _SlotSnapshot: return tuple( ( @@ -1787,7 +1781,9 @@ def _optimizer_state( for key, record in zip(keys, records, strict=True) } full = module._adapter_weight(tensors, suffix=suffix) - components[component].append(_localized(module, full, parameter)) + components[component].append( + module._localized_weight(full, into=parameter).contiguous() + ) key_steps = {source.manifest["steps"][key] for key in keys} if len(key_steps) != 1: raise RuntimeError(f"Optimizer steps differ for {keys}") @@ -1842,27 +1838,6 @@ def _validate_base_model( ) -def _rollback_load( - trainer: TrainerRank, - snapshot: _SlotSnapshot, - temporary: str, - name: str, - previous: object, - group: dist.ProcessGroup | None, -) -> None: - def rollback() -> None: - _restore_slots(snapshot) - trainer._checkpoint_slots.pop(temporary, None) - if previous is None: - trainer._checkpoint_slots.pop(name, None) - else: - from art.trainer_rank._impl import _CheckpointSlot - - trainer._checkpoint_slots[name] = cast(_CheckpointSlot, previous) - - _phase(rollback, "roll back checkpoint load", group) - - def load_checkpoint( trainer: TrainerRank, source: PreparedCheckpoint, @@ -1986,7 +1961,16 @@ def commit() -> None: _phase(commit, "commit checkpoint", group) except BaseException: - _rollback_load(trainer, snapshot, temporary, name, previous, group) + + def rollback() -> None: + _restore_slots(snapshot) + trainer._checkpoint_slots.pop(temporary, None) + if previous is None: + trainer._checkpoint_slots.pop(name, None) + else: + trainer._checkpoint_slots[name] = previous + + _phase(rollback, "roll back checkpoint load", group) raise diff --git a/src/art/trainer_rank/_commands.py b/src/art/trainer_rank/_commands.py index 5457f7864..dbcc5ab58 100644 --- a/src/art/trainer_rank/_commands.py +++ b/src/art/trainer_rank/_commands.py @@ -1192,6 +1192,9 @@ async def run_rank_callback_stream( if not (inspect.isgenerator(iterator) or inspect.isasyncgen(iterator)): raise TypeError("Stream callback must return a generator") sent = None + exhausted = ( + StopAsyncIteration if inspect.isasyncgen(iterator) else StopIteration + ) while True: try: value = ( @@ -1199,11 +1202,7 @@ async def run_rank_callback_stream( if inspect.isasyncgen(iterator) else iterator.send(sent) ) - except ( - StopAsyncIteration - if inspect.isasyncgen(iterator) - else StopIteration - ): + except exhausted: return sent = yield RankCallbackResult( 0 if mode == "zero" else executor.dp_rank, value From 869925f2247151fdc19bf81271423b0ab912ad2f Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 22:59:48 +0000 Subject: [PATCH 065/150] Coordinate checkpoint slot initialization failures --- src/art/trainer_rank/_checkpoint.py | 40 ++++++------- tests/unit/test_trainer_rank_validation.py | 70 +++++++++++++++++++++- 2 files changed, 87 insertions(+), 23 deletions(-) diff --git a/src/art/trainer_rank/_checkpoint.py b/src/art/trainer_rank/_checkpoint.py index f3117c94c..a66b57b70 100644 --- a/src/art/trainer_rank/_checkpoint.py +++ b/src/art/trainer_rank/_checkpoint.py @@ -1912,26 +1912,26 @@ def load_checkpoint( "validate staged checkpoint", group, ) - if forward_only: - for param in params: - param.requires_grad_(False) - from art.trainer_rank._impl import _CheckpointSlot - - trainer._checkpoint_slots[temporary] = _CheckpointSlot( - params, - config, - custom_payload=( - _forward_custom_payload(source.custom) - if forward_only - else source.custom - ), - snapshot=forward_only, - ) - _phase( - lambda: trainer._validate_loaded_checkpoint_config(temporary, config), - "validate loaded checkpoint config", - group, - ) + + def validate_loaded() -> None: + if forward_only: + for param in params: + param.requires_grad_(False) + from art.trainer_rank._impl import _CheckpointSlot + + trainer._checkpoint_slots[temporary] = _CheckpointSlot( + params, + config, + custom_payload=( + _forward_custom_payload(source.custom) + if forward_only + else source.custom + ), + snapshot=forward_only, + ) + trainer._validate_loaded_checkpoint_config(temporary, config) + + _phase(validate_loaded, "validate loaded checkpoint config", group) if ( not forward_only and source.manifest is not None diff --git a/tests/unit/test_trainer_rank_validation.py b/tests/unit/test_trainer_rank_validation.py index 17f0268e8..aa18497fb 100644 --- a/tests/unit/test_trainer_rank_validation.py +++ b/tests/unit/test_trainer_rank_validation.py @@ -2276,6 +2276,12 @@ def _checkpoint_load_failure_worker( assert completed.item() == world_size return + parameter = torch.nn.Parameter(torch.tensor([rank + 1.0])) + gradient = torch.tensor([rank + 3.0]) + parameter.grad = gradient + retained = _CheckpointSlot((parameter,), revision=3, generation=7) + trainer._checkpoint_slots["retained"] = retained + optimizer = ( OptimizerConfig( learning_rate=1e-3, @@ -2312,6 +2318,43 @@ def _checkpoint_load_failure_worker( manifest, "digest", ) + copy_failure = RuntimeError("injected custom payload tensor-copy failure") + if phase == "custom-copy": + tensor = torch.ones(1) + custom = checkpoint_module.PreparedCustomPayload( + { + "p": { + "kind": "parameter", + "tensor_keys": ["p"], + "trainable_keys": ["p"], + "parameter_aliases": [["p"]], + "buffer_aliases": [], + "persistent_buffer_keys": [], + } + }, + {"p": tensor}, + {}, + ) + assert source.manifest is not None + source = replace( + source, + custom=custom, + manifest={ + **source.manifest, + "format_version": 3, + "custom_tensors": custom.records, + }, + ) + original_copy = torch.Tensor.__deepcopy__ + + def copy_tensor( + value: torch.Tensor, memo: dict[int, object] + ) -> torch.Tensor: + if rank == 0 and value is tensor: + raise copy_failure + return original_copy(value, memo) + + monkeypatch.setattr(torch.Tensor, "__deepcopy__", copy_tensor) monkeypatch.setattr( checkpoint_module, @@ -2349,18 +2392,39 @@ def _checkpoint_load_failure_worker( ), ) - with pytest.raises(RuntimeError, match="injected|Another rank failed"): - checkpoint_module.load_checkpoint(trainer, source, "student") + with pytest.raises( + RuntimeError, match="injected|Another rank failed" + ) as caught: + checkpoint_module.load_checkpoint( + trainer, source, "student", forward_only=phase == "custom-copy" + ) + if phase == "custom-copy": + if rank == 0: + assert caught.value is copy_failure + else: + assert "validate loaded checkpoint config" in str(caught.value) assert "student" not in trainer._checkpoint_slots assert not any( name.startswith("__art_loading_") for name in trainer._checkpoint_slots ) + assert set(trainer._checkpoint_slots) == {"retained"} + assert trainer._checkpoint_slots["retained"] is retained + assert (retained.generation, retained.revision) == (7, 3) + assert retained.params[0] is parameter + assert parameter.requires_grad + assert parameter.grad is gradient + torch.testing.assert_close( + parameter, torch.tensor([rank + 1.0]), atol=0, rtol=0 + ) + torch.testing.assert_close(gradient, torch.tensor([rank + 3.0]), atol=0, rtol=0) completed = torch.tensor(1) dist.all_reduce(completed) assert completed.item() == world_size -@pytest.mark.parametrize("phase", ("read", "optimizer", "commit", "export")) +@pytest.mark.parametrize( + "phase", ("read", "optimizer", "commit", "export", "custom-copy") +) def test_checkpoint_load_failure_is_collective_and_transactional( tmp_path: Path, phase: str ) -> None: From 7e5dabd87d6ed6e62fe89ef1e442ba5955faf6ae Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 23:06:35 +0000 Subject: [PATCH 066/150] Simplify checkpoint failure injection callbacks --- tests/unit/test_trainer_rank_validation.py | 52 +++++++--------------- 1 file changed, 17 insertions(+), 35 deletions(-) diff --git a/tests/unit/test_trainer_rank_validation.py b/tests/unit/test_trainer_rank_validation.py index aa18497fb..4dc68e0bb 100644 --- a/tests/unit/test_trainer_rank_validation.py +++ b/tests/unit/test_trainer_rank_validation.py @@ -2356,41 +2356,23 @@ def copy_tensor( monkeypatch.setattr(torch.Tensor, "__deepcopy__", copy_tensor) - monkeypatch.setattr( - checkpoint_module, - "_load_adapter", - ( - lambda *_args: ( - (_ for _ in ()).throw(RuntimeError("injected snapshot read")) - if phase == "read" and rank == 1 - else {} - ) - ), - ) - monkeypatch.setattr( - checkpoint_module, - "_optimizer_state", - ( - lambda *_args: ( - (_ for _ in ()).throw(RuntimeError("injected optimizer read")) - if phase == "optimizer" and rank == 1 - else LocalOptimizerState( - (), (), (), (), cast(OptimizerConfig, optimizer) - ) - ) - ), - ) - monkeypatch.setattr( - checkpoint_module, - "_commit_slot", - ( - lambda *_args: ( - (_ for _ in ()).throw(RuntimeError("injected rank-zero commit")) - if phase == "commit" and rank == 0 - else None - ) - ), - ) + def load_adapter(*_args: object) -> dict[str, torch.Tensor]: + if phase == "read" and rank == 1: + raise RuntimeError("injected snapshot read") + return {} + + def optimizer_state(*_args: object) -> LocalOptimizerState: + if phase == "optimizer" and rank == 1: + raise RuntimeError("injected optimizer read") + return LocalOptimizerState((), (), (), (), cast(OptimizerConfig, optimizer)) + + def commit_slot(*_args: object) -> None: + if phase == "commit" and rank == 0: + raise RuntimeError("injected rank-zero commit") + + monkeypatch.setattr(checkpoint_module, "_load_adapter", load_adapter) + monkeypatch.setattr(checkpoint_module, "_optimizer_state", optimizer_state) + monkeypatch.setattr(checkpoint_module, "_commit_slot", commit_slot) with pytest.raises( RuntimeError, match="injected|Another rank failed" From d12ee06373394194a37835ee4b0d7762944894b1 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 23:10:59 +0000 Subject: [PATCH 067/150] Check checkpoint communicator reuse after failed loads --- tests/unit/test_trainer_rank_validation.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/tests/unit/test_trainer_rank_validation.py b/tests/unit/test_trainer_rank_validation.py index 4dc68e0bb..0194eb7c4 100644 --- a/tests/unit/test_trainer_rank_validation.py +++ b/tests/unit/test_trainer_rank_validation.py @@ -2399,9 +2399,10 @@ def commit_slot(*_args: object) -> None: parameter, torch.tensor([rank + 1.0]), atol=0, rtol=0 ) torch.testing.assert_close(gradient, torch.tensor([rank + 3.0]), atol=0, rtol=0) - completed = torch.tensor(1) - dist.all_reduce(completed) - assert completed.item() == world_size + for group in (trainer._checkpoint_process_group, None): + completed = torch.tensor(1) + dist.all_reduce(completed, group=group) + assert completed.item() == world_size @pytest.mark.parametrize( From 3b04f05a7f3b55d2f449601c0207cc723d12d93f Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 23:19:36 +0000 Subject: [PATCH 068/150] Reuse fresh adapter configuration in checkpoint tests --- tests/unit/test_trainer_rank_validation.py | 52 +++++++--------------- 1 file changed, 16 insertions(+), 36 deletions(-) diff --git a/tests/unit/test_trainer_rank_validation.py b/tests/unit/test_trainer_rank_validation.py index 0194eb7c4..077c9da80 100644 --- a/tests/unit/test_trainer_rank_validation.py +++ b/tests/unit/test_trainer_rank_validation.py @@ -67,6 +67,7 @@ if TYPE_CHECKING: from art.megatron.lora import LoRASlotRef from art.megatron.train import TrainingRuntime + from art.trainer_rank._impl import _AdapterConfig class _Model: @@ -1459,6 +1460,15 @@ def load( snapshot_prepared_checkpoint(trainer, source, "loaded") +def _adapter_config(model: str = "test/model") -> _AdapterConfig: + return { + "base_model_name_or_path": model, + "r": 1, + "lora_alpha": 1, + "target_modules": [], + } + + def test_checkpoint_export_requires_retained_adapter_config() -> None: trainer = TrainerRank(_runtime()) with pytest.raises(TrainerRankSlotStateError, match="unloaded checkpoint"): @@ -1490,12 +1500,7 @@ def capture(*_args: object, **_kwargs: object) -> tuple[object, dict[str, float] ) trainer = TrainerRank(_runtime()) trainer._checkpoint_slots["student"] = _CheckpointSlot( - config={ - "base_model_name_or_path": "test", - "r": 1, - "lora_alpha": 1, - "target_modules": [], - }, + config=_adapter_config("test"), revision=7, ) @@ -1533,12 +1538,7 @@ def test_checkpoint_save_rejects_accumulated_gradients() -> None: trainer._checkpoint_slots.setdefault("student", _CheckpointSlot()).params = ( parameter, ) - trainer._checkpoint_slots["student"].config = { - "base_model_name_or_path": "test", - "r": 1, - "lora_alpha": 1, - "target_modules": [], - } + trainer._checkpoint_slots["student"].config = _adapter_config("test") with pytest.raises(TrainerRankSlotStateError, match="accumulated gradients"): _validate_save_state(trainer, "student") @@ -2177,12 +2177,7 @@ def fail_finish(*_: object) -> None: def test_checkpoint_prepare_preserves_foreign_reservation(tmp_path: Path) -> None: trainer = _save_state_trainer() trainer._checkpoint_slots["student"] = _CheckpointSlot( - config={ - "base_model_name_or_path": "test", - "r": 1, - "lora_alpha": 1, - "target_modules": [], - } + config=_adapter_config("test") ) output = tmp_path / "save" reservation = tmp_path / ".save.reserved" @@ -2203,12 +2198,7 @@ def test_checkpoint_prepare_reports_snapshot_cleanup_failure( trainer = _save_state_trainer() trainer._checkpoint_slots["student"] = _CheckpointSlot( - config={ - "base_model_name_or_path": "test", - "r": 1, - "lora_alpha": 1, - "target_modules": [], - } + config=_adapter_config("test") ) original = _checkpoint.shutil.rmtree @@ -2262,12 +2252,7 @@ def _checkpoint_load_failure_worker( if phase == "export": if rank == 1: trainer._checkpoint_slots["student"] = _CheckpointSlot( - config={ - "base_model_name_or_path": "test/model", - "r": 1, - "lora_alpha": 1, - "target_modules": [], - } + config=_adapter_config() ) with pytest.raises((ValueError, RuntimeError), match="Unknown|Another"): lora_export_module.export_lora(trainer, "/unused", "student") @@ -2308,12 +2293,7 @@ def _checkpoint_load_failure_worker( ) source = PreparedCheckpoint( Path("/unused"), - { - "base_model_name_or_path": "test/model", - "r": 1, - "lora_alpha": 1, - "target_modules": [], - }, + cast(dict[str, object], _adapter_config()), (), manifest, "digest", From 7cfa820c004a0e38e06c92dc1e73ccb7dd3aa836 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 23:25:00 +0000 Subject: [PATCH 069/150] Reuse coordinated phase handling for rank-zero checkpoint work --- src/art/trainer_rank/_checkpoint.py | 8 +------- 1 file changed, 1 insertion(+), 7 deletions(-) diff --git a/src/art/trainer_rank/_checkpoint.py b/src/art/trainer_rank/_checkpoint.py index a66b57b70..2919ae076 100644 --- a/src/art/trainer_rank/_checkpoint.py +++ b/src/art/trainer_rank/_checkpoint.py @@ -1121,13 +1121,7 @@ def _rank_zero_phase( phase: str, group: dist.ProcessGroup | None, ) -> None: - error: BaseException | None = None - if _rank() == 0: - try: - action() - except BaseException as exc: - error = exc - raise_distributed(error, phase, group) + _phase(action if _rank() == 0 else lambda: None, phase, group) def _finish(trainer: TrainerRank, prepared: _PreparedSave) -> None: From 70c0b3374998247891ed175b0be5019c366839c8 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 27 Sep 2026 00:13:51 +0000 Subject: [PATCH 070/150] Share exact checkpoint runtime constructor across trainer tests --- .../unit/test_trainer_rank_custom_tensors.py | 27 +------------ tests/unit/test_trainer_rank_validation.py | 34 +--------------- tests/unit/trainer_rank_test_support.py | 39 ++++++++++++++++++- 3 files changed, 40 insertions(+), 60 deletions(-) diff --git a/tests/unit/test_trainer_rank_custom_tensors.py b/tests/unit/test_trainer_rank_custom_tensors.py index ff638f06f..e8f180c8f 100644 --- a/tests/unit/test_trainer_rank_custom_tensors.py +++ b/tests/unit/test_trainer_rank_custom_tensors.py @@ -16,6 +16,7 @@ import torch import torch.distributed as dist import torch.multiprocessing as mp +from trainer_rank_test_support import checkpoint_runtime as _runtime from trainer_rank_test_support import gloo_group from art.trainer_rank import ( @@ -96,32 +97,6 @@ def __init__(self, *, persistent: bool) -> None: self.register_buffer("running", torch.ones(1), persistent=persistent) -def _runtime(model: torch.nn.Module | None = None) -> Any: - return SimpleNamespace( - model=[model or torch.nn.Linear(1, 1)], - optimizer=None, - provider=SimpleNamespace( - hidden_size=4, - num_layers=1, - kv_channels=2, - art_flex_sliding_windows=(16,), - ), - 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, - ) - - def _config() -> dict[str, object]: return { "base_model_name_or_path": "test/model", diff --git a/tests/unit/test_trainer_rank_validation.py b/tests/unit/test_trainer_rank_validation.py index 077c9da80..031acaa17 100644 --- a/tests/unit/test_trainer_rank_validation.py +++ b/tests/unit/test_trainer_rank_validation.py @@ -19,6 +19,7 @@ import pytest import torch import torch.distributed as dist +from trainer_rank_test_support import checkpoint_runtime as _runtime from trainer_rank_test_support import gloo_group, spawn_and_join from art.megatron.prefix_tree_packing import prefix_tree_pack @@ -66,7 +67,6 @@ if TYPE_CHECKING: from art.megatron.lora import LoRASlotRef - from art.megatron.train import TrainingRuntime from art.trainer_rank._impl import _AdapterConfig @@ -131,38 +131,6 @@ class _SlotRef: name: str | None -def _runtime( - model: torch.nn.Module | None = None, - *, - optimizer: object | None = None, -) -> "TrainingRuntime": - # Deliberately lightweight structural fake; importing/constructing the real - # Megatron runtime would make these CPU-only unit tests require Megatron. - return SimpleNamespace( - model=[model or torch.nn.Linear(1, 1)], - optimizer=optimizer, - provider=SimpleNamespace( - hidden_size=4, - num_layers=1, - kv_channels=2, - art_flex_sliding_windows=(16,), - ), - 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(name: str | None) -> "LoRASlotRef": return _SlotRef(name) # type: ignore diff --git a/tests/unit/trainer_rank_test_support.py b/tests/unit/trainer_rank_test_support.py index 6dbab5fef..0e55abcfd 100644 --- a/tests/unit/trainer_rank_test_support.py +++ b/tests/unit/trainer_rank_test_support.py @@ -1,15 +1,52 @@ -"""Real process groups and lightweight topology for trainer-rank contract tests.""" +"""Shared runtime construction and process groups for trainer-rank contract tests.""" from contextlib import contextmanager from datetime import timedelta import sys import time from types import ModuleType, SimpleNamespace +from typing import TYPE_CHECKING import pytest +import torch import torch.distributed as dist import torch.multiprocessing as mp +if TYPE_CHECKING: + from art.megatron.train import TrainingRuntime + + +def checkpoint_runtime( + model: torch.nn.Module | None = None, + *, + optimizer: object | None = None, +) -> "TrainingRuntime": + # Deliberately lightweight structural fake; importing/constructing the real + # Megatron runtime would make these CPU-only unit tests require Megatron. + return SimpleNamespace( + model=[model or torch.nn.Linear(1, 1)], + optimizer=optimizer, + provider=SimpleNamespace( + hidden_size=4, + num_layers=1, + kv_channels=2, + art_flex_sliding_windows=(16,), + ), + 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 + @contextmanager def process_group(rank, rendezvous, *, world_size=2, timeout=30, backend="gloo"): From 2099ee293c2989c04aa91b4295af8e9e5350068e Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 27 Sep 2026 00:17:05 +0000 Subject: [PATCH 071/150] Retain explicit Linear type in command transport worker --- tests/unit/test_trainer_command_transport.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/tests/unit/test_trainer_command_transport.py b/tests/unit/test_trainer_command_transport.py index aa50a31d8..94a18eb05 100644 --- a/tests/unit/test_trainer_command_transport.py +++ b/tests/unit/test_trainer_command_transport.py @@ -130,7 +130,8 @@ def _transport_worker(physical: int, rendezvous: str, cuda: bool) -> None: gloo_group(physical, f"file://{rendezvous}"), megatron_topology(physical, dp_size=1, tp_size=2), ): - runtime = _runtime(torch.nn.Linear(1, 1).to(device)) + model = torch.nn.Linear(1, 1).to(device) + runtime = _runtime(model) runtime.rank, runtime.world_size = physical, 2 rank: Any = TrainerRank(runtime) rank._dp_rank_and_size = lambda: (0, 1) @@ -179,7 +180,7 @@ def broadcast(sequence: int, operation: str, *args: Any) -> _Command: def forward(inputs: ForwardInput) -> ForwardOutput: assert inputs.input_tokens.device.type == "cpu" values = inputs.input_tokens.to(device=device, dtype=torch.float32) - return ForwardOutput(None, None, None, values * runtime.model[0].weight) + return ForwardOutput(None, None, None, values * model.weight) rank.forward = forward inputs = ForwardInput(input_tokens=torch.tensor([1, 2, 3], device=device)) From 0debb264a60971a8a51a18b26f90f86cae221281 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 27 Sep 2026 00:35:20 +0000 Subject: [PATCH 072/150] Share lightweight GPT model fixture across trainer tests --- tests/unit/test_trainer_rank_split.py | 16 +--------------- tests/unit/test_trainer_rank_topology.py | 12 +----------- tests/unit/test_trainer_rank_weird_shapes.py | 16 +--------------- tests/unit/trainer_rank_test_support.py | 15 +++++++++++++++ 4 files changed, 18 insertions(+), 41 deletions(-) diff --git a/tests/unit/test_trainer_rank_split.py b/tests/unit/test_trainer_rank_split.py index 376676f96..cfdb3e78e 100644 --- a/tests/unit/test_trainer_rank_split.py +++ b/tests/unit/test_trainer_rank_split.py @@ -41,6 +41,7 @@ import pytest import torch +from trainer_rank_test_support import _FakeGPT from art.trainer_rank import ( ForwardInput, @@ -66,21 +67,6 @@ from art.megatron.train import TrainingRuntime -class _FakeGPT(torch.nn.Module): - def __init__(self, *, hidden_size: int = 8, vocab_size: int = 32) -> None: - super().__init__() - self.weight = torch.nn.Parameter(torch.zeros((), dtype=torch.float16)) - self.config = SimpleNamespace( - hidden_size=hidden_size, - num_layers=4, - padded_vocab_size=vocab_size, - ) - self.decoder = object() - - def _preprocess(self, *args: object, **kwargs: object) -> None: - return None - - def _runtime() -> "TrainingRuntime": return SimpleNamespace( model=[_FakeGPT()], diff --git a/tests/unit/test_trainer_rank_topology.py b/tests/unit/test_trainer_rank_topology.py index 45ec291dd..713dfe060 100644 --- a/tests/unit/test_trainer_rank_topology.py +++ b/tests/unit/test_trainer_rank_topology.py @@ -15,6 +15,7 @@ import pytest import torch +from trainer_rank_test_support import _FakeGPT from art.trainer_rank import ( ForwardInput, @@ -26,17 +27,6 @@ from art.megatron.train import TrainingRuntime -class _FakeGPT(torch.nn.Module): - def __init__(self) -> None: - super().__init__() - self.weight = torch.nn.Parameter(torch.zeros((), dtype=torch.float16)) - self.config = SimpleNamespace(hidden_size=8, num_layers=4, padded_vocab_size=32) - self.decoder = object() - - def _preprocess(self, *args: object, **kwargs: object) -> None: - return None - - def _runtime(*, tp: int = 1, pp: int = 1, chunks: int = 1) -> "TrainingRuntime": return SimpleNamespace( model=[_FakeGPT() for _ in range(chunks)], diff --git a/tests/unit/test_trainer_rank_weird_shapes.py b/tests/unit/test_trainer_rank_weird_shapes.py index 929960ff4..3d382a168 100644 --- a/tests/unit/test_trainer_rank_weird_shapes.py +++ b/tests/unit/test_trainer_rank_weird_shapes.py @@ -6,6 +6,7 @@ import pytest import torch +from trainer_rank_test_support import _FakeGPT from art.megatron.prefix_tree_packing import ( estimate_prefix_tree_packed_tokens, @@ -32,21 +33,6 @@ from art.megatron.train import TrainingRuntime -class _FakeGPT(torch.nn.Module): - def __init__(self, *, hidden_size: int = 8, vocab_size: int = 32) -> None: - super().__init__() - self.weight = torch.nn.Parameter(torch.zeros((), dtype=torch.float16)) - self.config = SimpleNamespace( - hidden_size=hidden_size, - num_layers=4, - padded_vocab_size=vocab_size, - ) - self.decoder = object() - - def _preprocess(self, *args: object, **kwargs: object) -> None: - return None - - def _runtime() -> "TrainingRuntime": # Deliberately lightweight structural fake; importing/constructing the real # Megatron runtime would make these CPU-only unit tests require Megatron. diff --git a/tests/unit/trainer_rank_test_support.py b/tests/unit/trainer_rank_test_support.py index 0e55abcfd..4f7393ad6 100644 --- a/tests/unit/trainer_rank_test_support.py +++ b/tests/unit/trainer_rank_test_support.py @@ -16,6 +16,21 @@ from art.megatron.train import TrainingRuntime +class _FakeGPT(torch.nn.Module): + def __init__(self, *, hidden_size: int = 8, vocab_size: int = 32) -> None: + super().__init__() + self.weight = torch.nn.Parameter(torch.zeros((), dtype=torch.float16)) + self.config = SimpleNamespace( + hidden_size=hidden_size, + num_layers=4, + padded_vocab_size=vocab_size, + ) + self.decoder = object() + + def _preprocess(self, *args: object, **kwargs: object) -> None: + return None + + def checkpoint_runtime( model: torch.nn.Module | None = None, *, From 924000dc4454bda50a40579552eacdf8d8e6d430 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 27 Sep 2026 00:57:01 +0000 Subject: [PATCH 073/150] Simplify packed trainer test setup --- tests/unit/test_trainer_rank_split.py | 24 +---- tests/unit/test_trainer_rank_weird_shapes.py | 98 ++++---------------- tests/unit/trainer_rank_test_support.py | 32 +++++-- 3 files changed, 47 insertions(+), 107 deletions(-) diff --git a/tests/unit/test_trainer_rank_split.py b/tests/unit/test_trainer_rank_split.py index cfdb3e78e..1a8f6f74a 100644 --- a/tests/unit/test_trainer_rank_split.py +++ b/tests/unit/test_trainer_rank_split.py @@ -31,7 +31,6 @@ from __future__ import annotations -from collections.abc import Callable from contextlib import nullcontext from dataclasses import dataclass, replace import math @@ -41,7 +40,7 @@ import pytest import torch -from trainer_rank_test_support import _FakeGPT +from trainer_rank_test_support import _FakeGPT, _packed_budget from art.trainer_rank import ( ForwardInput, @@ -56,7 +55,6 @@ _PACKED_PRICED_LOGICAL_ROW_BYTES, Unset, _FlatForwardPlan, - _MemoryCheck, _MemoryProfile, _SplitForwardPlan, ) @@ -93,26 +91,6 @@ def _request(marker: int, length: int = 10) -> ForwardInput: return ForwardInput(input_tokens=tokens, target_tokens=tokens) -def _packed_budget( - monkeypatch: pytest.MonkeyPatch, - rank: TrainerRank, - available: int | Callable[[], int], -) -> None: - """Express memory purely in packed tokens, bypassing the live model.""" - - monkeypatch.setattr( - rank, - "_estimate_required_memory_bytes_from_values", - lambda *, packed_tokens, **_kwargs: packed_tokens, - ) - - def check(required: int, *, sync_across_dp: bool = False) -> _MemoryCheck: - limit = available if isinstance(available, int) else available() - return _MemoryCheck(required, limit, required <= limit) - - monkeypatch.setattr(rank, "_memory_check_required", check) - - def _recording_executor( monkeypatch: pytest.MonkeyPatch, rank: TrainerRank ) -> list[_FlatForwardPlan]: diff --git a/tests/unit/test_trainer_rank_weird_shapes.py b/tests/unit/test_trainer_rank_weird_shapes.py index 3d382a168..6bdff3e52 100644 --- a/tests/unit/test_trainer_rank_weird_shapes.py +++ b/tests/unit/test_trainer_rank_weird_shapes.py @@ -1,12 +1,13 @@ from __future__ import annotations -from collections.abc import Callable, Iterable +from collections.abc import Iterable from types import SimpleNamespace from typing import TYPE_CHECKING import pytest import torch from trainer_rank_test_support import _FakeGPT +from trainer_rank_test_support import _packed_budget as _set_packed_token_budget from art.megatron.prefix_tree_packing import ( estimate_prefix_tree_packed_tokens, @@ -44,6 +45,17 @@ def _runtime() -> "TrainingRuntime": ) # type: ignore +def _empty_executor(monkeypatch: pytest.MonkeyPatch, rank: TrainerRank) -> None: + monkeypatch.setattr( + rank, + "_run_flat_plan_with_memory_tracking", + lambda plan, **_kwargs: ( + [ForwardOutput(None, None, None, None) for _ in range(plan.request_count)], + None, + ), + ) + + def _tokens(*values: int) -> torch.Tensor: return torch.tensor(values, dtype=torch.long) @@ -75,24 +87,6 @@ def _target_request( ) -def _set_packed_token_budget( - monkeypatch: pytest.MonkeyPatch, - rank: TrainerRank, - available: int | Callable[[], int], -) -> None: - monkeypatch.setattr( - rank, - "_estimate_required_memory_bytes_from_values", - lambda *, packed_tokens, **_kwargs: packed_tokens, - ) - - def check(required: int, *, sync_across_dp: bool = False) -> _MemoryCheck: - limit = available if isinstance(available, int) else available() - return _MemoryCheck(required, limit, required <= limit) - - monkeypatch.setattr(rank, "_memory_check_required", check) - - def _ternary_tree_sequences() -> tuple[torch.Tensor, ...]: # Shape: shared root, two continuation branches, and terminal nodes at # several depths. This mirrors prompt -> continuation A/B -> terminal data. @@ -237,14 +231,7 @@ def test_forward_batches_preserves_nested_vineppo_groups( plan.packed_tokens, 10_000, True ), ) - monkeypatch.setattr( - rank, - "_run_flat_plan_with_memory_tracking", - lambda plan, **_kwargs: ( - [ForwardOutput(None, None, None, None) for _ in range(plan.request_count)], - None, - ), - ) + _empty_executor(monkeypatch, rank) groups = _vineppo_like_inputs() micro_batches = list(rank.forward_batches(groups)) @@ -268,14 +255,7 @@ def test_forward_batches_prewarms_next_wave_during_yield( limit = rank._estimate_flat_forward(inputs[:4]) assert limit is not None _set_packed_token_budget(monkeypatch, rank, lambda: limit[0]) - monkeypatch.setattr( - rank, - "_run_flat_plan_with_memory_tracking", - lambda plan, **_kwargs: ( - [ForwardOutput(None, None, None, None) for _ in range(plan.request_count)], - None, - ), - ) + _empty_executor(monkeypatch, rank) generator = rank.forward_batches(inputs) first = next(generator) @@ -322,14 +302,7 @@ def _prewarmed_rank( limit = rank._estimate_flat_forward(budget_rows) assert limit is not None _set_packed_token_budget(monkeypatch, rank, lambda: limit[0]) - monkeypatch.setattr( - rank, - "_run_flat_plan_with_memory_tracking", - lambda plan, **_kwargs: ( - [ForwardOutput(None, None, None, None) for _ in range(plan.request_count)], - None, - ), - ) + _empty_executor(monkeypatch, rank) return rank @@ -405,14 +378,7 @@ def test_width_search_lets_prefix_sharing_widen_the_wave( rank = TrainerRank(_runtime()) monkeypatch.setattr(rank, "_dp_rank_and_size", lambda: (0, 1)) monkeypatch.setattr(rank, "_all_ranks_have_memory_profile", lambda **_kwargs: True) - monkeypatch.setattr( - rank, - "_run_flat_plan_with_memory_tracking", - lambda plan, **_kwargs: ( - [ForwardOutput(None, None, None, None) for _ in range(plan.request_count)], - None, - ), - ) + _empty_executor(monkeypatch, rank) plan = rank._plan_flat_forward(inputs) assert plan.packed_tokens < 4_002, "planner must share the common prefix" # Budget fits the shared plan (2,002 packed) but not the no-sharing bound @@ -459,17 +425,7 @@ def plans(shared_len: int) -> tuple[TrainerRank, list, int, int]: monkeypatch.setattr( rank, "_all_ranks_have_memory_profile", lambda **_kwargs: True ) - monkeypatch.setattr( - rank, - "_run_flat_plan_with_memory_tracking", - lambda plan, **_kwargs: ( - [ - ForwardOutput(None, None, None, None) - for _ in range(plan.request_count) - ], - None, - ), - ) + _empty_executor(monkeypatch, rank) two = rank._plan_flat_forward(inputs[:2]).packed_tokens three = rank._plan_flat_forward(inputs).packed_tokens return rank, inputs, two, three @@ -637,14 +593,7 @@ def test_profiled_steady_state_keeps_the_wide_shared_wave( inputs = [_target_request(_tokens(*prompt, tail)) for tail in range(16)] rank = TrainerRank(_attention_runtime()) monkeypatch.setattr(rank, "_dp_rank_and_size", lambda: (0, 1)) - monkeypatch.setattr( - rank, - "_run_flat_plan_with_memory_tracking", - lambda plan, **_kwargs: ( - [ForwardOutput(None, None, None, None) for _ in range(plan.request_count)], - None, - ), - ) + _empty_executor(monkeypatch, rank) plan = rank._plan_flat_forward(inputs) assert plan.packed_tokens == 1_016, plan.packed_tokens # Steady state: a prior call profiled exactly this shape. @@ -678,14 +627,7 @@ def test_forward_preserves_caller_owned_nested_input_tensors( rank = TrainerRank(_runtime()) monkeypatch.setattr(rank, "_dp_rank_and_size", lambda: (0, 1)) monkeypatch.setattr(rank, "_all_ranks_have_memory_profile", lambda **_: True) - monkeypatch.setattr( - rank, - "_run_flat_plan_with_memory_tracking", - lambda plan, **_kwargs: ( - [ForwardOutput(None, None, None, None) for _ in range(plan.request_count)], - None, - ), - ) + _empty_executor(monkeypatch, rank) groups = _vineppo_like_inputs() tensors = [ (request, request.input_tokens, request.target_tokens) diff --git a/tests/unit/trainer_rank_test_support.py b/tests/unit/trainer_rank_test_support.py index 4f7393ad6..ddf8794cd 100644 --- a/tests/unit/trainer_rank_test_support.py +++ b/tests/unit/trainer_rank_test_support.py @@ -1,5 +1,6 @@ """Shared runtime construction and process groups for trainer-rank contract tests.""" +from collections.abc import Callable from contextlib import contextmanager from datetime import timedelta import sys @@ -14,17 +15,14 @@ if TYPE_CHECKING: from art.megatron.train import TrainingRuntime + from art.trainer_rank import TrainerRank class _FakeGPT(torch.nn.Module): - def __init__(self, *, hidden_size: int = 8, vocab_size: int = 32) -> None: + def __init__(self) -> None: super().__init__() self.weight = torch.nn.Parameter(torch.zeros((), dtype=torch.float16)) - self.config = SimpleNamespace( - hidden_size=hidden_size, - num_layers=4, - padded_vocab_size=vocab_size, - ) + self.config = SimpleNamespace(hidden_size=8, num_layers=4, padded_vocab_size=32) self.decoder = object() def _preprocess(self, *args: object, **kwargs: object) -> None: @@ -63,6 +61,28 @@ def checkpoint_runtime( ) # type: ignore +def _packed_budget( + monkeypatch: pytest.MonkeyPatch, + rank: "TrainerRank", + available: int | Callable[[], int], +) -> None: + """Express memory purely in packed tokens, bypassing the live model.""" + + from art.trainer_rank._impl import _MemoryCheck + + monkeypatch.setattr( + rank, + "_estimate_required_memory_bytes_from_values", + lambda *, packed_tokens, **_kwargs: packed_tokens, + ) + + def check(required: int, *, sync_across_dp: bool = False) -> _MemoryCheck: + limit = available if isinstance(available, int) else available() + return _MemoryCheck(required, limit, required <= limit) + + monkeypatch.setattr(rank, "_memory_check_required", check) + + @contextmanager def process_group(rank, rendezvous, *, world_size=2, timeout=30, backend="gloo"): dist.init_process_group( From 9f6724f049463e946c8680120098a77683706008 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 27 Sep 2026 01:23:43 +0000 Subject: [PATCH 074/150] Remove redundant correction kind state --- src/art/trainer_rank/_corrections.py | 7 +------ 1 file changed, 1 insertion(+), 6 deletions(-) diff --git a/src/art/trainer_rank/_corrections.py b/src/art/trainer_rank/_corrections.py index 7feacb602..76040c4e8 100644 --- a/src/art/trainer_rank/_corrections.py +++ b/src/art/trainer_rank/_corrections.py @@ -118,7 +118,6 @@ def correct_logprob_cotangent( @dataclass(frozen=True) class _OutputCorrection: index: int - kind: str original_logprobs: torch.Tensor token_index: int | None = None original_tokens: torch.Tensor | None = None @@ -309,7 +308,6 @@ def leaves(value: Any) -> Iterator[Any]: index = indices[id(tensor)] entry = _OutputCorrection( index=index, - kind=kind, original_logprobs=tensor.detach().to("cpu", copy=True), token_index=indices[id(output.top_k.tokens)] if kind == "top_k" @@ -321,10 +319,7 @@ def leaves(value: Any) -> Iterator[Any]: if kind == "top_k" and output.logits is not None else None, ) - if index in entries and ( - entries[index].kind != kind - or entries[index].token_index != entry.token_index - ): + if index in entries and entries[index].token_index != entry.token_index: raise ValueError( "an aliased output tensor has ambiguous correction semantics" ) From 70588afc487e6d4b34f30283c1deb157d252cb5e Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 27 Sep 2026 01:35:55 +0000 Subject: [PATCH 075/150] Remove unused native module owner state --- src/art/trainer_rank/_heads.py | 13 +++---------- src/art/trainer_rank/_impl.py | 4 ++-- 2 files changed, 5 insertions(+), 12 deletions(-) diff --git a/src/art/trainer_rank/_heads.py b/src/art/trainer_rank/_heads.py index 86d611c34..af3eea98c 100644 --- a/src/art/trainer_rank/_heads.py +++ b/src/art/trainer_rank/_heads.py @@ -7,7 +7,6 @@ from copy import deepcopy from dataclasses import dataclass, replace from typing import TYPE_CHECKING, Any, Literal, SupportsIndex, cast -import weakref import torch @@ -464,13 +463,7 @@ def __reduce_ex__(self, protocol: SupportsIndex) -> tuple[Any, ...]: class _NativeModuleState: - def __init__( - self, - trainer: TrainerRank, - tracker: _CustomTensorTracker, - module: torch.nn.Module, - ): - self.trainer = weakref.ref(trainer) + def __init__(self, tracker: _CustomTensorTracker, module: torch.nn.Module): self.tracker = tracker self.module = module @@ -527,9 +520,9 @@ def _stage_local_buffers( def native_module_handle( - trainer: TrainerRank, custom: _CustomObject, tracker: _CustomTensorTracker + custom: _CustomObject, tracker: _CustomTensorTracker ) -> ModuleHandle: - state = _NativeModuleState(trainer, tracker, cast(torch.nn.Module, custom.value)) + state = _NativeModuleState(tracker, cast(torch.nn.Module, custom.value)) return ModuleHandle(state.module, state.capture, state.publish) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 9fecf27e2..17a9fbde8 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -10644,8 +10644,8 @@ def _track_custom_object( getattr(child, f"_{kind}s")[key] = replacements[identity] from ._heads import native_module_handle - trainer = tracker.validate() - return replace(custom, handle=native_module_handle(trainer, custom, tracker)) + tracker.validate() + return replace(custom, handle=native_module_handle(custom, tracker)) def _custom_layout( From a4187dcbeec7b71316fe070d87f8e6243bfb34b6 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 27 Sep 2026 02:49:19 +0000 Subject: [PATCH 076/150] Simplify inherited checkpoint finalization arbitration Rely on the existing finalizer lock and preserve the guarded FIFO check, cleanup outcomes, and retry protocol. Remove the redundant transient map/wait protocol and assert lock release in the existing retry control. --- src/art/trainer_rank/_checkpoint.py | 53 +++++----------------- src/art/trainer_rank/_impl.py | 1 - tests/unit/test_trainer_rank_validation.py | 2 +- 3 files changed, 13 insertions(+), 43 deletions(-) diff --git a/src/art/trainer_rank/_checkpoint.py b/src/art/trainer_rank/_checkpoint.py index 2919ae076..7dcd63685 100644 --- a/src/art/trainer_rank/_checkpoint.py +++ b/src/art/trainer_rank/_checkpoint.py @@ -987,7 +987,6 @@ def prepare_checkpoint_save( trainer._prepared_checkpoint_saves[output_dir] = prepared trainer._finalized_checkpoint_saves.pop(output_dir, None) trainer._checkpoint_preparing_saves.discard(output_dir) - trainer._checkpoint_save_condition.notify_all() def _read_snapshot( @@ -1299,7 +1298,6 @@ def _advance_save_queue(trainer: TrainerRank, sequence: int) -> None: while trainer._checkpoint_save_next in trainer._checkpoint_save_skipped: trainer._checkpoint_save_skipped.remove(trainer._checkpoint_save_next) trainer._checkpoint_save_next += 1 - trainer._checkpoint_save_condition.notify_all() def _cleanup_paths(paths: Iterable[Path]) -> BaseException | None: @@ -1314,38 +1312,6 @@ def _cleanup_paths(paths: Iterable[Path]) -> BaseException | None: return BaseExceptionGroup("checkpoint cleanup failed", errors) if errors else None -def _claim_finalization( - trainer: TrainerRank, - output_dir: str, - action: Literal["finish", "abort"], -) -> _PreparedSave | None: - with trainer._checkpoint_save_condition: - while True: - prepared = trainer._prepared_checkpoint_saves.get(output_dir) - if prepared is None: - if output_dir in trainer._finalized_checkpoint_saves: - return None - if action == "abort": - return None - raise RuntimeError(f"Checkpoint save was not prepared: {output_dir}") - outcome = trainer._checkpoint_save_outcomes.get(output_dir) - if outcome is not None and outcome != action: - raise RuntimeError( - f"Checkpoint save was already {outcome}ed: {output_dir}" - ) - if output_dir in trainer._checkpoint_finalizing_saves: - trainer._checkpoint_save_condition.wait() - continue - if outcome is None and prepared.sequence != trainer._checkpoint_save_next: - raise RuntimeError( - "Checkpoint saves must be finalized in preparation order: " - f"expected sequence {trainer._checkpoint_save_next}, got " - f"{prepared.sequence}" - ) - trainer._checkpoint_finalizing_saves[output_dir] = action - return prepared - - def _finalize_checkpoint_save( trainer: TrainerRank, output_dir: str, @@ -1387,12 +1353,19 @@ def _finalize_checkpoint_save( raise RuntimeError(f"Checkpoint save was already {outcome}ed: {output_dir}") if outcome is not None and outcome != action: raise RuntimeError(f"Checkpoint save was already {outcome}ed: {output_dir}") - prepared = ( - _claim_finalization(trainer, output_dir, action) - if finalized is None - else None - ) + prepared = local if finalized is None else None assert prepared is not None or finalized is not None + if prepared is not None: + with trainer._checkpoint_save_condition: + if ( + outcome is None + and prepared.sequence != trainer._checkpoint_save_next + ): + raise RuntimeError( + "Checkpoint saves must be finalized in preparation order: " + f"expected sequence {trainer._checkpoint_save_next}, got " + f"{prepared.sequence}" + ) error: BaseException | None = None cleanup_failed = True try: @@ -1464,7 +1437,6 @@ def _finalize_checkpoint_save( ) finally: with trainer._checkpoint_save_condition: - trainer._checkpoint_finalizing_saves.pop(output_dir, None) if not cleanup_failed: trainer._prepared_checkpoint_saves.pop(output_dir, None) trainer._checkpoint_save_outcomes.pop(output_dir, None) @@ -1472,7 +1444,6 @@ def _finalize_checkpoint_save( trainer._finalized_checkpoint_saves[output_dir] = _FinalizedSave( sequence, outcome ) - trainer._checkpoint_save_condition.notify_all() def finish_checkpoint_save(trainer: TrainerRank, output_dir: str) -> None: diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 17a9fbde8..06473b655 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -2095,7 +2095,6 @@ def memory_field(name: str, default: Any = None) -> Any: self._checkpoint_save_next = 0 self._checkpoint_save_skipped: set[int] = set() self._checkpoint_preparing_saves: set[str] = set() - self._checkpoint_finalizing_saves: dict[str, Literal["finish", "abort"]] = {} self._checkpoint_save_outcomes: dict[str, Literal["finish", "abort"]] = {} self._prepared_checkpoint_saves: dict[str, _PreparedSave] = {} self._finalized_checkpoint_saves: dict[str, _FinalizedSave] = {} diff --git a/tests/unit/test_trainer_rank_validation.py b/tests/unit/test_trainer_rank_validation.py index 031acaa17..d934047f6 100644 --- a/tests/unit/test_trainer_rank_validation.py +++ b/tests/unit/test_trainer_rank_validation.py @@ -2076,7 +2076,7 @@ def fail_once( monkeypatch.setattr(_checkpoint, "_gather", fail_once) with pytest.raises(RuntimeError, match="cleanup gather"): finish_checkpoint_save(trainer, "save") - assert "save" not in trainer._checkpoint_finalizing_saves + assert not trainer._checkpoint_finalize_lock.locked() finish_checkpoint_save(trainer, "save") assert "save" not in trainer._prepared_checkpoint_saves From 5d61a6b60402f08d92cf5f4dc0bf88a8828bc3f1 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 27 Sep 2026 03:17:21 +0000 Subject: [PATCH 077/150] Allow cleanup-only abort after checkpoint finish Recover retained native temporary files without rolling back the finished outcome or replaying serialization. Keep invalid terminal states and finish-after-abort rejected. Cover cleanup retries, terminal guards, committed bytes, and asymmetric finalization state. --- src/art/trainer_rank/_checkpoint.py | 5 +- tests/unit/test_trainer_rank_validation.py | 82 +++++++++++++++++++--- 2 files changed, 75 insertions(+), 12 deletions(-) diff --git a/src/art/trainer_rank/_checkpoint.py b/src/art/trainer_rank/_checkpoint.py index 7dcd63685..c2c0ed2ec 100644 --- a/src/art/trainer_rank/_checkpoint.py +++ b/src/art/trainer_rank/_checkpoint.py @@ -1347,12 +1347,13 @@ def _finalize_checkpoint_save( return raise RuntimeError(f"Checkpoint save was not prepared: {output_dir}") finalized_ranks = _gather(finalized is not None, group) + # Abort may finish cleanup without rolling back a committed save. + if outcome not in (None, action, "finish"): + raise RuntimeError(f"Checkpoint save was already {outcome}ed: {output_dir}") if all(finalized_ranks): if outcome == "finish" or action == "abort": return raise RuntimeError(f"Checkpoint save was already {outcome}ed: {output_dir}") - if outcome is not None and outcome != action: - raise RuntimeError(f"Checkpoint save was already {outcome}ed: {output_dir}") prepared = local if finalized is None else None assert prepared is not None or finalized is not None if prepared is not None: diff --git a/tests/unit/test_trainer_rank_validation.py b/tests/unit/test_trainer_rank_validation.py index d934047f6..fa447dd75 100644 --- a/tests/unit/test_trainer_rank_validation.py +++ b/tests/unit/test_trainer_rank_validation.py @@ -2012,11 +2012,16 @@ def finish() -> None: assert calls == 1 -@pytest.mark.parametrize("action", ("finish", "abort")) +@pytest.mark.parametrize( + "action,retry", (("finish", "finish"), ("finish", "abort"), ("abort", "abort")) +) +@pytest.mark.parametrize("failed_path", ("snapshot", "reservation")) def test_checkpoint_cleanup_failure_can_be_retried( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, action: str, + retry: str, + failed_path: str, ) -> None: from art.trainer_rank import _checkpoint @@ -2028,28 +2033,77 @@ def test_checkpoint_cleanup_failure_can_be_retried( def finalize(_trainer: TrainerRank, _prepared: _PreparedSave) -> None: nonlocal finalizations finalizations += 1 + prepared.destination.mkdir() + (prepared.destination / "committed").write_bytes(b"saved state") original = _checkpoint.shutil.rmtree failed = False + failure = OSError("injected cleanup failure") def fail_once(path: Path, ignore_errors: bool = False, **_: object) -> None: nonlocal failed - if Path(path) == prepared.snapshot and not failed: + if Path(path) == getattr(prepared, failed_path) and not failed: failed = True - raise OSError("injected cleanup failure") + raise failure original(path, ignore_errors=ignore_errors) monkeypatch.setattr(_checkpoint, "_finish", finalize) monkeypatch.setattr(_checkpoint.shutil, "rmtree", fail_once) operation = finish_checkpoint_save if action == "finish" else abort_checkpoint_save - with pytest.raises(BaseExceptionGroup, match="cleanup failed"): + with pytest.raises(BaseExceptionGroup, match="cleanup failed") as raised: operation(trainer, "save") - operation(trainer, "save") + assert raised.value.exceptions == (failure,) + assert trainer._checkpoint_save_outcomes["save"] == action + assert trainer._checkpoint_save_next == 1 + recovery = finish_checkpoint_save if retry == "finish" else abort_checkpoint_save + recovery(trainer, "save") + abort_checkpoint_save(trainer, "save") assert finalizations == (1 if action == "finish" else 0) + assert trainer._finalized_checkpoint_saves["save"].outcome == action + assert trainer._checkpoint_save_next == 1 assert "save" not in trainer._prepared_checkpoint_saves + assert "save" not in trainer._checkpoint_save_outcomes assert not prepared.snapshot.exists() assert not prepared.reservation.exists() + if action == "finish": + assert (prepared.destination / "committed").read_bytes() == b"saved state" + else: + assert not prepared.destination.exists() + + +@pytest.mark.parametrize("finalized", (False, True)) +@pytest.mark.parametrize( + "outcome,action", (("abort", "finish"), ("invalid", "finish"), ("invalid", "abort")) +) +def test_checkpoint_terminal_outcome_rejects_invalid_recovery( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, + finalized: bool, + outcome: str, + action: str, +) -> None: + from art.trainer_rank import _checkpoint + + trainer = _save_state_trainer() + prepared = _prepared_save(tmp_path, 0) + if finalized: + trainer._finalized_checkpoint_saves["save"] = _FinalizedSave( + 0, cast(Any, outcome) + ) + else: + trainer._prepared_checkpoint_saves["save"] = prepared + trainer._checkpoint_save_outcomes["save"] = cast(Any, outcome) + monkeypatch.setattr(_checkpoint, "_finish", lambda *_: pytest.fail("reran finish")) + monkeypatch.setattr( + _checkpoint, "_cleanup_paths", lambda *_: pytest.fail("cleanup") + ) + operation = finish_checkpoint_save if action == "finish" else abort_checkpoint_save + with pytest.raises(RuntimeError, match=f"already {outcome}ed"): + operation(trainer, "save") + assert trainer._checkpoint_save_next == 0 + assert prepared.snapshot.exists() and prepared.reservation.exists() + assert not prepared.destination.exists() def test_checkpoint_cleanup_gather_failure_releases_finalizer( @@ -2081,8 +2135,9 @@ def fail_once( assert "save" not in trainer._prepared_checkpoint_saves +@pytest.mark.parametrize("action", ("finish", "abort")) def test_checkpoint_asymmetric_cleanup_gather_can_converge( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, action: str ) -> None: from art.trainer_rank import _checkpoint @@ -2094,6 +2149,7 @@ def test_checkpoint_asymmetric_cleanup_gather_can_converge( completed._finalized_checkpoint_saves["save"] = _FinalizedSave(0, "finish") retained._prepared_checkpoint_saves["save"] = retained_save retained._checkpoint_save_outcomes["save"] = "finish" + completed._checkpoint_save_next = retained._checkpoint_save_next = 1 def mixed( value: object, _group: dist.ProcessGroup | None = None @@ -2103,11 +2159,17 @@ def mixed( return (value, value) monkeypatch.setattr(_checkpoint, "_gather", mixed) - finish_checkpoint_save(completed, "save") - finish_checkpoint_save(retained, "save") - assert "save" in completed._finalized_checkpoint_saves - assert "save" in retained._finalized_checkpoint_saves + monkeypatch.setattr(_checkpoint, "_finish", lambda *_: pytest.fail("reran finish")) + operation = finish_checkpoint_save if action == "finish" else abort_checkpoint_save + operation(completed, "save") + operation(retained, "save") + assert completed._finalized_checkpoint_saves["save"].outcome == "finish" + assert retained._finalized_checkpoint_saves["save"].outcome == "finish" + assert completed._checkpoint_save_next == retained._checkpoint_save_next == 1 assert "save" not in retained._prepared_checkpoint_saves + assert ( + not retained_save.snapshot.exists() and not retained_save.reservation.exists() + ) def test_checkpoint_cleanup_gather_preserves_finish_error( From c57cb1c96ce41b4f2ff1b8399bc5d2db89ef9594 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 08:41:10 +0000 Subject: [PATCH 078/150] Share fixed fake rank constructors in memory tests --- tests/unit/test_trainer_rank_active_memory.py | 26 ++++----- .../test_trainer_rank_checkpoint_memory.py | 15 ++---- tests/unit/test_trainer_rank_moe_memory.py | 53 +++++++------------ .../unit/test_trainer_rank_pending_memory.py | 13 +---- tests/unit/test_trainer_rank_tp_floor.py | 15 ++---- tests/unit/trainer_rank_test_support.py | 12 ++++- 6 files changed, 46 insertions(+), 88 deletions(-) diff --git a/tests/unit/test_trainer_rank_active_memory.py b/tests/unit/test_trainer_rank_active_memory.py index b00c3d78d..001f33651 100644 --- a/tests/unit/test_trainer_rank_active_memory.py +++ b/tests/unit/test_trainer_rank_active_memory.py @@ -3,10 +3,10 @@ import builtins from dataclasses import replace from types import SimpleNamespace -from typing import Any, cast import pytest import torch +from trainer_rank_test_support import fake_rank from art.trainer_rank import ( ForwardInput, @@ -41,22 +41,14 @@ def _preprocess(self, *args, **kwargs): def _rank(): - return TrainerRank( - cast( - Any, - SimpleNamespace( - model=[_Model()], - optimizer=None, - provider=SimpleNamespace( - hidden_size=8, - num_layers=4, - recompute_granularity="full", - recompute_method="uniform", - recompute_num_layers=1, - ), - model_support_handler=SimpleNamespace(build_gdn_execution_spec=False), - ), - ) + return fake_rank( + TrainerRank, + [_Model()], + hidden_size=8, + num_layers=4, + recompute_granularity="full", + recompute_method="uniform", + recompute_num_layers=1, ) diff --git a/tests/unit/test_trainer_rank_checkpoint_memory.py b/tests/unit/test_trainer_rank_checkpoint_memory.py index c6e33ee9d..6df7c9403 100644 --- a/tests/unit/test_trainer_rank_checkpoint_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_memory.py @@ -2,10 +2,11 @@ from dataclasses import replace from types import SimpleNamespace -from typing import Any, cast +from typing import Any import pytest import torch +from trainer_rank_test_support import fake_rank from art.trainer_rank import ForwardInput, TrainerRank from art.trainer_rank._impl import Unset, _ForwardRefusal, _MemoryProfile @@ -40,17 +41,7 @@ def rank(): model.config = block.config model.decoder = block model._preprocess = lambda: None - result = TrainerRank( - cast( - Any, - SimpleNamespace( - model=[model], - optimizer=None, - provider=SimpleNamespace(hidden_size=2048, num_layers=40), - model_support_handler=SimpleNamespace(build_gdn_execution_spec=False), - ), - ) - ) + result = fake_rank(TrainerRank, [model], hidden_size=2048, num_layers=40) result._moe_output_bytes_per_token = 188416 result._moe_checkpoint_grad_bytes_per_token = 188416 return result diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index a537ccd8a..09afc48ef 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -2,11 +2,12 @@ from dataclasses import replace from types import SimpleNamespace -from typing import Any, cast +from typing import Any import weakref import pytest import torch +from trainer_rank_test_support import fake_rank from art.trainer_rank import ForwardInput, ForwardOutput, TrainerRank from art.trainer_rank._impl import ( @@ -70,22 +71,14 @@ def module(cls): def _rank(layer=None): model = layer if layer is not None else torch.nn.Linear(1, 1).bfloat16() - return TrainerRank( - cast( - Any, - SimpleNamespace( - model=[model], - optimizer=None, - provider=SimpleNamespace( - hidden_size=2048, - num_layers=40, - recompute_granularity="full", - recompute_method="uniform", - recompute_num_layers=1, - ), - model_support_handler=SimpleNamespace(build_gdn_execution_spec=False), - ), - ) + return fake_rank( + TrainerRank, + [model], + hidden_size=2048, + num_layers=40, + recompute_granularity="full", + recompute_method="uniform", + recompute_num_layers=1, ) @@ -678,24 +671,14 @@ def hybrid_checkpoint_rank(layer, monkeypatch): # Only distributed topology is mocked. Real constructor metadata selects # the HybridEP coefficient, without loading a model or initializing CUDA. monkeypatch.setattr(TrainerRank, "_topology_key", lambda self: (1, 1, 2, 1)) - rank = TrainerRank( - cast( - Any, - SimpleNamespace( - model=[model], - optimizer=None, - provider=SimpleNamespace( - hidden_size=2048, - num_layers=40, - expert_model_parallel_size=2, - expert_tensor_parallel_size=1, - num_moe_experts=256, - ), - model_support_handler=SimpleNamespace( - build_gdn_execution_spec=False - ), - ), - ) + rank = fake_rank( + TrainerRank, + [model], + hidden_size=2048, + num_layers=40, + expert_model_parallel_size=2, + expert_tensor_parallel_size=1, + num_moe_experts=256, ) assert rank._moe_output_bytes_per_token == 282624 assert rank._moe_memory_supported diff --git a/tests/unit/test_trainer_rank_pending_memory.py b/tests/unit/test_trainer_rank_pending_memory.py index 359de8daa..89fba1b02 100644 --- a/tests/unit/test_trainer_rank_pending_memory.py +++ b/tests/unit/test_trainer_rank_pending_memory.py @@ -8,6 +8,7 @@ from test_trainer_rank_moe_memory import _enclosing_moe from test_trainer_rank_moe_memory import layer as layer import torch +from trainer_rank_test_support import fake_rank from art.megatron.prefix_tree_packing import prefix_tree_pack from art.trainer_rank import ForwardInput, TrainerRank @@ -84,17 +85,7 @@ def rank_with_moe(moe_layer, *, install_hooks=False): from art.megatron.gdn.operator import install_gdn_island_hooks install_gdn_island_hooks([model]) - r: Any = TrainerRank( - cast( - Any, - SimpleNamespace( - model=[model], - optimizer=None, - provider=SimpleNamespace(hidden_size=2048, num_layers=40), - model_support_handler=SimpleNamespace(build_gdn_execution_spec=False), - ), - ) - ) + r: Any = fake_rank(TrainerRank, [model], hidden_size=2048, num_layers=40) r._dp_rank_and_size = lambda: (0, 1) # Uninitialized MCore has no CPU DP group. return r, gd diff --git a/tests/unit/test_trainer_rank_tp_floor.py b/tests/unit/test_trainer_rank_tp_floor.py index 6d5f4d00d..6d1f571df 100644 --- a/tests/unit/test_trainer_rank_tp_floor.py +++ b/tests/unit/test_trainer_rank_tp_floor.py @@ -8,10 +8,11 @@ from dataclasses import replace from types import SimpleNamespace -from typing import Any, cast +from typing import Any import pytest import torch +from trainer_rank_test_support import fake_rank from art.trainer_rank import TrainerRank from art.trainer_rank._impl import _MemorySignature @@ -55,17 +56,7 @@ def tp_rank(layers=LAYERS, *, ffn=F, topology=TP4, sequence_parallel=True, **con model.config = block.config model.decoder = block model._preprocess = lambda: None - r: Any = TrainerRank( - cast( - Any, - SimpleNamespace( - model=[model], - optimizer=None, - provider=SimpleNamespace(hidden_size=H, num_layers=layers), - model_support_handler=SimpleNamespace(build_gdn_execution_spec=False), - ), - ) - ) + r: Any = fake_rank(TrainerRank, [model], hidden_size=H, num_layers=layers) # Qwen3.8-27B: gated attention every fourth layer, GDN otherwise. r._geometry = replace( r._geometry, diff --git a/tests/unit/trainer_rank_test_support.py b/tests/unit/trainer_rank_test_support.py index ddf8794cd..7df9fbb1c 100644 --- a/tests/unit/trainer_rank_test_support.py +++ b/tests/unit/trainer_rank_test_support.py @@ -6,7 +6,7 @@ import sys import time from types import ModuleType, SimpleNamespace -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any import pytest import torch @@ -29,6 +29,16 @@ def _preprocess(self, *args: object, **kwargs: object) -> None: return None +def fake_rank(rank_type: type["TrainerRank"], model, **provider) -> "TrainerRank": + runtime: Any = SimpleNamespace( + model=model, + optimizer=None, + provider=SimpleNamespace(**provider), + model_support_handler=SimpleNamespace(build_gdn_execution_spec=False), + ) + return rank_type(runtime) + + def checkpoint_runtime( model: torch.nn.Module | None = None, *, From e4545fc6dda68982385aa3403e90857e00a24dce Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 09:15:20 +0000 Subject: [PATCH 079/150] Share full recompute config in memory fixtures --- .../test_trainer_rank_checkpoint_memory.py | 19 ++--------------- tests/unit/test_trainer_rank_moe_memory.py | 19 ++--------------- .../unit/test_trainer_rank_pending_memory.py | 19 ++--------------- tests/unit/test_trainer_rank_tp_floor.py | 21 ++----------------- tests/unit/trainer_rank_test_support.py | 20 ++++++++++++++++++ 5 files changed, 28 insertions(+), 70 deletions(-) diff --git a/tests/unit/test_trainer_rank_checkpoint_memory.py b/tests/unit/test_trainer_rank_checkpoint_memory.py index 6df7c9403..b9d7dd68f 100644 --- a/tests/unit/test_trainer_rank_checkpoint_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_memory.py @@ -6,7 +6,7 @@ import pytest import torch -from trainer_rank_test_support import fake_rank +from trainer_rank_test_support import fake_rank, full_recompute_config from art.trainer_rank import ForwardInput, TrainerRank from art.trainer_rank._impl import Unset, _ForwardRefusal, _MemoryProfile @@ -17,22 +17,7 @@ def rank(): block = TransformerBlock.__new__(TransformerBlock) torch.nn.Module.__init__(block) - block.config = SimpleNamespace( - hidden_size=2048, - num_layers=40, - padded_vocab_size=32, - params_dtype=torch.bfloat16, - recompute_granularity="full", - recompute_method="uniform", - recompute_num_layers=1, - distribute_saved_activations=False, - sequence_parallel=False, - fp32_residual_connection=False, - cpu_offloading=False, - cuda_graph_impl="none", - fp8=None, - fp4=None, - ) + block.config = full_recompute_config(2048, 40, False) block.layers = torch.nn.ModuleList( [torch.nn.Linear(1, 1).bfloat16() for _ in range(40)] ) diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index 09afc48ef..6b81a7ebe 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -7,7 +7,7 @@ import pytest import torch -from trainer_rank_test_support import fake_rank +from trainer_rank_test_support import fake_rank, full_recompute_config from art.trainer_rank import ForwardInput, ForwardOutput, TrainerRank from art.trainer_rank._impl import ( @@ -645,22 +645,7 @@ def hybrid_checkpoint_rank(layer, monkeypatch): torch.empty(128, 8, outputs, dtype=torch.bfloat16) ) decoder = module(TransformerBlock) - decoder.config = SimpleNamespace( - hidden_size=2048, - num_layers=40, - padded_vocab_size=32, - params_dtype=torch.bfloat16, - recompute_granularity="full", - recompute_method="uniform", - recompute_num_layers=1, - distribute_saved_activations=False, - sequence_parallel=False, - fp32_residual_connection=False, - cpu_offloading=False, - cuda_graph_impl="none", - fp8=None, - fp4=None, - ) + decoder.config = full_recompute_config(2048, 40, False) decoder.num_layers_per_pipeline_rank = 40 decoder.layers = torch.nn.ModuleList( [moe] + [torch.nn.Linear(1, 1).bfloat16() for _ in range(39)] diff --git a/tests/unit/test_trainer_rank_pending_memory.py b/tests/unit/test_trainer_rank_pending_memory.py index 89fba1b02..2fdb3ac7b 100644 --- a/tests/unit/test_trainer_rank_pending_memory.py +++ b/tests/unit/test_trainer_rank_pending_memory.py @@ -8,7 +8,7 @@ from test_trainer_rank_moe_memory import _enclosing_moe from test_trainer_rank_moe_memory import layer as layer import torch -from trainer_rank_test_support import fake_rank +from trainer_rank_test_support import fake_rank, full_recompute_config from art.megatron.prefix_tree_packing import prefix_tree_pack from art.trainer_rank import ForwardInput, TrainerRank @@ -31,22 +31,7 @@ def rank_with_moe(moe_layer, *, install_hooks=False): from art.megatron.lora import LoRA, SelfAttentionLinearProjLoRA decoder = module(TransformerBlock) - decoder.config = SimpleNamespace( - hidden_size=2048, - num_layers=40, - padded_vocab_size=32, - params_dtype=torch.bfloat16, - recompute_granularity="full", - recompute_method="uniform", - recompute_num_layers=1, - distribute_saved_activations=False, - sequence_parallel=False, - fp32_residual_connection=False, - cpu_offloading=False, - cuda_graph_impl="none", - fp8=None, - fp4=None, - ) + decoder.config = full_recompute_config(2048, 40, False) decoder.layers = torch.nn.ModuleList( [torch.nn.Linear(1, 1).bfloat16() for _ in range(40)] ) diff --git a/tests/unit/test_trainer_rank_tp_floor.py b/tests/unit/test_trainer_rank_tp_floor.py index 6d1f571df..0500b75ef 100644 --- a/tests/unit/test_trainer_rank_tp_floor.py +++ b/tests/unit/test_trainer_rank_tp_floor.py @@ -7,12 +7,11 @@ """ from dataclasses import replace -from types import SimpleNamespace from typing import Any import pytest import torch -from trainer_rank_test_support import fake_rank +from trainer_rank_test_support import fake_rank, full_recompute_config from art.trainer_rank import TrainerRank from art.trainer_rank._impl import _MemorySignature @@ -31,23 +30,7 @@ def tp_rank(layers=LAYERS, *, ffn=F, topology=TP4, sequence_parallel=True, **con block = TransformerBlock.__new__(TransformerBlock) torch.nn.Module.__init__(block) - block.config = SimpleNamespace( - hidden_size=H, - num_layers=layers, - padded_vocab_size=32, - params_dtype=torch.bfloat16, - recompute_granularity="full", - recompute_method="uniform", - recompute_num_layers=1, - distribute_saved_activations=False, - sequence_parallel=sequence_parallel, - fp32_residual_connection=False, - cpu_offloading=False, - cuda_graph_impl="none", - fp8=None, - fp4=None, - **config, - ) + block.config = full_recompute_config(H, layers, sequence_parallel, **config) block.layers = torch.nn.ModuleList( [torch.nn.Linear(1, 1).bfloat16() for _ in range(layers)] ) diff --git a/tests/unit/trainer_rank_test_support.py b/tests/unit/trainer_rank_test_support.py index 7df9fbb1c..9d902c491 100644 --- a/tests/unit/trainer_rank_test_support.py +++ b/tests/unit/trainer_rank_test_support.py @@ -39,6 +39,26 @@ def fake_rank(rank_type: type["TrainerRank"], model, **provider) -> "TrainerRank return rank_type(runtime) +def full_recompute_config(hidden_size, num_layers, sequence_parallel, /, **config): + return SimpleNamespace( + hidden_size=hidden_size, + num_layers=num_layers, + padded_vocab_size=32, + params_dtype=torch.bfloat16, + recompute_granularity="full", + recompute_method="uniform", + recompute_num_layers=1, + distribute_saved_activations=False, + sequence_parallel=sequence_parallel, + fp32_residual_connection=False, + cpu_offloading=False, + cuda_graph_impl="none", + fp8=None, + fp4=None, + **config, + ) + + def checkpoint_runtime( model: torch.nn.Module | None = None, *, From 4a3da1169b44db2b0718b91648d644e268518213 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 11:16:16 +0000 Subject: [PATCH 080/150] Share real transformer model shells in memory fixtures --- .../test_trainer_rank_checkpoint_memory.py | 15 ++------------- tests/unit/test_trainer_rank_moe_memory.py | 13 ++----------- .../unit/test_trainer_rank_pending_memory.py | 15 +++------------ tests/unit/test_trainer_rank_tp_floor.py | 14 ++------------ tests/unit/trainer_rank_test_support.py | 19 +++++++++++++++++++ 5 files changed, 28 insertions(+), 48 deletions(-) diff --git a/tests/unit/test_trainer_rank_checkpoint_memory.py b/tests/unit/test_trainer_rank_checkpoint_memory.py index b9d7dd68f..4b5c4fac7 100644 --- a/tests/unit/test_trainer_rank_checkpoint_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_memory.py @@ -2,11 +2,10 @@ from dataclasses import replace from types import SimpleNamespace -from typing import Any import pytest import torch -from trainer_rank_test_support import fake_rank, full_recompute_config +from trainer_rank_test_support import fake_rank, recompute_model from art.trainer_rank import ForwardInput, TrainerRank from art.trainer_rank._impl import Unset, _ForwardRefusal, _MemoryProfile @@ -15,17 +14,7 @@ def rank(): from megatron.core.transformer.transformer_block import TransformerBlock - block = TransformerBlock.__new__(TransformerBlock) - torch.nn.Module.__init__(block) - block.config = full_recompute_config(2048, 40, False) - block.layers = torch.nn.ModuleList( - [torch.nn.Linear(1, 1).bfloat16() for _ in range(40)] - ) - block.num_layers_per_pipeline_rank = 40 - model: Any = torch.nn.Module() - model.config = block.config - model.decoder = block - model._preprocess = lambda: None + model = recompute_model(TransformerBlock, 2048, 40, False) result = fake_rank(TrainerRank, [model], hidden_size=2048, num_layers=40) result._moe_output_bytes_per_token = 188416 result._moe_checkpoint_grad_bytes_per_token = 188416 diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index 6b81a7ebe..541f4c2e4 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -7,7 +7,7 @@ import pytest import torch -from trainer_rank_test_support import fake_rank, full_recompute_config +from trainer_rank_test_support import fake_rank, recompute_model from art.trainer_rank import ForwardInput, ForwardOutput, TrainerRank from art.trainer_rank._impl import ( @@ -629,7 +629,6 @@ def child(workspace: int, growth: int) -> _SubforwardCost: def hybrid_checkpoint_rank(layer, monkeypatch): from megatron.core.transformer.transformer_block import TransformerBlock from test_trainer_rank_converted_memory import weights - from test_trainer_rank_pending_memory import module with torch.device("meta"): moe = _hybridep(weights(layer, 8), 2) @@ -644,15 +643,7 @@ def hybrid_checkpoint_rank(layer, monkeypatch): fc.lora.B_T = torch.nn.Parameter( torch.empty(128, 8, outputs, dtype=torch.bfloat16) ) - decoder = module(TransformerBlock) - decoder.config = full_recompute_config(2048, 40, False) - decoder.num_layers_per_pipeline_rank = 40 - decoder.layers = torch.nn.ModuleList( - [moe] + [torch.nn.Linear(1, 1).bfloat16() for _ in range(39)] - ) - model: Any = torch.nn.Module() - model.config, model.decoder = decoder.config, decoder - model._preprocess = lambda: None + model = recompute_model(TransformerBlock, 2048, 40, False, layers=(moe,)) # Only distributed topology is mocked. Real constructor metadata selects # the HybridEP coefficient, without loading a model or initializing CUDA. monkeypatch.setattr(TrainerRank, "_topology_key", lambda self: (1, 1, 2, 1)) diff --git a/tests/unit/test_trainer_rank_pending_memory.py b/tests/unit/test_trainer_rank_pending_memory.py index 2fdb3ac7b..c41a97229 100644 --- a/tests/unit/test_trainer_rank_pending_memory.py +++ b/tests/unit/test_trainer_rank_pending_memory.py @@ -8,7 +8,7 @@ from test_trainer_rank_moe_memory import _enclosing_moe from test_trainer_rank_moe_memory import layer as layer import torch -from trainer_rank_test_support import fake_rank, full_recompute_config +from trainer_rank_test_support import fake_rank, recompute_model from art.megatron.prefix_tree_packing import prefix_tree_pack from art.trainer_rank import ForwardInput, TrainerRank @@ -30,12 +30,7 @@ def rank_with_moe(moe_layer, *, install_hooks=False): from art.megatron.gdn.operator import _prefix_tree_forward from art.megatron.lora import LoRA, SelfAttentionLinearProjLoRA - decoder = module(TransformerBlock) - decoder.config = full_recompute_config(2048, 40, False) - decoder.layers = torch.nn.ModuleList( - [torch.nn.Linear(1, 1).bfloat16() for _ in range(40)] - ) - decoder.num_layers_per_pipeline_rank = 40 + model = recompute_model(TransformerBlock, 2048, 40, False) layer = torch.nn.Module() layer.mlp = moe_layer gd = module(GatedDeltaNet) @@ -61,11 +56,7 @@ def rank_with_moe(moe_layer, *, install_hooks=False): torch.empty(1, 2048, dtype=torch.bfloat16) ) layer.self_attention = gd - decoder.layers[38] = layer - model: Any = torch.nn.Module() - model.config = decoder.config - model.decoder = decoder - model._preprocess = lambda: None + model.decoder.layers[38] = layer if install_hooks: from art.megatron.gdn.operator import install_gdn_island_hooks diff --git a/tests/unit/test_trainer_rank_tp_floor.py b/tests/unit/test_trainer_rank_tp_floor.py index 0500b75ef..c77c67524 100644 --- a/tests/unit/test_trainer_rank_tp_floor.py +++ b/tests/unit/test_trainer_rank_tp_floor.py @@ -11,7 +11,7 @@ import pytest import torch -from trainer_rank_test_support import fake_rank, full_recompute_config +from trainer_rank_test_support import fake_rank, recompute_model from art.trainer_rank import TrainerRank from art.trainer_rank._impl import _MemorySignature @@ -28,17 +28,7 @@ def tp_rank(layers=LAYERS, *, ffn=F, topology=TP4, sequence_parallel=True, **config): from megatron.core.transformer.transformer_block import TransformerBlock - block = TransformerBlock.__new__(TransformerBlock) - torch.nn.Module.__init__(block) - block.config = full_recompute_config(H, layers, sequence_parallel, **config) - block.layers = torch.nn.ModuleList( - [torch.nn.Linear(1, 1).bfloat16() for _ in range(layers)] - ) - block.num_layers_per_pipeline_rank = layers - model: Any = torch.nn.Module() - model.config = block.config - model.decoder = block - model._preprocess = lambda: None + model = recompute_model(TransformerBlock, H, layers, sequence_parallel, **config) r: Any = fake_rank(TrainerRank, [model], hidden_size=H, num_layers=layers) # Qwen3.8-27B: gated attention every fourth layer, GDN otherwise. r._geometry = replace( diff --git a/tests/unit/trainer_rank_test_support.py b/tests/unit/trainer_rank_test_support.py index 9d902c491..5a4a1576f 100644 --- a/tests/unit/trainer_rank_test_support.py +++ b/tests/unit/trainer_rank_test_support.py @@ -59,6 +59,25 @@ def full_recompute_config(hidden_size, num_layers, sequence_parallel, /, **confi ) +def recompute_model( + block_type, hidden_size, num_layers, sequence_parallel, /, *, layers=(), **config +): + block = block_type.__new__(block_type) + torch.nn.Module.__init__(block) + block.config = full_recompute_config( + hidden_size, num_layers, sequence_parallel, **config + ) + block.layers = torch.nn.ModuleList( + list(layers) + + [torch.nn.Linear(1, 1).bfloat16() for _ in range(num_layers - len(layers))] + ) + block.num_layers_per_pipeline_rank = num_layers + model: Any = torch.nn.Module() + model.config, model.decoder = block.config, block + model._preprocess = lambda: None + return model + + def checkpoint_runtime( model: torch.nn.Module | None = None, *, From 067b203f2470cf2388857c52df609a1b81109580 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 11:44:19 +0000 Subject: [PATCH 081/150] Simplify checkpoint memory test cases and config setup --- .../test_trainer_rank_checkpoint_memory.py | 71 ++++++++----------- tests/unit/trainer_rank_test_support.py | 18 ++--- 2 files changed, 34 insertions(+), 55 deletions(-) diff --git a/tests/unit/test_trainer_rank_checkpoint_memory.py b/tests/unit/test_trainer_rank_checkpoint_memory.py index 4b5c4fac7..45e0e4b93 100644 --- a/tests/unit/test_trainer_rank_checkpoint_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_memory.py @@ -122,28 +122,39 @@ def test_no_grad_enclosure_empty_and_unsupported(): @pytest.mark.parametrize( - "field,value", + "mode,field,value", [ - ("recompute_granularity", "selective"), - ("recompute_method", "block"), - ("recompute_num_layers", 2), - ("distribute_saved_activations", True), - ("sequence_parallel", True), - ("fp32_residual_connection", True), - ("cpu_offloading", True), - ("cuda_graph_impl", "local"), - ("params_dtype", torch.float32), - ("fp8", "hybrid"), - ("fp4", True), - ("num_layers", 39), - ("hidden_size", 1024), + ("grad", "recompute_granularity", "selective"), + ("grad", "recompute_method", "block"), + ("grad", "recompute_num_layers", 2), + ("grad", "distribute_saved_activations", True), + ("grad", "sequence_parallel", True), + ("grad", "fp32_residual_connection", True), + ("grad", "cpu_offloading", True), + ("grad", "cuda_graph_impl", "local"), + ("grad", "params_dtype", torch.float32), + ("grad", "fp8", "hybrid"), + ("grad", "fp4", True), + ("grad", "num_layers", 39), + ("grad", "hidden_size", 1024), + ("cold-grad", "recompute_num_layers", True), + ("cold-grad", "cpu_offloading", 0), + ("no-grad", "recompute_granularity", None), + ("no-grad", "recompute_granularity", "selective"), + ("no-grad", "recompute_method", "block"), + ("no-grad", "recompute_num_layers", True), + ("no-grad", "cpu_offloading", True), + ("no-grad", "params_dtype", torch.float32), ], + ids=str, ) -def test_actual_config_revalidated(field, value): +def test_actual_config_revalidated(mode, field, value): r = rank() - assert r._checkpoint_memory_floor(((10, True),))[0] > 0 + if mode == "grad": + assert r._checkpoint_memory_floor(((10, True),))[0] > 0 setattr(r.runtime.model[0].decoder.config, field, value) - assert r._checkpoint_memory_floor(((10, True),)) == (0, 0) + groups = ((11, False),) if mode == "no-grad" else ((10, True),) + assert r._checkpoint_memory_floor(groups) == (0, 0) @pytest.mark.parametrize("axis", [1, 3]) @@ -329,15 +340,6 @@ def test_optimistic_split_profile_cliff_preserves_checkpoint_floor(): assert cost.retained == int((full.output_bytes + retained) * 1.1) -@pytest.mark.parametrize( - "field,value", [("recompute_num_layers", True), ("cpu_offloading", 0)] -) -def test_malformed_flag_types_do_not_claim_supported_schedule(field, value): - r = rank() - setattr(r.runtime.model[0].decoder.config, field, value) - assert r._checkpoint_memory_floor(((10, True),)) == (0, 0) - - @pytest.mark.parametrize("profile_rate", [None, 1, 1_000_000]) def test_no_grad_enclosure_exact_lower_and_profile(profile_rate): r = rank() @@ -445,20 +447,3 @@ def test_reference_prefix_search_agrees_with_mixed_demand(fits): == r._plan_cost(reference_plan).required ) assert not r._memory_profiles - - -@pytest.mark.parametrize( - "field,value", - [ - ("recompute_granularity", None), - ("recompute_granularity", "selective"), - ("recompute_method", "block"), - ("recompute_num_layers", True), - ("cpu_offloading", True), - ("params_dtype", torch.float32), - ], -) -def test_no_grad_enclosure_config_guard(field, value): - r = rank() - setattr(r.runtime.model[0].decoder.config, field, value) - assert r._checkpoint_memory_floor(((11, False),)) == (0, 0) diff --git a/tests/unit/trainer_rank_test_support.py b/tests/unit/trainer_rank_test_support.py index 5a4a1576f..b5cbfda55 100644 --- a/tests/unit/trainer_rank_test_support.py +++ b/tests/unit/trainer_rank_test_support.py @@ -39,8 +39,12 @@ def fake_rank(rank_type: type["TrainerRank"], model, **provider) -> "TrainerRank return rank_type(runtime) -def full_recompute_config(hidden_size, num_layers, sequence_parallel, /, **config): - return SimpleNamespace( +def recompute_model( + block_type, hidden_size, num_layers, sequence_parallel, /, *, layers=(), **config +): + block = block_type.__new__(block_type) + torch.nn.Module.__init__(block) + block.config = SimpleNamespace( hidden_size=hidden_size, num_layers=num_layers, padded_vocab_size=32, @@ -57,16 +61,6 @@ def full_recompute_config(hidden_size, num_layers, sequence_parallel, /, **confi fp4=None, **config, ) - - -def recompute_model( - block_type, hidden_size, num_layers, sequence_parallel, /, *, layers=(), **config -): - block = block_type.__new__(block_type) - torch.nn.Module.__init__(block) - block.config = full_recompute_config( - hidden_size, num_layers, sequence_parallel, **config - ) block.layers = torch.nn.ModuleList( list(layers) + [torch.nn.Linear(1, 1).bfloat16() for _ in range(num_layers - len(layers))] From f89d83b593b321c0b26d2941b6bfb867d74b2028 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 12:09:29 +0000 Subject: [PATCH 082/150] Consolidate checkpoint profile and TP pricing test scenarios --- ...trainer_rank_checkpoint_gradient_memory.py | 2 + .../test_trainer_rank_checkpoint_memory.py | 24 --------- tests/unit/test_trainer_rank_tp_floor.py | 50 +++++++------------ 3 files changed, 19 insertions(+), 57 deletions(-) diff --git a/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py b/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py index d6d9ded56..f67dd0b78 100644 --- a/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py @@ -157,6 +157,8 @@ def test_lower_bound_profile_cliff_preserves_separate_peak_component(): lower = r._split_chunk_lower_cost( req, tuple(q.input_tokens for q in req), checkpoint=Unset ) + retained = 128 * 40 * 2048 * 2 + assert lower.retained == int((full.output_bytes + retained) * 1.1) assert lower.checkpoint_input_gradient == 128 * 40 * 4096 assert lower.required <= r._plan_cost(full).required assert lower.checkpoint_retained == full.output_bytes + 128 * 40 * 4096 diff --git a/tests/unit/test_trainer_rank_checkpoint_memory.py b/tests/unit/test_trainer_rank_checkpoint_memory.py index 45e0e4b93..59f5fb34d 100644 --- a/tests/unit/test_trainer_rank_checkpoint_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_memory.py @@ -316,30 +316,6 @@ def test_split_keeps_complete_order_and_checks_each_new_subforward(): assert all(a is b for a, b in zip(restored, req, strict=True)) -def test_optimistic_split_profile_cliff_preserves_checkpoint_floor(): - r = rank() - req = [ - ForwardInput( - input_tokens=torch.arange(128), - target_tokens=torch.arange(128), - no_grad=False, - ) - for _ in range(16) - ] - full = r._plan_flat_forward(req, memory_minimal=True) - r._memory_profiles[full.signature] = _MemoryProfile( - bytes_per_token=1, - packed_tokens=256, - logical_per_packed=1, - retained_compute_bytes_per_token=1, - ) - cost = r._split_chunk_lower_cost( - req, tuple(x.input_tokens for x in req), checkpoint=Unset - ) - retained = 128 * 40 * 2048 * 2 - assert cost.retained == int((full.output_bytes + retained) * 1.1) - - @pytest.mark.parametrize("profile_rate", [None, 1, 1_000_000]) def test_no_grad_enclosure_exact_lower_and_profile(profile_rate): r = rank() diff --git a/tests/unit/test_trainer_rank_tp_floor.py b/tests/unit/test_trainer_rank_tp_floor.py index c77c67524..b79d6acaa 100644 --- a/tests/unit/test_trainer_rank_tp_floor.py +++ b/tests/unit/test_trainer_rank_tp_floor.py @@ -102,42 +102,26 @@ def test_rows_are_sharded_with_ceiling_and_only_gradient_groups_save_them(): @pytest.mark.parametrize( - "case", + "case,kwargs", [ - "tp2", - "tp8", - "cp2", - "pp2", - "no_sequence_parallel", - "sequence_parallel_at_tp1", - "selective_recompute", - "moe", - "moe_geometry", - "replicated_qkv", - "missing_attention_geometry", - "missing_conv_kernel", - "shallow", - "wide_ffn", + ("tp2", dict(topology=(1, 2, 1, 1))), + ("tp8", dict(topology=(1, 8, 1, 1))), + ("cp2", dict(topology=(1, 4, 2, 1))), + ("pp2", dict(topology=(1, 4, 1, 2))), + ("no_sequence_parallel", dict(sequence_parallel=False)), + ("sequence_parallel_at_tp1", dict(topology=(1, 1, 1, 1))), + ("selective_recompute", dict()), + ("moe", dict()), + ("moe_geometry", dict()), + ("replicated_qkv", dict()), + ("missing_attention_geometry", dict()), + ("missing_conv_kernel", dict()), + ("shallow", dict(layers=48)), + ("wide_ffn", dict(ffn=4 * F)), ], ) -def test_unproven_shapes_keep_todays_pricing(case): - shapes = { - "tp2": dict(topology=(1, 2, 1, 1)), - "tp8": dict(topology=(1, 8, 1, 1)), - "cp2": dict(topology=(1, 4, 2, 1)), - "pp2": dict(topology=(1, 4, 1, 2)), - "no_sequence_parallel": dict(sequence_parallel=False), - "sequence_parallel_at_tp1": dict(topology=(1, 1, 1, 1)), - "selective_recompute": dict(), - "moe": dict(), - "moe_geometry": dict(), - "replicated_qkv": dict(), - "missing_attention_geometry": dict(), - "missing_conv_kernel": dict(), - "shallow": dict(layers=48), - "wide_ffn": dict(ffn=4 * F), - } - r = tp_rank(**shapes[case]) +def test_unproven_shapes_keep_todays_pricing(case, kwargs): + r = tp_rank(**kwargs) if case == "selective_recompute": r.runtime.model[0].decoder.config.recompute_granularity = "selective" if case == "moe": From 63b91ca90566909447114689880298bdb1988f97 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 12:17:59 +0000 Subject: [PATCH 083/150] Normalize packed GPT-OSS exports in the native builder --- src/art/trainer_rank/_lora_export.py | 27 +++++++++++++++++++++++++++ 1 file changed, 27 insertions(+) diff --git a/src/art/trainer_rank/_lora_export.py b/src/art/trainer_rank/_lora_export.py index 0bb01707e..c73892be0 100644 --- a/src/art/trainer_rank/_lora_export.py +++ b/src/art/trainer_rank/_lora_export.py @@ -215,6 +215,33 @@ def _build_vllm_lora_tensors_from_inputs( packed_expert_metadata=inputs.packed_expert_metadata, packed_expert_tensors_by_owner_key=inputs.packed_expert_tensors_by_owner_key, ) + interleaved_keys = frozenset( + meta.key + for meta in inputs.packed_expert_metadata + if meta.pack_layout == "interleaved_gate_up_rank_major_expert_cols" + ) + if getattr(inputs.handler, "key", None) == "gpt_oss_moe" and interleaved_keys: + from art.megatron.model_support.handlers.gpt_oss import ( + _gpt_oss_padding_sizes_from_adapter_config, + _trim_gpt_oss_interleaved_gate_up_last, + ) + + sizes = _gpt_oss_padding_sizes_from_adapter_config(inputs.adapter_config) + assert sizes is not None + _, _, logical, internal = sizes + for key in interleaved_keys: + tensor = merged_tensors[key] + if tensor.ndim != 2 or tensor.shape[0] not in { + 2 * logical, + 2 * internal, + }: + raise ValueError("GPT-OSS packed gate/up LoRA has an invalid shape") + if tensor.shape[0] != 2 * logical: + # Packed producers interleave gate/up before trimming; the + # regular handler's half-split trim is for canonical tensors. + merged_tensors[key] = _trim_gpt_oss_interleaved_gate_up_last( + tensor.T, logical=logical, internal=internal + ).T.contiguous() return inputs.handler.to_vllm_lora_tensors( merged_tensors, adapter_config=inputs.adapter_config, From 46c911011c1943a3756e1faec558d07b445f18c2 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 12:54:10 +0000 Subject: [PATCH 084/150] Compose native export correction and checkpoint test reuse --- .../model_support/handlers/gpt_oss.py | 6 +- .../megatron/lora/test_lora_disk_codecs.py | 63 +++++++++---------- 2 files changed, 31 insertions(+), 38 deletions(-) diff --git a/src/art/megatron/model_support/handlers/gpt_oss.py b/src/art/megatron/model_support/handlers/gpt_oss.py index 320d223f7..9c169a097 100644 --- a/src/art/megatron/model_support/handlers/gpt_oss.py +++ b/src/art/megatron/model_support/handlers/gpt_oss.py @@ -1299,11 +1299,11 @@ def _trim_gpt_oss_lora_for_vllm( if key.endswith(".base_layer.lora_B.weight"): if int(tensor.shape[0]) == 2 * logical_ffn: return tensor.contiguous() - return _trim_gpt_oss_gate_up_dim0( - tensor, + return _trim_gpt_oss_interleaved_gate_up_last( + tensor.T, logical=logical_ffn, internal=internal_ffn, - ) + ).T.contiguous() if key.endswith(".lora_A.weight"): return _trim_dim_right(tensor, dim=-1, size=logical_ffn) if key.endswith(".lora_B.weight"): diff --git a/tests/integration/megatron/lora/test_lora_disk_codecs.py b/tests/integration/megatron/lora/test_lora_disk_codecs.py index 800143467..4e3db7039 100644 --- a/tests/integration/megatron/lora/test_lora_disk_codecs.py +++ b/tests/integration/megatron/lora/test_lora_disk_codecs.py @@ -1584,12 +1584,9 @@ def synchronize(error: BaseException | None, _phase: str, group: object) -> None assert synchronized_groups == [failure_group] -def test_trainer_rank_publishes_named_checkpoint_slot_without_mutating_base( - tmp_path: Path, -): - prefix = "base_model.model.model.layers.0.self_attn.q_proj" - lora = LoRA(prefix, 3, 4, 2, 2, torch.float32, torch.device("cpu")) - baseline = (lora.A_T.detach().clone(), lora.B_T.detach().clone()) +def _named_lora_checkpoint( + prefix: str, lora: LoRA +) -> tuple[TrainerRank, dict[str, torch.Tensor], dict[str, Any]]: adapter = { f"{prefix}.lora_A.weight": torch.arange(6, dtype=torch.float32).reshape(2, 3), f"{prefix}.lora_B.weight": torch.arange(8, dtype=torch.float32).reshape(4, 2), @@ -1615,6 +1612,16 @@ def test_trainer_rank_publishes_named_checkpoint_slot_without_mutating_base( tuple(trainer._iter_slot_parameters(trainer._slot_ref("student"))), cast(_AdapterConfig, config), ) + return trainer, adapter, config + + +def test_trainer_rank_publishes_named_checkpoint_slot_without_mutating_base( + tmp_path: Path, +): + prefix = "base_model.model.model.layers.0.self_attn.q_proj" + lora = LoRA(prefix, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + baseline = (lora.A_T.detach().clone(), lora.B_T.detach().clone()) + trainer, adapter, config = _named_lora_checkpoint(prefix, lora) output_dir = tmp_path / "checkpoint" assert trainer.export_lora(str(output_dir), "student") == 0 @@ -1631,31 +1638,7 @@ def test_trainer_rank_publishes_named_checkpoint_slot_without_mutating_base( def test_prepared_lora_export_is_immutable_and_abortable(tmp_path: Path): prefix = "base_model.model.model.layers.0.self_attn.q_proj" lora = LoRA(prefix, 3, 4, 2, 2, torch.float32, torch.device("cpu")) - adapter = { - f"{prefix}.lora_A.weight": torch.arange(6, dtype=torch.float32).reshape(2, 3), - f"{prefix}.lora_B.weight": torch.arange(8, dtype=torch.float32).reshape(4, 2), - } - trainer = TrainerRank.__new__(TrainerRank) - trainer.runtime = SimpleNamespace( - model=[lora], - model_support_handler=DEFAULT_DENSE_HANDLER, - rank=0, - world_size=1, - ) - trainer._slot_stack = [] - trainer._pending_slot_graphs = {} - trainer._checkpoint_slots = {} - trainer._skipped_forward_waves = {} - trainer._snapshot_checkpoint_names = set() - trainer._checkpoint_prefetch_sources = {} - trainer._checkpoint_prefetch_lock = threading.Lock() - trainer._checkpoint_mutation_lock = threading.RLock() - config = _config("Qwen/Qwen3-8B", rank=2, alpha=2) - assert trainer._load_checkpoint_slot("student", adapter, alpha=2) == 1 - trainer._checkpoint_slots["student"] = _CheckpointSlot( - tuple(trainer._iter_slot_parameters(trainer._slot_ref("student"))), - cast(_AdapterConfig, config), - ) + trainer, adapter, config = _named_lora_checkpoint(prefix, lora) revision, capture_timings = trainer._prepare_lora_export( "first", "student", owner_id="owner" @@ -1812,9 +1795,11 @@ def test_direct_3d_packed_expert_publish_matches_handler_vllm_exactly( ) +@pytest.mark.parametrize("internal_ffn", [128, 1024]) def test_direct_gpt_oss_packed_expert_publish_matches_handler_vllm_exactly( tmp_path: Path, monkeypatch, + internal_ffn: int, ): monkeypatch.setattr(lora_module.ps, "get_expert_model_parallel_rank", lambda: 0) monkeypatch.setattr(lora_module.ps, "get_expert_data_parallel_rank", lambda: 0) @@ -1834,7 +1819,7 @@ def test_direct_gpt_oss_packed_expert_publish_matches_handler_vllm_exactly( gate_up_lora = LoRA( adapter_model_prefix=f"{group_prefix}.{{expert}}.gate_up_proj", in_features=hidden, - out_features=2 * intermediate, + out_features=2 * internal_ffn, rank=rank, alpha=rank, dtype=torch.float32, @@ -1843,7 +1828,7 @@ def test_direct_gpt_oss_packed_expert_publish_matches_handler_vllm_exactly( ) down_lora = LoRA( adapter_model_prefix=f"{group_prefix}.{{expert}}.down_proj", - in_features=intermediate, + in_features=internal_ffn, out_features=hidden, rank=rank, alpha=rank, @@ -1857,10 +1842,18 @@ def test_direct_gpt_oss_packed_expert_publish_matches_handler_vllm_exactly( full[f"{expert_prefix}.gate_up_proj.lora_A.weight"].T ) gate_up_lora.B_T.data[expert].copy_( - full[f"{expert_prefix}.gate_up_proj.lora_B.weight"].T + torch.nn.functional.pad( + full[f"{expert_prefix}.gate_up_proj.lora_B.weight"].T.reshape( + rank, 2, intermediate + ), + (0, internal_ffn - intermediate), + ).flatten(1) ) down_lora.A_T.data[expert].copy_( - full[f"{expert_prefix}.down_proj.lora_A.weight"].T + torch.nn.functional.pad( + full[f"{expert_prefix}.down_proj.lora_A.weight"].T, + (0, 0, 0, internal_ffn - intermediate), + ) ) down_lora.B_T.data[expert].copy_( full[f"{expert_prefix}.down_proj.lora_B.weight"].T From 9bfdc4e941c936a0d4109521f46e2ac5239cf292 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 12:56:46 +0000 Subject: [PATCH 085/150] Remove unused GPT-OSS interleaved padding helper --- .../model_support/handlers/gpt_oss.py | 28 ------------------- 1 file changed, 28 deletions(-) diff --git a/src/art/megatron/model_support/handlers/gpt_oss.py b/src/art/megatron/model_support/handlers/gpt_oss.py index 9c169a097..9f88e2014 100644 --- a/src/art/megatron/model_support/handlers/gpt_oss.py +++ b/src/art/megatron/model_support/handlers/gpt_oss.py @@ -172,34 +172,6 @@ def _gate_up_from_etp_shard_order(tensor: torch.Tensor, etp_size: int) -> torch. ) -def _pad_gpt_oss_interleaved_gate_up_last( - tensor: torch.Tensor, - *, - logical: int, - internal: int, -) -> torch.Tensor: - if logical == internal: - return tensor.contiguous() - if int(tensor.shape[-1]) != 2 * logical: - raise RuntimeError( - "Expected GPT OSS interleaved gate/up logical dim " - f"{2 * logical}, got {tuple(tensor.shape)}" - ) - gate = tensor[..., 0::2] - up = tensor[..., 1::2] - return ( - torch.stack( - [ - _pad_dim_right(gate, dim=-1, size=internal), - _pad_dim_right(up, dim=-1, size=internal), - ], - dim=-1, - ) - .flatten(-2) - .contiguous() - ) - - def _trim_gpt_oss_interleaved_gate_up_last( tensor: torch.Tensor, *, From b745db5a99d372041eda6795a0ffe62a17b68046 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 13:34:25 +0000 Subject: [PATCH 086/150] Compose checkpoint export synchronization and accepted private simplifications --- .github/workflows/prek.yml | 3 +- .../model_support/handlers/gpt_oss.py | 19 +-- src/art/trainer_rank/_impl.py | 4 +- src/art/trainer_rank/_lora_export.py | 54 ++++---- .../megatron/lora/test_lora_disk_codecs.py | 116 ++++++++++++++++++ 5 files changed, 148 insertions(+), 48 deletions(-) diff --git a/.github/workflows/prek.yml b/.github/workflows/prek.yml index 0e3b3266f..bc7145649 100644 --- a/.github/workflows/prek.yml +++ b/.github/workflows/prek.yml @@ -213,7 +213,7 @@ jobs: - name: Run Megatron lightweight tests run: | megatron_runtime/.venv/bin/python -c "import megatron.core.packed_seq_params" - megatron_runtime/.venv/bin/python -m pytest --nbval --current-env --tb=short \ + megatron_runtime/.venv/bin/python -m pytest -v --nbval --current-env --tb=short \ tests/unit/test_megatron_reference_logprobs.py \ tests/unit/test_preprocessing_tokenize.py::test_gemma4_normalizes_json_tool_arguments_for_mapping_template \ tests/unit/test_moe_routing_replay.py \ @@ -249,6 +249,7 @@ jobs: tests/unit/test_trainer_rank_cache_recovery.py::test_dense_cp_exact_demand_fits_after_recovery \ tests/unit/test_trainer_rank_cache_recovery.py::test_dense_cp_exact_demand_refuses_after_recovery \ tests/acceptance/trainer_rank_planner \ + tests/integration/megatron/lora/test_lora_disk_codecs.py::test_export_preparation_failure_is_collective_and_retryable \ tests/integration/megatron/test_sft_packing.py::test_sft_packing_preserves_training_targets \ tests/integration/megatron/model_support/test_dispatcher_graph_retention.py \ tests/integration/megatron/gdn_shared_prefix/test_gdn_planner_runtime_model.py \ diff --git a/src/art/megatron/model_support/handlers/gpt_oss.py b/src/art/megatron/model_support/handlers/gpt_oss.py index 9f88e2014..d03823d0e 100644 --- a/src/art/megatron/model_support/handlers/gpt_oss.py +++ b/src/art/megatron/model_support/handlers/gpt_oss.py @@ -296,7 +296,7 @@ def _gpt_oss_config_dict(base_model_name_or_path: str) -> dict[str, Any]: def _gpt_oss_padding_sizes_from_adapter_config( adapter_config: dict[str, Any], -) -> tuple[int, int, int, int] | None: +) -> tuple[int, int, int, int]: base_model = adapter_config.get("base_model_name_or_path") if not isinstance(base_model, str) or not base_model: raise RuntimeError("GPT OSS LoRA conversion requires base_model_name_or_path") @@ -1244,8 +1244,6 @@ def _trim_gpt_oss_lora_for_vllm( adapter_config: dict[str, Any], ) -> torch.Tensor: sizes = _gpt_oss_padding_sizes_from_adapter_config(adapter_config) - if sizes is None: - return tensor.contiguous() logical_hidden, internal_hidden, logical_ffn, internal_ffn = sizes match = _ART_MOE_EXPERT_KEY_RE.match(key) if match is not None: @@ -1290,8 +1288,6 @@ def _pad_gpt_oss_lora_from_vllm( adapter_config: dict[str, Any], ) -> torch.Tensor: sizes = _gpt_oss_padding_sizes_from_adapter_config(adapter_config) - if sizes is None: - return tensor.contiguous() _logical_hidden, internal_hidden, _logical_ffn, internal_ffn = sizes match = _ART_MOE_EXPERT_KEY_RE.match(key) if match is not None: @@ -1309,19 +1305,6 @@ def _pad_gpt_oss_lora_from_vllm( return _pad_dim_right(tensor, dim=-1, size=internal_ffn) if module == "down_proj" and lora == "lora_B": return _pad_dim_right(tensor, dim=0, size=internal_hidden) - if _ART_PACKED_MOE_KEY_RE.match(key): - if key.endswith(".base_layer.lora_A.weight"): - return _pad_dim_right(tensor, dim=-1, size=internal_hidden) - if key.endswith(".base_layer.lora_B.weight"): - return _pad_gpt_oss_gate_up_dim0( - tensor, - logical=tensor.shape[0] // 2, - internal=internal_ffn, - ) - if key.endswith(".lora_A.weight"): - return _pad_dim_right(tensor, dim=-1, size=internal_ffn) - if key.endswith(".lora_B.weight"): - return _pad_dim_right(tensor, dim=0, size=internal_hidden) return tensor.contiguous() diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index ba0bd9fde..dcd4e429a 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -113,7 +113,7 @@ _FinalizedSave, _PreparedSave, ) - from art.trainer_rank._lora_export import _PreparedLoraExport + from art.trainer_rank._lora_export import _VllmLoraPublishInputs from ._heads import ModuleHandle @@ -2080,7 +2080,7 @@ def memory_field(name: str, default: Any = None) -> Any: self._slot_stack: list[LoRASlotRef] = [] self._checkpoint_slots: dict[str, _CheckpointSlot] = {} self._snapshot_checkpoint_names: set[str] = set() - self._prepared_lora_exports: dict[str, tuple[str, _PreparedLoraExport]] = {} + self._prepared_lora_exports: dict[str, tuple[str, _VllmLoraPublishInputs]] = {} self._checkpoint_prefetches: dict[str, Future[PreparedCheckpoint]] = {} self._checkpoint_prefetch_sources: dict[str, str] = {} self._checkpoint_prefetch_lock = threading.Lock() diff --git a/src/art/trainer_rank/_lora_export.py b/src/art/trainer_rank/_lora_export.py index c73892be0..0913389a4 100644 --- a/src/art/trainer_rank/_lora_export.py +++ b/src/art/trainer_rank/_lora_export.py @@ -18,11 +18,6 @@ _K = TypeVar("_K") -@dataclass(frozen=True) -class _PreparedLoraExport: - inputs: _VllmLoraPublishInputs - - @dataclass(frozen=True) class _VllmLoraPublishPlan: rank: int @@ -161,14 +156,12 @@ def _prepare_vllm_lora_publish( packed_expert_groups=packed_expert_groups, slot_ref=slot_ref, ) - all_packed_metadata = lora_publish._canonical_global_metadata(local_packed_metadata) - all_metadata = lora_publish._canonical_global_metadata(local_metadata) return _VllmLoraPublishPlan( rank=rank, device=device, - metadata=all_metadata, + metadata=local_metadata, local_tensors=local_tensors, - packed_expert_metadata=all_packed_metadata, + packed_expert_metadata=local_packed_metadata, local_packed_expert_tensors=local_packed_tensors, handler=handler, adapter_config=dict(adapter_config), @@ -223,11 +216,9 @@ def _build_vllm_lora_tensors_from_inputs( if getattr(inputs.handler, "key", None) == "gpt_oss_moe" and interleaved_keys: from art.megatron.model_support.handlers.gpt_oss import ( _gpt_oss_padding_sizes_from_adapter_config, - _trim_gpt_oss_interleaved_gate_up_last, ) sizes = _gpt_oss_padding_sizes_from_adapter_config(inputs.adapter_config) - assert sizes is not None _, _, logical, internal = sizes for key in interleaved_keys: tensor = merged_tensors[key] @@ -236,12 +227,6 @@ def _build_vllm_lora_tensors_from_inputs( 2 * internal, }: raise ValueError("GPT-OSS packed gate/up LoRA has an invalid shape") - if tensor.shape[0] != 2 * logical: - # Packed producers interleave gate/up before trimming; the - # regular handler's half-split trim is for canonical tensors. - merged_tensors[key] = _trim_gpt_oss_interleaved_gate_up_last( - tensor.T, logical=logical, internal=internal - ).T.contiguous() return inputs.handler.to_vllm_lora_tensors( merged_tensors, adapter_config=inputs.adapter_config, @@ -253,7 +238,7 @@ def _capture_lora_publish_inputs( checkpoint_name: str, adapter_config: dict[str, object], group: torch.distributed.ProcessGroup | None, -) -> tuple[_PreparedLoraExport | None, dict[str, float]]: +) -> tuple[_VllmLoraPublishInputs | None, dict[str, float]]: from art.trainer_rank import _checkpoint timings: dict[str, float] = {} @@ -282,6 +267,23 @@ def _capture_lora_publish_inputs( "plan LoRA publish", group, ) + # Every rank must finish local collection before metadata or tensor exchange. + from art.megatron.weights import lora_publish + + packed_metadata = _checkpoint._phase( + lambda: lora_publish._canonical_global_metadata(plan.packed_expert_metadata), + "gather packed LoRA metadata", + group, + ) + plan = _checkpoint._phase( + lambda: replace( + plan, + packed_expert_metadata=packed_metadata, + metadata=lora_publish._canonical_global_metadata(plan.metadata), + ), + "gather LoRA metadata", + group, + ) timings["plan_collect"] = time.monotonic() - started started = time.monotonic() @@ -290,7 +292,7 @@ def _capture_lora_publish_inputs( started = time.monotonic() - def stage() -> _PreparedLoraExport | None: + def stage() -> _VllmLoraPublishInputs | None: if inputs is not None: stager = _PinnedCpuStager() staged = replace( @@ -303,7 +305,7 @@ def stage() -> _PreparedLoraExport | None: ), ) stager.finish() - return _PreparedLoraExport(staged) + return staged return None prepared = _checkpoint._phase(stage, "stage LoRA publish tensors", group) @@ -312,14 +314,12 @@ def stage() -> _PreparedLoraExport | None: def _save_lora_publish_inputs( - output_dir: str, prepared: _PreparedLoraExport + output_dir: str, prepared: _VllmLoraPublishInputs ) -> dict[str, float]: from art.megatron.model_support.lora_disk import save_vllm_lora_tensors started = time.monotonic() - vllm_tensors, published_config = _build_vllm_lora_tensors_from_inputs( - prepared.inputs - ) + vllm_tensors, published_config = _build_vllm_lora_tensors_from_inputs(prepared) timings = {"convert": time.monotonic() - started} started = time.monotonic() save_vllm_lora_tensors(output_dir, vllm_tensors, published_config) @@ -338,7 +338,7 @@ def prepare_lora_export( started = time.monotonic() group = _checkpoint._ensure_group(trainer) - snapshots: dict[str, tuple[str, _PreparedLoraExport]] = getattr( + snapshots: dict[str, tuple[str, _VllmLoraPublishInputs]] = getattr( trainer, "_prepared_lora_exports", {} ) duplicate = ( @@ -387,7 +387,7 @@ def prepare_lora_export( def finish_lora_export( trainer: TrainerRank, export_id: str, output_dir: str, *, owner_id: str ) -> dict[str, float]: - snapshots: dict[str, tuple[str, _PreparedLoraExport]] = getattr( + snapshots: dict[str, tuple[str, _VllmLoraPublishInputs]] = getattr( trainer, "_prepared_lora_exports", {} ) try: @@ -401,7 +401,7 @@ def finish_lora_export( def abort_lora_export(trainer: TrainerRank, export_id: str, *, owner_id: str) -> None: - snapshots: dict[str, tuple[str, _PreparedLoraExport]] = getattr( + snapshots: dict[str, tuple[str, _VllmLoraPublishInputs]] = getattr( trainer, "_prepared_lora_exports", {} ) if (prepared := snapshots.get(export_id)) is not None and prepared[0] == owner_id: diff --git a/tests/integration/megatron/lora/test_lora_disk_codecs.py b/tests/integration/megatron/lora/test_lora_disk_codecs.py index 4e3db7039..0978768ca 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 datetime import timedelta import json import os from pathlib import Path @@ -10,6 +12,9 @@ import pytest from safetensors.torch import load_file, save_file import torch +import torch.distributed as dist + +from tests.unit.trainer_rank_test_support import gloo_group, spawn_and_join pytest.importorskip("megatron.bridge.models.gpt_provider") @@ -1584,6 +1589,117 @@ def synchronize(error: BaseException | None, _phase: str, group: object) -> None assert synchronized_groups == [failure_group] +def _export_preparation_failure_worker( + rank: int, init_method: str, failure_site: str +) -> None: + with ( + gloo_group(rank, init_method, timeout=10), + pytest.MonkeyPatch.context() as monkeypatch, + ): + monkeypatch.setattr(lora_module.ps, "get_data_parallel_rank", lambda **_: rank) + monkeypatch.setattr(lora_module.ps, "get_tensor_model_parallel_rank", lambda: 0) + monkeypatch.setattr( + lora_module.ps, "get_tensor_model_parallel_world_size", lambda: 1 + ) + monkeypatch.setattr(lora_module.ps, "get_expert_model_parallel_rank", lambda: 0) + prefix = "base_model.model.model.layers.0.self_attn.q_proj" + lora = LoRA(prefix, 3, 4, 2, 2, torch.float32, torch.device("cpu")) + trainer, adapter, _config = _named_lora_checkpoint(prefix, lora) + trainer.runtime.rank, trainer.runtime.world_size = rank, 2 + group = dist.new_group(backend="gloo", timeout=timedelta(seconds=10)) + trainer._checkpoint_process_group = group + trainer._checkpoint_finalize_process_group = group + failure = ( + asyncio.CancelledError("injected local collection cancellation") + if failure_site == "packed" + else RuntimeError("injected export preparation failure") + ) + metadata_calls: list[object] = [] + exchanges: list[object] = [] + canonical = lora_publish._canonical_global_metadata + exchange = _lora_export._exchange_vllm_lora_publish + + def metadata(local: list[Any]) -> list[Any]: + metadata_calls.append(local) + result = canonical(local) + if rank == 0 and failure_site == "metadata": + raise failure + return result + + def exchange_tensors(plan: _lora_export._VllmLoraPublishPlan): + exchanges.append(plan) + return exchange(plan) + + with pytest.MonkeyPatch.context() as inject: + if failure_site != "metadata": + collector = ( + "collect_local_lora_entries" + if failure_site == "dense" + else "collect_local_packed_expert_entries" + ) + collect = getattr(lora_publish, collector) + + def collect_or_fail(*args: Any, **kwargs: Any): + result = collect(*args, **kwargs) + if rank == 0: + raise failure + return result + + inject.setattr(lora_publish, collector, collect_or_fail) + inject.setattr(lora_publish, "_canonical_global_metadata", metadata) + inject.setattr( + _lora_export, "_exchange_vllm_lora_publish", exchange_tensors + ) + with pytest.raises(BaseException, match="injected") as caught: + trainer._prepare_lora_export("retry", "student", owner_id="owner") + if rank == 0: + assert caught.value is failure + else: + assert isinstance(caught.value, RuntimeError) + assert "Another rank failed" in str(caught.value) + assert len(metadata_calls) == (1 if failure_site == "metadata" else 0) + assert exchanges == [] + assert not getattr(trainer, "_prepared_lora_exports", {}) + + revision, timings = trainer._prepare_lora_export( + "retry", "student", owner_id="owner" + ) + assert revision == 0 + assert set(timings) == { + "slot_validation", + "runtime_validation", + "plan_collect", + "exchange", + "d2h", + } + if rank == 0: + owner, prepared = trainer._prepared_lora_exports["retry"] + assert owner == "owner" + _assert_tensors_equal( + _lora_export._build_vllm_lora_tensors_from_inputs(prepared)[0], + adapter, + ) + trainer._abort_lora_export("retry", owner_id="owner") + assert not getattr(trainer, "_prepared_lora_exports", {}) + for reuse_group in (group, None): + completed = torch.tensor(1) + dist.all_reduce(completed, group=reuse_group) + assert completed.item() == 2 + + +@pytest.mark.parametrize("failure_site", ("dense", "packed", "metadata")) +def test_export_preparation_failure_is_collective_and_retryable( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, failure_site: str +) -> None: + monkeypatch.setenv("CUDA_VISIBLE_DEVICES", "") + spawn_and_join( + _export_preparation_failure_worker, + args=(f"file://{tmp_path / 'export'}", failure_site), + timeout=90, + failure=f"collective export {failure_site} failure test hung", + ) + + def _named_lora_checkpoint( prefix: str, lora: LoRA ) -> tuple[TrainerRank, dict[str, torch.Tensor], dict[str, Any]]: From 3777c0bcc1f13b666483aca7161bfba31cf7336c Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 13:37:17 +0000 Subject: [PATCH 087/150] Share LoRA publication runtime validation --- src/art/megatron/weights/lora_publish.py | 37 ++++++++++++++---------- src/art/trainer_rank/_lora_export.py | 18 +----------- 2 files changed, 23 insertions(+), 32 deletions(-) diff --git a/src/art/megatron/weights/lora_publish.py b/src/art/megatron/weights/lora_publish.py index 63044560e..b9bfc442e 100644 --- a/src/art/megatron/weights/lora_publish.py +++ b/src/art/megatron/weights/lora_publish.py @@ -285,6 +285,27 @@ def _rank_and_device() -> tuple[int, torch.device]: ) +def _validate_vllm_lora_publish_runtime( + rank: int, world_size: int +) -> tuple[int, torch.device]: + 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: + raise RuntimeError( + "Non-distributed LoRA publish requires rank=0 and world_size=1, " + f"got rank={rank} world_size={world_size}" + ) + rank = 0 + return rank, device + + def _metadata_by_owner_dtype( metadata: Sequence[Any], ) -> dict[tuple[int, str], list[Any]]: @@ -646,21 +667,7 @@ def _build_merged_lora_tensors_from_model( world_size: int, slot_ref: LoRASlotRef | None = None, ) -> dict[str, torch.Tensor] | None: - 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: - raise RuntimeError( - "Non-distributed LoRA publish requires rank=0 and world_size=1, " - f"got rank={rank} world_size={world_size}" - ) - rank = 0 + rank, device = _validate_vllm_lora_publish_runtime(rank, world_size) packed_expert_groups = tuple(handler.expert_packed_lora_groups()) local_tensors, local_metadata = collect_local_lora_entries( model, diff --git a/src/art/trainer_rank/_lora_export.py b/src/art/trainer_rank/_lora_export.py index 0913389a4..18973593e 100644 --- a/src/art/trainer_rank/_lora_export.py +++ b/src/art/trainer_rank/_lora_export.py @@ -101,23 +101,7 @@ def _validate_vllm_lora_publish_runtime( ) -> tuple[int, torch.device]: from art.megatron.weights import lora_publish - actual_rank, device = lora_publish._rank_and_device() - if lora_publish._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})" - ) - else: - if 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 - return rank, device + return lora_publish._validate_vllm_lora_publish_runtime(rank, world_size) def _prepare_vllm_lora_publish( From 0fe6f91bb47ec878e1af053918d349668d67e384 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 13:59:08 +0000 Subject: [PATCH 088/150] Require validated runtime for LoRA publication preparation --- src/art/trainer_rank/_lora_export.py | 12 ++---------- 1 file changed, 2 insertions(+), 10 deletions(-) diff --git a/src/art/trainer_rank/_lora_export.py b/src/art/trainer_rank/_lora_export.py index 18973593e..9e97fb7e9 100644 --- a/src/art/trainer_rank/_lora_export.py +++ b/src/art/trainer_rank/_lora_export.py @@ -110,18 +110,12 @@ def _prepare_vllm_lora_publish( adapter_dtypes: dict[str, torch.dtype], handler: Any, adapter_config: dict[str, Any], - rank: int, - world_size: int, + runtime: tuple[int, torch.device], slot_ref: LoRASlotRef | None = None, - runtime: tuple[int, torch.device] | None = None, ) -> _VllmLoraPublishPlan: from art.megatron.weights import lora_publish - rank, device = ( - _validate_vllm_lora_publish_runtime(rank, world_size) - if runtime is None - else runtime - ) + rank, device = runtime packed_expert_groups = tuple(handler.expert_packed_lora_groups()) local_tensors, local_metadata = lora_publish.collect_local_lora_entries( model, @@ -243,8 +237,6 @@ def _capture_lora_publish_inputs( adapter_dtypes={}, handler=trainer.runtime.model_support_handler, adapter_config=adapter_config, - rank=trainer.runtime.rank, - world_size=trainer.runtime.world_size, slot_ref=trainer._slot_ref(checkpoint_name), runtime=runtime, ), From e2852cafafe2ad986815655419f323971e629d2a Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 14:03:50 +0000 Subject: [PATCH 089/150] Test second export metadata failure on rank one --- .../megatron/lora/test_lora_disk_codecs.py | 31 +++++++++++++------ 1 file changed, 22 insertions(+), 9 deletions(-) diff --git a/tests/integration/megatron/lora/test_lora_disk_codecs.py b/tests/integration/megatron/lora/test_lora_disk_codecs.py index 0978768ca..352da5f00 100644 --- a/tests/integration/megatron/lora/test_lora_disk_codecs.py +++ b/tests/integration/megatron/lora/test_lora_disk_codecs.py @@ -1590,8 +1590,9 @@ def synchronize(error: BaseException | None, _phase: str, group: object) -> None def _export_preparation_failure_worker( - rank: int, init_method: str, failure_site: str + rank: int, init_method: str, failure_case: tuple[str, int, int] ) -> None: + failure_site, failure_call, failing_rank = failure_case with ( gloo_group(rank, init_method, timeout=10), pytest.MonkeyPatch.context() as monkeypatch, @@ -1622,7 +1623,11 @@ def _export_preparation_failure_worker( def metadata(local: list[Any]) -> list[Any]: metadata_calls.append(local) result = canonical(local) - if rank == 0 and failure_site == "metadata": + if ( + rank == failing_rank + and failure_site == "metadata" + and len(metadata_calls) == failure_call + ): raise failure return result @@ -1641,7 +1646,7 @@ def exchange_tensors(plan: _lora_export._VllmLoraPublishPlan): def collect_or_fail(*args: Any, **kwargs: Any): result = collect(*args, **kwargs) - if rank == 0: + if rank == failing_rank: raise failure return result @@ -1652,12 +1657,14 @@ def collect_or_fail(*args: Any, **kwargs: Any): ) with pytest.raises(BaseException, match="injected") as caught: trainer._prepare_lora_export("retry", "student", owner_id="owner") - if rank == 0: + if rank == failing_rank: assert caught.value is failure else: assert isinstance(caught.value, RuntimeError) assert "Another rank failed" in str(caught.value) - assert len(metadata_calls) == (1 if failure_site == "metadata" else 0) + assert len(metadata_calls) == ( + failure_call if failure_site == "metadata" else 0 + ) assert exchanges == [] assert not getattr(trainer, "_prepared_lora_exports", {}) @@ -1687,16 +1694,22 @@ def collect_or_fail(*args: Any, **kwargs: Any): assert completed.item() == 2 -@pytest.mark.parametrize("failure_site", ("dense", "packed", "metadata")) +@pytest.mark.parametrize( + "failure_case", + [("dense", 1, 0), ("packed", 1, 0), ("metadata", 1, 0), ("metadata", 2, 1)], + ids=["dense", "packed", "metadata", "metadata-2-rank1"], +) def test_export_preparation_failure_is_collective_and_retryable( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch, failure_site: str + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, + failure_case: tuple[str, int, int], ) -> None: monkeypatch.setenv("CUDA_VISIBLE_DEVICES", "") spawn_and_join( _export_preparation_failure_worker, - args=(f"file://{tmp_path / 'export'}", failure_site), + args=(f"file://{tmp_path / 'export'}", failure_case), timeout=90, - failure=f"collective export {failure_site} failure test hung", + failure=f"collective export {failure_case[0]} failure test hung", ) From 8496d798f90f4f9d8f1edafe74c9c29fd467b904 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 14:39:43 +0000 Subject: [PATCH 090/150] Qualify joined-reasoning test fixture import --- tests/unit/test_reasoning_parser_joined_branches.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/unit/test_reasoning_parser_joined_branches.py b/tests/unit/test_reasoning_parser_joined_branches.py index 4c0aa0a01..47e3b1aca 100644 --- a/tests/unit/test_reasoning_parser_joined_branches.py +++ b/tests/unit/test_reasoning_parser_joined_branches.py @@ -1,11 +1,11 @@ from jinja2.sandbox import ImmutableSandboxedEnvironment import pytest -from test_literal_reasoning_content import _TEMPLATE from art_inference.chat_template import ( _QWEN_INLINE_REASONING, _without_inline_reasoning_parser, ) +from tests.unit.test_literal_reasoning_content import _TEMPLATE @pytest.mark.parametrize("scope", ["top", "macro"]) From bcb3cc8349968c269986a8721796730feb183551 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 15:09:29 +0000 Subject: [PATCH 091/150] Remove obsolete thinking workaround and strengthen branch coverage --- src/art/trajectories/_tokenize.py | 53 ---- src/art_inference/chat_template.py | 24 +- tests/unit/test_literal_reasoning_content.py | 238 +++++++----------- .../test_reasoning_parser_joined_branches.py | 50 +++- .../trajectories/test_literal_thinking_off.py | 47 ++-- 5 files changed, 164 insertions(+), 248 deletions(-) diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index f14834238..e9627324e 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -4532,58 +4532,6 @@ def _source_covers_complete_sampled_message( ) == normalize_chat_message(projected[0]) -def _preserve_literal_thinking_off_content( - history: ChatCompletionsHistory, - messages: list[dict[str, Any]], - template: object, - kwargs: Mapping[str, object], -) -> None: - # This Qwen3.5 template treats any in unstructured content as a - # reasoning separator, even with thinking disabled. Restrict the render-copy - # adaptation to its exact preserved template; other templates may interpret - # an empty reasoning_content field differently. - if ( - not isinstance(template, str) - or sha256(template.encode()).hexdigest() - != "098047d425a6673b1fe1a82a197a481616e53a283beaa8cb76cbb74d38ca6644" - or kwargs.get("enable_thinking") is not False - or kwargs.get("preserve_thinking") is not True - ): - return - for message, source in zip(messages, history.message_sources, strict=True): - if ( - source is None - or not isinstance(source.exchange, ChatCompletionsExchange) - or source.choice_index is None - or message.get("role") != "assistant" - or not isinstance(content := message.get("content"), str) - or "" not in content - ): - continue - request_kwargs = source.exchange.request.get("chat_template_kwargs") - if ( - not isinstance(request_kwargs, Mapping) - or request_kwargs.get("enable_thinking") is not False - ): - continue - choice = _chat_choice(source) - # Visible-only histories may omit structured reasoning present in the - # source response. Preserve both that source and normalized aliases. - if any( - value is not None and not (isinstance(value, str) and value == "") - for value in ( - message.get("reasoning"), - message.get("reasoning_content"), - _field(choice.message, "reasoning"), - _field(choice.message, "reasoning_content"), - ) - ): - continue - prompt, output, _ = _chat_choice_tokens(choice, source.exchange.response) - if prompt is not None and output is not None: - message["reasoning_content"] = "" - - class _ChatViewTokenizer: """Tokenize one chat-completions history view, stage by stage. @@ -4645,7 +4593,6 @@ def __init__( **default_chat_template_kwargs_for_template(template), **explicit_kwargs, } - _preserve_literal_thinking_off_content(history, messages, template, kwargs) ends_with_assistant = bool(messages) and messages[-1].get("role") == "assistant" segmented = False self.history = history diff --git a/src/art_inference/chat_template.py b/src/art_inference/chat_template.py index 430ee69f4..6a63b8f86 100644 --- a/src/art_inference/chat_template.py +++ b/src/art_inference/chat_template.py @@ -339,7 +339,7 @@ def apply_edits() -> str: selected = set() shared = set() - def writes_content(node: nodes.Assign | nodes.AssignBlock) -> bool: + def writes_content(node: nodes.Assign) -> bool: target = node.target return ( isinstance(target, nodes.Name) @@ -384,18 +384,9 @@ def remember_stores(node: nodes.Node, initialized: set[str] | None) -> None: initialized.update( n.name for n in node.find_all(nodes.Name) if n.ctx == "store" ) - for child in ( - node, - *node.find_all((nodes.Macro, nodes.Import, nodes.FromImport)), - ): + for child in (node, *node.find_all(nodes.Macro)): if isinstance(child, nodes.Macro): initialized.add(child.name) - elif isinstance(child, nodes.Import): - initialized.add(child.target) - elif isinstance(child, nodes.FromImport): - initialized.update( - name if isinstance(name, str) else name[1] for name in child.names - ) def visit( body: Sequence[nodes.Node], @@ -456,7 +447,7 @@ def visit( # object whose destructor observes the new value. shared.update(bindings) else: - # Macro/loop/with/block bodies have independent bindings. + # Macro/loop bodies have independent bindings. for _, value in node.iter_fields(): if isinstance(value, list) and all( isinstance(n, nodes.Node) for n in value @@ -477,21 +468,12 @@ def visit( ) if isinstance(node, nodes.Macro): local_names = {arg.name for arg in node.args} - elif isinstance(node, nodes.With): - local_names = { - bound.name - for target in node.targets - for bound in (target, *target.find_all(nodes.Name)) - if isinstance(bound, nodes.Name) - } visit( value, set(), local_names, isinstance(node, nodes.For) and value is node.body, ) - if isinstance(node, nodes.AssignBlock) and writes_content(node): - bindings.clear() return bindings visit(tree.body, set()) diff --git a/tests/unit/test_literal_reasoning_content.py b/tests/unit/test_literal_reasoning_content.py index d346d25d3..70898624a 100644 --- a/tests/unit/test_literal_reasoning_content.py +++ b/tests/unit/test_literal_reasoning_content.py @@ -35,6 +35,16 @@ ) +_TRIM = "{% set content = render_content(message.content, true)|trim %}" +_RENDERER = "{% macro render_content(content, count) %}{{ content }}{% endmacro %}" + + +def _fixture_parser() -> str: + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + return match.group() + + @pytest.mark.parametrize( "middle", [ @@ -56,13 +66,10 @@ ) @pytest.mark.parametrize("content", [" answer ", " beforexafter "]) def test_implicit_context_consumers_keep_original_shared_trim(middle, content): - match = _QWEN_INLINE_REASONING.search(_TEMPLATE) - assert match is not None - parser = match.group() - trim = "{% set content = render_content(message.content, true)|trim %}" + parser = _fixture_parser() template = ( "{% macro render_content(content, count) %}{{ content }}{% endmacro %}" - "{% set probe = bind() %}" + trim + middle + parser + "[{{ content }}]" + "{% set probe = bind() %}" + _TRIM + middle + parser + "[{{ content }}]" ) seen = [] @@ -137,18 +144,16 @@ def finalize(context, value): assert env.from_string(fixed).render(message=message, bind=bind) == expected assert seen == before assert all(value == content.strip() for value in seen) - assert trim in fixed + assert _TRIM in fixed def test_role_guard_is_not_proof_after_message_reassignment(): - match = _QWEN_INLINE_REASONING.search(_TEMPLATE) - assert match is not None - trim = "{% set content = render_content(message.content, true)|trim %}" + parser = _fixture_parser() template = ( "{% macro render_content(content, count) %}{{ content }}{% endmacro %}" - + trim + + _TRIM + "{% set message = none %}{% if message.role == 'assistant' %}" - + match.group() + + parser + "{% else %}[{{ content }}]{% endif %}" ) fixed = chat_template_with_preserved_thinking(template) @@ -156,21 +161,19 @@ def test_role_guard_is_not_proof_after_message_reassignment(): env = ImmutableSandboxedEnvironment() message = {"role": "assistant", "content": " answer "} assert env.from_string(fixed).render(message=message) == "[answer]" - assert trim in fixed + assert _TRIM in fixed @pytest.mark.parametrize("prior_binding", [False, True]) def test_only_fresh_loop_local_initialization_is_transparent(prior_binding): - match = _QWEN_INLINE_REASONING.search(_TEMPLATE) - assert match is not None - trim = "{% set content = render_content(message.content, true)|trim %}" + parser = _fixture_parser() template = ( "{% macro render_content(content, count) %}{{ content }}{% endmacro %}" "{% for message in messages %}" + ("{% set probe = make_probe() %}" if prior_binding else "") - + trim + + _TRIM + "{% set probe = none %}" - + match.group() + + parser + "[{{ content }}]{% endfor %}" ) destroyed = [] @@ -188,7 +191,7 @@ def __del__(self): ) assert actual == "[" + (content.strip() if prior_binding else content) + "]" assert bool(destroyed) is prior_binding - assert (trim in fixed) is prior_binding + assert (_TRIM in fixed) is prior_binding @pytest.mark.parametrize( @@ -200,8 +203,7 @@ def __del__(self): ], ) def test_role_pruning_keeps_prior_store_history(branch): - match = _QWEN_INLINE_REASONING.search(_TEMPLATE) - assert match is not None + parser = _fixture_parser() seen = [] class Probe: @@ -211,7 +213,6 @@ def __init__(self, reader): def __del__(self): seen.append(self.reader()) - trim = "{% set content = render_content(message.content, true)|trim %}" template = ( "{% macro render_content(content, count) %}{{ content }}{% endmacro %}" "{% for message in messages %}" @@ -221,9 +222,9 @@ def __del__(self): "{% set probe = make_probe(read_content) %}" "{% set message = {'role': 'assistant', 'content': message.content} %}", ) - + trim + + _TRIM + "{% set probe = none %}" - + match.group() + + parser + "{% endfor %}{{ seen|join('|') }}" ) env = ImmutableSandboxedEnvironment() @@ -237,7 +238,7 @@ def __del__(self): fixed = _without_inline_reasoning_parser(template) assert env.from_string(fixed).render(**kwargs) == "answer" assert seen == ["answer"] - assert trim in fixed + assert _TRIM in fixed @pytest.mark.parametrize( @@ -256,14 +257,12 @@ def __del__(self): ) @pytest.mark.parametrize("mutate_role", [False, True]) def test_unknown_renderer_effects_keep_trim_and_call_order(prefix, mutate_role): - match = _QWEN_INLINE_REASONING.search(_TEMPLATE) - assert match is not None - trim = "{% set content = render_content(message.content, true)|trim %}" + parser = _fixture_parser() source = ( prefix - + trim + + _TRIM + "{% if message.role == 'assistant' %}" - + match.group() + + parser + "{% endif %}[{{ content }}]" ) env = ImmutableSandboxedEnvironment() @@ -295,21 +294,19 @@ def __getitem__(self, key): fixed = _without_inline_reasoning_parser(source) assert not _QWEN_INLINE_REASONING.search(fixed) - assert trim in fixed - assert render(fixed) == render(source.replace(match.group(), "")) + assert _TRIM in fixed + assert render(fixed) == render(source.replace(parser, "")) assert render(fixed) == ("[beforexafter]", ["assistant"]) assert chat_template_with_preserved_thinking(source) == fixed assert chat_template_with_preserved_thinking(fixed) == fixed def test_renderer_declared_after_use_does_not_prove_role_stability(): - match = _QWEN_INLINE_REASONING.search(_TEMPLATE) - assert match is not None - trim = "{% set content = render_content(message.content, true)|trim %}" + parser = _fixture_parser() source = ( - trim + _TRIM + "{% if message.role == 'assistant' %}" - + match.group() + + parser + "{% else %}[{{ content }}]{% endif %}" "{% macro render_content(content, count) %}{{ content }}{% endmacro %}" ) @@ -327,21 +324,19 @@ def render_content(value, count): ) fixed = _without_inline_reasoning_parser(source) - assert trim in fixed + assert _TRIM in fixed assert render(source) == render(fixed) == "[answer]" def test_inherited_renderer_does_not_prove_role_stability(): - match = _QWEN_INLINE_REASONING.search(_TEMPLATE) - assert match is not None - trim = "{% set content = render_content(message.content, true)|trim %}" + parser = _fixture_parser() source = ( "{% extends 'parent' %}" "{% macro render_content(content,count) %}{{ content }}{% endmacro %}" "{% block body %}" - + trim + + _TRIM + "{% if message.role == 'assistant' %}" - + match.group() + + parser + "{% endif %}[{{ content }}]{% endblock %}" ) parent = ( @@ -355,7 +350,7 @@ def mutate(value, count): return value fixed = _without_inline_reasoning_parser(source) - assert trim in fixed + assert _TRIM in fixed assert ( ImmutableSandboxedEnvironment(loader=DictLoader({"parent": parent})) .from_string(fixed) @@ -375,18 +370,16 @@ def mutate(value, count): "text", [" answer ", " beforeliteralafter "] ) def test_renderer_cannot_mutate_parameter_namespace_aliases(mutation, text): - match = _QWEN_INLINE_REASONING.search(_TEMPLATE) - assert match is not None - trim = "{% set content = render_content(message.content, true)|trim %}" + parser = _fixture_parser() source = ( "{% macro render_content(content,count) %}" + mutation + "{{ content.text }}{% endmacro %}" + "{% set message = namespace(role=message.role, text=message.content) %}" + "{% set message.content = message %}" - + trim + + _TRIM + "{% if message.role == 'assistant' %}" - + match.group() + + parser + "{% else %}[{{ content }}]{% endif %}" ) env = ImmutableSandboxedEnvironment() @@ -395,25 +388,23 @@ def test_renderer_cannot_mutate_parameter_namespace_aliases(mutation, text): fixed = _without_inline_reasoning_parser(source) assert original == "[" + text.strip() + "]" assert env.from_string(fixed).render(**kwargs) == original - assert trim in fixed + assert _TRIM in fixed assert not _QWEN_INLINE_REASONING.search(fixed) assert _without_inline_reasoning_parser(fixed) == fixed @pytest.mark.parametrize("text", [" answer ", " literalafter "]) def test_private_renderer_counter_cannot_be_rebound_to_message(text): - match = _QWEN_INLINE_REASONING.search(_TEMPLATE) - assert match is not None - trim = "{% set content = render_content(message.content, true)|trim %}" + parser = _fixture_parser() source = ( "{% set counter = namespace(value=0) %}" "{% macro render_content(content,count) %}" "{% set counter.role = 'user' %}{{ content }}{% endmacro %}" "{% set message = namespace(role=message.role, content=message.content) %}" "{% set counter = message %}" - + trim + + _TRIM + "{% if message.role == 'assistant' %}" - + match.group() + + parser + "{% else %}[{{ content }}]{% endif %}" ) env = ImmutableSandboxedEnvironment() @@ -422,7 +413,7 @@ def test_private_renderer_counter_cannot_be_rebound_to_message(text): fixed = _without_inline_reasoning_parser(source) assert original == "[" + text.strip() + "]" assert env.from_string(fixed).render(**kwargs) == original - assert trim in fixed + assert _TRIM in fixed assert not _QWEN_INLINE_REASONING.search(fixed) assert _without_inline_reasoning_parser(fixed) == fixed @@ -438,15 +429,13 @@ def test_private_renderer_counter_cannot_be_rebound_to_message(text): ], ) def test_counter_free_renderer_rejects_implicit_macro_mutation(environment, middle): - match = _QWEN_INLINE_REASONING.search(_TEMPLATE) - assert match is not None - trim = "{% set content = render_content(message.content, true)|trim %}" + parser = _fixture_parser() source = ( "{% macro render_content(content,count) %}{{ content }}{% endmacro %}" + middle - + trim + + _TRIM + "{% if message.role == 'assistant' %}" - + match.group() + + parser + "{% endif %}[{{ content }}]" ) @@ -473,9 +462,9 @@ def replacement(content, count): expected = "[answer]", ["inject", "render:assistant"], "user" fixed = _without_inline_reasoning_parser(source) - assert render(source) == render(source.replace(match.group(), "")) == expected + assert render(source) == render(source.replace(parser, "")) == expected assert render(fixed) == expected - assert trim in fixed + assert _TRIM in fixed assert not _QWEN_INLINE_REASONING.search(fixed) assert _without_inline_reasoning_parser(fixed) == fixed @@ -490,16 +479,14 @@ def replacement(content, count): ], ) def test_counter_renderer_rejects_implicit_context_exports(middle): - match = _QWEN_INLINE_REASONING.search(_TEMPLATE) - assert match is not None - trim = "{% set content = render_content(message.content, true)|trim %}" + parser = _fixture_parser() source = ( "{% set counter=namespace(value=0) %}" "{% macro render_content(content,count) %}{% set counter.value=counter.value+1 %}{{ content }}{% endmacro %}" + middle - + trim + + _TRIM + "{% if message.role == 'assistant' %}" - + match.group() + + parser + "{% endif %}[{{ content }}]" ) message = {"role": "assistant", "content": " answer "} @@ -519,7 +506,7 @@ def inject(context, value=None): ) env.filters["inject"] = env.tests["inject"] = inject fixed = _without_inline_reasoning_parser(source) - assert trim in fixed + assert _TRIM in fixed assert ( env.from_string(fixed).render(message=message, inject=inject, probe=Probe()) == "[answer]" @@ -528,8 +515,7 @@ def inject(context, value=None): @pytest.mark.parametrize("scope", ["top", "loop", "macro", "with"]) def test_replacing_content_keeps_trim_for_its_own_observers(scope): - match = _QWEN_INLINE_REASONING.search(_TEMPLATE) - assert match is not None + parser = _fixture_parser() seen = [] class Probe: @@ -545,7 +531,7 @@ def __del__(self): "{% if message.role == 'user' %}{% set content = make_probe(read_content) %}" "{% set message = {'role': 'assistant', 'content': message.content} %}{% endif %}" + trim - + match.group() + + parser ) if scope == "loop": body = "{% for message in messages %}" + body + "{% endfor %}" @@ -581,9 +567,7 @@ def __del__(self): def test_disabling_inline_parser_preserves_outer_whitespace( trim_blocks, lstrip_blocks, left, right, newline ): - match = _QWEN_INLINE_REASONING.search(_TEMPLATE) - assert match is not None - operation = match.group() + operation = _fixture_parser() operation = "{%" + left + operation[3:] operation = operation[:-2].rstrip("-+") + right + "%}" template = ("HEADER \n\t" + operation + "\n \tTAIL{{ content }}").replace( @@ -784,15 +768,13 @@ def test_unconfigured_template_receives_the_same_correction(): @pytest.mark.parametrize("indirect", [False, True]) @pytest.mark.parametrize("content", [" answer ", " beforexafter "]) def test_macro_capture_keeps_shared_content_trim(indirect, content): - parser = _QWEN_INLINE_REASONING.search(_TEMPLATE) - assert parser is not None template = ( "{% macro render_content(content, count) %}{{ content }}{% endmacro %}" "{% macro preview() %}[{{ content }}]{% endmacro %}" "{% macro wrapper() %}{{ preview() }}{% endmacro %}" "{% set content = render_content(message.content, true)|trim %}" + ("{{ wrapper() }}" if indirect else "{{ preview() }}") - + parser.group() + + _fixture_parser() + "[{{ content }}]" ) fixed = chat_template_with_preserved_thinking(template) @@ -834,11 +816,7 @@ def test_extension_fallback_keeps_prior_structured_content_whitespace(preserve): @pytest.mark.parametrize("wrapper", [("{% raw %}", "{% endraw %}"), ("{#", "#}")]) def test_inline_operation_as_raw_or_comment_text_is_not_rewritten(wrapper): - from art_inference.chat_template import _QWEN_INLINE_REASONING - - match = _QWEN_INLINE_REASONING.search(_TEMPLATE) - assert match is not None - operation = match.group() + operation = _fixture_parser() template = wrapper[0] + operation + wrapper[1] assert chat_template_with_preserved_thinking(template) == template @@ -858,11 +836,7 @@ def test_other_structured_reasoning_condition_is_not_rewritten(): def test_inline_operation_inside_quoted_expression_is_literal(): - from art_inference.chat_template import _QWEN_INLINE_REASONING - - match = _QWEN_INLINE_REASONING.search(_TEMPLATE) - assert match is not None - operation = match.group().replace("\n", " ") + operation = _fixture_parser().replace("\n", " ") template = '{{ "' + operation + '" }}' fixed = chat_template_with_preserved_thinking(template) assert fixed == template @@ -873,11 +847,7 @@ def test_inline_operation_inside_quoted_expression_is_literal(): "wrapper", [("{% raw %}", "{% endraw %}"), ("{#", "#}"), ('{{ "', '" }}')] ) def test_mixed_executable_and_literal_operations_only_changes_executable(wrapper): - from art_inference.chat_template import _QWEN_INLINE_REASONING - - match = _QWEN_INLINE_REASONING.search(_TEMPLATE) - assert match is not None - operation = match.group().replace("\n", " ") + operation = _fixture_parser().replace("\n", " ") literal = wrapper[0] + operation + wrapper[1] template = _TEMPLATE + literal fixed = chat_template_with_preserved_thinking(template) @@ -914,9 +884,7 @@ def test_newline_lexing_preserves_literal_content(newline, prefix): @pytest.mark.parametrize("newline", ["\r\n", "\r"]) @pytest.mark.parametrize("wrapper", ["comment", "raw", "quoted"]) def test_newline_parser_spelling_in_nonexecutable_token_unchanged(newline, wrapper): - match = _QWEN_INLINE_REASONING.search(_TEMPLATE) - assert match is not None - operation = match.group().replace("\n", newline) + operation = _fixture_parser().replace("\n", newline) if wrapper == "comment": template = "{#" + newline + operation + newline + "#}" elif wrapper == "raw": @@ -959,9 +927,7 @@ def test_equivalent_inline_operations_preserve_literal_and_structured_fields(spe @pytest.mark.parametrize("wrapper", ["raw", "comment", "quoted"]) def test_equivalent_operation_as_literal_data_is_not_edited(wrapper): - match = _QWEN_INLINE_REASONING.search(_TEMPLATE) - assert match is not None - operation = match.group().replace("'", '"') + operation = _fixture_parser().replace("'", '"') if wrapper == "raw": literal = "{% raw %}" + operation + "{% endraw %}" elif wrapper == "comment": @@ -975,9 +941,8 @@ def test_equivalent_operation_as_literal_data_is_not_edited(wrapper): @pytest.mark.parametrize("change", ["different_split", "side_effect", "different_gate"]) def test_distinct_custom_content_operations_are_not_inferred(change): - match = _QWEN_INLINE_REASONING.search(_TEMPLATE) - assert match is not None - operation = match.group() + parser = _fixture_parser() + operation = parser if change == "different_split": operation = operation.replace( "content.split('')[-1]", "content.split('')[0]" @@ -990,7 +955,7 @@ def test_distinct_custom_content_operations_are_not_inferred(change): operation = operation.replace( "if '' in content", "if custom and '' in content" ) - assert operation != match.group() + assert operation != parser assert _without_inline_reasoning_parser(operation) == operation @@ -1018,36 +983,32 @@ def test_plain_block_whitespace_keeps_prior_inline_parser_coverage(separator, qu def test_custom_macro_calls_keep_trim_while_removing_recognized_parser( preserve, newline, layout, inline_structured_reasoning ): - match = _QWEN_INLINE_REASONING.search(_TEMPLATE) - assert match is not None - parser = match.group() + parser = _fixture_parser() if inline_structured_reasoning: parser += "{{ reasoning_content|trim }}" - trim = "{% set content = render_content(message.content, true)|trim %}" - render = "{% macro render_content(content, count) %}{{ content }}{% endmacro %}" - preview = "{% macro preview(message) %}" + trim + "[{{ content }}]{% endmacro %}" + preview = "{% macro preview(message) %}" + _TRIM + "[{{ content }}]{% endmacro %}" main = ( - "{% macro answer(message) %}" + trim + parser + "[{{ content }}]{% endmacro %}" + "{% macro answer(message) %}" + _TRIM + parser + "[{{ content }}]{% endmacro %}" ) if layout == "branches": main = ( "{% macro answer(message) %}{% if message.role == 'assistant' %}" - + trim + + _TRIM + parser + "[{{ content }}]{% else %}" - + trim + + _TRIM + "[{{ content }}]{% endif %}{% endmacro %}" ) separator = "" if layout == "same_line" else newline template = separator.join( - [render, preview, main, "{{ answer(message) }}|{{ preview(message) }}"] + [_RENDERER, preview, main, "{{ answer(message) }}|{{ preview(message) }}"] ) fixed = chat_template_with_preserved_thinking(template) assert isinstance(fixed, str) assert preview in fixed # Other macro calls lack a closed callable-custody proof. Retain their # original trims while still removing only the recognized parser. - assert trim in fixed + assert _TRIM in fixed assert not _QWEN_INLINE_REASONING.search(fixed) env = ImmutableSandboxedEnvironment(trim_blocks=True, lstrip_blocks=True) kwargs = dict( @@ -1103,30 +1064,26 @@ def test_qwen_content_preview_macro_is_unchanged(preserve, raw): @pytest.mark.parametrize("shadow", ["assignment", "conditional", "scope"]) def test_content_trim_requires_a_proven_local_binding(shadow): - match = _QWEN_INLINE_REASONING.search(_TEMPLATE) - assert match is not None - parser = match.group() - trim = "{% set content = render_content(message.content, true)|trim %}" - render = "{% macro render_content(content, count) %}{{ content }}{% endmacro %}" + parser = _fixture_parser() if shadow == "assignment": - template = render + trim + "{% set content = 'replacement' %}" + parser + template = _RENDERER + _TRIM + "{% set content = 'replacement' %}" + parser elif shadow == "conditional": template = ( - render - + trim + _RENDERER + + _TRIM + "{% if custom %}{% set content = 'replacement' %}{% endif %}" + parser ) else: template = ( - render - + trim + _RENDERER + + _TRIM + "{% macro nested() %}" + parser + "{{ content }}{% endmacro %}{{ nested() }}" ) fixed = _without_inline_reasoning_parser(template) - assert trim in fixed + assert _TRIM in fixed assert not _QWEN_INLINE_REASONING.search(fixed) assert _without_inline_reasoning_parser(fixed) == fixed @@ -1149,9 +1106,7 @@ def parse(self, parser): trim_blocks=True, lstrip_blocks=True, ) - match = _QWEN_INLINE_REASONING.search(_TEMPLATE) - assert match is not None - parser = match.group() + parser = _fixture_parser() prefix = ( "{% for item in [1] %}{% break %}{% endfor %}" if extension == "loopcontrols" @@ -1169,8 +1124,7 @@ def parse(self, parser): def test_tuple_scope_content_replacement_retains_trim(): - match = _QWEN_INLINE_REASONING.search(_TEMPLATE) - assert match is not None + parser = _fixture_parser() seen = [] class Probe: @@ -1191,7 +1145,7 @@ def bind(probe, reader): "{% macro read_content() %}{{ content }}{% endmacro %}" "{{ bind(content,read_content) }}" + trim - + match.group() + + parser + "{% endwith %}{{ seen|join('|') }}" ) fixed = _without_inline_reasoning_parser(source) @@ -1212,11 +1166,7 @@ def bind(probe, reader): @pytest.mark.parametrize("consumer", ["output", "alias", "condition", "other_branch"]) def test_shared_content_preview_prevents_ambiguous_trim_rewrite(consumer): - match = _QWEN_INLINE_REASONING.search(_TEMPLATE) - assert match is not None - parser = match.group() - trim = "{% set content = render_content(message.content, true)|trim %}" - render = "{% macro render_content(content, count) %}{{ content }}{% endmacro %}" + parser = _fixture_parser() if consumer == "output": body = "[{{ content }}]" + parser + "[{{ content }}]" elif consumer == "alias": @@ -1229,10 +1179,10 @@ def test_shared_content_preview_prevents_ambiguous_trim_rewrite(consumer): + parser + "[{{ content }}]{% endif %}" ) - template = render + trim + body + template = _RENDERER + _TRIM + body fixed = chat_template_with_preserved_thinking(template) assert isinstance(fixed, str) - assert trim in fixed + assert _TRIM in fixed assert not _QWEN_INLINE_REASONING.search(fixed) env = ImmutableSandboxedEnvironment(trim_blocks=True, lstrip_blocks=True) kwargs = dict(message={"role": "assistant", "content": " "}, preview_only=True) @@ -1242,28 +1192,24 @@ def test_shared_content_preview_prevents_ambiguous_trim_rewrite(consumer): def test_custom_macro_calls_keep_trim_and_unedited_parser(): - match = _QWEN_INLINE_REASONING.search(_TEMPLATE) - assert match is not None - parser = match.group() + parser = _fixture_parser() unedited = parser.replace("%}", "%}{# preserve this custom block #}", 1) - trim = "{% set content = render_content(message.content, true)|trim %}" - render = "{% macro render_content(content, count) %}{{ content }}{% endmacro %}" preview = ( "{% macro preview(message) %}" - + trim + + _TRIM + unedited + "[{{ content }}]{% endmacro %}" ) answer = ( - "{% macro answer(message) %}" + trim + parser + "[{{ content }}]{% endmacro %}" + "{% macro answer(message) %}" + _TRIM + parser + "[{{ content }}]{% endmacro %}" ) template = ( - render + preview + answer + "{{ answer(message) }}|{{ preview(message) }}" + _RENDERER + preview + answer + "{{ answer(message) }}|{{ preview(message) }}" ) fixed = _without_inline_reasoning_parser(template) assert preview in fixed assert answer not in fixed - assert trim in fixed + assert _TRIM in fixed env = ImmutableSandboxedEnvironment(trim_blocks=True, lstrip_blocks=True) rendered = env.from_string(fixed).render( message={"role": "assistant", "content": " HEADliteralTAIL "} diff --git a/tests/unit/test_reasoning_parser_joined_branches.py b/tests/unit/test_reasoning_parser_joined_branches.py index 47e3b1aca..9886efc34 100644 --- a/tests/unit/test_reasoning_parser_joined_branches.py +++ b/tests/unit/test_reasoning_parser_joined_branches.py @@ -1,3 +1,5 @@ +from pathlib import Path + from jinja2.sandbox import ImmutableSandboxedEnvironment import pytest @@ -5,10 +7,13 @@ _QWEN_INLINE_REASONING, _without_inline_reasoning_parser, ) -from tests.unit.test_literal_reasoning_content import _TEMPLATE + +_TEMPLATE = ( + Path(__file__).parents[1] / "fixtures/qwen35_preserved_thinking.jinja" +).read_text() -@pytest.mark.parametrize("scope", ["top", "macro"]) +@pytest.mark.parametrize("scope", ["top", "macro", "bound"]) @pytest.mark.parametrize("layout", ["if", "elif", "nested", "sequential"]) @pytest.mark.parametrize("content", [" answer ", " beforexafter "]) def test_joined_preview_retains_its_original_trim(scope, layout, content): @@ -34,6 +39,11 @@ def test_joined_preview_retains_its_original_trim(scope, layout, content): "{% if mode == 'other' %}" + parser + "{% endif %}" ) body = trim + branch + "[{{ content }}]" + if scope == "bound": + body = ( + "{% set preview_only = enable_thinking %}" + "{% set mode = message.mode %}{% set outer = message.outer %}" + body + ) if scope == "macro": body = ( "{% macro answer(message) %}" + body + "{% endmacro %}{{ answer(message) }}" @@ -44,7 +54,13 @@ def test_joined_preview_retains_its_original_trim(scope, layout, content): fixed = _without_inline_reasoning_parser(template) env = ImmutableSandboxedEnvironment(trim_blocks=True, lstrip_blocks=True) kwargs = dict( - message={"role": "assistant", "content": content}, + message={ + "role": "assistant", + "content": content, + "mode": "answer", + "outer": True, + }, + enable_thinking=True, preview_only=True, mode="answer", outer=True, @@ -56,13 +72,16 @@ def test_joined_preview_retains_its_original_trim(scope, layout, content): assert trim in fixed assert not _QWEN_INLINE_REASONING.search(fixed) # The parser path still treats the text as literal; its shared trim stays. - kwargs["preview_only"] = False + kwargs["preview_only"] = kwargs["enable_thinking"] = False assert env.from_string(fixed).render(**kwargs) == "[" + content.strip() + "]" assert _without_inline_reasoning_parser(fixed) == fixed @pytest.mark.parametrize("mode", ["a", "b", "c"]) -def test_unknown_comparison_keeps_shared_trim_even_when_all_paths_have_parser(mode): +@pytest.mark.parametrize("bound", [False, True]) +def test_unknown_comparison_keeps_shared_trim_even_when_all_paths_have_parser( + mode, bound +): match = _QWEN_INLINE_REASONING.search(_TEMPLATE) assert match is not None parser = match.group() @@ -77,13 +96,32 @@ def test_unknown_comparison_keeps_shared_trim_even_when_all_paths_have_parser(mo + parser + "{% endif %}[{{ content }}]" ) + if bound: + template = "{% set mode = message.mode %}" + template fixed = _without_inline_reasoning_parser(template) content = " beforexafter " env = ImmutableSandboxedEnvironment(trim_blocks=True, lstrip_blocks=True) assert ( env.from_string(fixed).render( - mode=mode, message={"role": "assistant", "content": content} + mode=mode, message={"role": "assistant", "content": content, "mode": mode} ) == "[" + content.strip() + "]" ) assert _without_inline_reasoning_parser(fixed) == fixed + + +def test_skipped_role_branch_keeps_unique_assistant_binding(): + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + template = ( + "{% macro render_content(content, count) %}{{ content }}{% endmacro %}" + "{% set content = render_content(message.content, true)|trim %}" + "{% if message.role != 'assistant' %}{% set content = 'other' %}{% endif %}" + + match.group() + + "[{{ content }}]" + ) + fixed = _without_inline_reasoning_parser(template) + render = ImmutableSandboxedEnvironment().from_string(fixed).render + message = {"role": "assistant", "content": " beforexafter "} + assert render(message=message) == "[ beforexafter ]" + assert render(message={**message, "role": "user"}) == "[other]" diff --git a/tests/unit/trajectories/test_literal_thinking_off.py b/tests/unit/trajectories/test_literal_thinking_off.py index c30c1a632..d094279ac 100644 --- a/tests/unit/trajectories/test_literal_thinking_off.py +++ b/tests/unit/trajectories/test_literal_thinking_off.py @@ -142,15 +142,11 @@ def test_native_thinking_off_retains_literal_content( ) -> None: history, tokenizer = _history(content=content) original = history.model_dump(mode="python") - # Render the original template without either the general parser correction - # or the old hash-specific workaround to retain the destructive baseline. + # Disable the general correction to retain the destructive original render. with monkeypatch.context() as patch: patch.setattr( _tokenize, "chat_template_with_preserved_thinking", lambda value: value ) - patch.setattr( - _tokenize, "_preserve_literal_thinking_off_content", lambda *args: None - ) _outcome(history, tokenizer) assert content not in tokenizer.rendered[0] tokenizer.calls.clear() @@ -186,9 +182,7 @@ def test_native_thinking_off_retains_literal_content( "visible_only", ], ) -def test_unrelated_histories_keep_original_rendering( - case: str, monkeypatch: pytest.MonkeyPatch -) -> None: +def test_unrelated_histories_keep_original_rendering(case: str) -> None: history, tokenizer = _history( thinking=True if case == "source_on" @@ -225,17 +219,29 @@ def test_unrelated_histories_keep_original_rendering( if case == "visible_only": cast(dict[str, Any], history.messages[-1]).pop("reasoning") original = history.model_dump(mode="python") - candidate = _outcome(history, tokenizer) - calls = deepcopy(tokenizer.calls) - tokenizer.calls.clear() - with monkeypatch.context() as patch: - patch.setattr( - _tokenize, "_preserve_literal_thinking_off_content", lambda *args: None - ) - baseline = _outcome(history, tokenizer) - assert candidate == baseline - assert len(calls) == len(tokenizer.calls) - assert calls[0] == tokenizer.calls[0] + tokenized = history.tokenize(tokenizer=tokenizer) + # Independent public transcript and source-field oracles, not a second run + # with an already-inert workaround disabled. + expected = ( + "<|im_start|>user\nPublic query.<|im_end|>\n" + "<|im_start|>assistant\n\n\n\n\n" + _LITERAL + ) + assert tokenizer.rendered[0] == expected + "<|im_end|>\n" + structured = case in {"structured", "alias"} + assert tokenizer.decode(tokenized.tokens) == expected + ( + "" if structured else "<|im_end|>\n" + ) + assert tokenizer.calls[0][-1] == { + "role": "assistant", + "content": _LITERAL, + **({"reasoning": "explicit reasoning"} if structured else {}), + } + sampled = [ + i for i, flag in enumerate(tokenized.flags) if flag & tr.TokenFlag.SAMPLED + ] + expected_sampled = "" if case in {"no_source", "request_source"} else _LITERAL + assert tokenizer.decode([tokenized.tokens[i] for i in sampled]) == expected_sampled + assert [tokenized.logprobs[i] for i in sampled] == [-0.5] * len(expected_sampled) assert history.model_dump(mode="python") == original @@ -410,9 +416,6 @@ def observe(*args: Any, **kwargs: Any) -> tr.TokenizedHistory | None: patch.setattr( _tokenize, "chat_template_with_preserved_thinking", lambda value: value ) - patch.setattr( - _tokenize, "_preserve_literal_thinking_off_content", lambda *args: None - ) _outcome(history, tokenizer) boundary, old_exact = observed[0] stored = list(boundary.tail + boundary.following) From 1a02ed72394857280ca9e242dcd5a679918734d0 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 16:25:48 +0000 Subject: [PATCH 092/150] Fix runtime annotations on extracted public trainer methods (cherry picked from commit 7233b46ca49c1da69c56c30d33c07fe768c137ee) --- src/art/trainer_rank/_micro_batch_planner.py | 6 +++--- src/art/trainer_rank/_optimizer.py | 4 ++-- src/art/trainer_rank/_slots.py | 6 +++--- 3 files changed, 8 insertions(+), 8 deletions(-) diff --git a/src/art/trainer_rank/_micro_batch_planner.py b/src/art/trainer_rank/_micro_batch_planner.py index 72925dd94..9f1e253e5 100644 --- a/src/art/trainer_rank/_micro_batch_planner.py +++ b/src/art/trainer_rank/_micro_batch_planner.py @@ -1660,7 +1660,7 @@ def _complete_planner_observation( _impl._planner_misses._warn("could not finish planner-miss observation") -def finish_planner_observation(self: TrainerRank) -> None: +def finish_planner_observation(self: _impl.TrainerRank) -> None: """Release execution context without sampling an unbounded caller peak. ART compares completed peaks at its existing profiling boundaries: @@ -1675,7 +1675,7 @@ def finish_planner_observation(self: TrainerRank) -> None: _impl._planner_misses._warn("could not finish planner-miss execution") -def report_planner_oom(self: TrainerRank, error: BaseException) -> None: +def report_planner_oom(self: _impl.TrainerRank, error: BaseException) -> None: """Persist a caught CUDA OOM before caller cleanup, then leave it alone. This does not suppress, retry, or recover the original failure. An OOM @@ -1767,7 +1767,7 @@ def replay() -> dict[str, Any]: _impl._planner_misses._warn("could not persist planner OOM report") -def discard_planner_observation(self: TrainerRank) -> None: +def discard_planner_observation(self: _impl.TrainerRank) -> None: try: observation = getattr(self, "_planner_observation", None) if observation is not None: diff --git a/src/art/trainer_rank/_optimizer.py b/src/art/trainer_rank/_optimizer.py index de98ca260..ae45bd45e 100644 --- a/src/art/trainer_rank/_optimizer.py +++ b/src/art/trainer_rank/_optimizer.py @@ -91,9 +91,9 @@ def _extend_dynamic_optimizer( def optim_step( - self: TrainerRank, + self: _impl.TrainerRank, *, - params: AdamParams | Mapping[str, AdamParams], + params: _impl.AdamParams | Mapping[str, _impl.AdamParams], scale_grads: float | Mapping[str, float] = 1.0, checkpoints: Sequence[str] | None = None, on_live_graphs: Literal["allow", "error"] = "allow", diff --git a/src/art/trainer_rank/_slots.py b/src/art/trainer_rank/_slots.py index bc02626d8..126f8e422 100644 --- a/src/art/trainer_rank/_slots.py +++ b/src/art/trainer_rank/_slots.py @@ -56,7 +56,7 @@ def _resolve_custom_checkpoint(self: TrainerRank, checkpoint: AdapterSelection) def prefetch_checkpoints( - self: TrainerRank, *checkpoints: str | MaterializedCheckpoint + self: _impl.TrainerRank, *checkpoints: str | _impl.MaterializedCheckpoint ) -> asyncio.Task[None]: futures = [] for checkpoint in checkpoints: @@ -182,7 +182,7 @@ def _ensure_checkpoint_slots(self: TrainerRank, checkpoints: Iterable[str]) -> N def load_checkpoint( - self: TrainerRank, checkpoint: str | MaterializedCheckpoint | None + self: _impl.TrainerRank, checkpoint: str | _impl.MaterializedCheckpoint | None ) -> None: self._guard_forward_collective("load_checkpoint") logical, source = self._checkpoint_source(checkpoint) @@ -232,7 +232,7 @@ def _push_checkpoint_sync( self._slot_stack.append(self._slot_ref(logical_path)) -def pop_checkpoint(self: TrainerRank) -> None: +def pop_checkpoint(self: _impl.TrainerRank) -> None: with self._checkpoint_mutation_lock: if not self._slot_stack: raise RuntimeError("No pushed checkpoint to pop") From 7ce1cc81a02056e436bc596cfdddc56bd81f5e75 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 16:30:30 +0000 Subject: [PATCH 093/150] Remove stale imports and repeated extraction notes (cherry picked from commit 821bcf1c70d231218b599d308c35f79bbbac73d5) --- src/art/trainer_rank/_impl.py | 21 +------------------- src/art/trainer_rank/_micro_batch_planner.py | 14 +------------ src/art/trainer_rank/_optimizer.py | 10 ---------- src/art/trainer_rank/_slots.py | 10 ---------- 4 files changed, 2 insertions(+), 53 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 62bf04a5d..c12a6755f 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -16,10 +16,9 @@ from contextlib import contextmanager, nullcontext from contextvars import ContextVar from copy import deepcopy -from dataclasses import asdict, dataclass, fields, is_dataclass, replace +from dataclasses import dataclass, fields, is_dataclass, replace from dataclasses import field as dataclass_field from functools import lru_cache, partial -import hashlib import logging import math import os @@ -73,7 +72,6 @@ resolve_forward_options, ) from art.trainer_rank._planner_cost import ( - COEFFICIENT_VERSION_FALLBACK, ModelGeometry, ParallelShape, select_scoring, @@ -82,9 +80,6 @@ from art.trainer_rank._prefix_tree_planner import ( CanonicalPrefixTree, PrefixTreeLayout, - build_canonical_prefix_tree, - prefix_tree_layout_candidates, - select_prefix_tree_layout, ) from art.trainer_rank._rng import TrainerRNG, caller_group from art.trainer_rank._telemetry import phase as _telemetry_phase @@ -107,7 +102,6 @@ from art.megatron.train import TrainingRuntime from art.trainer_rank._checkpoint import ( CustomOptimizerState, - LocalOptimizerState, PreparedCheckpoint, PreparedCustomPayload, _FinalizedSave, @@ -5750,9 +5744,6 @@ def _gather_tensor_parallel_logits(self, logits: torch.Tensor) -> torch.Tensor: tensor_parallel.gather_from_tensor_model_parallel_region(logits), ) - # Memory estimation, profiling and admission accounting live in - # ``_memory``; binding the functions here keeps ``self._x(...)`` dispatch - # and per-instance overrides (tests monkeypatch these) behaving as before. _split_required_memory = staticmethod(_memory._split_required_memory) _split_memory_key = staticmethod(_memory._split_memory_key) _record_split_memory_floor = _memory._record_split_memory_floor @@ -5782,10 +5773,6 @@ def _gather_tensor_parallel_logits(self, logits: torch.Tensor) -> torch.Tensor: _all_ranks_have_memory_profile = _memory._all_ranks_have_memory_profile _update_memory_profile = _memory._update_memory_profile - # Micro-batch planning, split search and admission live in - # ``_micro_batch_planner``; binding the functions here keeps ``self._x(...)`` - # dispatch and per-instance overrides (tests monkeypatch these) behaving as - # before. _forward_batches = _micro_batch_planner._forward_batches _plan_admissible_forward = _micro_batch_planner._plan_admissible_forward _find_admissible_forward = _micro_batch_planner._find_admissible_forward @@ -5823,9 +5810,6 @@ 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 @@ -5858,9 +5842,6 @@ 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 diff --git a/src/art/trainer_rank/_micro_batch_planner.py b/src/art/trainer_rank/_micro_batch_planner.py index 9f1e253e5..1262197b2 100644 --- a/src/art/trainer_rank/_micro_batch_planner.py +++ b/src/art/trainer_rank/_micro_batch_planner.py @@ -1,16 +1,4 @@ -"""TrainerRank micro-batch planning, split search and admission. - -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``, -``time``, ``_telemetry_phase``, sibling helpers, plan/cost types) are still -resolved through ``_impl`` at call time, so tests that patch ``_impl.dist`` and -friends keep intercepting them. Only pure stdlib helpers and prefix-tree / -planner-cost functions that nothing patches are imported here directly. -Referencing ``_impl`` as a module also lets the circular import resolve lazily. -""" +"""TrainerRank micro-batch planning, split search and admission.""" from __future__ import annotations diff --git a/src/art/trainer_rank/_optimizer.py b/src/art/trainer_rank/_optimizer.py index ae45bd45e..eaf412f8f 100644 --- a/src/art/trainer_rank/_optimizer.py +++ b/src/art/trainer_rank/_optimizer.py @@ -1,16 +1,6 @@ """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 diff --git a/src/art/trainer_rank/_slots.py b/src/art/trainer_rank/_slots.py index 126f8e422..b053e7e65 100644 --- a/src/art/trainer_rank/_slots.py +++ b/src/art/trainer_rank/_slots.py @@ -1,15 +1,5 @@ """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 From 7c0d0b92266c0b85a0f547e76871c1a0a9093864 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 16:38:05 +0000 Subject: [PATCH 094/150] Preserve literal renderer oracles across explicit override modes --- .../trajectories/test_literal_thinking_off.py | 25 ++++++++++++++++++- 1 file changed, 24 insertions(+), 1 deletion(-) diff --git a/tests/unit/trajectories/test_literal_thinking_off.py b/tests/unit/trajectories/test_literal_thinking_off.py index bf621d702..01ec7c654 100644 --- a/tests/unit/trajectories/test_literal_thinking_off.py +++ b/tests/unit/trajectories/test_literal_thinking_off.py @@ -234,7 +234,15 @@ def test_literal_content_is_not_inferred_from_source_thinking_mode(case: str) -> if case == "visible_only": cast(dict[str, Any], history.messages[-1]).pop("reasoning") original = history.model_dump(mode="python") - tokenized = history.tokenize(tokenizer=tokenizer, chat_template=_RENDER_OVERRIDE) + tokenized = _tokenize._tokenize_chat_view( + history, + base_model=None, + tokenizer=tokenizer, + chat_template=None, + chat_template_kwargs=None, + _projection_matches=_tokenize._history_render_state(history).projection_matches, + _recorded_boundaries=False, + ) # Independent public transcript and source-field oracles, not a second run # with an already-inert workaround disabled. expected = ( @@ -257,6 +265,21 @@ def test_literal_content_is_not_inferred_from_source_thinking_mode(case: str) -> expected_sampled = "" if case in {"no_source", "request_source"} else _LITERAL assert tokenizer.decode([tokenized.tokens[i] for i in sampled]) == expected_sampled assert [tokenized.logprobs[i] for i in sampled] == [-0.5] * len(expected_sampled) + assert history.model_dump(mode="python") == original + tokenizer.calls.clear() + tokenizer.rendered.clear() + assert _outcome(history, tokenizer, chat_template=_RENDER_OVERRIDE) == ( + ( + ValueError, + "Could not locate a sampled history message in the rendered history", + ) + if structured + else ( + tokenized.tokens, + tokenized.flags, + [None if x != x else x for x in tokenized.logprobs], + ) + ) # Exercise rendering explicitly even when complete native output can bypass # it. Plain content stays literal independently of recorded/current thinking mode # and whether the message has complete native token metadata. Structured From b13b5f5c7150110af96ad7cefe6731ca6adefb5f Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 16:44:57 +0000 Subject: [PATCH 095/150] Respect prompt provenance in the explicit template override control --- tests/unit/trajectories/test_literal_thinking_off.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/tests/unit/trajectories/test_literal_thinking_off.py b/tests/unit/trajectories/test_literal_thinking_off.py index 01ec7c654..959ddce74 100644 --- a/tests/unit/trajectories/test_literal_thinking_off.py +++ b/tests/unit/trajectories/test_literal_thinking_off.py @@ -276,7 +276,11 @@ def test_literal_content_is_not_inferred_from_source_thinking_mode(case: str) -> if structured else ( tokenized.tokens, - tokenized.flags, + # The override does not certify the recorded prompt as exact. + [tr.TokenFlag(0)] * (len(expected) - len(_LITERAL)) + + tokenized.flags[len(expected) - len(_LITERAL) :] + if case == "visible_only" + else tokenized.flags, [None if x != x else x for x in tokenized.logprobs], ) ) From 8711cc5c4fa739f3ba8a94125dcd7268f5262516 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 17:14:28 +0000 Subject: [PATCH 096/150] Share recorded boundary invocation setup in tests --- .../trajectories/test_recorded_boundaries.py | 38 +++++++++---------- .../test_recorded_boundary_source_guard.py | 29 ++------------ 2 files changed, 21 insertions(+), 46 deletions(-) diff --git a/tests/unit/trajectories/test_recorded_boundaries.py b/tests/unit/trajectories/test_recorded_boundaries.py index 3d56af540..e3d2a462d 100644 --- a/tests/unit/trajectories/test_recorded_boundaries.py +++ b/tests/unit/trajectories/test_recorded_boundaries.py @@ -829,6 +829,20 @@ def test_complete_messages_records_do_not_require_a_chat_projection( assert trajectory.model_dump_json() == before +def _recorded_boundaries( + history: module.ChatCompletionsHistory, + tokenizer: Any, + render: module._ChatRender, +) -> module.TokenizedHistory | None: + return module._tokenize_recorded_chat_boundaries( + history, + [dict(message) for message in history.messages], + tokenizer=tokenizer, + render=render, + _trace=None, + ) + + def _boundary_render(tokenizer: Any) -> module._ChatRender: def render( selected_messages: list[dict[str, Any]], *, add_generation_prompt: bool @@ -856,13 +870,7 @@ def limited(tokens, **kwargs): return decode(tokens, **kwargs) monkeypatch.setattr(tokenizer, "decode", limited) - result = module._tokenize_recorded_chat_boundaries( - history, - [dict(message) for message in history.messages], - tokenizer=tokenizer, - render=_boundary_render(tokenizer), - _trace=None, - ) + result = _recorded_boundaries(history, tokenizer, _boundary_render(tokenizer)) assert trailing and result is None @@ -890,13 +898,7 @@ def fail(*args, **kwargs): monkeypatch.setattr(tokenizer, "decode", fail) render = fail if stage == "render" else _boundary_render(tokenizer) with pytest.raises(type(error)) as caught: - module._tokenize_recorded_chat_boundaries( - history, - [dict(message) for message in history.messages], - tokenizer=tokenizer, - render=render, - _trace=None, - ) + _recorded_boundaries(history, tokenizer, render) assert calls == [True] and caught.value is error @@ -916,11 +918,5 @@ def render(*args, **kwargs): raise AssertionError("should not reach rendering") with pytest.raises(ValueError, match="token_ids"): - module._tokenize_recorded_chat_boundaries( - history, - [dict(message) for message in history.messages], - tokenizer=tokenizer, - render=render, - _trace=None, - ) + _recorded_boundaries(history, tokenizer, render) assert not called diff --git a/tests/unit/trajectories/test_recorded_boundary_source_guard.py b/tests/unit/trajectories/test_recorded_boundary_source_guard.py index 094a8077d..fb3605765 100644 --- a/tests/unit/trajectories/test_recorded_boundary_source_guard.py +++ b/tests/unit/trajectories/test_recorded_boundary_source_guard.py @@ -2,7 +2,7 @@ from typing import Any import pytest -from test_recorded_boundaries import _boundary_render +from test_recorded_boundaries import _boundary_render, _recorded_boundaries from test_tokenize import _character_template_history from art.trajectories import _tokenize as module @@ -33,13 +33,7 @@ def mutate(*args: Any, **kwargs: Any): monkeypatch.setattr(tokenizer, "decode", mutate) with pytest.raises(RuntimeError if outcome == "fatal" else ValueError) as caught: - module._tokenize_recorded_chat_boundaries( - history, - [dict(message) for message in history.messages], - tokenizer=tokenizer, - render=_boundary_render(tokenizer), - _trace=None, - ) + _recorded_boundaries(history, tokenizer, _boundary_render(tokenizer)) if outcome == "fatal": assert caught.value is failure else: @@ -131,16 +125,7 @@ class Masquerading(metaclass=Meta): def forbidden(*args, **kwargs): pytest.fail("an unproved boundary must decline before rendering") - assert ( - module._tokenize_recorded_chat_boundaries( - history, - [dict(message) for message in history.messages], - tokenizer=tokenizer, - render=forbidden, - _trace=None, - ) - is None - ) + assert _recorded_boundaries(history, tokenizer, forbidden) is None assert not equality_calls if isinstance(context, list): context.clear() @@ -187,13 +172,7 @@ def render(selected_messages, *, add_generation_prompt): monkeypatch.setattr(ProjectedToolTokenizer, "__call__", observed) try: - value = module._tokenize_recorded_chat_boundaries( - history, - [dict(message) for message in history.messages], - tokenizer=tokenizer, - render=render, - _trace=None, - ) + value = _recorded_boundaries(history, tokenizer, render) except RuntimeError as error: assert behavior == "fatal" and error is fatal except ValueError as error: From 24bdb522ac35e022b5707c6a6373ca4bbf388bfe Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 21:47:27 +0000 Subject: [PATCH 097/150] Release registered forward graphs when correction setup fails --- src/art/trainer_rank/_impl.py | 111 +++++++++--------- .../test_trainer_rank_slot_graph_lifetime.py | 66 ++++++++++- 2 files changed, 123 insertions(+), 54 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index c12a6755f..540366d1f 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -3963,62 +3963,67 @@ def execute(captured: _ForwardGroupPlan) -> tuple[torch.Tensor, ...]: output_device=output_device, execution_peak_bytes=getattr(placement, "execution_peak_bytes", 0), ) - if topology.cp > 1 and retention != "replay": - residual = cache.state(handle).non_offloadable_bytes - if residual is not None: - profiles = getattr(self, "_graph_residency", None) - if profiles is None: - self._graph_residency = profiles = OrderedDict() - key = self._graph_residency_key(group) - profiles[key] = max(residual, profiles.pop(key, 0)) - if len(profiles) > 256: - profiles.popitem(last=False) - assert spec is not None - if version is not None: + try: + if topology.cp > 1 and retention != "replay": + residual = cache.state(handle).non_offloadable_bytes + if residual is not None: + profiles = getattr(self, "_graph_residency", None) + if profiles is None: + self._graph_residency = profiles = OrderedDict() + key = self._graph_residency_key(group) + profiles[key] = max(residual, profiles.pop(key, 0)) + if len(profiles) > 256: + profiles.popitem(last=False) + assert spec is not None + if version is not None: - @contextmanager - def current_context(): - current = self._capture_lora_version( - ref, options.max_gradient_staleness, origin=version.version - ) - assert current is not None - previous_storages = storages.copy() - storages.update( - (parameter.device, parameter.untyped_storage().data_ptr()) - for slot in current.slots.values() - for parameter in slot.parameters() + @contextmanager + def current_context(): + current = self._capture_lora_version( + ref, options.max_gradient_staleness, origin=version.version + ) + assert current is not None + previous_storages = storages.copy() + storages.update( + (parameter.device, parameter.untyped_storage().data_ptr()) + for slot in current.slots.values() + for parameter in slot.parameters() + ) + try: + with use_lora_slot(ref, version=current): + yield + finally: + storages.clear() + storages.update(previous_storages) + + cache.set_corrections( + handle, + capture_forward_corrections( + unflatten_tensors(spec, tensors), tensors, options + ), + is_stale=lambda: ( + self._capture_checkpoint_version( + version.version.checkpoint + ).revision + != version.weight_version.revision + ), + current_context_factory=current_context, ) - try: - with use_lora_slot(ref, version=current): - yield - finally: - storages.clear() - storages.update(previous_storages) - - cache.set_corrections( - handle, - capture_forward_corrections( - unflatten_tensors(spec, tensors), tensors, options - ), - is_stale=lambda: ( - self._capture_checkpoint_version( - version.version.checkpoint - ).revision - != version.weight_version.revision - ), - current_context_factory=current_context, + packet = TensorPacket( + handle, spec, tensors, tuple(tensor.requires_grad for tensor in tensors) ) - packet = TensorPacket( - handle, spec, tensors, tuple(tensor.requires_grad for tensor in tensors) - ) - outputs = self._forward_cotangent_collector().attach( - packet, - managed=output_device is not None, - on_release=partial(cache.release, handle), - ) - # Track the caller graph outside saved-state hooks, which detach markers. - # Consumption also ends the lifetime of unused sibling outputs. - return self._track_slot_graph_outputs(ref, outputs) + outputs = self._forward_cotangent_collector().attach( + packet, + managed=output_device is not None, + on_release=partial(cache.release, handle), + ) + # Track the caller graph outside saved-state hooks, which detach markers. + # Consumption also ends the lifetime of unused sibling outputs. + return self._track_slot_graph_outputs(ref, outputs) + except BaseException: + # No caller owns a failed handoff; a partial bridge may also release. + cache.release(handle) + raise def _forward_output_metadata( self, diff --git a/tests/unit/test_trainer_rank_slot_graph_lifetime.py b/tests/unit/test_trainer_rank_slot_graph_lifetime.py index ec872a982..f0be446fc 100644 --- a/tests/unit/test_trainer_rank_slot_graph_lifetime.py +++ b/tests/unit/test_trainer_rank_slot_graph_lifetime.py @@ -1,8 +1,11 @@ from __future__ import annotations +import asyncio from contextlib import nullcontext import gc +import traceback from typing import Literal +import weakref import pytest from test_trainer_rank_custom_tensors import _trainer @@ -43,7 +46,7 @@ def pack(tensor): yield saved if request.param == "detach" else None -def _forward(monkeypatch, retention, output_device, *, grad_enabled=True): +def _prepare_forward(monkeypatch, retention, output_device, *, grad_enabled=True): trainer, _ = _trainer("student") ref = trainer._slot_ref("student") weight = torch.nn.Parameter(torch.tensor(2.0)) @@ -77,10 +80,71 @@ def _forward(monkeypatch, retention, output_device, *, grad_enabled=True): ) ], ) + return trainer, ref, weight, group + + +def _forward(monkeypatch, retention, output_device, *, grad_enabled=True): + trainer, ref, weight, group = _prepare_forward( + monkeypatch, retention, output_device, grad_enabled=grad_enabled + ) output = trainer._execute_graph_group(group)[0] return trainer, ref, weight, output +@pytest.mark.parametrize("failure_type", [MemoryError, asyncio.CancelledError]) +def test_failed_correction_capture_releases_unattached_graph(monkeypatch, failure_type): + trainer, ref, weight, group = _prepare_forward(monkeypatch, "cpu", "cpu") + monkeypatch.setattr( + trainer, + "_forward_packed", + lambda items, prepared: [ForwardOutput(weight.square(), None, None, None)], + ) + cache = trainer._forward_graph_cache() + primary = failure_type("correction CPU snapshot failed") + primary.__cause__ = cause = RuntimeError("original cause") + saved, outputs, versions = [], [], [] + to = torch.Tensor.to + + def fail_snapshot(value, *args, **kwargs): + if args == ("cpu",) and kwargs.get("copy") and cache.handles(): + (handle,) = cache.handles() + record = cache._records[handle] + saved.extend(record.saved or ()) + outputs.extend(weakref.ref(output) for output in record.outputs or ()) + versions.extend( + weakref.ref(v) for v in trainer._version_state().lora.values() + ) + assert saved and outputs and len(versions) == 1 + assert record.checkpoint_versions == ( + trainer._capture_checkpoint_version("student"), + ) + raise primary + return to(value, *args, **kwargs) + + with monkeypatch.context() as allocation: + allocation.setattr(torch.Tensor, "to", fail_snapshot) + with pytest.raises(failure_type) as caught: + trainer._execute_graph_group(group) + assert caught.value is primary and primary.__cause__ is cause + assert not cache.handles() + assert all(reference() is None for reference in (*saved, *outputs)) + assert not trainer._has_live_slot_graph(ref) + # The original traceback owns function locals; the cache must not keep the + # checkpoint capture alive once those independent references are released. + traceback.clear_frames(primary.__traceback__) + gc.collect() + assert all(reference() is None for reference in versions) + assert not trainer._version_state().lora + trainer._guard_slot_can_load(ref) + assert weight.grad is None + output = trainer._execute_graph_group(group)[0].target_logprobs + assert output is not None and output.item() == 4 + assert len(cache.handles()) == 1 + trainer.backward(output) + torch.testing.assert_close(weight.grad, torch.tensor(4.0)) + assert not cache.handles() + + @pytest.mark.parametrize("retention", ["gpu", "cpu", "replay"]) def test_native_cached_slot_guard_tracks_retained_and_consumed_graph( monkeypatch, retention: Literal["gpu", "cpu", "replay"], output_device From 35498dea5c088ac46d5b0aa89276144ee983ee36 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 22:50:23 +0000 Subject: [PATCH 098/150] Compose trainer handoff and callback cleanup fixes --- .github/workflows/prek.yml | 2 + src/art/trainer_rank/_commands.py | 50 +++-- src/art/trainer_rank/_impl.py | 128 +++++++++---- tests/unit/test_trainer_rank_backward_work.py | 3 + tests/unit/test_trainer_rank_commands.py | 50 +++++ .../unit/test_trainer_rank_forward_handoff.py | 174 ++++++++++++++++++ tests/unit/test_trainer_rank_live_heads.py | 93 +++++++++- 7 files changed, 454 insertions(+), 46 deletions(-) create mode 100644 tests/unit/test_trainer_rank_forward_handoff.py diff --git a/.github/workflows/prek.yml b/.github/workflows/prek.yml index bc7145649..ab82f37cd 100644 --- a/.github/workflows/prek.yml +++ b/.github/workflows/prek.yml @@ -230,6 +230,7 @@ jobs: tests/unit/test_trainer_rank_validation.py \ tests/unit/test_trainer_rank_weird_shapes.py \ tests/unit/test_trainer_rank_slot_graph_lifetime.py \ + tests/unit/test_trainer_rank_forward_handoff.py \ tests/unit/test_trainer_rank_commands.py::test_gloo_dp2_tp2_participation_and_gradients \ tests/unit/test_trainer_rank_admission_inputs.py \ tests/unit/test_trainer_rank_checkpoint_memory.py \ @@ -281,6 +282,7 @@ jobs: --ignore=tests/unit/test_trainer_rank_validation.py \ --ignore=tests/unit/test_trainer_rank_weird_shapes.py \ --ignore=tests/unit/test_trainer_rank_slot_graph_lifetime.py \ + --ignore=tests/unit/test_trainer_rank_forward_handoff.py \ --ignore=tests/unit/test_trainer_rank_admission_inputs.py \ --ignore=tests/unit/test_trainer_rank_checkpoint_memory.py \ --ignore=tests/unit/test_trainer_rank_profile_warm.py \ diff --git a/src/art/trainer_rank/_commands.py b/src/art/trainer_rank/_commands.py index dbcc5ab58..849fb4434 100644 --- a/src/art/trainer_rank/_commands.py +++ b/src/art/trainer_rank/_commands.py @@ -1137,6 +1137,18 @@ def reduce(self, tensor: torch.Tensor, **kwargs: Any) -> None: ) +@contextmanager +def _preserve_callback_error(primary: BaseException | None) -> Iterator[None]: + try: + yield + except BaseException as cleanup: + if primary is None: + raise + _impl.TrainerRank._memory_error_with_reduction_note( + primary, cleanup, operation="callback cleanup" + ) + + async def run_rank_callback( rank: TrainerRank, callback: Callable[[Any], Any], *, mode: Mode = "rank" ) -> RankCallbackResult: @@ -1149,6 +1161,7 @@ async def run_rank_callback( raise cancelled return RankCallbackResult(None) view = _view(executor) + primary: BaseException | None = None try: if cancelled is not None: raise cancelled @@ -1159,11 +1172,15 @@ async def run_rank_callback( if inspect.isgenerator(result) or inspect.isasyncgen(result): raise TypeError("Use run_rank_callback_stream for generator callbacks") return RankCallbackResult(0 if mode == "zero" else executor.dp_rank, result) + except BaseException as error: + primary = error + raise finally: - try: - view._flush_heads() - finally: - executor.stop() + with _preserve_callback_error(primary): + try: + view._flush_heads() + finally: + executor.stop() async def run_rank_callback_stream( @@ -1180,6 +1197,7 @@ async def run_rank_callback_stream( yield RankCallbackResult(None) return iterator = value = None + primary: BaseException | None = None view = _view(executor) async with executor.release_on_exit(): try: @@ -1207,13 +1225,19 @@ async def run_rank_callback_stream( sent = yield RankCallbackResult( 0 if mode == "zero" else executor.dp_rank, value ) + except GeneratorExit: + raise + except BaseException as error: + primary = error + raise finally: - try: - if inspect.isasyncgen(iterator): - await iterator.aclose() - elif inspect.isgenerator(iterator): - iterator.close() - view._flush_heads() - finally: - value = None - executor.stop() + with _preserve_callback_error(primary): + try: + if inspect.isasyncgen(iterator): + await iterator.aclose() + elif inspect.isgenerator(iterator): + iterator.close() + view._flush_heads() + finally: + value = None + executor.stop() diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 540366d1f..4a00d2828 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -3048,6 +3048,49 @@ def _execute_admitted_plan( self._complete_planner_observation(phase="forward") return outputs + @contextmanager + def _forward_handoff(self, *, advance: bool) -> Iterator[None]: + # Only intermediate TP/CP frontiers can race the next model collective. + group = caller_group() if advance else None + if group is None or dist.get_world_size(group) == 1: + yield + return + error: BaseException | None = None + try: + yield + except BaseException as exc: + error = exc + try: + (failed,) = self._recovery_reduce( + [float(error is not None)], op="MAX", sync_across_dp=False + ) + except BaseException as exchange_error: + if error is None: + raise + self._memory_error_with_reduction_note( + error, exchange_error, operation="forward handoff" + ) + else: + if error is None and failed: + raise RuntimeError("Forward handoff failed on another rank") + if error is not None: + raise error + + def _discard_forward_graphs( + self, previous: tuple[str, ...], error: BaseException + ) -> None: + cache = getattr(self, "_graph_cache", None) + if cache is not None: + previous_handles = set(previous) + for handle in cache.handles(): + if handle not in previous_handles: + try: + cache.release(handle) + except BaseException as cleanup_error: + self._memory_error_with_reduction_note( + error, cleanup_error, operation="forward graph release" + ) + @_backward_region def _execute_split_plan_with_memory_tracking( self, plan: _SplitForwardPlan, *, check: _MemoryCheck, context: str @@ -3055,40 +3098,51 @@ def _execute_split_plan_with_memory_tracking( state = self._recovery_state() work_before = state.work self._begin_planner_observation(plan, check) + previous = self._graph_cache.handles() if hasattr(self, "_graph_cache") else () + outputs: list[AnyForwardOutput] = [] + output: AnyForwardOutput | None = None + merged: list[AnyForwardOutput | None] = [] self._planner_observing_split = True try: baseline, peak = None, 0 - merged: list[AnyForwardOutput | None] = [None] * plan.request_count + merged = [None] * plan.request_count for ordinal, (subforward, indices) in enumerate( zip(plan.subforwards, plan.request_indices, strict=True) ): - try: - outputs, child_baseline = self._run_flat_plan_with_memory_tracking( - subforward, check=check, context=context - ) - if child_baseline is not None: - if baseline is None: - baseline = child_baseline - peak = max( - peak, int(torch.cuda.max_memory_allocated(self.device)) + with self._forward_handoff(advance=ordinal + 1 < plan.subforward_count): + try: + outputs, child_baseline = ( + self._run_flat_plan_with_memory_tracking( + subforward, check=check, context=context + ) ) - except TrainerRankMemoryError as error: - # Model execution already began, so no replanning is possible - # and the caller must not mistake this for an up-front refusal. - raise TrainerRankPartialExecutionError( - f"{context}: subforward {ordinal + 1} of " - f"{plan.subforward_count} failed during execution " - f"({ordinal} of {plan.subforward_count} completed). {error}", - predicted_peak_bytes=error.predicted_peak_bytes, - usable_limit_bytes=error.usable_limit_bytes, - suggestion=error.suggestion, - ) from error - for index, output in zip(indices, outputs, strict=True): - merged[index] = output + if child_baseline is not None: + if baseline is None: + baseline = child_baseline + peak = max( + peak, int(torch.cuda.max_memory_allocated(self.device)) + ) + except TrainerRankMemoryError as error: + # Model execution already began, so no replanning is possible + # and the caller must not mistake this for an up-front refusal. + raise TrainerRankPartialExecutionError( + f"{context}: subforward {ordinal + 1} of " + f"{plan.subforward_count} failed during execution " + f"({ordinal} of {plan.subforward_count} completed). {error}", + predicted_peak_bytes=error.predicted_peak_bytes, + usable_limit_bytes=error.usable_limit_bytes, + suggestion=error.suggestion, + ) from error + for index, output in zip(indices, outputs, strict=True): + merged[index] = output if any(output is None for output in merged): raise AssertionError("split execution did not cover every request") return cast(list[AnyForwardOutput], merged), baseline, peak - except BaseException: + except BaseException as error: + outputs.clear() + merged.clear() + output = None + self._discard_forward_graphs(previous, error) state.work = work_before raise finally: @@ -3811,16 +3865,26 @@ def _execute_flat_plan(self, plan: _FlatForwardPlan) -> list[AnyForwardOutput]: if plan.groups else None ) + previous = self._graph_cache.handles() if hasattr(self, "_graph_cache") else () + item_outputs: list[AnyForwardOutput] = [] + output: AnyForwardOutput | None = None try: for group_index, group in enumerate(plan.groups): - if hybridep is not None: - self._set_hybridep_rows(hybridep[0][group_index]) - with torch.set_grad_enabled(group.grad_enabled): - item_outputs = self._execute_graph_group(group) - for index, output in zip( - group.request_indices, item_outputs, strict=True - ): - outputs[index] = output + with self._forward_handoff(advance=group_index + 1 < len(plan.groups)): + if hybridep is not None: + self._set_hybridep_rows(hybridep[0][group_index]) + with torch.set_grad_enabled(group.grad_enabled): + item_outputs = self._execute_graph_group(group) + for index, output in zip( + group.request_indices, item_outputs, strict=True + ): + outputs[index] = output + except BaseException as error: + outputs.clear() + item_outputs.clear() + output = None + self._discard_forward_graphs(previous, error) + raise finally: if hybridep is not None: self._set_hybridep_rows(hybridep[1]) diff --git a/tests/unit/test_trainer_rank_backward_work.py b/tests/unit/test_trainer_rank_backward_work.py index 02641dbf2..b7647e247 100644 --- a/tests/unit/test_trainer_rank_backward_work.py +++ b/tests/unit/test_trainer_rank_backward_work.py @@ -168,6 +168,8 @@ def actual_rank(self): "_release_cached_memory_for_backward", "_memory_error_with_reduction_note", "_execute_split_plan_with_memory_tracking", + "_forward_handoff", + "_discard_forward_graphs", "_begin_planner_observation", } selected = [ @@ -191,6 +193,7 @@ def actual_rank(self): dataclass_field=field, threading=threading, contextmanager=contextmanager, + caller_group=lambda: None, traceback=traceback, BackwardWork=self.module.BackwardWork, _backward_region=self.module.region, diff --git a/tests/unit/test_trainer_rank_commands.py b/tests/unit/test_trainer_rank_commands.py index 4c32e0744..bc1de179a 100644 --- a/tests/unit/test_trainer_rank_commands.py +++ b/tests/unit/test_trainer_rank_commands.py @@ -812,3 +812,53 @@ def test_nested_aggregate_outputs_admit_before_any_model_copy(): request = replace(request, options=ForwardOptions(output_device="model")) with pytest.raises(MemoryError): view._place_outputs([(output, [[request]]) for output in outputs]) + + +@pytest.mark.parametrize("asynchronous", [False, True]) +@pytest.mark.parametrize("error_type", [ValueError, asyncio.CancelledError, None]) +def test_stream_close_keeps_primary_and_stops_executor(asynchronous, error_type): + rank: Any = _Rank() + primary = error_type("consumer failure") if error_type is not None else None + cleanup = LookupError("iterator close failed") + views, flushed = [], [] + + def prepare(view): + views.append(view) + view._flush_heads = lambda: flushed.append(True) + + def callback(view): + prepare(view) + try: + yield 1 + finally: + raise cleanup + + async def async_callback(view): + prepare(view) + try: + yield 1 + finally: + raise cleanup + + async def run(): + stream = run_rank_callback_stream( + rank, async_callback if asynchronous else callback + ) + assert (await anext(stream)).value == 1 + with pytest.raises( + error_type if primary is not None else LookupError + ) as caught: + if primary is None: + await stream.aclose() + else: + await stream.athrow(primary) + assert caught.value is (cleanup if primary is None else primary) + if primary is not None: + assert any("iterator close failed" in note for note in primary.__notes__) + assert views[0]._executor.stopped and not flushed + await stream.aclose() + assert ( + await run_rank_callback(rank, lambda view: view.optim_step()) + ).value == {"steps": 1} + + asyncio.run(run()) diff --git a/tests/unit/test_trainer_rank_forward_handoff.py b/tests/unit/test_trainer_rank_forward_handoff.py new file mode 100644 index 000000000..b67bf9add --- /dev/null +++ b/tests/unit/test_trainer_rank_forward_handoff.py @@ -0,0 +1,174 @@ +"""A local handoff failure must precede the next model-parallel forward.""" + +import asyncio +from dataclasses import replace +import gc +import traceback +import weakref + +import pytest +from test_trainer_rank_slot_graph_lifetime import _prepare_forward +import torch +import torch.distributed as dist +from trainer_rank_test_support import gloo_group, megatron_topology, spawn_and_join + +from art.trainer_rank import ForwardOutput +from art.trainer_rank._impl import ( + _FlatForwardPlan, + _MemoryCheck, + _MemorySignature, + _SplitForwardPlan, +) + + +@pytest.mark.parametrize("split", [False, True], ids=["groups", "split"]) +def test_failed_handoff_precedes_next_physical_forward(tmp_path, monkeypatch, split): + monkeypatch.setenv("CUDA_VISIBLE_DEVICES", "") + spawn_and_join( + _handoff_worker, + (f"file://{tmp_path / 'handoff'}", split), + timeout=180, + failure="Handoff failure stranded a peer in the next physical forward", + ) + + +def _handoff_worker(physical, rendezvous, split): + with pytest.MonkeyPatch.context() as patch: + # Load the real native LoRA types before the callback topology shim. + trainer, ref, weight, group = _prepare_forward(patch, "cpu", "cpu") + with ( + gloo_group(physical, rendezvous, timeout=10), + megatron_topology(physical, dp_size=1, tp_size=2), + ): + calls = 0 + + def forward(items, prepared): + nonlocal calls + calls += 1 + # This real collective models the next TP/CP model boundary. + value = torch.ones(1) + dist.all_reduce(value) + assert value.item() == 2 + return [ForwardOutput(weight.square(), None, None, None)] + + patch.setattr(trainer, "_forward_packed", forward) + signature = _MemorySignature( + (1, 2, 1, 1), (0, None), 3, (), True, (True,) * 3 + ) + flat = _FlatForwardPlan( + 3, + (("student", False),) * 3, + tuple(replace(group, request_indices=(i,)) for i in range(3)), + 6, + 6, + 12, + signature, + ) + plan = ( + _SplitForwardPlan( + tuple( + replace( + flat, + request_count=1, + output_metadata=(("student", False),), + groups=(group,), + ) + for _ in range(3) + ), + ((0,), (1,), (2,)), + 3, + ) + if split + else flat + ) + patch.setattr( + trainer, + "_plan_admissible_forward", + lambda *a, **k: (plan, _MemoryCheck(12, 100, True)), + ) + inputs = [group.items[0].request] * 3 + cache = trainer._forward_graph_cache() + original_to = torch.Tensor.to + for fail_at, kind in ((1, MemoryError), (2, asyncio.CancelledError)): + preserved = trainer._execute_graph_group(group)[0].target_logprobs + assert preserved is not None + original_handles = cache.handles() + calls = 0 + primary = kind("injected correction copy failure") + primary.__cause__ = cause = ValueError("original cause") + references = [] + observed = set() + + def copy(value, *args, **kwargs): + if args == ("cpu",) and kwargs.get("copy"): + for handle in ( + set(cache.handles()) - set(original_handles) - observed + ): + record = cache._records[handle] + references.extend(record.saved or ()) + references.extend( + weakref.ref(t) for t in record.outputs or () + ) + observed.add(handle) + if physical == 1 and calls == fail_at: + raise primary + return original_to(value, *args, **kwargs) + + with patch.context() as allocation: + allocation.setattr(torch.Tensor, "to", copy) + with pytest.raises( + kind if physical == 1 else RuntimeError + ) as caught: + trainer.forward(inputs) + assert calls == fail_at, (calls, fail_at, repr(caught.value)) + if physical == 1: + assert caught.value is primary and primary.__cause__ is cause + else: + assert "another rank" in str(caught.value).lower(), repr( + caught.value + ) + assert cache.handles() == original_handles + assert trainer._has_live_slot_graph(ref) + assert len(trainer._slot_graphs()[ref]) == 1 + assert observed and all(reference() is None for reference in references) + traceback.clear_frames(caught.value.__traceback__) + gc.collect() + assert weight.grad is None + calls = 0 + outputs = trainer.forward(inputs) + assert calls == 3 + targets = [output.target_logprobs for output in outputs] + assert all(value is not None and value.item() == 4 for value in targets) + trainer.backward( + torch.stack([value for value in targets if value is not None]).sum() + ) + torch.testing.assert_close(weight.grad, torch.tensor(12.0)) + assert cache.handles() == original_handles + trainer.backward(preserved) + torch.testing.assert_close(weight.grad, torch.tensor(16.0)) + trainer.zero_grad() + # The fixture's analytic parameter is independent of checkpoint slots. + weight.grad = None + assert not cache.handles() and not trainer._has_live_slot_graph(ref) + + single = replace( + flat, + request_count=1, + output_metadata=(("student", False),), + groups=(group,), + ) + patch.setattr( + trainer, + "_plan_admissible_forward", + lambda *a, **k: (single, _MemoryCheck(4, 100, True)), + ) + patch.setattr( + trainer, + "_recovery_reduce", + lambda *a, **k: pytest.fail("single forward added a handoff reduction"), + ) + output = trainer.forward(inputs[:1])[0].target_logprobs + assert output is not None and output.item() == 4 + trainer.backward(output) + torch.testing.assert_close(weight.grad, torch.tensor(4.0)) + assert not cache.handles() diff --git a/tests/unit/test_trainer_rank_live_heads.py b/tests/unit/test_trainer_rank_live_heads.py index eb87e7d89..a0f4d1b06 100644 --- a/tests/unit/test_trainer_rank_live_heads.py +++ b/tests/unit/test_trainer_rank_live_heads.py @@ -12,7 +12,12 @@ from torch.utils.checkpoint import checkpoint from trainer_rank_test_support import gloo_group, megatron_topology, spawn_and_join -from art.trainer_rank import AdamParams, ModuleHandle, run_rank_callback +from art.trainer_rank import ( + AdamParams, + ModuleHandle, + run_rank_callback, + run_rank_callback_stream, +) from art.trainer_rank._commands import join_rank_callback_release from art.trainer_rank._heads import ( HeadRegistration, @@ -1394,3 +1399,89 @@ def test_cuda_buffer_sync_stages_cpu_authority_before_comparison(tmp_path): synchronize_head_buffers(trainer) torch.testing.assert_close(buffer, torch.ones(2, device="cuda")) assert export_head(trainer, "student", "mean").buffer_revision == 0 + + +@pytest.mark.parametrize( + "style,error_type", + [ + (style, error_type) + for style in ("callback", "sync", "async") + for error_type in (ValueError, asyncio.CancelledError, None) + ] + + [("close_sync", None), ("close_async", None)], +) +def test_callback_primary_survives_rejected_buffer_publication(style, error_type): + async def run(): + factory = lambda: torch.tensor(1.0) + trainer, native = _native_head("buffer", "head", factory) + primary = error_type("callback primary") if error_type is not None else None + cause, ambient = KeyError("original cause"), LookupError("ambient exception") + views, retained, closed = [], [], [] + + def callback(view): + views.append(view) + logical = view.buffer("head", factory, checkpoint="student") + retained.append(logical) + logical.add_(1) + native.add_(5) + if primary is not None: + raise primary from cause + return 7 + + def sync_stream(view): + try: + yield callback(view) + finally: + closed.append(True) + + async def async_stream(view): + try: + yield callback(view) + finally: + closed.append(True) + + try: + raise ambient + except LookupError: + expected = error_type if primary is not None else RuntimeError + with pytest.raises(expected) as caught: + if style == "callback": + await run_rank_callback(trainer, callback) + else: + stream = run_rank_callback_stream( + trainer, + async_stream if style.endswith("async") else sync_stream, + ) + try: + assert (await anext(stream)).value == 7 + if style.startswith("close"): + await stream.aclose() + else: + await anext(stream) + finally: + await stream.aclose() + if primary is not None: + assert caught.value is primary + assert primary.__cause__ is cause and primary.__context__ is ambient + assert any( + "Secondary callback cleanup failure" in note + and "changed before publication" in note + for note in primary.__notes__ + ) + else: + assert "changed before publication" in str(caught.value) + assert views[0]._executor.stopped + assert closed == ([] if style == "callback" else [True]) + assert native.item() == 6 + with pytest.raises(RuntimeError, match="publication failed"): + retained[0].item() + + def recover(view): + fresh = view.buffer("head", factory, checkpoint="student") + assert fresh is not retained[0] and fresh.item() == 6 + fresh.add_(2) + + await run_rank_callback(trainer, recover) + assert native.item() == 8 + + asyncio.run(run()) From f002b100c445868a842470dac6e07b9ecec52501 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 23:33:10 +0000 Subject: [PATCH 099/150] Preserve placed report completeness and failed graph ownership --- src/art/trainer_rank/_graphs.py | 46 +++++++---- src/art/trainer_rank/_micro_batch_planner.py | 2 + tests/unit/test_grouped_planner_replay.py | 69 ++++++++++++++++ tests/unit/test_trainer_rank_graphs.py | 87 ++++++++++++++++++++ 4 files changed, 189 insertions(+), 15 deletions(-) diff --git a/src/art/trainer_rank/_graphs.py b/src/art/trainer_rank/_graphs.py index e6aa39acf..72d616853 100644 --- a/src/art/trainer_rank/_graphs.py +++ b/src/art/trainer_rank/_graphs.py @@ -232,6 +232,17 @@ class _ForwardRecord: ) execution_peak_bytes: int = 0 + def release(self) -> None: + # Saved-variable hooks can outlive their Python outputs. Break all + # ownership edges even when a caller retains a failure traceback. + self.outputs = self.saved = self.resident = None + self.restored.clear() + self.inputs = self.corrections = None + self.execute = lambda _: () + self.context_factory = nullcontext + self.validate_backward = self.current_context_factory = None + self.is_stale = self.keep_on_device = None + def run( self, *, @@ -370,12 +381,25 @@ def run( for value in physical ) # clone: a detached view may still pin a much larger model activation. - detached = tuple( - value.detach() - .to(device=output_device, copy=True) - .requires_grad_(value.requires_grad) - for value in physical - ) + copied: list[torch.Tensor] = [] + value = None + try: + for value in physical: + copied.append( + value.detach() + .to(device=output_device, copy=True) + .requires_grad_(value.requires_grad) + ) + detached = tuple(copied) + except BaseException: + record.release() + copied.clear() + # Unlike a generator expression, these output aliases can be + # cleared while preserving the caller's original traceback. + del record, physical, value + del execute, inputs, context_factory, validate_backward + del checkpoint_versions, options, rng_tracker, keep_on_device + raise handle = uuid4().hex self._records[handle] = record if retention == "replay": @@ -476,15 +500,7 @@ def evict( def release(self, handle: ForwardHandle) -> None: if (record := self._records.pop(handle, None)) is not None: - # Saved-variable hooks can outlive their Python outputs. Break all - # ownership edges even when a caller retains a failure traceback. - record.outputs = record.saved = record.resident = None - record.restored.clear() - record.inputs = record.corrections = None - record.execute = lambda _: () - record.context_factory = nullcontext - record.validate_backward = record.current_context_factory = None - record.is_stale = record.keep_on_device = None + record.release() def backward( self, diff --git a/src/art/trainer_rank/_micro_batch_planner.py b/src/art/trainer_rank/_micro_batch_planner.py index 29f653ffc..c84f0b3f0 100644 --- a/src/art/trainer_rank/_micro_batch_planner.py +++ b/src/art/trainer_rank/_micro_batch_planner.py @@ -1409,6 +1409,8 @@ def _fill_planner_snapshot( missing = [ f"runtime_facts_unavailable:{_planner_replay.refusal_reason(error)}" ] + if any(g.memory_placement is not None for g in child.groups): + missing.append("graph_placement_admission_unavailable") estimates.append( { "signature": asdict(child.signature), diff --git a/tests/unit/test_grouped_planner_replay.py b/tests/unit/test_grouped_planner_replay.py index 575fafa54..9d318b0c5 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -77,6 +77,75 @@ def test_selected_grouped_cost_recomputed( assert not torch.cuda.is_initialized() +@pytest.mark.parametrize("split", [False, True]) +def test_placement_admission_with_runtime_facts_stays_incomplete( + split, monkeypatch, tmp_path +): + from art.trainer_rank import ForwardOptions, _planner_evidence + + rank = head_rank() + rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) + monkeypatch.setattr(rank, "_available_memory_bytes", lambda *_args: 10**12) + monkeypatch.setattr(rank, "_available_cpu_memory_bytes", lambda: 10**12) + options = ForwardOptions( + backward_state="replay" if split else "gpu", + output_device="cpu" if split else "model", + ) + items = [replace(request(rows, grad=True), options=options) for rows in (65, 33)] + plan = ( + tr._SplitForwardPlan( + tuple(rank._plan_flat_forward([item]) for item in items), ((0,), (1,)), 2 + ) + if split + else rank._plan_flat_forward(items) + ) + children = plan.subforwards if isinstance(plan, tr._SplitForwardPlan) else (plan,) + unplaced = rank._split_required_memory([rank._plan_cost(p) for p in children]) + with _planner_evidence.scope( + _planner_evidence.Decision("forward", sync_across_dp=False, owner=rank) + ): + plan, check = rank._admit_graph_memory(plan) + assert check.fits and check.cpu_fits and check.cpu_required_bytes > 0 + assert check.sample is not None + local = check.sample.local_required_bytes + assert local is not None + assert local == check.estimated_required_bytes and local != unplaced + rank._begin_planner_observation( + plan, replace(check, estimated_required_bytes=local + 123) + ) + observation = rank._planner_observation + assert observation is not None and observation["comparable"] + path = rank._planner_reporter.report( + predicted_peak_bytes=observation["predicted"], + observed_peak_bytes=local * 2, + phase="forward", + admission_peak_bytes=local + 123, + replay_factory=observation["replay"], + ) + assert path is not None + rank.finish_planner_observation() + report = reports.validate_report(path.read_bytes()) + payload = report["replay"] + assert payload["local_admission_peak_bytes"] == local + assert payload["reduced_admission_peak_bytes"] == local + 123 + assert report["predicted_peak_bytes"] == round(local / tr._MEMORY_SAFETY_FACTOR) + assert payload["requests"][0]["options"]["backward_state"] == options.backward_state + assert payload["requests"][0]["options"]["output_device"] == options.output_device + assert [r["target_tokens"] for r in payload["requests"]] == [ + list(range(65)), + list(range(33)), + ] + assert all( + e["runtime_facts"] is not None for e in payload["memory_replay"]["estimates"] + ) + # Without the placement boundary, unchanged-source replay reconstructs every + # model estimate but incorrectly claims it can reproduce placement admission. + assert not report["replay_complete"], reports.replay(report) + assert report["incomplete_reasons"] == ["graph_placement_admission_unavailable"] + with pytest.raises(ValueError, match="graph_placement_admission_unavailable"): + reports.replay(report) + + def test_runtime_dimensions_change_recomputed_cost_not_expected_answer(tmp_path): rank = head_rank() rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) diff --git a/tests/unit/test_trainer_rank_graphs.py b/tests/unit/test_trainer_rank_graphs.py index 7a7658021..3c3806ea1 100644 --- a/tests/unit/test_trainer_rank_graphs.py +++ b/tests/unit/test_trainer_rank_graphs.py @@ -1,5 +1,6 @@ from __future__ import annotations +import asyncio from contextlib import contextmanager, nullcontext import gc from types import SimpleNamespace @@ -506,3 +507,89 @@ def failing(x): assert snapshot.grad is None assert not trainer._version_state()._origins assert cache.handles() == () + + +@pytest.mark.parametrize("failure_type", [MemoryError, asyncio.CancelledError]) +@pytest.mark.parametrize("retention", ["gpu", "cpu", "replay"]) +@pytest.mark.parametrize("copy_index", [1, 2]) +def test_initial_output_copy_failure_releases_only_failed_graph( + monkeypatch, failure_type, retention, copy_index +): + from art.trainer_rank._graphs import _ForwardRecord + + trainer = TrainerRank.__new__(TrainerRank) + parameter = torch.nn.Parameter(torch.tensor(2.0)) + trainer._checkpoint_slots = {"student": _CheckpointSlot(params=(parameter,))} + trainer.runtime = SimpleNamespace(model=[], optimizer=None) + cache = GraphCache() + older = torch.nn.Parameter(torch.tensor(3.0)) + old_handle, (old_output,) = cache.run(lambda _: (older.square(),), None) + old_record = weakref.ref(cache._records[old_handle]) + records, physical, saved, snapshots, versions, copies = [], [], [], [], [], [] + primary = failure_type("initial output copy failed") + primary.__cause__ = cause = RuntimeError("original cause") + run, to = _ForwardRecord.run, torch.Tensor.to + attempted = 0 + + def arguments(): + version = trainer._capture_checkpoint_version("student") + snapshot = trainer._snapshot_parameter(parameter, version) + snapshots.append(weakref.ref(snapshot)) + versions.append(weakref.ref(version)) + return dict( + execute=lambda x: (snapshot * x, snapshot.square()), + inputs=torch.tensor(3.0), + context_factory=lambda: nullcontext(snapshot), + validate_backward=lambda: trainer._version_state().validate(version), + checkpoint_versions=(version,), + keep_on_device=lambda value: value is snapshot, + retention=retention, + ) + + def observe(record): + outputs = run(record) + records.append(weakref.ref(record)) + physical.extend(weakref.ref(output) for output in outputs) + saved.extend(record.saved or ()) + return outputs + + def fail_copy(value, *args, **kwargs): + nonlocal attempted + if kwargs.get("copy") and "device" in kwargs: + attempted += 1 + if attempted == copy_index: + # The injected hook must not add its own tensor owner to the + # retained traceback; only ART's failed attempt is under test. + del value + raise primary + result = to(value, *args, **kwargs) + copies.append(weakref.ref(result)) + return result + return to(value, *args, **kwargs) + + with monkeypatch.context() as failure: + failure.setattr(_ForwardRecord, "run", observe) + failure.setattr(torch.Tensor, "to", fail_copy) + with pytest.raises(failure_type) as caught: + cache.run(**arguments()) + assert caught.value is primary and primary.__cause__ is cause + assert primary.__traceback__ is not None and attempted == copy_index + assert len(records) == len(snapshots) == len(versions) == 1 + assert len(physical) == 2 and saved and len(copies) == copy_index - 1 + assert cache.handles() == (old_handle,) + assert cache._records[old_handle] is old_record() + assert all( + reference() is None + for reference in (*records, *physical, *saved, *snapshots, *versions, *copies) + ) + assert parameter.grad is None + assert old_output.item() == 9 and old_output.requires_grad + cache.backward(old_handle, (torch.tensor(1.0),)) + torch.testing.assert_close(older.grad, torch.tensor(6.0)) + + handle, outputs = cache.run(**arguments()) + assert tuple(output.item() for output in outputs) == (6, 4) + with trainer._gradient_transaction(): + cache.backward(handle, (torch.tensor(1.0), torch.tensor(1.0))) + torch.testing.assert_close(parameter.grad, torch.tensor(7.0)) + assert cache.handles() == () From d758c0a994eecc769a3d11d6c73d785d56b3bd2e Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 00:19:37 +0000 Subject: [PATCH 100/150] Release failed native forward capture aliases --- src/art/trainer_rank/_impl.py | 44 ++++--- .../test_trainer_rank_slot_graph_lifetime.py | 114 +++++++++++++++++- 2 files changed, 140 insertions(+), 18 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 4a00d2828..3dc2e4706 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -4011,23 +4011,24 @@ def execute(captured: _ForwardGroupPlan) -> tuple[torch.Tensor, ...]: else torch.cuda.current_device(), ) cache = self._forward_graph_cache() - handle, tensors = cache.run( - execute, - group, - context_factory=lambda: use_lora_slot(ref, version=version), - validate_backward=None if version is None else version.validate, - retention=retention, - checkpoint_versions=() if version is None else (version.version,), - options=options, - cuda_devices=devices, - rng_tracker=tracker, - keep_on_device=lambda tensor: ( - (tensor.device, tensor.untyped_storage().data_ptr()) in storages - ), - output_device=output_device, - execution_peak_bytes=getattr(placement, "execution_peak_bytes", 0), - ) + handle = None try: + handle, tensors = cache.run( + execute, + group, + context_factory=lambda: use_lora_slot(ref, version=version), + validate_backward=None if version is None else version.validate, + retention=retention, + checkpoint_versions=() if version is None else (version.version,), + options=options, + cuda_devices=devices, + rng_tracker=tracker, + keep_on_device=lambda tensor: ( + (tensor.device, tensor.untyped_storage().data_ptr()) in storages + ), + output_device=output_device, + execution_peak_bytes=getattr(placement, "execution_peak_bytes", 0), + ) if topology.cp > 1 and retention != "replay": residual = cache.state(handle).non_offloadable_bytes if residual is not None: @@ -4086,7 +4087,16 @@ def current_context(): return self._track_slot_graph_outputs(ref, outputs) except BaseException: # No caller owns a failed handoff; a partial bridge may also release. - cache.release(handle) + try: + if handle is not None: + cache.release(handle) + finally: + # A retained traceback must not own this failed call's captures. + # Clear its closure cell, never the shared version or live graphs. + del version + packet = outputs = None + tensors = () + parameters.clear() raise def _forward_output_metadata( diff --git a/tests/unit/test_trainer_rank_slot_graph_lifetime.py b/tests/unit/test_trainer_rank_slot_graph_lifetime.py index f0be446fc..9515724f4 100644 --- a/tests/unit/test_trainer_rank_slot_graph_lifetime.py +++ b/tests/unit/test_trainer_rank_slot_graph_lifetime.py @@ -8,7 +8,7 @@ import weakref import pytest -from test_trainer_rank_custom_tensors import _trainer +from test_trainer_rank_custom_tensors import _real_lora_trainer, _trainer import torch from art.megatron.context_parallel.types import ParallelTopology @@ -145,6 +145,118 @@ def fail_snapshot(value, *args, **kwargs): assert not cache.handles() +@pytest.mark.parametrize( + "phase,shared", + [ + (phase, shared) + for phase in ("copy", "correction", "handoff") + for shared in (False, True) + ] + + [("cancel", False)], +) +def test_native_failed_handoff_drops_only_its_capture_owners( + monkeypatch, phase, shared +): + trainer, ref, _, group = _prepare_forward(monkeypatch, "cpu", "cpu") + real, _ = _real_lora_trainer() + from art.megatron.lora import LoRA + + trainer.runtime, trainer._checkpoint_slots = real.runtime, real._checkpoint_slots + lora = trainer.runtime.model[0] + assert isinstance(lora, LoRA) + current = lora.lora_slot_params(ref) + with torch.no_grad(): + current[0].fill_(1) + current[1].fill_(2) + versions, snapshots, delivered = [], [], [] + + def forward(items, prepared): + active = lora.active_lora_tensors() + assert active is not None + a, b, _ = active + assert a is not current[0] and b is not current[1] + versions.extend(weakref.ref(v) for v in trainer._version_state().lora.values()) + snapshots.extend((weakref.ref(a), weakref.ref(b))) + return [ForwardOutput(a.square().sum() + b.square().sum(), None, None, None)] + + monkeypatch.setattr(trainer, "_forward_packed", forward) + cache = trainer._forward_graph_cache() + older = torch.nn.Parameter(torch.tensor(3.0)) + if shared: + old_output = trainer._execute_graph_group(group)[0].target_logprobs + (old_handle,) = cache.handles() + else: + old_handle, (old_output,) = cache.run(lambda _: (older.square(),), None) + old_record = weakref.ref(cache._records[old_handle]) + versions.clear() + snapshots.clear() + error_type = asyncio.CancelledError if phase == "cancel" else MemoryError + primary = error_type("native handoff failed") + primary.__cause__ = cause = RuntimeError("original cause") + to = torch.Tensor.to + + def fail_copy(value, *args, **kwargs): + if kwargs.get("copy") and ( + (phase in ("copy", "cancel") and "device" in kwargs) + or ( + phase == "correction" and args == ("cpu",) and len(cache.handles()) == 2 + ) + ): + del value + raise primary + return to(value, *args, **kwargs) + + def fail_handoff(_ref, outputs): + delivered.extend(weakref.ref(output.target_logprobs) for output in outputs) + del outputs + raise primary + + with monkeypatch.context() as failure: + failure.setattr(torch.Tensor, "to", fail_copy) + if phase == "handoff": + failure.setattr(trainer, "_track_slot_graph_outputs", fail_handoff) + with pytest.raises(error_type) as caught: + trainer._execute_graph_group(group) + assert caught.value is primary and primary.__cause__ is cause + assert primary.__traceback__ is not None + assert ( + cache.handles() == (old_handle,) and cache._records[old_handle] is old_record() + ) + assert len(versions) == 1 and len(snapshots) == 2 + assert all( + (reference() is not None) == shared for reference in (*versions, *snapshots) + ), ( + "retained version and snapshots", + tuple(reference() is not None for reference in (*versions, *snapshots)), + ) + assert all(reference() is None for reference in delivered) + assert all(parameter.grad is None for parameter in current) + if shared: + assert old_output is not None and old_output.item() == 38 + cache.evict(old_handle) + trainer.backward(old_output) + torch.testing.assert_close(current[0].grad, torch.full_like(current[0], 2)) + torch.testing.assert_close(current[1].grad, torch.full_like(current[1], 4)) + for parameter in current: + parameter.grad = None + else: + assert old_output is not None and old_output.item() == 9 + cache.backward(old_handle, (torch.tensor(1.0),)) + torch.testing.assert_close(older.grad, torch.tensor(6.0)) + assert all(reference() is None for reference in (*versions, *snapshots)) + assert not trainer._version_state().lora + trainer._guard_slot_can_load(ref) + + output = trainer._execute_graph_group(group)[0].target_logprobs + assert output is not None and output.item() == 38 + (handle,) = cache.handles() + cache.evict(handle) + trainer.backward(output) + torch.testing.assert_close(current[0].grad, torch.full_like(current[0], 2)) + torch.testing.assert_close(current[1].grad, torch.full_like(current[1], 4)) + assert not cache.handles() + + @pytest.mark.parametrize("retention", ["gpu", "cpu", "replay"]) def test_native_cached_slot_guard_tracks_retained_and_consumed_graph( monkeypatch, retention: Literal["gpu", "cpu", "replay"], output_device From 4fc438ce62021a1a1554837665c2acae482e5a91 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 00:55:20 +0000 Subject: [PATCH 101/150] Release failed correction capture copies --- src/art/trainer_rank/_corrections.py | 66 +++++++++++-------- .../test_trainer_rank_slot_graph_lifetime.py | 45 +++++++++---- 2 files changed, 70 insertions(+), 41 deletions(-) diff --git a/src/art/trainer_rank/_corrections.py b/src/art/trainer_rank/_corrections.py index 76040c4e8..a8ca52336 100644 --- a/src/art/trainer_rank/_corrections.py +++ b/src/art/trainer_rank/_corrections.py @@ -294,34 +294,42 @@ def leaves(value: Any) -> Iterator[Any]: if len(indices) != len(tensors): raise ValueError("correction capture requires deduplicated flat tensors") entries: dict[int, _OutputCorrection] = {} - for output in leaves(outputs): - for kind, tensor in ( - ("target_logprobs", output.target_logprobs), - ("top_k", None if output.top_k is None else output.top_k.logprobs), - ): - if ( - tensor is None - or not tensor.requires_grad - or (correction is None and kind != "top_k") + try: + for output in leaves(outputs): + for kind, tensor in ( + ("target_logprobs", output.target_logprobs), + ("top_k", None if output.top_k is None else output.top_k.logprobs), ): - continue - index = indices[id(tensor)] - entry = _OutputCorrection( - index=index, - original_logprobs=tensor.detach().to("cpu", copy=True), - token_index=indices[id(output.top_k.tokens)] - if kind == "top_k" - else None, - original_tokens=output.top_k.tokens.detach().to("cpu", copy=True) - if kind == "top_k" - else None, - logits_index=indices[id(output.logits)] - if kind == "top_k" and output.logits is not None - else None, - ) - if index in entries and entries[index].token_index != entry.token_index: - raise ValueError( - "an aliased output tensor has ambiguous correction semantics" + if ( + tensor is None + or not tensor.requires_grad + or (correction is None and kind != "top_k") + ): + continue + index = indices[id(tensor)] + entry = _OutputCorrection( + index=index, + original_logprobs=tensor.detach().to("cpu", copy=True), + token_index=indices[id(output.top_k.tokens)] + if kind == "top_k" + else None, + original_tokens=output.top_k.tokens.detach().to("cpu", copy=True) + if kind == "top_k" + else None, + logits_index=indices[id(output.logits)] + if kind == "top_k" and output.logits is not None + else None, ) - entries.setdefault(index, entry) - return ForwardCorrectionContext(len(tensors), correction, tuple(entries.values())) + if index in entries and entries[index].token_index != entry.token_index: + raise ValueError( + "an aliased output tensor has ambiguous correction semantics" + ) + entries.setdefault(index, entry) + return ForwardCorrectionContext( + len(tensors), correction, tuple(entries.values()) + ) + except BaseException: + entries.clear() + del outputs, tensors + output = tensor = entry = None + raise diff --git a/tests/unit/test_trainer_rank_slot_graph_lifetime.py b/tests/unit/test_trainer_rank_slot_graph_lifetime.py index 9515724f4..9ff8e850a 100644 --- a/tests/unit/test_trainer_rank_slot_graph_lifetime.py +++ b/tests/unit/test_trainer_rank_slot_graph_lifetime.py @@ -3,7 +3,6 @@ import asyncio from contextlib import nullcontext import gc -import traceback from typing import Literal import weakref @@ -13,7 +12,7 @@ from art.megatron.context_parallel.types import ParallelTopology from art.megatron.prefix_tree_packing import prefix_tree_pack -from art.trainer_rank import ForwardInput, ForwardOptions, ForwardOutput +from art.trainer_rank import ForwardInput, ForwardOptions, ForwardOutput, TopK from art.trainer_rank._impl import ( TrainerRankSlotStateError, _ForwardGroupPlan, @@ -97,17 +96,31 @@ def test_failed_correction_capture_releases_unattached_graph(monkeypatch, failur monkeypatch.setattr( trainer, "_forward_packed", - lambda items, prepared: [ForwardOutput(weight.square(), None, None, None)], + lambda items, prepared: [ + ForwardOutput( + weight.square(), + TopK(weight.pow(3).reshape(1), torch.tensor([0])), + None, + None, + ) + ], ) cache = trainer._forward_graph_cache() + older = torch.nn.Parameter(torch.tensor(3.0)) + old_handle, (old_output,) = cache.run(lambda _: (older.square(),), None) + old_record = weakref.ref(cache._records[old_handle]) primary = failure_type("correction CPU snapshot failed") primary.__cause__ = cause = RuntimeError("original cause") - saved, outputs, versions = [], [], [] + saved, outputs, versions, detached, copies = [], [], [], [], [] to = torch.Tensor.to def fail_snapshot(value, *args, **kwargs): - if args == ("cpu",) and kwargs.get("copy") and cache.handles(): - (handle,) = cache.handles() + if args == ("cpu",) and kwargs.get("copy") and len(cache.handles()) == 2: + if len(copies) < 2: + copied = to(value, *args, **kwargs) + copies.append(weakref.ref(copied)) + return copied + handle = next(handle for handle in cache.handles() if handle != old_handle) record = cache._records[handle] saved.extend(record.saved or ()) outputs.extend(weakref.ref(output) for output in record.outputs or ()) @@ -118,25 +131,33 @@ def fail_snapshot(value, *args, **kwargs): assert record.checkpoint_versions == ( trainer._capture_checkpoint_version("student"), ) + del value raise primary - return to(value, *args, **kwargs) + copied = to(value, *args, **kwargs) + if kwargs.get("copy") and "device" in kwargs: + detached.append(weakref.ref(copied)) + return copied with monkeypatch.context() as allocation: allocation.setattr(torch.Tensor, "to", fail_snapshot) with pytest.raises(failure_type) as caught: trainer._execute_graph_group(group) assert caught.value is primary and primary.__cause__ is cause - assert not cache.handles() + assert primary.__traceback__ is not None + assert cache.handles() == (old_handle,) + assert cache._records[old_handle] is old_record() + assert len(detached) == 3 and len(copies) == 2 assert all(reference() is None for reference in (*saved, *outputs)) + assert all(reference() is None for reference in (*detached, *copies)) assert not trainer._has_live_slot_graph(ref) - # The original traceback owns function locals; the cache must not keep the - # checkpoint capture alive once those independent references are released. - traceback.clear_frames(primary.__traceback__) - gc.collect() assert all(reference() is None for reference in versions) assert not trainer._version_state().lora trainer._guard_slot_can_load(ref) assert weight.grad is None + assert old_output.item() == 9 + cache.backward(old_handle, (torch.tensor(1.0),)) + torch.testing.assert_close(older.grad, torch.tensor(6.0)) + assert not cache.handles() output = trainer._execute_graph_group(group)[0].target_logprobs assert output is not None and output.item() == 4 assert len(cache.handles()) == 1 From 43b65bb9bd658716b8e76d0f95e7549ceca59e69 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 01:01:14 +0000 Subject: [PATCH 102/150] Release failed forward captures before graph registration --- src/art/trainer_rank/_graphs.py | 17 +++++---- tests/unit/test_trainer_rank_graphs.py | 49 +++++++++++++++++++++----- 2 files changed, 51 insertions(+), 15 deletions(-) diff --git a/src/art/trainer_rank/_graphs.py b/src/art/trainer_rank/_graphs.py index 72d616853..4989a9ceb 100644 --- a/src/art/trainer_rank/_graphs.py +++ b/src/art/trainer_rank/_graphs.py @@ -302,6 +302,10 @@ def pack(tensor: torch.Tensor) -> _SavedTensor: for reference in saved: if (cell := reference()) is not None: cell.tensor = torch.empty(0) + # These frame aliases otherwise outlive record.release(). + del self + context_factory = keep_on_device = None + cell = None raise finally: copies.clear() @@ -375,15 +379,16 @@ def run( for device in ("cpu", "cuda") ) record.execution_peak_bytes = execution_peak_bytes - physical = record.run() - record.metadata = tuple( - (value.shape, value.dtype, value.device, value.requires_grad) - for value in physical - ) - # clone: a detached view may still pin a much larger model activation. copied: list[torch.Tensor] = [] + physical: tuple[torch.Tensor, ...] = () value = None try: + physical = record.run() + record.metadata = tuple( + (value.shape, value.dtype, value.device, value.requires_grad) + for value in physical + ) + # clone: a detached view may still pin a much larger model activation. for value in physical: copied.append( value.detach() diff --git a/tests/unit/test_trainer_rank_graphs.py b/tests/unit/test_trainer_rank_graphs.py index 3c3806ea1..55179f438 100644 --- a/tests/unit/test_trainer_rank_graphs.py +++ b/tests/unit/test_trainer_rank_graphs.py @@ -2,6 +2,7 @@ import asyncio from contextlib import contextmanager, nullcontext +from functools import partial import gc from types import SimpleNamespace import weakref @@ -511,7 +512,7 @@ def failing(x): @pytest.mark.parametrize("failure_type", [MemoryError, asyncio.CancelledError]) @pytest.mark.parametrize("retention", ["gpu", "cpu", "replay"]) -@pytest.mark.parametrize("copy_index", [1, 2]) +@pytest.mark.parametrize("copy_index", [pytest.param(0, id="forward"), 1, 2]) def test_initial_output_copy_failure_releases_only_failed_graph( monkeypatch, failure_type, retention, copy_index ): @@ -530,6 +531,17 @@ def test_initial_output_copy_failure_releases_only_failed_graph( primary.__cause__ = cause = RuntimeError("original cause") run, to = _ForwardRecord.run, torch.Tensor.to attempted = 0 + fail_forward = copy_index == 0 + input_snapshots = [] + + def execute(snapshot, value): + outputs = (snapshot * value, snapshot.square()) + if fail_forward: + physical.extend(weakref.ref(output) for output in outputs) + # Attribute ownership only to ART, not this injected model frame. + del snapshot, value, outputs + raise primary + return outputs def arguments(): version = trainer._capture_checkpoint_version("student") @@ -537,21 +549,27 @@ def arguments(): snapshots.append(weakref.ref(snapshot)) versions.append(weakref.ref(version)) return dict( - execute=lambda x: (snapshot * x, snapshot.square()), + execute=partial(execute, snapshot), inputs=torch.tensor(3.0), context_factory=lambda: nullcontext(snapshot), - validate_backward=lambda: trainer._version_state().validate(version), + validate_backward=torch.nn.ParameterList([snapshot]).zero_grad + if copy_index == 0 + else lambda: trainer._version_state().validate(version), checkpoint_versions=(version,), keep_on_device=lambda value: value is snapshot, retention=retention, ) def observe(record): - outputs = run(record) records.append(weakref.ref(record)) - physical.extend(weakref.ref(output) for output in outputs) - saved.extend(record.saved or ()) - return outputs + input_snapshots.append(weakref.ref(record.inputs.value)) + try: + outputs = run(record) + physical.extend(weakref.ref(output) for output in outputs) + saved.extend(record.saved or ()) + return outputs + finally: + del record def fail_copy(value, *args, **kwargs): nonlocal attempted @@ -575,18 +593,31 @@ def fail_copy(value, *args, **kwargs): assert caught.value is primary and primary.__cause__ is cause assert primary.__traceback__ is not None and attempted == copy_index assert len(records) == len(snapshots) == len(versions) == 1 - assert len(physical) == 2 and saved and len(copies) == copy_index - 1 + assert len(physical) == 2 and len(input_snapshots) == 1 + if copy_index: + assert saved and len(copies) == copy_index - 1 + else: + assert not copies assert cache.handles() == (old_handle,) assert cache._records[old_handle] is old_record() assert all( reference() is None - for reference in (*records, *physical, *saved, *snapshots, *versions, *copies) + for reference in ( + *records, + *physical, + *saved, + *snapshots, + *versions, + *copies, + *input_snapshots, + ) ) assert parameter.grad is None assert old_output.item() == 9 and old_output.requires_grad cache.backward(old_handle, (torch.tensor(1.0),)) torch.testing.assert_close(older.grad, torch.tensor(6.0)) + fail_forward = False handle, outputs = cache.run(**arguments()) assert tuple(output.item() for output in outputs) == (6, 4) with trainer._gradient_transaction(): From 468e2b50a4e2a3bbcdecd7aeb231a45aa17879bd Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 03:06:40 +0000 Subject: [PATCH 103/150] Restore reviewed ART backward-state cleanup from durable patch --- src/art/trainer_rank/_graphs.py | 53 ++++--- .../test_trainer_rank_offload_lifetime.py | 142 ++++++++++++++++++ 2 files changed, 172 insertions(+), 23 deletions(-) create mode 100644 tests/unit/test_trainer_rank_offload_lifetime.py diff --git a/src/art/trainer_rank/_graphs.py b/src/art/trainer_rank/_graphs.py index 4989a9ceb..91ffb73b5 100644 --- a/src/art/trainer_rank/_graphs.py +++ b/src/art/trainer_rank/_graphs.py @@ -113,18 +113,30 @@ class _TransferStats: restore_max_bytes: int = 0 @torch.compiler.disable - def copy(self, tensor: torch.Tensor, device: torch.device | str) -> torch.Tensor: - start = perf_counter() - if tensor.device.type == "cuda" and torch.device(device).type == "cpu": - # Pin the only host copy of each storage. Keep copies blocking so - # offload releases GPU ownership before returning, even on user - # streams, and unpack never exposes an unfinished restore. - result = torch.empty_like(tensor, device="cpu", pin_memory=True) - result.copy_(tensor) - else: - result = tensor.to(device, copy=True) + def copy(self, cell: _SavedTensor, device: torch.device | str) -> torch.Tensor: + tensor = cell.tensor + storage = tensor.untyped_storage() + raw = result = None + try: + raw = torch.empty(0, dtype=torch.uint8, device=tensor.device).set_( + storage, 0, (storage.nbytes(),), (1,) + ) + start = perf_counter() + if tensor.device.type == "cuda" and torch.device(device).type == "cpu": + # Pin the only host copy of each storage. Keep copies blocking so + # offload releases GPU ownership before returning, even on user + # streams, and unpack never exposes an unfinished restore. + result = torch.empty_like(raw, device="cpu", pin_memory=True) + result.copy_(raw) + else: + result = raw.to(device, copy=True) + except BaseException: + # The disabled wrapper retains its arguments on failure: pass a cell + # that failed-forward cleanup can empty, never a borrowed raw tensor. + del tensor, storage, raw, result + raise elapsed = perf_counter() - start - size = tensor.numel() * tensor.element_size() + size = storage.nbytes() if result.device.type == "cpu": self.offload_bytes += size self.offload_seconds += elapsed @@ -160,16 +172,11 @@ def view(self, storage, device) -> torch.Tensor: def offload(self, copies: dict[StorageWeakRef, torch.Tensor]) -> None: if self.managed and self.tensor.device.type != "cpu": - tensor = self.tensor - storage = tensor.untyped_storage() # Weak storage identity prevents allocator address reuse from # confusing distinct activations without pinning CUDA storage. key = self.source if key not in copies: - raw = torch.empty(0, dtype=torch.uint8, device=tensor.device).set_( - storage, 0, (storage.nbytes(),), (1,) - ) - copies[key] = self.transfer_stats.copy(raw, "cpu") + copies[key] = self.transfer_stats.copy(self, "cpu") self.tensor = self.view(copies[key].untyped_storage(), "cpu") def unpack(self) -> torch.Tensor: @@ -181,11 +188,7 @@ def unpack(self) -> torch.Tensor: return self.tensor key = (self.device, self.source) if key not in self.restored: - storage = self.tensor.untyped_storage() - raw = torch.empty(0, dtype=torch.uint8).set_( - storage, 0, (storage.nbytes(),), (1,) - ) - self.restored[key] = self.transfer_stats.copy(raw, self.device) + self.restored[key] = self.transfer_stats.copy(self, self.device) return self.view(self.restored[key].untyped_storage(), self.device) @@ -235,6 +238,9 @@ class _ForwardRecord: def release(self) -> None: # Saved-variable hooks can outlive their Python outputs. Break all # ownership edges even when a caller retains a failure traceback. + for reference in self.saved or (): + if (cell := reference()) is not None: + cell.tensor = torch.empty(0) self.outputs = self.saved = self.resident = None self.restored.clear() self.inputs = self.corrections = None @@ -275,9 +281,10 @@ def pack(tensor: torch.Tensor) -> _SavedTensor: restored, transfer_stats, ) + saved.append(weakref.ref(cell)) + del tensor if retention == "cpu": cell.offload(copies) - saved.append(weakref.ref(cell)) return cell with ( diff --git a/tests/unit/test_trainer_rank_offload_lifetime.py b/tests/unit/test_trainer_rank_offload_lifetime.py new file mode 100644 index 000000000..534f05359 --- /dev/null +++ b/tests/unit/test_trainer_rank_offload_lifetime.py @@ -0,0 +1,142 @@ +"""CPU storage-ownership probes; synthetic device labels do not qualify CUDA IO.""" + +import asyncio +import gc +from typing import Any, cast + +import pytest +import torch +from torch.multiprocessing.reductions import StorageWeakRef + +from art.trainer_rank import _graphs as graphs + + +@pytest.mark.parametrize("failure_type", [MemoryError, asyncio.CancelledError]) +@pytest.mark.parametrize("stage", ["allocation", "copy"]) +@pytest.mark.parametrize("late", [False, True], ids=["initial", "later"]) +def test_offload_failure_storage_lifetime(monkeypatch, failure_type, stage, late): + class SyntheticCuda(torch.Tensor): + __torch_function__ = cast(Any, torch._C._disabled_torch_function_impl) + + @property + def device(self): + return torch.device("cuda") + + class Destination(torch.Tensor): + __torch_function__ = cast(Any, torch._C._disabled_torch_function_impl) + + def copy_( + self, other: torch.Tensor, non_blocking: bool = False + ) -> torch.Tensor: + if failing and attempts == (2 if late else 1) and stage == "copy": + del self, other + raise error from cause + return super().copy_(other, non_blocking=non_blocking) + + class SavedTensor(graphs._SavedTensor): + def __init__(self, tensor, *args): + super().__init__(tensor, *args) + if self.managed: + self.tensor = tensor.as_subclass(SyntheticCuda) + sources.append(StorageWeakRef(tensor.untyped_storage())) + + original_empty, original_like = torch.empty, torch.empty_like + sources, destinations = [], [] + attempts = 0 + failing = True + error, cause = failure_type("injected offload failure"), ValueError("copy cause") + + def empty(*args, **kwargs): + if kwargs.get("device") == torch.device("cuda"): + kwargs["device"] = "cpu" + return original_empty(*args, **kwargs).as_subclass(SyntheticCuda) + return original_empty(*args, **kwargs) + + def empty_like(tensor, **kwargs): + nonlocal attempts + if kwargs.pop("pin_memory", False): + attempts += 1 + if failing and attempts == (2 if late else 1) and stage == "allocation": + del tensor + raise error from cause + result = original_like(tensor, **kwargs).as_subclass(Destination) + destinations.append(StorageWeakRef(result.untyped_storage())) + return result + return original_like(tensor, **kwargs) + + weight = torch.nn.Parameter(torch.tensor(2.0)) + borrowed = StorageWeakRef(weight.untyped_storage()) + + def execute(value): + activation = None + try: + activation = weight * value + 1 + return (activation.square(),) + finally: + del activation, value + + cache = graphs.GraphCache() + older_weight = torch.nn.Parameter(torch.tensor(3.0)) + older, (older_output,) = cache.run(lambda _: (older_weight.square(),), ()) + monkeypatch.setattr(graphs, "_SavedTensor", SavedTensor) + monkeypatch.setattr(torch, "empty", empty) + monkeypatch.setattr(torch, "empty_like", empty_like) + inputs = torch.arange(1.0, 4.0, requires_grad=True) + + def keep_on_device(tensor): + return tensor.data_ptr() == weight.data_ptr() + + was_enabled = gc.isenabled() + gc.disable() + try: + handle = None + if late: + handle, (output,) = cache.run( + execute, inputs, keep_on_device=keep_on_device + ) + with pytest.raises(failure_type) as failure: + if handle is None: + cache.run( + execute, inputs, retention="cpu", keep_on_device=keep_on_device + ) + else: + cache.offload(handle) + assert failure.value is error and error.__cause__ is cause + assert error.__traceback__ is not None + assert cache.transfer_stats.offload_count == int(late) + assert cache.transfer_stats.restore_count == 0 + assert not borrowed.expired() and weight.item() == 2 + if late: + assert handle is not None + assert cache.handles() == (older, handle) + # A valid partial graph owns the earlier copy and the failed source. + assert not destinations[0].expired() + assert not sources[-1].expired() + failing = False + cache.backward(handle, (torch.ones_like(output),)) + torch.testing.assert_close(weight.grad, torch.tensor(68.0)) + weight.grad = None + assert cache.handles() == (older,) + assert sources and all(storage.expired() for storage in sources) + assert all(storage.expired() for storage in destinations) + failing = False + handle, (output,) = cache.run( + execute, inputs, retention="cpu", keep_on_device=keep_on_device + ) + torch.testing.assert_close(output, torch.tensor([9.0, 25.0, 49.0])) + cache.backward(handle, (torch.ones_like(output),), retain_graph=True) + torch.testing.assert_close(weight.grad, torch.tensor(68.0)) + weight.grad = None + assert cache.handles() == (older, handle) + cache.backward(handle, (torch.ones_like(output),)) + torch.testing.assert_close(weight.grad, torch.tensor(68.0)) + assert cache.handles() == (older,) + assert all(storage.expired() for storage in sources + destinations) + assert not borrowed.expired() and weight.item() == 2 + torch.testing.assert_close(older_output, torch.tensor(9.0)) + cache.backward(older, (torch.ones_like(older_output),)) + torch.testing.assert_close(older_weight.grad, torch.tensor(6.0)) + assert cache.handles() == () + finally: + if was_enabled: + gc.enable() From 7a64114e7ad2db71c52f15b8875e244f299c7f28 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 03:39:57 +0000 Subject: [PATCH 104/150] Release saved graph storage on eviction and replay rejection --- src/art/trainer_rank/_graphs.py | 10 +++-- tests/unit/test_trainer_rank_graphs.py | 40 ++++++++++++++++--- .../test_trainer_rank_offload_lifetime.py | 10 ++++- 3 files changed, 49 insertions(+), 11 deletions(-) diff --git a/src/art/trainer_rank/_graphs.py b/src/art/trainer_rank/_graphs.py index 91ffb73b5..90ce0bd04 100644 --- a/src/art/trainer_rank/_graphs.py +++ b/src/art/trainer_rank/_graphs.py @@ -235,7 +235,7 @@ class _ForwardRecord: ) execution_peak_bytes: int = 0 - def release(self) -> None: + def release_physical(self) -> None: # Saved-variable hooks can outlive their Python outputs. Break all # ownership edges even when a caller retains a failure traceback. for reference in self.saved or (): @@ -243,6 +243,9 @@ def release(self) -> None: cell.tensor = torch.empty(0) self.outputs = self.saved = self.resident = None self.restored.clear() + + def release(self) -> None: + self.release_physical() self.inputs = self.corrections = None self.execute = lambda _: () self.context_factory = nullcontext @@ -506,8 +509,7 @@ def evict( if replay_with_current and record.current_context_factory is None: raise ValueError("Current-weight replay requires a version context") record.replay_with_current = replay_with_current - record.outputs = record.saved = record.resident = None - record.restored.clear() + record.release_physical() record.retention = "replay" def release(self, handle: ForwardHandle) -> None: @@ -673,8 +675,8 @@ def _prepare_backward(record, gradients, stale, prepared): (value.shape, value.dtype, value.device, value.requires_grad) for value in physical ) + del physical if metadata != record.metadata: - record.outputs = record.saved = None raise RuntimeError( "Replayed output metadata differs from original forward" ) diff --git a/tests/unit/test_trainer_rank_graphs.py b/tests/unit/test_trainer_rank_graphs.py index 55179f438..e65f72a63 100644 --- a/tests/unit/test_trainer_rank_graphs.py +++ b/tests/unit/test_trainer_rank_graphs.py @@ -9,6 +9,7 @@ import pytest import torch +from torch.multiprocessing.reductions import StorageWeakRef from torch.utils.checkpoint import checkpoint from art.trainer_rank import ForwardOutput, TopK, TrainerRank @@ -478,7 +479,12 @@ def test_replay_restores_original_autocast_context(): torch.testing.assert_close(parameter.grad, torch.full_like(parameter, 2.0)) -def test_replay_failure_discards_transaction_and_releases_participating_records(): +def test_replay_failure_discards_transaction_and_releases_participating_records( + request, +): + if gc.isenabled(): + request.addfinalizer(gc.enable) + gc.disable() trainer = TrainerRank.__new__(TrainerRank) parameter = torch.nn.Parameter(torch.tensor(2.0)) trainer._checkpoint_slots = {"student": _CheckpointSlot(params=(parameter,))} @@ -487,26 +493,50 @@ def test_replay_failure_discards_transaction_and_releases_participating_records( parameter, trainer._capture_checkpoint_version("student") ) cache = GraphCache() + older_weight = torch.nn.Parameter(torch.tensor(3.0)) + older, (older_output,) = cache.run(lambda _: (older_weight.square(),), ()) + borrowed = StorageWeakRef(parameter.untyped_storage()) + storages = [] executions = [0] def failing(x): executions[0] += 1 - result = snapshot * x + activation = snapshot * x + 1 + result = activation.square() + storages.extend( + StorageWeakRef(value.untyped_storage()) for value in (activation, result) + ) return (result if executions[0] == 1 else result.expand(2),) first, _ = cache.run( lambda x: (snapshot * x,), torch.tensor(3.0), retention="replay" ) second, _ = cache.run(failing, torch.tensor(4.0), retention="replay") - parameter.grad = torch.tensor(7.0) - with pytest.raises(RuntimeError, match="metadata differs"): + parameter.grad = prior_gradient = torch.tensor(7.0) + with pytest.raises(RuntimeError, match="metadata differs") as failure: with trainer._gradient_transaction(): cache.backward_many( [(first, (torch.tensor(1.0),)), (second, (torch.tensor(1.0),))] ) - assert parameter.grad.item() == 7 + assert parameter.grad is prior_gradient and parameter.grad.item() == 7 assert snapshot.grad is None assert not trainer._version_state()._origins + assert failure.value.__traceback__ is not None and failure.value.__cause__ is None + assert len(storages) == 4 and all(storage.expired() for storage in storages) + assert not borrowed.expired() and parameter.item() == snapshot.item() == 2 + assert cache.handles() == (older,) + torch.testing.assert_close(older_output, torch.tensor(9.0)) + cache.backward(older, (torch.ones_like(older_output),)) + torch.testing.assert_close(older_weight.grad, torch.tensor(6.0)) + retry, (output,) = cache.run( + lambda x: ((snapshot * x + 1).square(),), + torch.arange(1.0, 4.0), + retention="replay", + ) + torch.testing.assert_close(output, torch.tensor([9.0, 25.0, 49.0])) + with trainer._gradient_transaction(): + cache.backward(retry, (torch.ones_like(output),)) + assert parameter.grad.item() == 75 and snapshot.grad is None assert cache.handles() == () diff --git a/tests/unit/test_trainer_rank_offload_lifetime.py b/tests/unit/test_trainer_rank_offload_lifetime.py index 534f05359..abade529e 100644 --- a/tests/unit/test_trainer_rank_offload_lifetime.py +++ b/tests/unit/test_trainer_rank_offload_lifetime.py @@ -13,7 +13,9 @@ @pytest.mark.parametrize("failure_type", [MemoryError, asyncio.CancelledError]) @pytest.mark.parametrize("stage", ["allocation", "copy"]) -@pytest.mark.parametrize("late", [False, True], ids=["initial", "later"]) +@pytest.mark.parametrize( + "late", [False, True, "evicted"], ids=["initial", "later", "evicted"] +) def test_offload_failure_storage_lifetime(monkeypatch, failure_type, stage, late): class SyntheticCuda(torch.Tensor): __torch_function__ = cast(Any, torch._C._disabled_torch_function_impl) @@ -103,7 +105,7 @@ def keep_on_device(tensor): cache.offload(handle) assert failure.value is error and error.__cause__ is cause assert error.__traceback__ is not None - assert cache.transfer_stats.offload_count == int(late) + assert cache.transfer_stats.offload_count == int(bool(late)) assert cache.transfer_stats.restore_count == 0 assert not borrowed.expired() and weight.item() == 2 if late: @@ -113,6 +115,10 @@ def keep_on_device(tensor): assert not destinations[0].expired() assert not sources[-1].expired() failing = False + if late == "evicted": + cache.evict(handle) + assert all(storage.expired() for storage in sources + destinations) + assert cache.state(handle).retention == "replay" cache.backward(handle, (torch.ones_like(output),)) torch.testing.assert_close(weight.grad, torch.tensor(68.0)) weight.grad = None From ec7f2c117b8503398fd8a679167db7632829e447 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 04:09:26 +0000 Subject: [PATCH 105/150] Release correction staging and backward aliases on failure --- src/art/trainer_rank/_corrections.py | 311 ++++++++++++++----------- src/art/trainer_rank/_graphs.py | 148 ++++++------ tests/unit/test_trainer_rank_graphs.py | 164 ++++++++++++- 3 files changed, 407 insertions(+), 216 deletions(-) diff --git a/src/art/trainer_rank/_corrections.py b/src/art/trainer_rank/_corrections.py index a8ca52336..174cb0fae 100644 --- a/src/art/trainer_rank/_corrections.py +++ b/src/art/trainer_rank/_corrections.py @@ -25,40 +25,48 @@ def importance_weights( silently replaced. Computation uses at least float32 and promotes to float64 for float64 inputs or clipping bounds outside the float32 normal range. """ - if original_logprobs.shape != current_logprobs.shape: - raise ValueError("correction logprob shapes must match exactly") - if original_logprobs.device != current_logprobs.device: - raise ValueError("correction logprobs must be on the same device") - if ( - not original_logprobs.is_floating_point() - or not current_logprobs.is_floating_point() - ): - raise TypeError("correction logprobs must be floating-point tensors") - if not bool(torch.isfinite(original_logprobs).all()): - raise ValueError("original logprobs must be finite (positive sampling support)") - if bool((torch.isnan(current_logprobs) | torch.isposinf(current_logprobs)).any()): - raise ValueError("current logprobs must be finite or negative infinity") - dtype = ( - torch.float64 - if torch.float64 in (original_logprobs.dtype, current_logprobs.dtype) - or correction.clip_high > torch.finfo(torch.float32).max - or any( - 0 < bound < torch.finfo(torch.float32).tiny - for bound in (correction.clip_low, correction.clip_high) + try: + if original_logprobs.shape != current_logprobs.shape: + raise ValueError("correction logprob shapes must match exactly") + if original_logprobs.device != current_logprobs.device: + raise ValueError("correction logprobs must be on the same device") + if ( + not original_logprobs.is_floating_point() + or not current_logprobs.is_floating_point() + ): + raise TypeError("correction logprobs must be floating-point tensors") + if not bool(torch.isfinite(original_logprobs).all()): + raise ValueError( + "original logprobs must be finite (positive sampling support)" + ) + if bool( + (torch.isnan(current_logprobs) | torch.isposinf(current_logprobs)).any() + ): + raise ValueError("current logprobs must be finite or negative infinity") + dtype = ( + torch.float64 + if torch.float64 in (original_logprobs.dtype, current_logprobs.dtype) + or correction.clip_high > torch.finfo(torch.float32).max + or any( + 0 < bound < torch.finfo(torch.float32).tiny + for bound in (correction.clip_low, correction.clip_high) + ) + else torch.float32 + ) + log_ratio = current_logprobs.detach().to(dtype) - original_logprobs.detach().to( + dtype + ) + if correction.clip_high == 0: + return torch.zeros_like(log_ratio) + log_low = math.log(correction.clip_low) if correction.clip_low else -math.inf + return ( + log_ratio.clamp(log_low, math.log(correction.clip_high)) + .exp() + .clamp(correction.clip_low, correction.clip_high) ) - else torch.float32 - ) - log_ratio = current_logprobs.detach().to(dtype) - original_logprobs.detach().to( - dtype - ) - if correction.clip_high == 0: - return torch.zeros_like(log_ratio) - log_low = math.log(correction.clip_low) if correction.clip_low else -math.inf - return ( - log_ratio.clamp(log_low, math.log(correction.clip_high)) - .exp() - .clamp(correction.clip_low, correction.clip_high) - ) + finally: + del original_logprobs, current_logprobs + log_ratio = None def correct_logprob_cotangent( @@ -77,42 +85,49 @@ def correct_logprob_cotangent( their identity/order. Values are full-vocabulary logprobs, not probabilities renormalized over top-k. No forward is performed here. """ - if cotangent.shape != original_logprobs.shape: - raise ValueError("cotangent and correction logprob shapes must match exactly") - active = cotangent != 0 - if not bool(active.any()): - return cotangent - if current_logprobs is None: - if correction.policy == "always": - raise RuntimeError( - "importance sampling correction requires current logprobs" - ) - return cotangent - if (original_tokens is None) != (current_tokens is None): - raise ValueError("correction requires both original and current token IDs") - if original_tokens is not None and current_tokens is not None: - if ( - original_tokens.shape != original_logprobs.shape - or current_tokens.shape != current_logprobs.shape - ): - raise ValueError("correction token IDs must match logprob shapes") - if not torch.equal(original_tokens, current_tokens): + try: + if cotangent.shape != original_logprobs.shape: raise ValueError( - "correction must compare the same token IDs in the same order" + "cotangent and correction logprob shapes must match exactly" ) - if current_logprobs.shape != original_logprobs.shape: - raise ValueError("correction logprob shapes must match exactly") - if current_logprobs.device != original_logprobs.device: - raise ValueError("correction logprobs must be on the same device") - selected = active.to(original_logprobs.device) - weights = importance_weights( - original_logprobs[selected], current_logprobs[selected], correction - ) - corrected = cotangent.clone() - corrected[active] = (cotangent[active] * weights.to(cotangent.device)).to( - cotangent.dtype - ) - return corrected + active = cotangent != 0 + if not bool(active.any()): + return cotangent + if current_logprobs is None: + if correction.policy == "always": + raise RuntimeError( + "importance sampling correction requires current logprobs" + ) + return cotangent + if (original_tokens is None) != (current_tokens is None): + raise ValueError("correction requires both original and current token IDs") + if original_tokens is not None and current_tokens is not None: + if ( + original_tokens.shape != original_logprobs.shape + or current_tokens.shape != current_logprobs.shape + ): + raise ValueError("correction token IDs must match logprob shapes") + if not torch.equal(original_tokens, current_tokens): + raise ValueError( + "correction must compare the same token IDs in the same order" + ) + if current_logprobs.shape != original_logprobs.shape: + raise ValueError("correction logprob shapes must match exactly") + if current_logprobs.device != original_logprobs.device: + raise ValueError("correction logprobs must be on the same device") + selected = active.to(original_logprobs.device) + weights = importance_weights( + original_logprobs[selected], current_logprobs[selected], correction + ) + corrected = cotangent.clone() + corrected[active] = (cotangent[active] * weights.to(cotangent.device)).to( + cotangent.dtype + ) + return corrected + finally: + del cotangent, original_logprobs, current_logprobs + del original_tokens, current_tokens + active = selected = weights = corrected = None @dataclass(frozen=True) @@ -178,62 +193,72 @@ def correct( current_tensors: Sequence[torch.Tensor] | None = None, ) -> tuple[torch.Tensor | None, ...]: """Stage corrected cotangents without mutating gradients or model state.""" - self.requires_current(gradients) - if current_tensors is not None and len(current_tensors) != self.output_count: - raise ValueError("current tensors must match the captured output count") - corrected = list(gradients) - if self.correction is None: - return tuple(corrected) - for output in self.outputs: - gradient = gradients[output.index] - original = output.original_logprobs - if gradient is None or not bool((gradient != 0).any()): - continue - current = ( - None - if current_tensors is None - else current_tensors[output.index].detach() - ) - if current is not None and output.original_tokens is not None: - assert current_tensors is not None and output.token_index is not None - tokens = current_tensors[output.token_index] - original_tokens = output.original_tokens.to(tokens.device) - if tokens.shape != original_tokens.shape: - raise ValueError( - "current top-k token shape must match original top-k" + try: + self.requires_current(gradients) + if ( + current_tensors is not None + and len(current_tensors) != self.output_count + ): + raise ValueError("current tensors must match the captured output count") + corrected = list(gradients) + if self.correction is None: + return tuple(corrected) + for output in self.outputs: + gradient = gradients[output.index] + original = output.original_logprobs + if gradient is None or not bool((gradient != 0).any()): + continue + current = ( + None + if current_tensors is None + else current_tensors[output.index].detach() + ) + if current is not None and output.original_tokens is not None: + assert ( + current_tensors is not None and output.token_index is not None ) - if not torch.equal(tokens, original_tokens): - # A changed top-k ordering can still contain every old ID. - sorted_tokens, order = tokens.sort(dim=-1) - positions = torch.searchsorted( - sorted_tokens.contiguous(), original_tokens.contiguous() - ).clamp_max(tokens.shape[-1] - 1) - matched = sorted_tokens.gather(-1, positions) == original_tokens - if bool((matched | (gradient == 0).to(matched.device)).all()): - current = current.gather(-1, order.gather(-1, positions)) - elif output.logits_index is not None: - logits = current_tensors[output.logits_index].detach() - dtype = ( - torch.float64 - if logits.dtype == torch.float64 - else torch.float32 + tokens = current_tensors[output.token_index] + original_tokens = output.original_tokens.to(tokens.device) + if tokens.shape != original_tokens.shape: + raise ValueError( + "current top-k token shape must match original top-k" ) - logits = logits.to(dtype) - current = logits.gather( - -1, output.original_tokens.to(logits.device) - ) - logits.logsumexp(-1, keepdim=True) - else: - # The new top-k lacks original events: no ratio is available. - current = None - corrected[output.index] = correct_logprob_cotangent( - gradient, - original_logprobs=original - if current is None - else original.to(current.device), - current_logprobs=current, - correction=self.correction, - ) - return tuple(corrected) + if not torch.equal(tokens, original_tokens): + # A changed top-k ordering can still contain every old ID. + sorted_tokens, order = tokens.sort(dim=-1) + positions = torch.searchsorted( + sorted_tokens.contiguous(), original_tokens.contiguous() + ).clamp_max(tokens.shape[-1] - 1) + matched = sorted_tokens.gather(-1, positions) == original_tokens + if bool((matched | (gradient == 0).to(matched.device)).all()): + current = current.gather(-1, order.gather(-1, positions)) + elif output.logits_index is not None: + logits = current_tensors[output.logits_index].detach() + dtype = ( + torch.float64 + if logits.dtype == torch.float64 + else torch.float32 + ) + logits = logits.to(dtype) + current = logits.gather( + -1, output.original_tokens.to(logits.device) + ) - logits.logsumexp(-1, keepdim=True) + else: + # The new top-k lacks original events: no ratio is available. + current = None + corrected[output.index] = correct_logprob_cotangent( + gradient, + original_logprobs=original + if current is None + else original.to(current.device), + current_logprobs=current, + correction=self.correction, + ) + return tuple(corrected) + finally: + del self, gradients, current_tensors + output = gradient = original = current = tokens = original_tokens = None + sorted_tokens = order = positions = matched = logits = corrected = None def validate_replay( self, @@ -246,26 +271,30 @@ def validate_replay( Physical current-weight replay cannot feed original-position cotangents to a Jacobian whose selected token at that position has changed. """ - self.requires_current(gradients) - if len(current_tensors) != self.output_count: - raise ValueError("current tensors must match the captured output count") - for output in self.outputs: - gradient = gradients[output.index] - if output.original_tokens is None or gradient is None: - continue - active = gradient != 0 - if not bool(active.any()): - continue - assert output.token_index is not None - tokens = current_tensors[output.token_index] - original = output.original_tokens.to(tokens.device) - if tokens.shape != original.shape or bool( - ((tokens != original) & active.to(tokens.device)).any() - ): - raise RuntimeError( - "current replay changed active top-k token identities; " - "replay the original weights instead" - ) + try: + self.requires_current(gradients) + if len(current_tensors) != self.output_count: + raise ValueError("current tensors must match the captured output count") + for output in self.outputs: + gradient = gradients[output.index] + if output.original_tokens is None or gradient is None: + continue + active = gradient != 0 + if not bool(active.any()): + continue + assert output.token_index is not None + tokens = current_tensors[output.token_index] + original = output.original_tokens.to(tokens.device) + if tokens.shape != original.shape or bool( + ((tokens != original) & active.to(tokens.device)).any() + ): + raise RuntimeError( + "current replay changed active top-k token identities; " + "replay the original weights instead" + ) + finally: + del self, gradients, current_tensors + output = gradient = active = tokens = original = None def capture_forward_corrections( diff --git a/src/art/trainer_rank/_graphs.py b/src/art/trainer_rank/_graphs.py index 90ce0bd04..8c0ac8e37 100644 --- a/src/art/trainer_rank/_graphs.py +++ b/src/art/trainer_rank/_graphs.py @@ -603,37 +603,40 @@ def _backward_validated( } # The coordinator uses this logical rank's TP/CP group, never all DP # ranks: different DP owners may have different numbers of graphs. - for handle, gradients in packets: - record = self._records[handle] - prepared[handle] = coordinate( - lambda: self._prepare_correction(record, gradients, stale[handle]) - ) - for handle, gradients in packets: - record = self._records[handle] - pairs = coordinate( - lambda: self._prepare_backward( - record, gradients, stale[handle], prepared[handle] + try: + for handle, gradients in packets: + record = self._records[handle] + prepared[handle] = coordinate( + lambda: self._prepare_correction(record, gradients, stale[handle]) ) - ) - try: - coordinate( - lambda: ( - torch.autograd.backward( - [output for output, _ in pairs], - [gradient for _, gradient in pairs], - retain_graph=retain_graph, - ) - if pairs - else None + for handle, gradients in packets: + record = self._records[handle] + pairs = coordinate( + lambda: self._prepare_backward( + record, gradients, stale[handle], prepared[handle] ) ) - finally: - record.restored.clear() - del pairs - if not retain_graph: - self.release(handle) - elif record.retention == "replay": - self.evict(handle) + try: + coordinate( + lambda: ( + torch.autograd.backward( + [output for output, _ in pairs], + [gradient for _, gradient in pairs], + retain_graph=retain_graph, + ) + if pairs + else None + ) + ) + finally: + record.restored.clear() + pairs.clear() + if not retain_graph: + self.release(handle) + elif record.retention == "replay": + self.evict(handle) + finally: + prepared.clear() @staticmethod def _prepare_correction(record, gradients, stale): @@ -646,50 +649,57 @@ def _prepare_correction(record, gradients, stale): return None # Explicit always may add a no-grad forward. Stage every correction # before physical backward; changed cotangents wait on CPU. - with record.rng.replay(record.rng_tracker): - current = record.run( - context_factory=record.current_context_factory, - store=False, - grad_enabled=False, + try: + with record.rng.replay(record.rng_tracker): + current = record.run( + context_factory=record.current_context_factory, + store=False, + grad_enabled=False, + ) + corrected = record.corrections.correct(gradients, current) + return tuple( + value if value is None or value is original else value.to("cpu") + for value, original in zip(corrected, gradients, strict=True) ) - corrected = record.corrections.correct(gradients, current) - return tuple( - value.to("cpu") if value is not None and value is not original else value - for value, original in zip(corrected, gradients, strict=True) - ) + finally: + current = corrected = None @staticmethod def _prepare_backward(record, gradients, stale, prepared): - gradients = gradients if prepared is None else prepared - if not any(gradient is not None for gradient in gradients): - return [] - current_replay = stale and record.replay_with_current - if record.outputs is None: - with record.rng.replay(record.rng_tracker): - physical = record.run( - context_factory=record.current_context_factory - if current_replay - else None + try: + gradients = gradients if prepared is None else prepared + if not any(gradient is not None for gradient in gradients): + return [] + current_replay = stale and record.replay_with_current + if record.outputs is None: + with record.rng.replay(record.rng_tracker): + physical = record.run( + context_factory=record.current_context_factory + if current_replay + else None + ) + metadata = tuple( + (value.shape, value.dtype, value.device, value.requires_grad) + for value in physical ) - metadata = tuple( - (value.shape, value.dtype, value.device, value.requires_grad) - for value in physical - ) - del physical - if metadata != record.metadata: - raise RuntimeError( - "Replayed output metadata differs from original forward" + del physical + if metadata != record.metadata: + raise RuntimeError( + "Replayed output metadata differs from original forward" + ) + record.replay_count += 1 + if prepared is None and stale and record.corrections is not None: + if current_replay: + record.corrections.validate_replay(gradients, record.outputs) + gradients = record.corrections.correct( + gradients, record.outputs if current_replay else None ) - record.replay_count += 1 - if prepared is None and stale and record.corrections is not None: - if current_replay: - record.corrections.validate_replay(gradients, record.outputs) - gradients = record.corrections.correct( - gradients, record.outputs if current_replay else None - ) - assert record.outputs is not None - return [ - (output, gradient.to(output.device)) - for output, gradient in zip(record.outputs, gradients, strict=True) - if gradient is not None - ] + assert record.outputs is not None + return [ + (output, gradient.to(output.device)) + for output, gradient in zip(record.outputs, gradients, strict=True) + if gradient is not None + ] + finally: + del gradients, prepared + physical = None diff --git a/tests/unit/test_trainer_rank_graphs.py b/tests/unit/test_trainer_rank_graphs.py index e65f72a63..e1dac63fa 100644 --- a/tests/unit/test_trainer_rank_graphs.py +++ b/tests/unit/test_trainer_rank_graphs.py @@ -283,10 +283,12 @@ def test_backward_uses_creation_order_not_wire_handle_order(): assert seen == [0, 1, 2] -def _corrected_cache(*, retention="gpu", policy="when_available", stale=True): +def _corrected_cache( + *, retention="gpu", policy="when_available", stale=True, cache=None +): from art.trainer_rank._corrections import capture_forward_corrections - cache = GraphCache() + cache = cache or GraphCache() original = torch.nn.Parameter(torch.tensor(-1.0)) current = torch.nn.Parameter(torch.tensor(-0.5)) selected = [original] @@ -344,6 +346,141 @@ def test_always_correction_uses_no_grad_current_evaluation_and_old_jacobian(): assert executions == [(False, True), (True, False)] +@pytest.mark.parametrize("phase", ["always", "coordinator", "metadata"]) +def test_correction_and_coordinator_failures_release_owned_storage( + monkeypatch, request, phase +): + from art.trainer_rank import _corrections as corrections + from art.trainer_rank import _graphs as graphs + + if gc.isenabled(): + request.addfinalizer(gc.enable) + gc.disable() + cache = GraphCache() + older_weight = torch.nn.Parameter(torch.tensor(3.0)) + older, (older_output,) = cache.run(lambda _: (older_weight.square(),), ()) + storages, errors = [], [] + + def observe(function, *, inputs=False): + def call(*args, **kwargs): + try: + if inputs: + storages.extend( + StorageWeakRef(value.untyped_storage()) for value in args[:2] + ) + result = function(*args, **kwargs) + for value in result if isinstance(result, tuple) else (result,): + if isinstance(value, torch.Tensor): + storages.append(StorageWeakRef(value.untyped_storage())) + return result + except BaseException as error: + errors.append(error) + raise + finally: + args = kwargs = result = value = None + + return call + + monkeypatch.setattr( + graphs._ForwardRecord, "run", observe(graphs._ForwardRecord.run) + ) + monkeypatch.setattr( + GraphCache, + "_prepare_correction", + staticmethod(observe(GraphCache._prepare_correction)), + ) + monkeypatch.setattr( + GraphCache, + "_prepare_backward", + staticmethod(observe(GraphCache._prepare_backward)), + ) + monkeypatch.setattr( + corrections, + "importance_weights", + observe(corrections.importance_weights, inputs=True), + ) + packets = [] + transaction = nullcontext() + if phase == "metadata": + trainer = TrainerRank.__new__(TrainerRank) + parameter = torch.nn.Parameter(torch.tensor(2.0)) + trainer._checkpoint_slots = {"student": _CheckpointSlot(params=(parameter,))} + trainer.runtime = SimpleNamespace(model=[], optimizer=None) + snapshot = trainer._snapshot_parameter( + parameter, trainer._capture_checkpoint_version("student") + ) + first, _ = cache.run( + lambda x: (snapshot * x,), torch.tensor(3.0), retention="replay" + ) + packets.append((first, (torch.tensor(1.0),))) + parameter.grad = prior_gradient = torch.tensor(7.0) + transaction = trainer._gradient_transaction() + cache, handle, original, current, executions, _ = _corrected_cache( + policy="always", + retention="replay" if phase == "metadata" else "gpu", + cache=cache, + ) + if phase == "metadata": + execute = cache._records[handle].execute + cache._records[handle].execute = lambda x: tuple( + value.expand(2) if torch.is_grad_enabled() else value + for value in execute(x) + ) + gradient = torch.tensor(1.0) + borrowed = [ + StorageWeakRef(value.untyped_storage()) + for value in (original, current, gradient) + ] + primary, cause = RuntimeError("peer backward failed"), ValueError("peer cause") + calls = 0 + + def coordinate(function): + nonlocal calls + calls += 1 + result = function() + if phase == "coordinator" and calls == 3: + raise primary from cause + return result + + if phase == "always": + with torch.no_grad(): + current.fill_(float("nan")) + packets.append((handle, (gradient,))) + with pytest.raises(ValueError if phase == "always" else RuntimeError) as failure: + with transaction: + cache.backward_many(packets, coordinate=coordinate) + if phase != "coordinator": + assert errors and all(error is failure.value for error in errors) + assert ( + "metadata differs" if phase == "metadata" else "current logprobs" + ) in str(failure.value) + assert original.grad is None + if phase == "metadata": + assert parameter.grad is prior_gradient and parameter.grad.item() == 7 + assert snapshot.grad is None and not trainer._version_state()._origins + else: + assert failure.value is primary and primary.__cause__ is cause and calls == 3 + torch.testing.assert_close( + original.grad, torch.tensor(2.0) * torch.tensor(0.75).exp() + ) + assert failure.value.__traceback__ is not None + assert len(storages) >= 4 and all(storage.expired() for storage in storages) + assert all(not storage.expired() for storage in borrowed) + assert gradient.item() == 1 and original.item() == -1 and current.grad is None + assert executions == [(False, True), (True, False)] + ( + [(False, True)] if phase == "metadata" else [] + ) + assert cache.handles() == (older,) + cache.backward(older, (torch.ones_like(older_output),)) + torch.testing.assert_close(older_weight.grad, torch.tensor(6.0)) + cache, retry, retry_weight, _, _, _ = _corrected_cache(policy="always", cache=cache) + cache.backward(retry, (gradient,)) + torch.testing.assert_close( + retry_weight.grad, torch.tensor(2.0) * torch.tensor(0.75).exp() + ) + assert all(storage.expired() for storage in storages) and not cache.handles() + + def test_newer_replay_opportunistically_corrects_current_jacobian(): cache, handle, original, current, executions, _ = _corrected_cache() cache.evict(handle, replay_with_current=True) @@ -354,15 +491,24 @@ def test_newer_replay_opportunistically_corrects_current_jacobian(): @pytest.mark.parametrize("corrections", [False, True]) -def test_current_replay_rejects_changed_selected_token_events(corrections): +def test_current_replay_rejects_changed_selected_token_events(corrections, request): from art.trainer_rank._corrections import capture_forward_corrections + if gc.isenabled(): + request.addfinalizer(gc.enable) + gc.disable() cache = GraphCache() parameter = torch.nn.Parameter(torch.tensor([1.0, 2.0])) tokens = [torch.tensor([0, 1])] + storages = [] def execute(_): - return parameter.log_softmax(-1)[tokens[0]], tokens[0] + logprobs = parameter.log_softmax(-1) + output = logprobs[tokens[0]] + storages.extend( + StorageWeakRef(value.untyped_storage()) for value in (logprobs, output) + ) + return output, tokens[0] handle, outputs = cache.run(execute, None) context = capture_forward_corrections( @@ -379,8 +525,14 @@ def execute(_): ) cache.evict(handle, replay_with_current=True) tokens[0] = torch.tensor([1, 0]) - with pytest.raises(RuntimeError, match="token identit"): - cache.backward(handle, (torch.ones(2), None)) + gradient = torch.ones(2) + with pytest.raises(RuntimeError, match="token identit") as failure: + cache.backward(handle, (gradient, None)) + assert failure.value.__traceback__ is not None and failure.value.__cause__ is None + assert len(storages) == 4 and all(storage.expired() for storage in storages) + torch.testing.assert_close(parameter, torch.tensor([1.0, 2.0])) + torch.testing.assert_close(gradient, torch.ones(2)) + torch.testing.assert_close(tokens[0], torch.tensor([1, 0])) assert parameter.grad is None assert not cache.handles() From e156bb7251a7c8b0c8f1e6908c5399ba0216ba36 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 04:37:00 +0000 Subject: [PATCH 106/150] Release successful preparation results on coordinated failure --- src/art/trainer_rank/_commands.py | 29 ++-- tests/unit/test_trainer_rank_graphs.py | 181 ++++++++++++++----------- 2 files changed, 118 insertions(+), 92 deletions(-) diff --git a/src/art/trainer_rank/_commands.py b/src/art/trainer_rank/_commands.py index 849fb4434..2f71c1c20 100644 --- a/src/art/trainer_rank/_commands.py +++ b/src/art/trainer_rank/_commands.py @@ -28,19 +28,22 @@ def _coordinate_call(call: Callable[[], T], *, group: dist.ProcessGroup | None) -> T: result, error = None, None try: - result = call() - except BaseException as exc: - error = exc - failures = [None if error is None else f"{type(error).__name__}: {error}"] - if dist.is_initialized(): - local = failures[0] - failures = [None] * dist.get_world_size(group) - dist.all_gather_object(failures, local, group=group) - if any(failures): - if error is not None: - raise error - raise RuntimeError(f"Physical trainer preflight failed: {failures}") - return cast(T, result) + try: + result = call() + except BaseException as exc: + error = exc + failures = [None if error is None else f"{type(error).__name__}: {error}"] + if dist.is_initialized(): + local = failures[0] + failures = [None] * dist.get_world_size(group) + dist.all_gather_object(failures, local, group=group) + if any(failures): + if error is not None: + raise error + raise RuntimeError(f"Physical trainer preflight failed: {failures}") + return cast(T, result) + finally: + result = None @dataclass(frozen=True) diff --git a/tests/unit/test_trainer_rank_graphs.py b/tests/unit/test_trainer_rank_graphs.py index e1dac63fa..3e48398e0 100644 --- a/tests/unit/test_trainer_rank_graphs.py +++ b/tests/unit/test_trainer_rank_graphs.py @@ -283,11 +283,23 @@ def test_backward_uses_creation_order_not_wire_handle_order(): assert seen == [0, 1, 2] +def _logprob_corrections(outputs, policy): + from art.trainer_rank._corrections import capture_forward_corrections + + return capture_forward_corrections( + ForwardOutput(outputs[0], None, None, None), + outputs, + ResolvedForwardOptions( + stale_gradient_corrections=( + ImportanceSamplingGradientCorrection(policy=policy), + ) + ), + ) + + def _corrected_cache( *, retention="gpu", policy="when_available", stale=True, cache=None ): - from art.trainer_rank._corrections import capture_forward_corrections - cache = cache or GraphCache() original = torch.nn.Parameter(torch.tensor(-1.0)) current = torch.nn.Parameter(torch.tensor(-0.5)) @@ -307,15 +319,7 @@ def execute(x): return (selected[0].square() * x,) handle, outputs = cache.run(execute, torch.tensor(-1.0), retention=retention) - context = capture_forward_corrections( - ForwardOutput(outputs[0], None, None, None), - outputs, - ResolvedForwardOptions( - stale_gradient_corrections=( - ImportanceSamplingGradientCorrection(policy=policy), - ) - ), - ) + context = _logprob_corrections(outputs, policy) cache.set_corrections( handle, context, @@ -346,38 +350,37 @@ def test_always_correction_uses_no_grad_current_evaluation_and_old_jacobian(): assert executions == [(False, True), (True, False)] -@pytest.mark.parametrize("phase", ["always", "coordinator", "metadata"]) -def test_correction_and_coordinator_failures_release_owned_storage( - monkeypatch, request, phase -): +def _observe_backward_storage(monkeypatch, request): from art.trainer_rank import _corrections as corrections from art.trainer_rank import _graphs as graphs if gc.isenabled(): request.addfinalizer(gc.enable) gc.disable() - cache = GraphCache() - older_weight = torch.nn.Parameter(torch.tensor(3.0)) - older, (older_output,) = cache.run(lambda _: (older_weight.square(),), ()) - storages, errors = [], [] + storages, errors, borrowed = [], [], [] + + def watch(value): + if isinstance(value, torch.Tensor): + storage = StorageWeakRef(value.untyped_storage()) + if storage not in borrowed: + storages.append(storage) + elif isinstance(value, (tuple, list)): + for child in value: + watch(child) def observe(function, *, inputs=False): def call(*args, **kwargs): try: if inputs: - storages.extend( - StorageWeakRef(value.untyped_storage()) for value in args[:2] - ) + watch(args[:2]) result = function(*args, **kwargs) - for value in result if isinstance(result, tuple) else (result,): - if isinstance(value, torch.Tensor): - storages.append(StorageWeakRef(value.untyped_storage())) + watch(result) return result except BaseException as error: errors.append(error) raise finally: - args = kwargs = result = value = None + args = kwargs = result = None return call @@ -399,45 +402,56 @@ def call(*args, **kwargs): "importance_weights", observe(corrections.importance_weights, inputs=True), ) - packets = [] - transaction = nullcontext() - if phase == "metadata": - trainer = TrainerRank.__new__(TrainerRank) - parameter = torch.nn.Parameter(torch.tensor(2.0)) - trainer._checkpoint_slots = {"student": _CheckpointSlot(params=(parameter,))} - trainer.runtime = SimpleNamespace(model=[], optimizer=None) - snapshot = trainer._snapshot_parameter( - parameter, trainer._capture_checkpoint_version("student") - ) - first, _ = cache.run( - lambda x: (snapshot * x,), torch.tensor(3.0), retention="replay" - ) - packets.append((first, (torch.tensor(1.0),))) - parameter.grad = prior_gradient = torch.tensor(7.0) - transaction = trainer._gradient_transaction() + return storages, errors, borrowed + + +@pytest.mark.parametrize( + "phase", ["always", "coordinator", "correction_peer", "backward_peer"] +) +def test_correction_and_coordinator_failures_release_owned_storage( + monkeypatch, request, phase +): + from art.trainer_rank import _commands as commands + + cache = GraphCache() + older_weight = torch.nn.Parameter(torch.tensor(3.0)) + older, (older_output,) = cache.run(lambda _: (older_weight.square(),), ()) + storages, errors, borrowed = _observe_backward_storage(monkeypatch, request) cache, handle, original, current, executions, _ = _corrected_cache( - policy="always", - retention="replay" if phase == "metadata" else "gpu", - cache=cache, + policy="always", cache=cache ) - if phase == "metadata": - execute = cache._records[handle].execute - cache._records[handle].execute = lambda x: tuple( - value.expand(2) if torch.is_grad_enabled() else value - for value in execute(x) - ) gradient = torch.tensor(1.0) - borrowed = [ + borrowed[:] = [ StorageWeakRef(value.untyped_storage()) for value in (original, current, gradient) ] primary, cause = RuntimeError("peer backward failed"), ValueError("peer cause") calls = 0 + peer_phase = phase in ("correction_peer", "backward_peer") + fail_at = 1 if phase == "correction_peer" else 2 + + def exchange(failures, local, *, group): + nonlocal calls + calls += 1 + failures[:] = [local, "peer prepare failed" if calls == fail_at else None] + + if peer_phase: + monkeypatch.setattr( + commands, + "dist", + SimpleNamespace( + is_initialized=lambda: True, + get_world_size=lambda group: 2, + all_gather_object=exchange, + ), + ) def coordinate(function): nonlocal calls + if peer_phase: + return commands._coordinate_call(function, group=None) calls += 1 - result = function() + result = commands._coordinate_call(function, group=None) if phase == "coordinator" and calls == 3: raise primary from cause return result @@ -445,19 +459,16 @@ def coordinate(function): if phase == "always": with torch.no_grad(): current.fill_(float("nan")) - packets.append((handle, (gradient,))) with pytest.raises(ValueError if phase == "always" else RuntimeError) as failure: - with transaction: - cache.backward_many(packets, coordinate=coordinate) - if phase != "coordinator": + cache.backward_many(((handle, (gradient,)),), coordinate=coordinate) + if peer_phase: + assert "Physical trainer preflight failed" in str(failure.value) + assert calls == fail_at and original.grad is None + assert failure.value.__cause__ is None + elif phase != "coordinator": assert errors and all(error is failure.value for error in errors) - assert ( - "metadata differs" if phase == "metadata" else "current logprobs" - ) in str(failure.value) + assert "current logprobs" in str(failure.value) assert original.grad is None - if phase == "metadata": - assert parameter.grad is prior_gradient and parameter.grad.item() == 7 - assert snapshot.grad is None and not trainer._version_state()._origins else: assert failure.value is primary and primary.__cause__ is cause and calls == 3 torch.testing.assert_close( @@ -467,14 +478,12 @@ def coordinate(function): assert len(storages) >= 4 and all(storage.expired() for storage in storages) assert all(not storage.expired() for storage in borrowed) assert gradient.item() == 1 and original.item() == -1 and current.grad is None - assert executions == [(False, True), (True, False)] + ( - [(False, True)] if phase == "metadata" else [] - ) + assert executions == [(False, True), (True, False)] assert cache.handles() == (older,) cache.backward(older, (torch.ones_like(older_output),)) torch.testing.assert_close(older_weight.grad, torch.tensor(6.0)) cache, retry, retry_weight, _, _, _ = _corrected_cache(policy="always", cache=cache) - cache.backward(retry, (gradient,)) + cache.backward_many(((retry, (gradient,)),), coordinate=coordinate) torch.testing.assert_close( retry_weight.grad, torch.tensor(2.0) * torch.tensor(0.75).exp() ) @@ -631,12 +640,10 @@ def test_replay_restores_original_autocast_context(): torch.testing.assert_close(parameter.grad, torch.full_like(parameter, 2.0)) +@pytest.mark.parametrize("always_prepass", [False, True], ids=["uncorrected", "always"]) def test_replay_failure_discards_transaction_and_releases_participating_records( - request, + monkeypatch, request, always_prepass ): - if gc.isenabled(): - request.addfinalizer(gc.enable) - gc.disable() trainer = TrainerRank.__new__(TrainerRank) parameter = torch.nn.Parameter(torch.tensor(2.0)) trainer._checkpoint_slots = {"student": _CheckpointSlot(params=(parameter,))} @@ -647,7 +654,11 @@ def test_replay_failure_discards_transaction_and_releases_participating_records( cache = GraphCache() older_weight = torch.nn.Parameter(torch.tensor(3.0)) older, (older_output,) = cache.run(lambda _: (older_weight.square(),), ()) - borrowed = StorageWeakRef(parameter.untyped_storage()) + observed, errors, borrowed = _observe_backward_storage(monkeypatch, request) + gradient = torch.tensor(1.0) + borrowed[:] = [ + StorageWeakRef(value.untyped_storage()) for value in (parameter, gradient) + ] storages = [] executions = [0] @@ -658,24 +669,36 @@ def failing(x): storages.extend( StorageWeakRef(value.untyped_storage()) for value in (activation, result) ) - return (result if executions[0] == 1 else result.expand(2),) + return ( + result + if executions[0] == 1 or not torch.is_grad_enabled() + else result.expand(2), + ) first, _ = cache.run( lambda x: (snapshot * x,), torch.tensor(3.0), retention="replay" ) - second, _ = cache.run(failing, torch.tensor(4.0), retention="replay") + second, outputs = cache.run(failing, torch.tensor(4.0), retention="replay") + if always_prepass: + cache.set_corrections( + second, + _logprob_corrections(outputs, "always"), + is_stale=lambda: True, + current_context_factory=nullcontext, + ) parameter.grad = prior_gradient = torch.tensor(7.0) with pytest.raises(RuntimeError, match="metadata differs") as failure: with trainer._gradient_transaction(): - cache.backward_many( - [(first, (torch.tensor(1.0),)), (second, (torch.tensor(1.0),))] - ) + cache.backward_many([(first, (gradient,)), (second, (gradient,))]) assert parameter.grad is prior_gradient and parameter.grad.item() == 7 assert snapshot.grad is None assert not trainer._version_state()._origins assert failure.value.__traceback__ is not None and failure.value.__cause__ is None - assert len(storages) == 4 and all(storage.expired() for storage in storages) - assert not borrowed.expired() and parameter.item() == snapshot.item() == 2 + assert errors and all(error is failure.value for error in errors) + assert len(storages) == 4 + 2 * always_prepass + assert all(storage.expired() for storage in (*storages, *observed)) + assert all(not storage.expired() for storage in borrowed) + assert parameter.item() == snapshot.item() == 2 and gradient.item() == 1 assert cache.handles() == (older,) torch.testing.assert_close(older_output, torch.tensor(9.0)) cache.backward(older, (torch.ones_like(older_output),)) From 1f594e73754126d5f98e5812b887a74f4881e2d7 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 04:53:27 +0000 Subject: [PATCH 107/150] Simplify backward orchestration and distributed test coordination --- src/art/trainer_rank/_graphs.py | 108 +++++++++----------- tests/unit/test_trainer_rank_graph_order.py | 12 +-- tests/unit/test_trainer_rank_versions.py | 13 +-- 3 files changed, 53 insertions(+), 80 deletions(-) diff --git a/src/art/trainer_rank/_graphs.py b/src/art/trainer_rank/_graphs.py index 8c0ac8e37..6f3657a3f 100644 --- a/src/art/trainer_rank/_graphs.py +++ b/src/art/trainer_rank/_graphs.py @@ -572,11 +572,55 @@ def backward_many( ) -> None: self.validate_many(packets) try: - self._backward_validated( - packets, - retain_graph=retain_graph, - coordinate=coordinate or (lambda function: function()), - ) + coordinate = coordinate or (lambda function: function()) + # Random wire handles differ across physical ranks. Creation order is + # shared by TP/CP peers and therefore fixes backward collective order. + ordinals = {handle: index for index, handle in enumerate(self._records)} + ordered_packets = sorted(packets, key=lambda packet: ordinals[packet[0]]) + prepared = {} + stale = { + handle: record.is_stale is not None and record.is_stale() + for handle, _ in ordered_packets + for record in (self._records[handle],) + } + # The coordinator uses this logical rank's TP/CP group, never all DP + # ranks: different DP owners may have different numbers of graphs. + try: + for handle, gradients in ordered_packets: + record = self._records[handle] + prepared[handle] = coordinate( + lambda: self._prepare_correction( + record, gradients, stale[handle] + ) + ) + for handle, gradients in ordered_packets: + record = self._records[handle] + pairs = coordinate( + lambda: self._prepare_backward( + record, gradients, stale[handle], prepared[handle] + ) + ) + try: + coordinate( + lambda: ( + torch.autograd.backward( + [output for output, _ in pairs], + [gradient for _, gradient in pairs], + retain_graph=retain_graph, + ) + if pairs + else None + ) + ) + finally: + record.restored.clear() + pairs.clear() + if not retain_graph: + self.release(handle) + elif record.retention == "replay": + self.evict(handle) + finally: + prepared.clear() except BaseException: # A replay/backward failure consumes the operation. The enclosing # checkpoint transaction discards unpublished optimizer gradients. @@ -584,60 +628,6 @@ def backward_many( self.release(handle) raise - def _backward_validated( - self, - packets: Sequence[tuple[ForwardHandle, Sequence[torch.Tensor | None]]], - *, - retain_graph: bool, - coordinate: Callable[[Callable[[], Any]], Any], - ) -> None: - # Random wire handles differ across physical ranks. Creation order is - # shared by TP/CP peers and therefore fixes backward collective order. - ordinals = {handle: index for index, handle in enumerate(self._records)} - packets = sorted(packets, key=lambda packet: ordinals[packet[0]]) - prepared = {} - stale = { - handle: record.is_stale is not None and record.is_stale() - for handle, _ in packets - for record in (self._records[handle],) - } - # The coordinator uses this logical rank's TP/CP group, never all DP - # ranks: different DP owners may have different numbers of graphs. - try: - for handle, gradients in packets: - record = self._records[handle] - prepared[handle] = coordinate( - lambda: self._prepare_correction(record, gradients, stale[handle]) - ) - for handle, gradients in packets: - record = self._records[handle] - pairs = coordinate( - lambda: self._prepare_backward( - record, gradients, stale[handle], prepared[handle] - ) - ) - try: - coordinate( - lambda: ( - torch.autograd.backward( - [output for output, _ in pairs], - [gradient for _, gradient in pairs], - retain_graph=retain_graph, - ) - if pairs - else None - ) - ) - finally: - record.restored.clear() - pairs.clear() - if not retain_graph: - self.release(handle) - elif record.retention == "replay": - self.evict(handle) - finally: - prepared.clear() - @staticmethod def _prepare_correction(record, gradients, stale): if not ( diff --git a/tests/unit/test_trainer_rank_graph_order.py b/tests/unit/test_trainer_rank_graph_order.py index 95110ccc7..f2863224e 100644 --- a/tests/unit/test_trainer_rank_graph_order.py +++ b/tests/unit/test_trainer_rank_graph_order.py @@ -9,6 +9,7 @@ from trainer_rank_test_support import gloo_group from art.trainer_rank import TrainerRank, _graphs +from art.trainer_rank._commands import _coordinate_call from art.trainer_rank._impl import _CheckpointSlot @@ -64,16 +65,7 @@ def execute(x, snapshot=snapshot, tag=tag, calls=calls): packets.append((handle, (torch.tensor(1.0),))) def coordinate(function): - result, error = None, None - try: - result = function() - except Exception as exc: - error = str(exc) - errors = [None, None] - dist.all_gather_object(errors, error) - if any(errors): - raise RuntimeError(str(errors)) - return result + return _coordinate_call(function, group=None) for parameter in parameters: parameter.grad = torch.tensor(7.0) diff --git a/tests/unit/test_trainer_rank_versions.py b/tests/unit/test_trainer_rank_versions.py index 1a89fa519..d93455b36 100644 --- a/tests/unit/test_trainer_rank_versions.py +++ b/tests/unit/test_trainer_rank_versions.py @@ -13,6 +13,7 @@ from trainer_rank_test_support import gloo_group, spawn_and_join from art.trainer_rank import TrainerRank, TrainerRankSlotStateError +from art.trainer_rank._commands import _coordinate_call from art.trainer_rank._impl import _CheckpointSlot @@ -425,17 +426,7 @@ def _transaction_exit_failure_worker(rank: int, rendezvous: str, nested: bool) - def coordinate(validate) -> None: calls.append(True) - error = None - try: - validate() - except BaseException as exc: - error = exc - errors = [None, None] - dist.all_gather_object(errors, None if error is None else str(error)) - if error is not None: - raise error - if any(errors): - raise RuntimeError(f"another rank failed: {errors}") + _coordinate_call(validate, group=None) saved_error = None reference = None From e84a8578b8d754ed6c4f3461b442129f26d2a833 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 04:58:11 +0000 Subject: [PATCH 108/150] Narrow filtered checkpoint slot modules for static typing --- src/art/trainer_rank/_checkpoint.py | 10 +++++++--- 1 file changed, 7 insertions(+), 3 deletions(-) diff --git a/src/art/trainer_rank/_checkpoint.py b/src/art/trainer_rank/_checkpoint.py index 34f214e61..f19c035b4 100644 --- a/src/art/trainer_rank/_checkpoint.py +++ b/src/art/trainer_rank/_checkpoint.py @@ -1600,6 +1600,12 @@ def _load_adapter( def _slot_snapshot(trainer: TrainerRank) -> _SlotSnapshot: + modules = ( + module + for chunk in trainer.runtime.model + for module in chunk.modules() + if hasattr(module, "_slot_keys") and hasattr(module, "_slot_modules") + ) return tuple( ( module, @@ -1607,9 +1613,7 @@ def _slot_snapshot(trainer: TrainerRank) -> _SlotSnapshot: dict(module._slot_modules.items()), {key: getattr(slot, "ref") for key, slot in module._slot_modules.items()}, ) - for chunk in trainer.runtime.model - for module in chunk.modules() - if hasattr(module, "_slot_keys") and hasattr(module, "_slot_modules") + for module in cast("Iterable[LoRA]", modules) ) From 2b980b226950c3a0f4cbe845211babe0f35fb903 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 04:54:59 +0000 Subject: [PATCH 109/150] Release undelivered executor results after peer failure --- src/art/trainer_rank/_commands.py | 1 + .../unit/test_trainer_rank_output_delivery.py | 66 +++++++++++++++++++ 2 files changed, 67 insertions(+) diff --git a/src/art/trainer_rank/_commands.py b/src/art/trainer_rank/_commands.py index 2f71c1c20..c8b4daa3f 100644 --- a/src/art/trainer_rank/_commands.py +++ b/src/art/trainer_rank/_commands.py @@ -403,6 +403,7 @@ def _execute(self, command: _Command) -> Any: None if error is None else f"{type(error).__name__}: {error}" ) if any(errors): + result = None self.state.graphs.pop( f"{self.mode}:{command.sequence}:dp:{self.dp_rank}", None ) diff --git a/tests/unit/test_trainer_rank_output_delivery.py b/tests/unit/test_trainer_rank_output_delivery.py index 4b63d43bd..cdb4e6b80 100644 --- a/tests/unit/test_trainer_rank_output_delivery.py +++ b/tests/unit/test_trainer_rank_output_delivery.py @@ -3,12 +3,14 @@ from collections.abc import Generator from dataclasses import replace from functools import partial +import gc from typing import Any import weakref import pytest from test_trainer_rank_commands import _input, _Rank import torch +from torch.multiprocessing.reductions import StorageWeakRef from art.trainer_rank import ( ForwardInput, @@ -124,6 +126,70 @@ def test_failed_delivery_releases_registered_graph_and_native_cache( assert not executor.state.graphs and not rank.cache.handles() +@pytest.mark.parametrize("operation", ["forward", "next", "batches_next"]) +def test_peer_rejection_releases_undelivered_packet_storage( + monkeypatch, request, operation +): + if gc.isenabled(): + request.addfinalizer(gc.enable) + gc.disable() + rank: Any = _CachedRank() + executor = _Executor(rank, "zero") + view = _view(executor) + previous = view.forward(_input(7)) + graphs, handles = set(executor.state.graphs), rank.cache.handles() + borrowed = StorageWeakRef(rank.weight.untyped_storage()) + peer_rank: Any = _Rank(1, 2) + peer = _Executor(peer_rank, "zero") + monkeypatch.setattr(peer, "_available_host_memory", lambda: 0) + with pytest.raises(MemoryError, match="output snapshot") as rejected: + peer.invoke("forward", _input(5)) + assert not peer.state.graphs + argument = ( + _input(3) + if operation == "forward" + else executor.invoke( + "batches" if operation == "next" else "batches_open", [_input(3)] + ) + ) + storages, packets, exchanges = [], [], [] + packet = executor._packet + + def observe(*args): + value = packet(*args) + packets.append(weakref.ref(value)) + storages.extend( + StorageWeakRef(t.untyped_storage()) for t in value.packet.tensors + ) + return value + + def gather(error): + exchanges.append(error) + return [error, f"MemoryError: {rejected.value}"] + + monkeypatch.setattr(executor, "_packet", observe) + with monkeypatch.context() as patch: + patch.setattr(executor, "_gather", gather) + with pytest.raises(RuntimeError, match="output snapshot") as failure: + executor.invoke(operation, argument) + if operation != "forward": + executor.invoke("close" if operation == "next" else "batches_close", argument) + assert exchanges == [None] # This leader completed its real forward and packet. + assert failure.value.__traceback__ is not None + assert rejected.value.__traceback__ is not None + assert set(executor.state.graphs) == graphs and rank.cache.handles() == handles + assert storages and all(storage.expired() for storage in storages) + assert all(packet() is None for packet in packets) + assert not borrowed.expired() and rank.weight.item() == 2 + view.backward(previous.hidden_states.sum()) + assert rank.weight.grad.item() == 7 + result = executor.invoke("forward", _input(11)) + assert result[0] is packets[-1]() + executor.invoke("release", (result[0].packet.handle,)) + assert not executor.state.graphs and not rank.cache.handles() + assert not storages[-1].expired() and result[0].packet.tensors[0].item() == 22 + + @pytest.mark.parametrize("delivery", ["forward", "iterator", "persistent"]) def test_later_wave_failure_keeps_only_previously_delivered_outputs( monkeypatch, delivery From b5765d1ced4d052e8e2b5683370d1a459ebc2a81 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 05:04:05 +0000 Subject: [PATCH 110/150] Release undelivered output aliases across executor failures --- src/art/trainer_rank/_commands.py | 187 ++++++++++-------- .../unit/test_trainer_rank_output_delivery.py | 67 +++++-- 2 files changed, 156 insertions(+), 98 deletions(-) diff --git a/src/art/trainer_rank/_commands.py b/src/art/trainer_rank/_commands.py index c8b4daa3f..919494115 100644 --- a/src/art/trainer_rank/_commands.py +++ b/src/art/trainer_rank/_commands.py @@ -395,84 +395,93 @@ def _gather(self, value: Any) -> list[Any]: def _execute(self, command: _Command) -> Any: result, error = None, None try: - with torch.set_grad_enabled(command.grad_enabled): - result = self._dispatch(command) - except BaseException as exc: - error = exc - errors = self._gather( - None if error is None else f"{type(error).__name__}: {error}" - ) - if any(errors): - result = None - self.state.graphs.pop( - f"{self.mode}:{command.sequence}:dp:{self.dp_rank}", None - ) - if error is not None: - raise error - raise RuntimeError( - f"Physical trainer command {command.operation!r} failed: {errors}" - ) - if command.operation in ("forward", "next", "batches_next"): - value = ( - result if get_rank_callback_metadata(self.rank) is not None else None - ) try: - return self._gather_outputs(value) - except BaseException: + with torch.set_grad_enabled(command.grad_enabled): + result = self._dispatch(command) + except BaseException as exc: + error = exc + errors = self._gather( + None if error is None else f"{type(error).__name__}: {error}" + ) + if any(errors): self.state.graphs.pop( f"{self.mode}:{command.sequence}:dp:{self.dp_rank}", None ) - raise - return result + if error is not None: + raise error + raise RuntimeError( + f"Physical trainer command {command.operation!r} failed: {errors}" + ) + if command.operation in ("forward", "next", "batches_next"): + try: + return self._gather_outputs( + result + if get_rank_callback_metadata(self.rank) is not None + else None + ) + except BaseException: + self.state.graphs.pop( + f"{self.mode}:{command.sequence}:dp:{self.dp_rank}", None + ) + raise + return result + finally: + result = None def _gather_outputs(self, value: Any) -> list[Any] | None: - if not self.distributed or len(self.members) == 1: - return [value] + try: + if not self.distributed or len(self.members) == 1: + return [value] - def admit_serialization() -> None: - from ._tensors import flatten_tensors + def admit_serialization() -> None: + from ._tensors import flatten_tensors - tensors, _ = flatten_tensors(value) - # Pickling creates storage bytes before gather admission can sample - # their size. Reserve tensor storage, a copy, and per-leaf metadata. - required = 2 * sum(t.numel() * t.element_size() for t in tensors) - required += 4096 * (len(tensors) + 1) - if required > self._available_host_memory(): - raise MemoryError( - f"Trainer output serialization requires {required} CPU bytes" - ) + tensors, _ = flatten_tensors(value) + # Pickling creates storage bytes before gather admission can sample + # their size. Reserve tensor storage, a copy, and per-leaf metadata. + try: + required = 2 * sum(t.numel() * t.element_size() for t in tensors) + required += 4096 * (len(tensors) + 1) + finally: + del tensors, _ + if required > self._available_host_memory(): + raise MemoryError( + f"Trainer output serialization requires {required} CPU bytes" + ) - self._coordinated_preflight(admit_serialization) - payload = self._coordinated_preflight(lambda: cloudpickle.dumps(value)) - sizes = self._gather(len(payload)) - - def admit() -> None: - available = self._available_host_memory() - # Gloo gather_object pads every sender to the largest serialized - # payload. Include receive storage and unpickling copies on leader. - padded = max(sizes) + 1024 - required = 2 * padded - if self.is_leader: - required += 2 * len(sizes) * padded + sum(sizes) - if required > available: - raise MemoryError( - f"Trainer output transfer requires {required} CPU bytes, " - f"but the per-process shared-host budget has {available}" - ) + self._coordinated_preflight(admit_serialization) + payload = self._coordinated_preflight(lambda: cloudpickle.dumps(value)) + sizes = self._gather(len(payload)) + + def admit() -> None: + available = self._available_host_memory() + # Gloo gather_object pads every sender to the largest serialized + # payload. Include receive storage and unpickling copies on leader. + padded = max(sizes) + 1024 + required = 2 * padded + if self.is_leader: + required += 2 * len(sizes) * padded + sum(sizes) + if required > available: + raise MemoryError( + f"Trainer output transfer requires {required} CPU bytes, " + f"but the per-process shared-host budget has {available}" + ) - self._coordinated_preflight(admit) - values: list[Any] | None = ( - [None] * len(self.members) if self.is_leader else None - ) - dist.gather_object(payload, values, dst=self.leader, group=self.group) + self._coordinated_preflight(admit) + values: list[Any] | None = ( + [None] * len(self.members) if self.is_leader else None + ) + dist.gather_object(payload, values, dst=self.leader, group=self.group) - def decode() -> list[Any] | None: - if values is None: - return None - decoded = [cloudpickle.loads(item) for item in values] - return [item for item in decoded if item is not None] + def decode() -> list[Any] | None: + if values is None: + return None + decoded = [cloudpickle.loads(item) for item in values] + return [item for item in decoded if item is not None] - return self._coordinated_preflight(decode) + return self._coordinated_preflight(decode) + finally: + value = payload = values = None def _available_host_memory(self) -> int: from ._memory_policy import host_memory_budget, local_rank_count @@ -488,20 +497,25 @@ def _available_host_memory(self) -> int: def _packet(self, tree: Any, sequence: int) -> Any: from ._tensors import ManagedTensor, detach_tree, flatten_tensors - handle = f"{self.mode}:{sequence}:dp:{self.dp_rank}" - tensors, _ = flatten_tensors(tree) - if any(tensor.requires_grad for tensor in tensors): - self.state.graphs[handle] = tuple(tensors) - if get_rank_callback_metadata(self.rank) is None: - return None - required = sum(tensor.numel() * tensor.element_size() for tensor in tensors) - if required > self._available_host_memory(): - raise MemoryError(f"Trainer output snapshot requires {required} CPU bytes") - return _OutputPacket( - detach_tree(handle, tree, device="cpu"), - tuple(tensor.device.type == "cpu" for tensor in tensors), - any(isinstance(tensor, ManagedTensor) for tensor in tensors), - ) + try: + handle = f"{self.mode}:{sequence}:dp:{self.dp_rank}" + tensors, _ = flatten_tensors(tree) + if any(tensor.requires_grad for tensor in tensors): + self.state.graphs[handle] = tuple(tensors) + if get_rank_callback_metadata(self.rank) is None: + return None + required = sum(tensor.numel() * tensor.element_size() for tensor in tensors) + if required > self._available_host_memory(): + raise MemoryError( + f"Trainer output snapshot requires {required} CPU bytes" + ) + return _OutputPacket( + detach_tree(handle, tree, device="cpu"), + tuple(tensor.device.type == "cpu" for tensor in tensors), + any(isinstance(tensor, ManagedTensor) for tensor in tensors), + ) + finally: + tree = tensors = _ = None def _dispatch(self, command: _Command) -> Any: op, args, kwargs = command.operation, command.args, command.kwargs @@ -524,11 +538,14 @@ def _dispatch(self, command: _Command) -> Any: ) if op in ("next", "batches_next"): batch = next(iterators[args[0]], None) - if batch is None: - return (None, None) - return replace(batch, inputs=[], outputs=[]), self._packet( - batch.outputs, command.sequence - ) + try: + if batch is None: + return (None, None) + return replace(batch, inputs=[], outputs=[]), self._packet( + batch.outputs, command.sequence + ) + finally: + batch = None iterator = iterators.pop(args[0], None) if op == "batches_close": self.state.batch_inputs.pop(args[0], None) diff --git a/tests/unit/test_trainer_rank_output_delivery.py b/tests/unit/test_trainer_rank_output_delivery.py index cdb4e6b80..9150dc525 100644 --- a/tests/unit/test_trainer_rank_output_delivery.py +++ b/tests/unit/test_trainer_rank_output_delivery.py @@ -126,9 +126,12 @@ def test_failed_delivery_releases_registered_graph_and_native_cache( assert not executor.state.graphs and not rank.cache.handles() +@pytest.mark.parametrize( + "phase", ["peer", "snapshot", "serialization", "transfer", "exchange"] +) @pytest.mark.parametrize("operation", ["forward", "next", "batches_next"]) def test_peer_rejection_releases_undelivered_packet_storage( - monkeypatch, request, operation + monkeypatch, request, operation, phase ): if gc.isenabled(): request.addfinalizer(gc.enable) @@ -156,29 +159,67 @@ def test_peer_rejection_releases_undelivered_packet_storage( packet = executor._packet def observe(*args): - value = packet(*args) - packets.append(weakref.ref(value)) - storages.extend( - StorageWeakRef(t.untyped_storage()) for t in value.packet.tensors - ) - return value + try: + if phase == "snapshot": + storages.extend( + StorageWeakRef(t.untyped_storage()) + for t in _tensors.flatten_tensors(args[0])[0] + ) + value = packet(*args) + packets.append(weakref.ref(value)) + storages.extend( + StorageWeakRef(t.untyped_storage()) for t in value.packet.tensors + ) + return value + finally: + args = () # Observation must not retain the failed call's input tree. - def gather(error): - exchanges.append(error) - return [error, f"MemoryError: {rejected.value}"] + exchange_error = RuntimeError("status exchange failed") + + def gather(value): + exchanges.append(value) + if phase == "exchange": + raise exchange_error + return [value, f"MemoryError: {rejected.value}" if phase == "peer" else value] monkeypatch.setattr(executor, "_packet", observe) with monkeypatch.context() as patch: patch.setattr(executor, "_gather", gather) - with pytest.raises(RuntimeError, match="output snapshot") as failure: + if phase == "snapshot": + patch.setattr(executor, "_available_host_memory", lambda: 0) + if phase in ("serialization", "transfer"): + budgets = iter( + [2**20, 0] if phase == "serialization" else [2**20, 2**20, 0] + ) + patch.setattr(executor, "_available_host_memory", lambda: next(budgets)) + patch.setattr(executor, "distributed", True) + patch.setattr(executor, "members", [0, 1]) + patch.setattr(executor, "_broadcast", lambda command: command) + message = ( + "output snapshot" + if phase == "peer" + else "status exchange" + if phase == "exchange" + else f"output {phase}" + ) + with pytest.raises((MemoryError, RuntimeError), match=message) as failure: executor.invoke(operation, argument) + if phase == "exchange": + assert failure.value is exchange_error + # A failed status collective has no recovery contract; release only this + # test's successful graph through the existing explicit release operation. + pending = set(executor.state.graphs) - graphs + assert len(pending) == 1 + executor.invoke("release", tuple(pending)) if operation != "forward": executor.invoke("close" if operation == "next" else "batches_close", argument) - assert exchanges == [None] # This leader completed its real forward and packet. + assert (exchanges[0] is None) is (phase != "snapshot") + assert len(exchanges) == (2 if phase == "transfer" else 1) assert failure.value.__traceback__ is not None assert rejected.value.__traceback__ is not None - assert set(executor.state.graphs) == graphs and rank.cache.handles() == handles + assert set(executor.state.graphs) == graphs assert storages and all(storage.expired() for storage in storages) + assert rank.cache.handles() == handles assert all(packet() is None for packet in packets) assert not borrowed.expired() and rank.weight.item() == 2 view.backward(previous.hidden_states.sum()) From aa55abe3fa8c25b3924926bbb67e3f0ab33411c1 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 05:23:53 +0000 Subject: [PATCH 111/150] Release detached output references after copy failures --- src/art/trainer_rank/_tensors.py | 25 ++++++---- .../unit/test_trainer_rank_output_delivery.py | 49 +++++++++++++++---- 2 files changed, 54 insertions(+), 20 deletions(-) diff --git a/src/art/trainer_rank/_tensors.py b/src/art/trainer_rank/_tensors.py index aec5cea3c..47ff44097 100644 --- a/src/art/trainer_rank/_tensors.py +++ b/src/art/trainer_rank/_tensors.py @@ -178,16 +178,21 @@ def detach_tree( handle: str, tree: Any, *, device: torch.device | str | None = None ) -> TensorPacket: """Snapshot supported output containers; reject opaque tensor-bearing objects.""" - tensors, spec = flatten_tensors(tree) - _validate_output_spec(spec) - return TensorPacket( - handle, - spec, - tuple( - _plain(tensor).detach().to(device=device, copy=True) for tensor in tensors - ), - tuple(tensor.requires_grad for tensor in tensors), - ) + copies: list[torch.Tensor] = [] + try: + tensors, spec = flatten_tensors(tree) + _validate_output_spec(spec) + for tensor in tensors: + copies.append(_plain(tensor).detach().to(device=device, copy=True)) + return TensorPacket( + handle, + spec, + tuple(copies), + tuple(tensor.requires_grad for tensor in tensors), + ) + finally: + tree = tensors = tensor = None + del copies class _OutputBridge(torch.autograd.Function): diff --git a/tests/unit/test_trainer_rank_output_delivery.py b/tests/unit/test_trainer_rank_output_delivery.py index 9150dc525..0170e14bb 100644 --- a/tests/unit/test_trainer_rank_output_delivery.py +++ b/tests/unit/test_trainer_rank_output_delivery.py @@ -127,7 +127,7 @@ def test_failed_delivery_releases_registered_graph_and_native_cache( @pytest.mark.parametrize( - "phase", ["peer", "snapshot", "serialization", "transfer", "exchange"] + "phase", ["peer", "snapshot", "copy", "serialization", "transfer", "exchange"] ) @pytest.mark.parametrize("operation", ["forward", "next", "batches_next"]) def test_peer_rejection_releases_undelivered_packet_storage( @@ -148,23 +148,25 @@ def test_peer_rejection_releases_undelivered_packet_storage( with pytest.raises(MemoryError, match="output snapshot") as rejected: peer.invoke("forward", _input(5)) assert not peer.state.graphs + inputs = [_input(3), _input(5)] if phase == "copy" else _input(3) argument = ( - _input(3) + inputs if operation == "forward" else executor.invoke( - "batches" if operation == "next" else "batches_open", [_input(3)] + "batches" if operation == "next" else "batches_open", [inputs] ) ) storages, packets, exchanges = [], [], [] + copy_targets, copy_calls, copy_error = set(), 0, None packet = executor._packet def observe(*args): try: - if phase == "snapshot": - storages.extend( - StorageWeakRef(t.untyped_storage()) - for t in _tensors.flatten_tensors(args[0])[0] - ) + if phase in ("snapshot", "copy"): + for tensor in _tensors.flatten_tensors(args[0])[0]: + storages.append(StorageWeakRef(tensor.untyped_storage())) + copy_targets.add(tensor.untyped_storage().data_ptr()) + tensor = None value = packet(*args) packets.append(weakref.ref(value)) storages.extend( @@ -174,6 +176,27 @@ def observe(*args): finally: args = () # Observation must not retain the failed call's input tree. + to = torch.Tensor.to + + def copy(tensor, *args, **kwargs): + nonlocal copy_calls, copy_error + try: + targeted = tensor.untyped_storage().data_ptr() in copy_targets + if targeted: + copy_calls += 1 + if copy_calls == 2: + # Real CPU Tensor.to failure after one successful owned copy. + kwargs["memory_format"] = torch.channels_last + copied = to(tensor, *args, **kwargs) + if targeted: + storages.append(StorageWeakRef(copied.untyped_storage())) + return copied + except RuntimeError as error: + copy_error = error + raise + finally: + tensor = None # The fault injector must not own the failing tensor. + exchange_error = RuntimeError("status exchange failed") def gather(value): @@ -185,6 +208,8 @@ def gather(value): monkeypatch.setattr(executor, "_packet", observe) with monkeypatch.context() as patch: patch.setattr(executor, "_gather", gather) + if phase == "copy": + patch.setattr(torch.Tensor, "to", copy) if phase == "snapshot": patch.setattr(executor, "_available_host_memory", lambda: 0) if phase in ("serialization", "transfer"): @@ -196,7 +221,9 @@ def gather(value): patch.setattr(executor, "members", [0, 1]) patch.setattr(executor, "_broadcast", lambda command: command) message = ( - "output snapshot" + "required rank 4" + if phase == "copy" + else "output snapshot" if phase == "peer" else "status exchange" if phase == "exchange" @@ -204,6 +231,8 @@ def gather(value): ) with pytest.raises((MemoryError, RuntimeError), match=message) as failure: executor.invoke(operation, argument) + if phase == "copy": + assert failure.value is copy_error and copy_calls == 2 and len(storages) == 3 if phase == "exchange": assert failure.value is exchange_error # A failed status collective has no recovery contract; release only this @@ -213,7 +242,7 @@ def gather(value): executor.invoke("release", tuple(pending)) if operation != "forward": executor.invoke("close" if operation == "next" else "batches_close", argument) - assert (exchanges[0] is None) is (phase != "snapshot") + assert (exchanges[0] is None) is (phase not in ("snapshot", "copy")) assert len(exchanges) == (2 if phase == "transfer" else 1) assert failure.value.__traceback__ is not None assert rejected.value.__traceback__ is not None From 4530fe09d1c6f49950586e7657fbce41d8033244 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 05:55:22 +0000 Subject: [PATCH 112/150] Bound asynchronous checkpoint captures by host headroom --- src/art/trainer_rank/_checkpoint.py | 95 +++++++- tests/unit/test_checkpoint_snapshot_spill.py | 205 ++++++++++++++---- ...t_checkpoint_snapshot_spill_distributed.py | 71 +++--- 3 files changed, 292 insertions(+), 79 deletions(-) diff --git a/src/art/trainer_rank/_checkpoint.py b/src/art/trainer_rank/_checkpoint.py index dae0cc5da..219e9063c 100644 --- a/src/art/trainer_rank/_checkpoint.py +++ b/src/art/trainer_rank/_checkpoint.py @@ -129,22 +129,36 @@ def __init__(self) -> None: tuple[Path, dict[str, dict[str, torch.Tensor]], Future[None]] ] = deque() self.thread: threading.Thread | None = None + self.workspace: dict[Future[None], int] = {} def submit( self, snapshot: Path, payloads: dict[str, dict[str, torch.Tensor]] ) -> Future[None]: result: Future[None] = Future() with self.lock: + self.workspace[result] = max( + ( + sum( + value.numel() + * value.element_size() + * (1 if value.is_contiguous() else 2) + for value in tensors.values() + ) + for tensors in payloads.values() + ), + default=0, + ) self.pending.append((snapshot, payloads, result)) if self.thread is None: - self.thread = threading.Thread( - target=self._run, name="checkpoint-snapshot" - ) try: + self.thread = threading.Thread( + target=self._run, name="checkpoint-snapshot" + ) self.thread.start() except BaseException: self.thread = None self.pending.pop() + self.workspace.pop(result) raise return result @@ -188,6 +202,8 @@ def _run(self) -> None: finally: payloads.clear() tensors = None + with self.lock: + self.workspace.pop(result) if error is None: result.set_result(None) else: @@ -940,6 +956,74 @@ def _local_state( return tuple(records), optimizer, _custom_snapshot(trainer, name, files) +def _admit_snapshot(trainer: TrainerRank, name: str) -> None: + """Estimate registered copies without executing user serialization hooks. + + Unregistered allocations, payload-expanding hooks and concurrent external + allocations are outside this estimate. + """ + from ._impl import _custom_named_parameters + + slot = trainer._checkpoint_slots[name] + custom_params = { + id(param) + for key, custom in slot.custom.items() + for _, param in _custom_named_parameters(key, custom) + } + tensors: list[torch.Tensor] = [ + param for param in slot.params if id(param) not in custom_params + ] + for custom in slot.custom.values(): + if custom.kind == "module": + for _, child in cast(torch.nn.Module, custom.value).named_modules( + remove_duplicate=False + ): + tensors.extend(p for p in child._parameters.values() if p is not None) + tensors.extend( + value + for key, value in child._buffers.items() + if value is not None + and key not in child._non_persistent_buffers_set + ) + else: + tensors.append(cast(torch.Tensor, custom.value)) + cached = slot.custom_payload + if cached is not None: + tensors.extend(cached.tensors.values()) + tensors.extend(cached.optimizer.values()) + size = sum(value.numel() * value.element_size() for value in tensors) + if slot.optimizer is not None: + step_bytes = torch.finfo(torch.get_default_dtype()).bits // 8 + # Three FP32 optimizer components and at most one step per expert element. + size += sum( + 3 * master.numel() * max(4, master.element_size()) + + max(1, master.numel()) * step_bytes + for master in slot.optimizer.master_params + ) + if cached is not None: + size += sum( + 12 * cached.tensors[key].numel() + step_bytes + for record in cached.records.values() + for key in record["trainable_keys"] + if f"master/{key}" not in cached.optimizer + ) + # Capture may overlap source copies/zeros and an older writer. Writing holds + # captured tensors, contiguous packing, and one file's serialized byte strings. + # Fresh headroom already excludes resident captures; do not charge them again. + workspace = 0 + spill = getattr(trainer, "_checkpoint_snapshot_spill", None) + if spill is not None: + with spill.lock: + workspace = max(spill.workspace.values(), default=0) + required = max(2 * size + workspace, 3 * size) + available = trainer._available_cpu_memory_bytes() + if required > available: + raise RuntimeError( + f"Cannot capture checkpoint: estimated host memory needs {required} additional " + f"bytes, {available} available; finish pending saves and retry" + ) + + def prepare_checkpoint_save( trainer: TrainerRank, output_dir: str, checkpoint_name: str ) -> None: @@ -972,6 +1056,11 @@ def prepare_checkpoint_save( ) from ._heads import synchronize_head_buffers + _phase( + lambda: _admit_snapshot(trainer, checkpoint_name), + "admit checkpoint host memory", + group, + ) synchronize_head_buffers(trainer, (checkpoint_name,)) config = deepcopy(_validate_save_state(trainer, checkpoint_name)) if any(value != config for value in _gather(config, group)): diff --git a/tests/unit/test_checkpoint_snapshot_spill.py b/tests/unit/test_checkpoint_snapshot_spill.py index 8cb551eac..99f9cb912 100644 --- a/tests/unit/test_checkpoint_snapshot_spill.py +++ b/tests/unit/test_checkpoint_snapshot_spill.py @@ -3,7 +3,9 @@ from concurrent.futures import Future from dataclasses import replace from pathlib import Path +import sys import threading +from types import SimpleNamespace import weakref import pytest @@ -15,17 +17,10 @@ from art.trainer_rank._impl import _CheckpointSlot, _CustomObject, _DynamicOptimizer -def test_captured_cpu_custom_optimizer_is_independent() -> None: +def _snapshot_trainer(monkeypatch, parameter=None): trainer = _save_state_trainer() - parameter = torch.nn.Parameter(torch.tensor([1.0, 2.0])) - buffer = torch.tensor([3.0]) - master = torch.nn.Parameter(parameter.detach().clone()) - optimizer = torch.optim.Adam((master,), lr=0.125) - optimizer.state[master] = { - "step": torch.tensor(7.0), - "exp_avg": torch.ones(2), - "exp_avg_sq": torch.full((2,), 2.0), - } + if parameter is None: + parameter = torch.nn.Parameter(torch.tensor([1.0])) trainer._checkpoint_slots["a"] = _CheckpointSlot( params=(parameter,), config={ @@ -34,12 +29,40 @@ def test_captured_cpu_custom_optimizer_is_independent() -> None: "lora_alpha": 1, "target_modules": ["q_proj"], }, - optimizer=_DynamicOptimizer(optimizer, (master,)), - custom={ - "p": _CustomObject("parameter", parameter, object()), - "b": _CustomObject("buffer", buffer, object()), - }, + custom={"p": _CustomObject("parameter", parameter, object())}, + ) + monkeypatch.setattr(trainer, "_slot_ref", lambda _: None) + monkeypatch.setitem( + sys.modules, + "art.megatron.lora", + SimpleNamespace(LoRA=type("UnusedLoRA", (), {})), + ) + monkeypatch.setitem( + sys.modules, + "art.megatron.weights.lora_publish", + SimpleNamespace(collect_local_lora_entries=lambda *a, **kw: ({}, [])), ) + return trainer + + +def test_captured_cpu_custom_optimizer_is_independent(monkeypatch) -> None: + parameter = torch.nn.Parameter(torch.tensor([1.0, 2.0])) + buffer = torch.tensor([3.0]) + master = torch.nn.Parameter(parameter.detach().clone()) + optimizer = torch.optim.Adam((master,), lr=0.125) + optimizer.state[master] = { + "step": torch.tensor(7.0), + "exp_avg": torch.ones(2), + "exp_avg_sq": torch.full((2,), 2.0), + } + trainer = _snapshot_trainer(monkeypatch, parameter) + slot = trainer._checkpoint_slots["a"] + slot.optimizer = _DynamicOptimizer(optimizer, (master,)) + slot.custom["b"] = _CustomObject("buffer", buffer, object()) + # Parameter/buffer copies alone fit; the optimizer capture does not. + monkeypatch.setattr(trainer, "_available_cpu_memory_bytes", lambda: 36) + with pytest.raises(RuntimeError, match="checkpoint.*host memory"): + cp._admit_snapshot(trainer, "a") payloads = {} records = cp._custom_snapshot(trainer, "a", payloads) before = { @@ -85,6 +108,7 @@ def write(tensors, path): assert entered.wait(3) assert all(not result.done() for result in results) assert len(spill.pending) == 7 + assert list(spill.workspace.values()) == [16] * 8 finally: release.set() assert worker is not None @@ -93,7 +117,7 @@ def write(tensors, path): for result in results: result.result() assert len(set(calls)) == 1 - assert spill.thread is None and not spill.pending + assert spill.thread is None and not spill.pending and not spill.workspace assert all(not values for values in payloads) assert all(ref() is None for ref in refs) @@ -167,35 +191,12 @@ def write(tensors, path): assert caught.value is error second.result(3) assert (tmp_path / "second/v.safetensors").is_file() + assert not spill.workspace def test_rank_prepare_returns_before_disk_and_owns_capture(tmp_path, monkeypatch): - import sys - from types import SimpleNamespace - - trainer = _save_state_trainer() - parameter = torch.nn.Parameter(torch.tensor([1.0])) - trainer._checkpoint_slots["a"] = _CheckpointSlot( - params=(parameter,), - config={ - "base_model_name_or_path": "test/model", - "r": 1, - "lora_alpha": 1, - "target_modules": ["q_proj"], - }, - custom={"p": _CustomObject("parameter", parameter, object())}, - ) - monkeypatch.setattr(trainer, "_slot_ref", lambda _: None) - monkeypatch.setitem( - sys.modules, - "art.megatron.lora", - SimpleNamespace(LoRA=type("UnusedLoRA", (), {})), - ) - monkeypatch.setitem( - sys.modules, - "art.megatron.weights.lora_publish", - SimpleNamespace(collect_local_lora_entries=lambda *a, **kw: ({}, [])), - ) + trainer = _snapshot_trainer(monkeypatch) + parameter = trainer._checkpoint_slots["a"].params[0] entered, release = threading.Event(), threading.Event() saved = {} original = safetensors.torch.save_file @@ -226,19 +227,20 @@ def write(tensors, path): assert not trainer._prepared_checkpoint_saves -def test_writer_start_failure_has_no_orphaned_backlog(tmp_path, monkeypatch): +@pytest.mark.parametrize("stage", ("__init__", "start")) +def test_writer_start_failure_has_no_orphaned_backlog(tmp_path, monkeypatch, stage): spill = cp._SnapshotSpill() error = RuntimeError("cannot start writer") - def fail(_self): + def fail(_self, *args, **kwargs): raise error with monkeypatch.context() as patch: - patch.setattr(threading.Thread, "start", fail) + patch.setattr(threading.Thread, stage, fail) with pytest.raises(RuntimeError) as caught: spill.submit(tmp_path / "failed", {"v.safetensors": {"v": torch.ones(1)}}) assert caught.value is error - assert spill.thread is None and not spill.pending + assert spill.thread is None and not spill.pending and not spill.workspace spill.submit(tmp_path / "next", {"v.safetensors": {"v": torch.ones(1)}}).result(3) assert (tmp_path / "next/v.safetensors").is_file() @@ -441,3 +443,114 @@ def save(tensors, path): assert threads == [worker] assert raw() is None and all(ref() is None for ref in packed) assert isinstance(result.exception(), OSError) if fails else result.result() is None + + +@pytest.mark.parametrize("backlog", ("empty", "active", "queued")) +def test_snapshot_capture_admission_precedes_allocation(tmp_path, monkeypatch, backlog): + parameter = torch.nn.Parameter(torch.arange(12.0).reshape(3, 4).T) + if backlog == "queued": + parameter.data = parameter.data.contiguous() + trainer = _snapshot_trainer(monkeypatch, parameter) + entered, release = threading.Event(), threading.Event() + available = 384 + calls = [] + original_state, original_write = cp._local_state, safetensors.torch.save_file + + def capture(*args): + calls.append(args[1]) + return original_state(*args) + + def write(tensors, path): + entered.set() + assert release.wait(3) + original_write(tensors, path) + + monkeypatch.setattr(trainer, "_available_cpu_memory_bytes", lambda: available) + monkeypatch.setattr(cp, "_local_state", capture) + monkeypatch.setattr(safetensors.torch, "save_file", write) + output = str(tmp_path / "refused") + try: + if backlog != "empty": + for index in range(2): + trainer.prepare_checkpoint_save(str(tmp_path / str(index)), "a") + parameter.data = parameter.data.T.contiguous().T + assert entered.wait(3) + assert len(trainer._checkpoint_snapshot_spill.pending) == 1 + # Headroom is additional allocation, already excluding resident captures. + available = 144 if backlog != "empty" else 0 + with pytest.raises(RuntimeError, match="checkpoint.*host memory"): + trainer.prepare_checkpoint_save(output, "a") + assert len(calls) == (2 if backlog != "empty" else 0) + assert len(trainer._prepared_checkpoint_saves) == ( + 2 if backlog != "empty" else 0 + ) + assert not trainer._checkpoint_preparing_saves + assert not list(tmp_path.glob(".refused.*")) + if backlog != "empty": + # Existing captures are already resident: do not charge them twice. + available = 192 + trainer.prepare_checkpoint_save(output, "a") + assert len(calls) == 3 + finally: + release.set() + for pending in list(trainer._prepared_checkpoint_saves): + trainer.abort_checkpoint_save(pending) + # Reservations must disappear on drain, and a refused destination is reusable. + available = 144 + trainer.prepare_checkpoint_save(output, "a") + trainer.abort_checkpoint_save(output) + assert not trainer._prepared_checkpoint_saves + torch.testing.assert_close(parameter, torch.arange(12.0).reshape(3, 4).T) + + +@pytest.mark.parametrize("has_optimizer", (False, True)) +def test_lazy_custom_snapshot_admission_counts_cached_state(monkeypatch, has_optimizer): + trainer = _snapshot_trainer(monkeypatch) + slot = trainer._checkpoint_slots["a"] + payloads = {} + records = cp._custom_snapshot(trainer, "a", payloads) + slot.custom.clear() + slot.params = () + cached_optimizer: dict[str, torch.Tensor] = ( + { + f"{key}/p": torch.ones(1) + for key in ("master", "exp_avg", "exp_avg_sq", "step") + } + if has_optimizer + else {} + ) + slot.custom_payload = cp.PreparedCustomPayload( + records, payloads["custom_tensors.safetensors"], cached_optimizer + ) + slot.optimizer = _DynamicOptimizer( + torch.optim.Adam((torch.nn.Parameter(torch.ones(1)),)), () + ) + # Both loaded optimizer data and synthesized missing state need admission. + monkeypatch.setattr(trainer, "_available_cpu_memory_bytes", lambda: 12) + with pytest.raises(RuntimeError, match="checkpoint.*host memory"): + cp._admit_snapshot(trainer, "a") + + +def test_module_snapshot_admission_reads_metadata_without_hooks(monkeypatch): + parameter = torch.nn.Parameter(torch.ones(2)) + trainer = _snapshot_trainer(monkeypatch, parameter) + module = torch.nn.Module() + module.register_parameter("left", parameter) + module.register_parameter("right", parameter) + module.register_buffer("scratch", torch.ones(2), persistent=False) + hooks = [] + module.register_state_dict_pre_hook(lambda *_: hooks.append(True)) + trainer._checkpoint_slots["a"].custom["p"] = _CustomObject( + "module", module, object() + ) + available = 47 + monkeypatch.setattr(trainer, "_available_cpu_memory_bytes", lambda: available) + with pytest.raises(RuntimeError, match="checkpoint.*host memory"): + cp._admit_snapshot(trainer, "a") + assert not hooks + # Two eight-byte saved keys, capture/packing/serialization; no scratch buffer. + available = 48 + cp._admit_snapshot(trainer, "a") + assert not hooks + cp._custom_snapshot(trainer, "a", {}) + assert hooks == [True] diff --git a/tests/unit/test_checkpoint_snapshot_spill_distributed.py b/tests/unit/test_checkpoint_snapshot_spill_distributed.py index 96b3eea0f..f6e8780e4 100644 --- a/tests/unit/test_checkpoint_snapshot_spill_distributed.py +++ b/tests/unit/test_checkpoint_snapshot_spill_distributed.py @@ -2,19 +2,17 @@ from datetime import timedelta from pathlib import Path -import sys import threading import time -from types import SimpleNamespace import pytest import safetensors.torch -from test_trainer_rank_validation import _save_state_trainer +from test_checkpoint_snapshot_spill import _snapshot_trainer import torch import torch.distributed as dist import torch.multiprocessing as mp -from art.trainer_rank._impl import _CheckpointSlot, _CustomObject +from art.trainer_rank import _checkpoint as cp def _worker(rank, directory, failure, action, prepared, released, finalized): @@ -25,18 +23,6 @@ def _worker(rank, directory, failure, action, prepared, released, finalized): init_method=f"file://{directory}/gloo", timeout=timedelta(seconds=15), ) - trainer = _save_state_trainer() - parameter = torch.nn.Parameter(torch.tensor([1.0])) - trainer._checkpoint_slots["a"] = _CheckpointSlot( - params=(parameter,), - config={ - "base_model_name_or_path": "test/model", - "r": 1, - "lora_alpha": 1, - "target_modules": ["q_proj"], - }, - custom={"p": _CustomObject("parameter", parameter, object())}, - ) output = str(Path(directory) / "failed") original_write, original_start = safetensors.torch.save_file, threading.Thread.start @@ -55,17 +41,34 @@ def start(thread): try: with pytest.MonkeyPatch.context() as patch: - patch.setattr(trainer, "_slot_ref", lambda _: None) - patch.setitem( - sys.modules, - "art.megatron.lora", - SimpleNamespace(LoRA=type("UnusedLoRA", (), {})), - ) - patch.setitem( - sys.modules, - "art.megatron.weights.lora_publish", - SimpleNamespace(collect_local_lora_entries=lambda *a, **kw: ({}, [])), - ) + trainer = _snapshot_trainer(patch) + parameter = trainer._checkpoint_slots["a"].params[0] + if failure == "admission": + with pytest.MonkeyPatch.context() as admission: + admission.setattr( + trainer, + "_available_cpu_memory_bytes", + lambda: 0 if rank == 0 else 12, + ) + admission.setattr( + cp, + "_local_state", + lambda *_: pytest.fail( + "capture ran before collective admission" + ), + ) + admission.setattr( + "art.trainer_rank._heads.synchronize_head_buffers", + lambda *_: pytest.fail("buffer copies ran before admission"), + ) + with pytest.raises(RuntimeError, match="checkpoint.*host memory"): + trainer.prepare_checkpoint_save(output, "a") + assert not trainer._checkpoint_preparing_saves + assert not trainer._prepared_checkpoint_saves + assert trainer._checkpoint_snapshot_spill is None + assert not list(Path(directory).glob(".failed.*")) + assert trainer._checkpoint_save_sequence == 0 + failure = "write" patch.setattr(safetensors.torch, "save_file", write) patch.setattr(threading.Thread, "start", start) trainer.prepare_checkpoint_save(output, "a") @@ -105,8 +108,16 @@ def start(thread): dist.destroy_process_group() -@pytest.mark.parametrize("action", ["finish", "abort"]) -@pytest.mark.parametrize("failure", ["start", "write"]) +@pytest.mark.parametrize( + "failure,action", + [ + ("start", "finish"), + ("start", "abort"), + ("write", "finish"), + ("write", "abort"), + ("admission", "finish"), + ], +) def test_asymmetric_snapshot_failure_does_not_block_capture(tmp_path, action, failure): context = mp.get_context("spawn") prepared = [context.Event() for _ in range(2)] @@ -119,7 +130,7 @@ def test_asymmetric_snapshot_failure_does_not_block_capture(tmp_path, action, fa ) for rank in range(2) ] - deadline = time.monotonic() + 40 + deadline = time.monotonic() + 60 try: for process in processes: process.start() From 77658b152404f788f4973897958ca2af7248552d Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 06:02:50 +0000 Subject: [PATCH 113/150] Include persistent buffer gather in snapshot admission --- src/art/trainer_rank/_checkpoint.py | 18 +++++++++--- tests/unit/test_checkpoint_snapshot_spill.py | 31 ++++++++++++++++++++ 2 files changed, 45 insertions(+), 4 deletions(-) diff --git a/src/art/trainer_rank/_checkpoint.py b/src/art/trainer_rank/_checkpoint.py index 219e9063c..de4cb4b66 100644 --- a/src/art/trainer_rank/_checkpoint.py +++ b/src/art/trainer_rank/_checkpoint.py @@ -959,8 +959,8 @@ def _local_state( def _admit_snapshot(trainer: TrainerRank, name: str) -> None: """Estimate registered copies without executing user serialization hooks. - Unregistered allocations, payload-expanding hooks and concurrent external - allocations are outside this estimate. + Unregistered allocations, unusually large serialization metadata, + payload-expanding hooks and concurrent allocations are outside this estimate. """ from ._impl import _custom_named_parameters @@ -973,20 +973,24 @@ def _admit_snapshot(trainer: TrainerRank, name: str) -> None: tensors: list[torch.Tensor] = [ param for param in slot.params if id(param) not in custom_params ] + buffers: list[torch.Tensor] = [] for custom in slot.custom.values(): if custom.kind == "module": for _, child in cast(torch.nn.Module, custom.value).named_modules( remove_duplicate=False ): tensors.extend(p for p in child._parameters.values() if p is not None) - tensors.extend( + buffers.extend( value for key, value in child._buffers.items() if value is not None and key not in child._non_persistent_buffers_set ) else: - tensors.append(cast(torch.Tensor, custom.value)) + (buffers if custom.kind == "buffer" else tensors).append( + cast(torch.Tensor, custom.value) + ) + tensors.extend(buffers) cached = slot.custom_payload if cached is not None: tensors.extend(cached.tensors.values()) @@ -1016,6 +1020,12 @@ def _admit_snapshot(trainer: TrainerRank, name: str) -> None: with spill.lock: workspace = max(spill.workspace.values(), default=0) required = max(2 * size + workspace, 3 * size) + if buffers and _distributed(): + # Buffer sync clones logical contents before pickling: no backing views. + # Allow one page per tensor for ordinary pickle metadata, the padded + # all-gather output, input, and cloning/serialization/deserialization copies. + sync = sum(value.numel() * value.element_size() + 4096 for value in buffers) + required = max(required, (dist.get_world_size() + 6) * sync + workspace) available = trainer._available_cpu_memory_bytes() if required > available: raise RuntimeError( diff --git a/tests/unit/test_checkpoint_snapshot_spill.py b/tests/unit/test_checkpoint_snapshot_spill.py index 99f9cb912..11ed47e0c 100644 --- a/tests/unit/test_checkpoint_snapshot_spill.py +++ b/tests/unit/test_checkpoint_snapshot_spill.py @@ -3,6 +3,7 @@ from concurrent.futures import Future from dataclasses import replace from pathlib import Path +import pickle import sys import threading from types import SimpleNamespace @@ -554,3 +555,33 @@ def test_module_snapshot_admission_reads_metadata_without_hooks(monkeypatch): assert not hooks cp._custom_snapshot(trainer, "a", {}) assert hooks == [True] + + +@pytest.mark.parametrize("world", (2, 8)) +@pytest.mark.parametrize("kind", ("buffer", "module")) +def test_snapshot_admission_counts_buffer_gather(monkeypatch, world, kind): + from art.trainer_rank._heads import _plain + + trainer = _snapshot_trainer(monkeypatch) + buffer = torch.ones(64)[1:3] + value: torch.Tensor | torch.nn.Module = buffer + if kind == "module": + value = torch.nn.Module() + value.register_buffer("saved", buffer) + value.register_buffer("scratch", torch.ones(64), persistent=False) + trainer._checkpoint_slots["a"].custom["b"] = _CustomObject(kind, value, object()) + payload = _plain(buffer).cpu() + assert payload.untyped_storage().nbytes() == 8 # Not the 256-byte backing view. + assert 8 < len(pickle.dumps({("a", "b"): (0, {"saved": payload})})) <= 4104 + trainer._checkpoint_snapshot_spill = SimpleNamespace( + lock=threading.Lock(), workspace={Future(): 512} + ) + monkeypatch.setattr(cp, "_distributed", lambda: True) + monkeypatch.setattr(cp.dist, "get_world_size", lambda: world) + # Padded gather plus cloning/serialization/deserialization and an old writer. + available = (world + 6) * 4104 + 512 - 1 + monkeypatch.setattr(trainer, "_available_cpu_memory_bytes", lambda: available) + with pytest.raises(RuntimeError, match="checkpoint.*host memory"): + cp._admit_snapshot(trainer, "a") + available += 1 + cp._admit_snapshot(trainer, "a") From 80297168e3a1c696cda811726fbfea0a28737585 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 06:12:09 +0000 Subject: [PATCH 114/150] Narrow command test head types without changing test behavior --- tests/unit/test_trainer_rank_commands.py | 16 ++++++++++++---- 1 file changed, 12 insertions(+), 4 deletions(-) diff --git a/tests/unit/test_trainer_rank_commands.py b/tests/unit/test_trainer_rank_commands.py index bc1de179a..31cd686e3 100644 --- a/tests/unit/test_trainer_rank_commands.py +++ b/tests/unit/test_trainer_rank_commands.py @@ -9,7 +9,7 @@ import sys import threading from types import SimpleNamespace -from typing import Any +from typing import Any, cast from unittest.mock import patch import weakref @@ -300,7 +300,7 @@ def forward(ctx, value): return value @staticmethod - def backward(ctx, gradient): + def backward(ctx, gradient): # ty: ignore[invalid-method-override] if dist.get_rank() == 1: raise RuntimeError("intentional physical backward failure") return gradient @@ -513,7 +513,13 @@ def forward(self, value): view.backward(head(torch.tensor(3.0))) asyncio.run(run_rank_callback(native, head_callback, mode="zero")) - weight = native._checkpoint_slots["student"].custom["head"].value.weight + weight = cast( + torch.nn.Parameter, + cast( + torch.nn.Module, + native._checkpoint_slots["student"].custom["head"].value, + ).weight, + ) gradient = ( torch.zeros_like(weight) if weight.grad is None else weight.grad.clone() ) @@ -644,6 +650,8 @@ def test_logical_native_head_backward_commits_after_local_autograd(): trainer, _ = _trainer("student") class LocalHead(torch.nn.Module): + count: torch.Tensor + def __init__(self): super().__init__() self.weight = torch.nn.Parameter(torch.tensor(2.0)) @@ -658,7 +666,7 @@ def callback(view): view.backward(head(torch.tensor(3.0))) asyncio.run(run_rank_callback(trainer, callback, mode="zero")) - native = trainer._checkpoint_slots["student"].custom["head"].value + native = cast(LocalHead, trainer._checkpoint_slots["student"].custom["head"].value) torch.testing.assert_close(native.weight.grad, torch.tensor(12.0)) assert native.count.item() == 1 From 9d02f3bdc64aa37e195fd9e451a6cfc03b6fcddb Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 06:24:57 +0000 Subject: [PATCH 115/150] Release partial checkpoint copies from failed capture frames --- src/art/trainer_rank/_checkpoint.py | 10 +++ tests/unit/test_checkpoint_moment_capture.py | 73 ++++++++++++++++++++ 2 files changed, 83 insertions(+) diff --git a/src/art/trainer_rank/_checkpoint.py b/src/art/trainer_rank/_checkpoint.py index de4cb4b66..d7f0dd8bc 100644 --- a/src/art/trainer_rank/_checkpoint.py +++ b/src/art/trainer_rank/_checkpoint.py @@ -956,6 +956,9 @@ def _local_state( return tuple(records), optimizer, _custom_snapshot(trainer, name, files) +_CAPTURE_FRAME_CODES = (_local_state.__code__, _custom_snapshot.__code__) + + def _admit_snapshot(trainer: TrainerRank, name: str) -> None: """Estimate registered copies without executing user serialization hooks. @@ -1112,6 +1115,13 @@ def prepare_checkpoint_save( ) except BaseException as exc: error = exc + # Only our completed capture frames own these partial snapshots; + # preserve active callers and foreign copy/hook traceback locals. + capture_tb = exc.__traceback__ + while capture_tb is not None: + if capture_tb.tb_frame.f_code in _CAPTURE_FRAME_CODES: + capture_tb.tb_frame.clear() + capture_tb = capture_tb.tb_next try: raise_distributed(error, "prepare checkpoint", group) if any(value != optimizer for value in _gather(optimizer, group)): diff --git a/tests/unit/test_checkpoint_moment_capture.py b/tests/unit/test_checkpoint_moment_capture.py index 50b6e41ad..bd3e420a7 100644 --- a/tests/unit/test_checkpoint_moment_capture.py +++ b/tests/unit/test_checkpoint_moment_capture.py @@ -1,12 +1,15 @@ """Tiny CPU captures through the real collector; no CUDA performance claim.""" +import gc from importlib.util import find_spec +import traceback from typing import Any, cast import pytest import safetensors.torch from test_trainer_rank_validation import _save_state_trainer import torch +from torch.multiprocessing.reductions import StorageWeakRef from art.trainer_rank import _checkpoint as cp from art.trainer_rank._impl import _CheckpointSlot, _CustomObject, _DynamicOptimizer @@ -164,3 +167,73 @@ def owned_cpu_zero(value, *args, **kwargs): before = saved[f"exp_avg_sq/{key}"].clone() saved[f"exp_avg/{key}"].add_(100) torch.testing.assert_close(saved[f"exp_avg_sq/{key}"], before, rtol=0, atol=0) + + +@pytest.mark.parametrize("captured_state", ("custom", "dense"), indirect=True) +def test_failed_capture_releases_partial_copies(captured_state, tmp_path, monkeypatch): + trainer, params, masters, _, _, expected, _ = captured_state + trainer._checkpoint_slots["a"].config = { + "base_model_name_or_path": "test/model", + "r": 2, + "lora_alpha": 2, + "target_modules": ["q_proj"], + } + borrowed = [ + StorageWeakRef(value.untyped_storage()) for value in (*params, *masters) + ] + sentinel = torch.tensor([13.0]) + copies = [] + calls = 0 + error, cause = OSError("CPU capture copy failed"), RuntimeError("copy cause") + original = torch.Tensor.to + + def copy(value, *args, **kwargs): + nonlocal calls + foreign_marker = sentinel + if kwargs.get("copy"): + calls += 1 + if calls == 2: + try: + raise cause + except RuntimeError: + raise error from cause + result = original(value, *args, **kwargs) + if kwargs.get("copy"): + copies.append(StorageWeakRef(result.untyped_storage())) + assert foreign_marker is sentinel + return result + + output = str(tmp_path / "reusable") + enabled = gc.isenabled() + gc.disable() + try: + with monkeypatch.context() as patch: + patch.setattr(torch.Tensor, "to", copy) + with pytest.raises(OSError) as caught: + trainer.prepare_checkpoint_save(output, "a") + assert caught.value is error and error.__cause__ is cause + for failure in (error, cause): + frame = next( + frame + for frame, _ in traceback.walk_tb(failure.__traceback__) + if frame.f_code is copy.__code__ + ) + assert frame.f_locals["foreign_marker"] is sentinel + assert calls == 2 and len(copies) == 1 and copies[0].expired() + assert all(not storage.expired() for storage in borrowed) + assert not trainer._checkpoint_preparing_saves + assert not trainer._prepared_checkpoint_saves + assert not list(tmp_path.glob(".reusable.*")) + trainer.prepare_checkpoint_save(output, "a") + prepared = trainer._prepared_checkpoint_saves[output] + assert prepared.writer is not None + prepared.writer.result(3) + for filename, tensors in expected.items(): + actual = safetensors.torch.load_file(prepared.snapshot / filename) + for key, reference in tensors.items(): + torch.testing.assert_close(actual[key], reference, rtol=0, atol=0) + finally: + if enabled: + gc.enable() + for pending in list(trainer._prepared_checkpoint_saves): + trainer.abort_checkpoint_save(pending) From 5c2a530a7ccf643d5408722774d0a0cfb42149a0 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 06:29:03 +0000 Subject: [PATCH 116/150] Match checkpoint capture frames by exact code identity --- src/art/trainer_rank/_checkpoint.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/src/art/trainer_rank/_checkpoint.py b/src/art/trainer_rank/_checkpoint.py index d7f0dd8bc..977db170b 100644 --- a/src/art/trainer_rank/_checkpoint.py +++ b/src/art/trainer_rank/_checkpoint.py @@ -1119,7 +1119,8 @@ def prepare_checkpoint_save( # preserve active callers and foreign copy/hook traceback locals. capture_tb = exc.__traceback__ while capture_tb is not None: - if capture_tb.tb_frame.f_code in _CAPTURE_FRAME_CODES: + frame_code = capture_tb.tb_frame.f_code + if any(frame_code is code for code in _CAPTURE_FRAME_CODES): capture_tb.tb_frame.clear() capture_tb = capture_tb.tb_next try: From 2e1281acd033875206d0270828ccc565017d3cb3 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 07:02:20 +0000 Subject: [PATCH 117/150] Reuse adapter config factory in checkpoint tests --- tests/unit/test_checkpoint_moment_capture.py | 11 ++++----- tests/unit/test_checkpoint_snapshot_spill.py | 20 ++++++---------- tests/unit/test_trainer_rank_validation.py | 24 +++++++++----------- 3 files changed, 22 insertions(+), 33 deletions(-) diff --git a/tests/unit/test_checkpoint_moment_capture.py b/tests/unit/test_checkpoint_moment_capture.py index 1020c80ff..e640a5ac7 100644 --- a/tests/unit/test_checkpoint_moment_capture.py +++ b/tests/unit/test_checkpoint_moment_capture.py @@ -7,7 +7,7 @@ import pytest import safetensors.torch -from test_trainer_rank_validation import _save_state_trainer +from test_trainer_rank_validation import _adapter_config, _save_state_trainer import torch from torch.multiprocessing.reductions import StorageWeakRef @@ -181,12 +181,9 @@ def pack_owned(value, *args, **kwargs): @pytest.mark.parametrize("captured_state", ("custom", "dense"), indirect=True) def test_failed_capture_releases_partial_copies(captured_state, tmp_path, monkeypatch): trainer, params, masters, _, _, expected, _ = captured_state - trainer._checkpoint_slots["a"].config = { - "base_model_name_or_path": "test/model", - "r": 2, - "lora_alpha": 2, - "target_modules": ["q_proj"], - } + trainer._checkpoint_slots["a"].config = _adapter_config( + rank=2, alpha=2, target_modules=("q_proj",) + ) borrowed = [ StorageWeakRef(value.untyped_storage()) for value in (*params, *masters) ] diff --git a/tests/unit/test_checkpoint_snapshot_spill.py b/tests/unit/test_checkpoint_snapshot_spill.py index 11ed47e0c..c1904162d 100644 --- a/tests/unit/test_checkpoint_snapshot_spill.py +++ b/tests/unit/test_checkpoint_snapshot_spill.py @@ -11,7 +11,11 @@ import pytest import safetensors.torch -from test_trainer_rank_validation import _prepared_save, _save_state_trainer +from test_trainer_rank_validation import ( + _adapter_config, + _prepared_save, + _save_state_trainer, +) import torch from art.trainer_rank import _checkpoint as cp @@ -24,12 +28,7 @@ def _snapshot_trainer(monkeypatch, parameter=None): parameter = torch.nn.Parameter(torch.tensor([1.0])) trainer._checkpoint_slots["a"] = _CheckpointSlot( params=(parameter,), - config={ - "base_model_name_or_path": "test/model", - "r": 1, - "lora_alpha": 1, - "target_modules": ["q_proj"], - }, + config=_adapter_config(target_modules=("q_proj",)), custom={"p": _CustomObject("parameter", parameter, object())}, ) monkeypatch.setattr(trainer, "_slot_ref", lambda _: None) @@ -252,12 +251,7 @@ def test_start_failure_is_owned_until_collective_finalization( ): trainer = _save_state_trainer() trainer._checkpoint_slots["a"] = _CheckpointSlot( - config={ - "base_model_name_or_path": "test/model", - "r": 1, - "lora_alpha": 1, - "target_modules": ["q_proj"], - } + config=_adapter_config(target_modules=("q_proj",)) ) monkeypatch.setattr(cp, "_validate_save_state", lambda *_: {}) refs = [] diff --git a/tests/unit/test_trainer_rank_validation.py b/tests/unit/test_trainer_rank_validation.py index fa447dd75..e6a6d17a3 100644 --- a/tests/unit/test_trainer_rank_validation.py +++ b/tests/unit/test_trainer_rank_validation.py @@ -1428,12 +1428,18 @@ def load( snapshot_prepared_checkpoint(trainer, source, "loaded") -def _adapter_config(model: str = "test/model") -> _AdapterConfig: +def _adapter_config( + model: str = "test/model", + *, + rank: int = 1, + alpha: float = 1, + target_modules: tuple[str, ...] = (), +) -> _AdapterConfig: return { "base_model_name_or_path": model, - "r": 1, - "lora_alpha": 1, - "target_modules": [], + "r": rank, + "lora_alpha": alpha, + "target_modules": list(target_modules), } @@ -2446,15 +2452,7 @@ def test_real_checkpoint_codec_round_trips_with_optional_optimizer( "get_data_parallel_rank", lambda **_kwargs: 0, ) - config = cast( - Any, - { - "base_model_name_or_path": "test/model", - "r": 2, - "lora_alpha": 2, - "target_modules": ["q_proj"], - }, - ) + config = _adapter_config(rank=2, alpha=2, target_modules=("q_proj",)) adapter = { "layer.q_proj.lora_A.weight": torch.tensor([[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]]), "layer.q_proj.lora_B.weight": torch.tensor( From 86d951d98c0848bf9b280ea4db1c89ac2dcd4750 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Wed, 30 Sep 2026 19:15:25 +0000 Subject: [PATCH 118/150] test: consolidate live-head coverage and reuse snapshot config Exercise the existing callback API in the reentrant-head rejection test and parameterize native/client buffer publication across CPU and CUDA. Preserve the error, gradient, publication, and read-only-view assertions. Reuse the existing adapter-config helper in the forward-only snapshot test. Validation: 30 distinct CPU cases pass across existing light/full environments; both affected CUDA cases are collected and skipped with CUDA hidden. Ruff and format checks pass. Scoped ty has no errors and the same 18 warnings as b574. Remove 43 net lines and 25 gross PR lines relative to main 229490db. --- tests/unit/test_trainer_rank_live_heads.py | 89 +++++++--------------- tests/unit/test_trainer_rank_validation.py | 10 +-- 2 files changed, 28 insertions(+), 71 deletions(-) diff --git a/tests/unit/test_trainer_rank_live_heads.py b/tests/unit/test_trainer_rank_live_heads.py index a0f4d1b06..39a47034a 100644 --- a/tests/unit/test_trainer_rank_live_heads.py +++ b/tests/unit/test_trainer_rank_live_heads.py @@ -3,7 +3,6 @@ import asyncio from copy import deepcopy from datetime import timedelta -from types import SimpleNamespace import pytest from test_trainer_rank_custom_tensors import _trainer, _use_local_gradients @@ -25,7 +24,6 @@ execute_head_operation, export_head, head_gradient_targets, - logical_register_head, ) from art.trainer_rank._options import ForwardOptions from art.trainer_rank._tensors import CotangentCollector, detach_tree @@ -1084,38 +1082,15 @@ def test_inplace_operation_snapshots_readonly_client_tensor(kind): def test_logical_callback_reentrant_head_rejects_before_gradient_publication(): trainer, _ = _trainer("student") - collector = CotangentCollector() - - def invoke(operation, kind, payload): - assert operation == "head" - return execute_head_operation(trainer, kind, payload) - - view = SimpleNamespace( - _rank=trainer, - _invoke=invoke, - _executor=SimpleNamespace( - state=SimpleNamespace(collector=collector), invoke=invoke - ), - device=torch.device("cpu"), - ) def callback(rank): - head = logical_register_head( - rank, "module", "head", lambda: TiedHead(True), checkpoint="student" - ) + head = rank.module("head", lambda: TiedHead(True), checkpoint="student") loss = head(torch.tensor(3.0, requires_grad=True)) # The logical executor submits packets only after local collection succeeds. - packets = collector.backward(loss) - trainer._commit_versioned_gradients( - [ - target - for packet in packets - for target in head_gradient_targets(trainer, packet) - ] - ) + rank.backward(loss) with pytest.raises(RuntimeError, match="nested remote backward is unsupported"): - callback(view) + asyncio.run(run_rank_callback(trainer, callback)) custom = trainer._checkpoint_slots["student"].custom["head"].value assert isinstance(custom, torch.nn.Module) assert all(parameter.grad is None for parameter in custom.parameters()) @@ -1213,20 +1188,39 @@ def test_live_buffer_view_mutation_rejects_without_silent_write( @pytest.mark.parametrize("client", (False, True)) -def test_functional_batchnorm_no_grad_publishes_buffer_changes(client): - trainer, native = _native_head("buffer", "mean", lambda: torch.zeros(2)) - live = _live_head(trainer, "mean", torch.zeros(2)) if client else None +@pytest.mark.parametrize("device", ("cpu", "cuda")) +def test_functional_batchnorm_no_grad_publishes_buffer_changes(client, device): + if device == "cuda" and not torch.cuda.is_available(): + pytest.skip("requires CUDA") + trainer, rank = _trainer("student") + if device == "cuda": + trainer.device = torch.device("cuda", 0) + native = rank.buffer("mean", lambda: torch.zeros(2), checkpoint="student") + live = ( + _live_head(trainer, "mean", torch.zeros(2, device=device)) if client else None + ) mean = native if live is None else _tensor(live) before = export_head(trainer, "student", "mean").buffer_revision + assert before == 0 + if device == "cuda": + with pytest.raises( + RuntimeError, match="Views of live checkpoint buffers are read-only" + ): + mean[:].fill_(7) with torch.no_grad(): torch.nn.functional.batch_norm( - torch.ones(4, 2), mean, torch.ones(2), training=True + torch.ones(4, 2, device=device), + mean, + torch.ones(2, device=device), + training=True, ) - torch.testing.assert_close(mean, torch.full((2,), 0.1)) + expected = torch.full((2,), 0.1, device=device) + torch.testing.assert_close(mean, expected) if live is not None: update = live.take_publication() assert update is not None execute_head_operation(trainer, "head_publish", (update,)) + torch.testing.assert_close(native, expected) assert export_head(trainer, "student", "mean").buffer_revision == before + 1 @@ -1359,35 +1353,6 @@ def test_stateful_function_on_buffer_snapshot_view_rejects(client): assert live.take_publication() is None -@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") -@pytest.mark.parametrize("client", (False, True)) -def test_cuda_live_buffer_views_and_functional_publication(client): - trainer, rank = _trainer("student") - trainer.device = torch.device("cuda", 0) - native = rank.buffer("mean", lambda: torch.zeros(2), checkpoint="student") - live = ( - _live_head(trainer, "mean", torch.zeros(2, device="cuda")) if client else None - ) - mean = native if live is None else _tensor(live) - with pytest.raises( - RuntimeError, match="Views of live checkpoint buffers are read-only" - ): - mean[:].fill_(7) - with torch.no_grad(): - torch.nn.functional.batch_norm( - torch.ones(4, 2, device="cuda"), - mean, - torch.ones(2, device="cuda"), - training=True, - ) - if live is not None: - update = live.take_publication() - assert update is not None - execute_head_operation(trainer, "head_publish", (update,)) - torch.testing.assert_close(native, torch.full((2,), 0.1, device="cuda")) - assert export_head(trainer, "student", "mean").buffer_revision == 1 - - @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") def test_cuda_buffer_sync_stages_cpu_authority_before_comparison(tmp_path): from art.trainer_rank._heads import synchronize_head_buffers diff --git a/tests/unit/test_trainer_rank_validation.py b/tests/unit/test_trainer_rank_validation.py index befc133be..dcffe5413 100644 --- a/tests/unit/test_trainer_rank_validation.py +++ b/tests/unit/test_trainer_rank_validation.py @@ -1356,15 +1356,7 @@ def test_forward_snapshot_is_independent_and_forward_only( ) trainer._checkpoint_slots["student"] = _CheckpointSlot( tuple(lora.lora_slot_params(source)), - cast( - Any, - { - "base_model_name_or_path": "test/model", - "r": 2, - "lora_alpha": 2, - "target_modules": ["q_proj"], - }, - ), + _adapter_config(rank=2, alpha=2, target_modules=("q_proj",)), ) assert trainer.snapshot_checkpoint("student", "saved") From 74ab4988c5be00060ca53459ceb1620331d32f49 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Wed, 30 Sep 2026 19:28:47 +0000 Subject: [PATCH 119/150] test: check reentrant head state before callback cleanup Catch the expected backward error and check native gradients and pending publication inside the callback, before run_rank_callback flushes heads. This preserves the original rejection assertions with the real runner. --- tests/unit/test_trainer_rank_live_heads.py | 24 +++++++++++----------- 1 file changed, 12 insertions(+), 12 deletions(-) diff --git a/tests/unit/test_trainer_rank_live_heads.py b/tests/unit/test_trainer_rank_live_heads.py index 39a47034a..d4cc2198b 100644 --- a/tests/unit/test_trainer_rank_live_heads.py +++ b/tests/unit/test_trainer_rank_live_heads.py @@ -1087,19 +1087,19 @@ def callback(rank): head = rank.module("head", lambda: TiedHead(True), checkpoint="student") loss = head(torch.tensor(3.0, requires_grad=True)) # The logical executor submits packets only after local collection succeeds. - rank.backward(loss) + with pytest.raises(RuntimeError, match="nested remote backward is unsupported"): + rank.backward(loss) + custom = trainer._checkpoint_slots["student"].custom["head"].value + assert isinstance(custom, torch.nn.Module) + assert all(parameter.grad is None for parameter in custom.parameters()) + assert ( + getattr(trainer, "_logical_head_handles")[ + ("student", "head") + ].take_publication() + is None + ) - with pytest.raises(RuntimeError, match="nested remote backward is unsupported"): - asyncio.run(run_rank_callback(trainer, callback)) - custom = trainer._checkpoint_slots["student"].custom["head"].value - assert isinstance(custom, torch.nn.Module) - assert all(parameter.grad is None for parameter in custom.parameters()) - assert ( - getattr(trainer, "_logical_head_handles")[ - ("student", "head") - ].take_publication() - is None - ) + asyncio.run(run_rank_callback(trainer, callback)) @pytest.mark.parametrize("mode", ("rank", "zero")) From 6e7aa16c660ec78ed077b9c75ad20508417941e8 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Wed, 30 Sep 2026 19:46:17 +0000 Subject: [PATCH 120/150] test: reuse process group lifetime in RNG workers Use the shared process_group helper with the existing backend, physical rank, 2 * dp_size world size and 90-second timeout. Keep subgroup setup, RNG and gradient assertions, and cleanup ordering unchanged. --- tests/unit/test_trainer_rank_rng.py | 18 +++++------------- 1 file changed, 5 insertions(+), 13 deletions(-) diff --git a/tests/unit/test_trainer_rank_rng.py b/tests/unit/test_trainer_rank_rng.py index 161144c74..a3b2313fb 100644 --- a/tests/unit/test_trainer_rank_rng.py +++ b/tests/unit/test_trainer_rank_rng.py @@ -2,7 +2,6 @@ import asyncio from contextlib import nullcontext -from datetime import timedelta from types import SimpleNamespace from typing import Any, cast @@ -11,7 +10,7 @@ import torch.distributed as dist import torch.multiprocessing as mp from torch.utils.checkpoint import checkpoint -from trainer_rank_test_support import gloo_group, megatron_topology, spawn_and_join +from trainer_rank_test_support import megatron_topology, process_group, spawn_and_join from art.trainer_rank import ( AdamParams, @@ -223,7 +222,7 @@ def test_failed_forward_keeps_command_collectives_aligned(tmp_path): def _failed_forward_worker(physical, rendezvous): with ( - gloo_group(physical, rendezvous, timeout=10), + process_group(physical, rendezvous, timeout=10), megatron_topology(physical, dp_size=1, tp_size=2), ): for primary_type, sync_failure, preflight in ( @@ -416,14 +415,9 @@ def _distributed_worker(rank, dp_size, parallelism, backend, init_method): device = torch.device("cpu" if backend == "gloo" else f"cuda:{rank}") if device.type == "cuda": torch.cuda.set_device(device) - dist.init_process_group( - backend, - init_method=init_method, - rank=rank, - world_size=2 * dp_size, - timeout=timedelta(seconds=90), - ) - try: + with process_group( + rank, init_method, world_size=2 * dp_size, timeout=90, backend=backend + ): replica_groups = [dist.new_group([2 * dp, 2 * dp + 1]) for dp in range(dp_size)] dp_groups = [ dist.new_group(list(range(replica, 2 * dp_size, 2))) for replica in range(2) @@ -467,8 +461,6 @@ def _distributed_worker(rank, dp_size, parallelism, backend, init_method): dp_group, parallelism, ) - finally: - dist.destroy_process_group() def _gradient_oracle( From 86505c98fb5091ee94c9744d72dd520441628114 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Wed, 30 Sep 2026 19:51:27 +0000 Subject: [PATCH 121/150] test: assert rejected callbacks do not publish head buffers --- tests/unit/test_trainer_rank_live_heads.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/unit/test_trainer_rank_live_heads.py b/tests/unit/test_trainer_rank_live_heads.py index d4cc2198b..0b1d74821 100644 --- a/tests/unit/test_trainer_rank_live_heads.py +++ b/tests/unit/test_trainer_rank_live_heads.py @@ -1098,6 +1098,7 @@ def callback(rank): ].take_publication() is None ) + assert export_head(trainer, "student", "head").buffer_revision == 0 asyncio.run(run_rank_callback(trainer, callback)) From 6ad8688b818e409e224068cf2c2de6ba571b052b Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Wed, 30 Sep 2026 19:54:01 +0000 Subject: [PATCH 122/150] test: keep routed graph checks independent of Megatron --- tests/unit/test_trainer_rank_routed_experts.py | 10 +++++++--- 1 file changed, 7 insertions(+), 3 deletions(-) diff --git a/tests/unit/test_trainer_rank_routed_experts.py b/tests/unit/test_trainer_rank_routed_experts.py index fa955a73e..6b99a474c 100644 --- a/tests/unit/test_trainer_rank_routed_experts.py +++ b/tests/unit/test_trainer_rank_routed_experts.py @@ -1,3 +1,4 @@ +from contextlib import nullcontext from dataclasses import replace from enum import Enum import sys @@ -77,13 +78,16 @@ def test_forward_routes_align_with_inputs_without_shift_and_group_separately(): def test_cached_forward_and_replay_preserve_owned_routes(monkeypatch, retention): from trainer_rank_test_support import checkpoint_runtime - from art.megatron.context_parallel.types import ParallelTopology - model = torch.nn.Linear(1, 1, bias=False) rank = TrainerRank(checkpoint_runtime(model)) binding, seen = router(), [] rank._routing_bindings = [(binding, 0)] - monkeypatch.setattr(rank, "_topology", lambda: ParallelTopology()) + monkeypatch.setattr(rank, "_resolve_slot_ref", lambda *_a, **_kw: None) + lora = SimpleNamespace(use_lora_slot=lambda _slot, **_kwargs: nullcontext()) + monkeypatch.setitem(sys.modules, "art.megatron.lora", lora) + monkeypatch.setattr( + rank, "_topology", lambda: SimpleNamespace(dp=1, tp=1, cp=1, pp=1) + ) monkeypatch.setattr(rank, "_dp_rank_and_size", lambda: (0, 1)) monkeypatch.setattr(rank, "_configure_hybridep", lambda *a, **kw: None) monkeypatch.setattr( From 87aa20da9815ff6f6fd27a1a8c91b6112677350e Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Wed, 30 Sep 2026 20:10:42 +0000 Subject: [PATCH 123/150] test: preserve pending head buffers until callback cleanup --- tests/unit/test_trainer_rank_live_heads.py | 10 ++++------ 1 file changed, 4 insertions(+), 6 deletions(-) diff --git a/tests/unit/test_trainer_rank_live_heads.py b/tests/unit/test_trainer_rank_live_heads.py index 0b1d74821..c0fde30d4 100644 --- a/tests/unit/test_trainer_rank_live_heads.py +++ b/tests/unit/test_trainer_rank_live_heads.py @@ -1085,6 +1085,7 @@ def test_logical_callback_reentrant_head_rejects_before_gradient_publication(): def callback(rank): head = rank.module("head", lambda: TiedHead(True), checkpoint="student") + head.offset.add_(1) loss = head(torch.tensor(3.0, requires_grad=True)) # The logical executor submits packets only after local collection succeeds. with pytest.raises(RuntimeError, match="nested remote backward is unsupported"): @@ -1092,15 +1093,12 @@ def callback(rank): custom = trainer._checkpoint_slots["student"].custom["head"].value assert isinstance(custom, torch.nn.Module) assert all(parameter.grad is None for parameter in custom.parameters()) - assert ( - getattr(trainer, "_logical_head_handles")[ - ("student", "head") - ].take_publication() - is None - ) assert export_head(trainer, "student", "head").buffer_revision == 0 asyncio.run(run_rank_callback(trainer, callback)) + state = export_head(trainer, "student", "head") + assert state.buffer_revision == 1 + assert state.buffers["offset"].item() == 2 @pytest.mark.parametrize("mode", ("rank", "zero")) From 7512bcefd2cc5200956567e20bfcb1a25f96c1b3 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Wed, 30 Sep 2026 20:26:12 +0000 Subject: [PATCH 124/150] refactor: reuse resolved trainer forward policy values --- src/art/trainer_rank/_heads.py | 6 +----- src/art/trainer_rank/_impl.py | 11 +---------- 2 files changed, 2 insertions(+), 15 deletions(-) diff --git a/src/art/trainer_rank/_heads.py b/src/art/trainer_rank/_heads.py index af3eea98c..65fa84661 100644 --- a/src/art/trainer_rank/_heads.py +++ b/src/art/trainer_rank/_heads.py @@ -1295,8 +1295,6 @@ def logical_register_head( *, checkpoint: Any = ..., ) -> Any: - from ._options import resolve_forward_options - if checkpoint is ...: from ._impl import Unset @@ -1345,9 +1343,7 @@ def logical_register_head( if isinstance(value, torch.nn.Module) else value.to(view.device) ) - maximum = resolve_forward_options( - getattr(view._rank, "_forward_options", None) - ).max_gradient_staleness + maximum = head_staleness(view._rank) head = LiveHead( state, value, view._executor.state.collector, max_gradient_staleness=maximum ) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index fcd906d51..45a153b6b 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -2877,22 +2877,13 @@ def _capture_forward_options( # lazy. Own the submitted tensor storage before returning an iterator. materialized = _snapshot(_materialize(inputs)) constructor = getattr(self, "_forward_options", None) - from dataclasses import fields def capture(value: ForwardInputs) -> ForwardInputs: if isinstance(value, ForwardInput): if constructor is None and options is None: return replace(value) resolved = resolve_forward_options(constructor, options, value.options) - return replace( - value, - options=ForwardOptions( - **{ - field.name: getattr(resolved, field.name) - for field in fields(resolved) - } - ), - ) + return replace(value, options=ForwardOptions(**vars(resolved))) return _rebuild_forward_tree(value, [capture(child) for child in value]) return capture(materialized) From 4a008d4b9367694556d3260b2cd1a367e6353ff3 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Wed, 30 Sep 2026 21:07:58 +0000 Subject: [PATCH 125/150] test: share native LoRA checkpoint setup --- .../megatron/lora/test_dynamic_lora_slots.py | 13 +++++++ .../megatron/lora/test_lora_versions.py | 39 ++++--------------- .../lora/test_trainer_v1_graph_cache.py | 8 +--- .../megatron/lora/test_trainer_v1_versions.py | 16 ++------ 4 files changed, 25 insertions(+), 51 deletions(-) diff --git a/tests/integration/megatron/lora/test_dynamic_lora_slots.py b/tests/integration/megatron/lora/test_dynamic_lora_slots.py index d902b6d2a..a66d58530 100644 --- a/tests/integration/megatron/lora/test_dynamic_lora_slots.py +++ b/tests/integration/megatron/lora/test_dynamic_lora_slots.py @@ -637,6 +637,19 @@ def canonicalize_loaded_lora_state( return state +@contextmanager +def _lora_checkpoint(seed=1, *, rng_seed=None): + """Construct the shared single-rank dense checkpoint used by version tests.""" + with _single_rank_model_parallel(): + if rng_seed is not None: + torch.manual_seed(rng_seed) + device = torch.device("cuda") + lora = LoRA("dense", 4, 5, 2, 32, torch.float32, device) + trainer = _trainer_for(lora, device) + _install_checkpoint(trainer, "A", _adapter("dense", rank=2, seed=seed)) + yield device, lora, trainer + + @contextmanager def _single_rank_model_parallel(): os.environ.setdefault("MASTER_ADDR", "127.0.0.1") diff --git a/tests/integration/megatron/lora/test_lora_versions.py b/tests/integration/megatron/lora/test_lora_versions.py index 07478d7b5..b2907edc6 100644 --- a/tests/integration/megatron/lora/test_lora_versions.py +++ b/tests/integration/megatron/lora/test_lora_versions.py @@ -8,19 +8,14 @@ torch = pytest.importorskip("torch") pytest.importorskip("megatron.core") -from art.megatron.lora import LoRA, LoRASlotRef, use_lora_slot # noqa: E402 +from art.megatron.lora import LoRASlotRef, use_lora_slot # noqa: E402 from art.trainer_rank import TrainerRankSlotStateError # noqa: E402 from art.trainer_rank._checkpoint import ( # noqa: E402 discard_snapshot_checkpoint, snapshot_checkpoint, ) -from .test_dynamic_lora_slots import ( # noqa: E402 - _adapter, - _install_checkpoint, - _single_rank_model_parallel, - _trainer_for, -) +from .test_dynamic_lora_slots import _lora_checkpoint # noqa: E402 @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") @@ -28,11 +23,7 @@ def test_native_capture_preserves_independent_parameter_trainability( train_a: bool, ) -> None: - with _single_rank_model_parallel(): - device = torch.device("cuda") - lora = LoRA("dense", 4, 5, 2, 32, torch.float32, device) - trainer = _trainer_for(lora, device) - _install_checkpoint(trainer, "A", _adapter("dense", rank=2, seed=1)) + with _lora_checkpoint() as (device, lora, trainer): ref = LoRASlotRef("checkpoint", "A") current = lora._slot(ref) assert current is not None @@ -68,11 +59,7 @@ def test_native_capture_preserves_independent_parameter_trainability( @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") def test_native_version_storage_reuse_accounting_and_checkpoint_lifetime() -> None: - with _single_rank_model_parallel(): - device = torch.device("cuda") - lora = LoRA("dense", 4, 5, 2, 32, torch.float32, device) - trainer = _trainer_for(lora, device) - _install_checkpoint(trainer, "A", _adapter("dense", rank=2, seed=1)) + with _lora_checkpoint() as (device, lora, trainer): ref = LoRASlotRef("checkpoint", "A") expected_bytes = (4 * 2 + 2 * 5) * 4 custom = torch.nn.Parameter(torch.ones(100, device=device)) @@ -117,11 +104,7 @@ def test_native_version_storage_reuse_accounting_and_checkpoint_lifetime() -> No @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") def test_native_capture_rejects_staleness_and_checkpoint_replacement() -> None: - with _single_rank_model_parallel(): - device = torch.device("cuda") - lora = LoRA("dense", 4, 5, 2, 32, torch.float32, device) - trainer = _trainer_for(lora, device) - _install_checkpoint(trainer, "A", _adapter("dense", rank=2, seed=1)) + with _lora_checkpoint() as (device, lora, trainer): ref = LoRASlotRef("checkpoint", "A") capture = trainer._capture_lora_version(ref, max_gradient_staleness=0) with use_lora_slot(ref, version=capture): @@ -141,11 +124,7 @@ def test_native_capture_rejects_staleness_and_checkpoint_replacement() -> None: @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") def test_newer_weight_replay_keeps_original_gradient_age() -> None: - with _single_rank_model_parallel(): - device = torch.device("cuda") - lora = LoRA("dense", 4, 5, 2, 32, torch.float32, device) - trainer = _trainer_for(lora, device) - _install_checkpoint(trainer, "A", _adapter("dense", rank=2, seed=1)) + with _lora_checkpoint() as (device, lora, trainer): ref = LoRASlotRef("checkpoint", "A") origin = trainer._capture_checkpoint_version("A") trainer._checkpoint_slots["A"].revision = origin.revision + 2 @@ -173,11 +152,7 @@ def test_newer_weight_replay_keeps_original_gradient_age() -> None: @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") def test_discarded_snapshot_name_cannot_reuse_old_capture() -> None: - with _single_rank_model_parallel(): - device = torch.device("cuda") - lora = LoRA("dense", 4, 5, 2, 32, torch.float32, device) - trainer = _trainer_for(lora, device) - _install_checkpoint(trainer, "A", _adapter("dense", rank=2, seed=1)) + with _lora_checkpoint() as (device, lora, trainer): trainer._checkpoint_slots["A"].config = { "base_model_name_or_path": "test/model", "r": 2, diff --git a/tests/integration/megatron/lora/test_trainer_v1_graph_cache.py b/tests/integration/megatron/lora/test_trainer_v1_graph_cache.py index fe419ea69..12da62285 100644 --- a/tests/integration/megatron/lora/test_trainer_v1_graph_cache.py +++ b/tests/integration/megatron/lora/test_trainer_v1_graph_cache.py @@ -22,6 +22,7 @@ from .test_dynamic_lora_slots import ( # noqa: E402 _adapter, _install_checkpoint, + _lora_checkpoint, _single_rank_model_parallel, _trainer_for, ) @@ -141,12 +142,7 @@ def test_native_stale_logprob_correction_keeps_original_gradient_age(mode, monke TrainerRankSlotStateError, ) - with _single_rank_model_parallel(): - torch.manual_seed(19) - device = torch.device("cuda") - lora = LoRA("dense", 4, 5, 2, 32, torch.float32, device) - trainer = _trainer_for(lora, device) - _install_checkpoint(trainer, "A", _adapter("dense", rank=2, seed=91)) + with _lora_checkpoint(seed=91, rng_seed=19) as (device, lora, trainer): ref = LoRASlotRef("checkpoint", "A") origin = trainer._capture_checkpoint_version("A") parameters = trainer._checkpoint_slots["A"].params diff --git a/tests/integration/megatron/lora/test_trainer_v1_versions.py b/tests/integration/megatron/lora/test_trainer_v1_versions.py index d80d82664..9e54fbb0d 100644 --- a/tests/integration/megatron/lora/test_trainer_v1_versions.py +++ b/tests/integration/megatron/lora/test_trainer_v1_versions.py @@ -14,15 +14,10 @@ torch = pytest.importorskip("torch") pytest.importorskip("megatron.core") -from art.megatron.lora import LoRA, LoRASlotRef, use_lora_slot # noqa: E402 +from art.megatron.lora import LoRASlotRef, use_lora_slot # noqa: E402 from art.trainer_rank import AdamParams # noqa: E402 -from .test_dynamic_lora_slots import ( # noqa: E402 - _adapter, - _install_checkpoint, - _single_rank_model_parallel, - _trainer_for, -) +from .test_dynamic_lora_slots import _lora_checkpoint # noqa: E402 def _coupled_loss(first, second): @@ -49,12 +44,7 @@ def _adam(parameters, gradients, moments, step, params): @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") @pytest.mark.parametrize("recompute", ["none", "torch", "reentrant", "megatron"]) def test_native_old_lora_graph_matches_matrix_and_adam_oracle(recompute, artifact_dir): - with _single_rank_model_parallel(): - torch.manual_seed(1709) - device = torch.device("cuda") - lora = LoRA("dense", 4, 5, 2, 32, torch.float32, device) - trainer = _trainer_for(lora, device) - _install_checkpoint(trainer, "A", _adapter("dense", rank=2, seed=91)) + with _lora_checkpoint(seed=91, rng_seed=1709) as (device, lora, trainer): ref = LoRASlotRef("checkpoint", "A") originals = tuple(trainer._checkpoint_slots["A"].params) reference = [p.detach().double().requires_grad_() for p in originals] From d3987147a8a182890aff617c9342688c3e7abd5d Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Wed, 30 Sep 2026 21:22:41 +0000 Subject: [PATCH 126/150] test: preserve types in shared LoRA setup --- .../megatron/lora/test_dynamic_lora_slots.py | 14 ++++++-------- 1 file changed, 6 insertions(+), 8 deletions(-) diff --git a/tests/integration/megatron/lora/test_dynamic_lora_slots.py b/tests/integration/megatron/lora/test_dynamic_lora_slots.py index a66d58530..987e86c47 100644 --- a/tests/integration/megatron/lora/test_dynamic_lora_slots.py +++ b/tests/integration/megatron/lora/test_dynamic_lora_slots.py @@ -1,6 +1,6 @@ from __future__ import annotations -from collections.abc import Sequence +from collections.abc import Iterator, Sequence from contextlib import contextmanager import os from pathlib import Path @@ -122,11 +122,7 @@ def test_dynamic_lora_slots_capture_recompute_context_and_step_independently() - @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required.") def test_trainer_rank_custom_objects_train_and_become_stale_on_cuda() -> None: - with _single_rank_model_parallel(): - device = torch.device("cuda") - lora = LoRA("dense", 4, 5, 2, 32, torch.float32, device) - trainer = _trainer_for(lora, device) - _install_checkpoint(trainer, "A", _adapter("dense", rank=2, seed=1)) + with _lora_checkpoint() as (device, _lora, trainer): head = trainer.module( "value_head", lambda: _CudaValueHead(4).to(device), checkpoint="A" ) @@ -638,8 +634,10 @@ def canonicalize_loaded_lora_state( @contextmanager -def _lora_checkpoint(seed=1, *, rng_seed=None): - """Construct the shared single-rank dense checkpoint used by version tests.""" +def _lora_checkpoint( + seed: int = 1, *, rng_seed: int | None = None +) -> Iterator[tuple[torch.device, LoRA, TrainerRank]]: + """Construct a dense LoRA checkpoint in a fresh single-rank context.""" with _single_rank_model_parallel(): if rng_seed is not None: torch.manual_seed(rng_seed) From 6e5ef7f760c5105e3cd4a0aadfdf08a6d03477d9 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Wed, 30 Sep 2026 23:28:33 +0000 Subject: [PATCH 127/150] test: share CP attention worker runtime --- .../cp_attn/test_cpu_offload_residency.py | 30 ++------------ .../cp_attn/test_retained_backward.py | 33 ++------------- tests/support/cp_attention.py | 40 ++++++++++++++++++- 3 files changed, 45 insertions(+), 58 deletions(-) diff --git a/tests/integration/megatron/cp_attn/test_cpu_offload_residency.py b/tests/integration/megatron/cp_attn/test_cpu_offload_residency.py index c80657e20..18fa0f4f7 100644 --- a/tests/integration/megatron/cp_attn/test_cpu_offload_residency.py +++ b/tests/integration/megatron/cp_attn/test_cpu_offload_residency.py @@ -1,11 +1,9 @@ """Actual CP2 graph residency and constrained complete-root placement.""" from dataclasses import asdict, replace -from datetime import timedelta import gc import json import os -from unittest.mock import patch import pytest import torch @@ -15,14 +13,12 @@ pytest.importorskip("megatron.core") from art.megatron.context_parallel import executor # noqa: E402 -from art.megatron.flex_attn import compiled # noqa: E402 -from art.megatron.runtime.compile_cache import configure_reusable_backward # noqa: E402 from art.trainer_rank._graphs import GraphCache # noqa: E402 from art.trainer_rank._memory_policy import ( # noqa: E402 ForwardMemoryCost, placement_cost, ) -from tests.support.cp_attention import prepare_cp2_attention # noqa: E402 +from tests.support.cp_attention import cp2_runtime, prepare_cp2_attention # noqa: E402 @pytest.mark.skipif( @@ -37,29 +33,9 @@ def test_cp_cpu_residency_constrains_complete_root(tmp_path): def _worker(rank, rendezvous): torch.set_num_threads(2) - torch.cuda.set_device(rank) device = torch.device("cuda", rank) - configure_reusable_backward() - dist.init_process_group( - "nccl", - init_method=rendezvous, - rank=rank, - world_size=2, - timeout=timedelta(seconds=120), - device_id=device, - ) - try: - with ( - patch.object(compiled, "_FORCED_FLEX_BACKEND", "TRITON"), - patch.object( - compiled, - "sparse_compiled_flex_attention", - compiled.triton_sparse_compiled_flex_attention, - ), - ): - _check(rank, device) - finally: - dist.destroy_process_group() + with cp2_runtime(rank, rendezvous, device, backend="TRITON", timeout=120): + _check(rank, device) def _check(rank, device): diff --git a/tests/integration/megatron/cp_attn/test_retained_backward.py b/tests/integration/megatron/cp_attn/test_retained_backward.py index 870d1adc5..1795e2ae3 100644 --- a/tests/integration/megatron/cp_attn/test_retained_backward.py +++ b/tests/integration/megatron/cp_attn/test_retained_backward.py @@ -1,6 +1,5 @@ """Actual CP collectives and compiled attention against a dense manual oracle.""" -from datetime import timedelta import os from pathlib import Path from unittest.mock import patch @@ -8,16 +7,13 @@ import pytest import torch -import torch.distributed as dist import torch.multiprocessing as mp pytest.importorskip("megatron.core") from art.megatron.context_parallel import executor # noqa: E402 -from art.megatron.flex_attn import compiled # noqa: E402 -from art.megatron.runtime.compile_cache import configure_reusable_backward # noqa: E402 from art.trainer_rank._graphs import GraphCache # noqa: E402 -from tests.support.cp_attention import prepare_cp2_attention # noqa: E402 +from tests.support.cp_attention import cp2_runtime, prepare_cp2_attention # noqa: E402 def test_cp_retained_failure_releases_original_records(monkeypatch): @@ -74,31 +70,8 @@ def test_cp_retained_backward_matches_dense_attention( def _worker(rank: int, init_method: str, backend: str, dim: int) -> None: # Exercise group ranks independently from CUDA device numbering. device = torch.device("cuda", 1 - rank) - torch.cuda.set_device(device) - configure_reusable_backward() - dist.init_process_group( - "nccl", - init_method=init_method, - rank=rank, - world_size=2, - timeout=timedelta(seconds=90), - device_id=device, - ) - try: - if backend == "TRITON": - with ( - patch.object(compiled, "_FORCED_FLEX_BACKEND", "TRITON"), - patch.object( - compiled, - "sparse_compiled_flex_attention", - compiled.triton_sparse_compiled_flex_attention, - ), - ): - _check_repeated_backward(rank, device, backend, dim) - else: - _check_repeated_backward(rank, device, backend, dim) - finally: - dist.destroy_process_group() + with cp2_runtime(rank, init_method, device, backend=backend, timeout=90): + _check_repeated_backward(rank, device, backend, dim) def _check_repeated_backward( diff --git a/tests/support/cp_attention.py b/tests/support/cp_attention.py index aa07831c2..d4de40102 100644 --- a/tests/support/cp_attention.py +++ b/tests/support/cp_attention.py @@ -1,6 +1,10 @@ -"""Shared CP2 layout setup for attention and graph-residency tests.""" +"""Shared CP2 runtime and layout setup for attention and graph-residency tests.""" +from collections.abc import Iterator +from contextlib import ExitStack, contextmanager +from datetime import timedelta from typing import cast +from unittest.mock import patch import torch import torch.distributed as dist @@ -15,6 +19,8 @@ ParallelTopology, RankRuntimePlan, ) +from art.megatron.flex_attn import compiled +from art.megatron.runtime.compile_cache import configure_reusable_backward from art.preprocessing.pack import PackedTensors @@ -52,3 +58,35 @@ def prepare_cp2_attention( assert indices.numel() == sum(plan.local_valid_lengths) executor.prepare_context_parallel_execution_state(state=state, device=device) return micro, state, plan, indices + + +@contextmanager +def cp2_runtime( + rank: int, init_method: str, device: torch.device, *, backend: str, timeout: int +) -> Iterator[None]: + torch.cuda.set_device(device) + configure_reusable_backward() + dist.init_process_group( + "nccl", + init_method=init_method, + rank=rank, + world_size=2, + timeout=timedelta(seconds=timeout), + device_id=device, + ) + try: + with ExitStack() as stack: + if backend == "TRITON": + stack.enter_context( + patch.object(compiled, "_FORCED_FLEX_BACKEND", "TRITON") + ) + stack.enter_context( + patch.object( + compiled, + "sparse_compiled_flex_attention", + compiled.triton_sparse_compiled_flex_attention, + ) + ) + yield + finally: + dist.destroy_process_group() From 2f6717f9984ae872477c1ea163b42ce155bda491 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 1 Oct 2026 00:33:56 +0000 Subject: [PATCH 128/150] refactor: simplify checkpoint slot snapshot traversal --- src/art/trainer_rank/_checkpoint.py | 10 +++------- 1 file changed, 3 insertions(+), 7 deletions(-) diff --git a/src/art/trainer_rank/_checkpoint.py b/src/art/trainer_rank/_checkpoint.py index 104fe4bac..db35a94a3 100644 --- a/src/art/trainer_rank/_checkpoint.py +++ b/src/art/trainer_rank/_checkpoint.py @@ -1800,12 +1800,6 @@ def _load_adapter( def _slot_snapshot(trainer: TrainerRank) -> _SlotSnapshot: - modules = ( - module - for chunk in trainer.runtime.model - for module in chunk.modules() - if hasattr(module, "_slot_keys") and hasattr(module, "_slot_modules") - ) return tuple( ( module, @@ -1813,7 +1807,9 @@ def _slot_snapshot(trainer: TrainerRank) -> _SlotSnapshot: dict(module._slot_modules.items()), {key: getattr(slot, "ref") for key, slot in module._slot_modules.items()}, ) - for module in cast("Iterable[LoRA]", modules) + for chunk in trainer.runtime.model + for module in cast("Iterable[LoRA]", chunk.modules()) + if hasattr(module, "_slot_keys") and hasattr(module, "_slot_modules") ) From b93343550124ed174252fb5641d76953bfa82cc3 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 1 Oct 2026 00:48:14 +0000 Subject: [PATCH 129/150] perf: skip empty buffer synchronization phases --- src/art/trainer_rank/_heads.py | 2 ++ tests/unit/test_trainer_rank_live_heads.py | 31 ++++++++++++++++++++++ 2 files changed, 33 insertions(+) diff --git a/src/art/trainer_rank/_heads.py b/src/art/trainer_rank/_heads.py index 65fa84661..aaecd5f82 100644 --- a/src/art/trainer_rank/_heads.py +++ b/src/art/trainer_rank/_heads.py @@ -1239,6 +1239,8 @@ def synchronize_head_buffers(trainer: TrainerRank, checkpoints: Any = None) -> N raise trainer._slot_state_error( "Custom buffer registrations differ across ranks" ) + if not targets: + return payload = _checkpoint._phase( lambda: ( { diff --git a/tests/unit/test_trainer_rank_live_heads.py b/tests/unit/test_trainer_rank_live_heads.py index c0fde30d4..09fb5f884 100644 --- a/tests/unit/test_trainer_rank_live_heads.py +++ b/tests/unit/test_trainer_rank_live_heads.py @@ -3,6 +3,7 @@ import asyncio from copy import deepcopy from datetime import timedelta +from unittest.mock import patch import pytest from test_trainer_rank_custom_tensors import _trainer, _use_local_gradients @@ -378,6 +379,36 @@ def test_constructor_staleness_applies_to_heads_before_mutating_gradients(): assert head.left.grad is None +def _empty_buffer_sync_worker(process_rank, init_method): + from art.trainer_rank import _checkpoint, _heads + + with gloo_group(process_rank, init_method, timeout=10): + trainer, rank = _trainer("student") + for parameter in (False, True): + if parameter: + rank.parameter( + "weight", lambda: torch.tensor(1.0), checkpoint="student" + ) + with patch.object( + _checkpoint, "_gather", wraps=_checkpoint._gather + ) as gather: + _heads.synchronize_head_buffers(trainer) + assert gather.call_count == 1 + assert gather.call_args.args[0] == {} + completed = torch.tensor(1) + dist.all_reduce(completed, group=trainer._checkpoint_group()) + assert completed.item() == 2 + + +def test_distributed_empty_buffer_sync_keeps_registration_agreement(tmp_path): + spawn_and_join( + _empty_buffer_sync_worker, + (f"file://{tmp_path / 'empty-buffer-sync'}",), + timeout=60, + failure="Empty buffer synchronization stranded a collective peer", + ) + + def _buffer_authority_worker(process_rank, init_method): from art.trainer_rank._heads import synchronize_head_buffers From 91915300641f33664ff0761b656248580cf26b44 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 1 Oct 2026 00:53:02 +0000 Subject: [PATCH 130/150] test: fold empty buffer checks into failure lifecycle --- tests/unit/test_trainer_rank_live_heads.py | 45 +++++++--------------- 1 file changed, 13 insertions(+), 32 deletions(-) diff --git a/tests/unit/test_trainer_rank_live_heads.py b/tests/unit/test_trainer_rank_live_heads.py index 09fb5f884..0c4bccfe0 100644 --- a/tests/unit/test_trainer_rank_live_heads.py +++ b/tests/unit/test_trainer_rank_live_heads.py @@ -3,7 +3,7 @@ import asyncio from copy import deepcopy from datetime import timedelta -from unittest.mock import patch +from unittest.mock import patch as mock_patch import pytest from test_trainer_rank_custom_tensors import _trainer, _use_local_gradients @@ -379,36 +379,6 @@ def test_constructor_staleness_applies_to_heads_before_mutating_gradients(): assert head.left.grad is None -def _empty_buffer_sync_worker(process_rank, init_method): - from art.trainer_rank import _checkpoint, _heads - - with gloo_group(process_rank, init_method, timeout=10): - trainer, rank = _trainer("student") - for parameter in (False, True): - if parameter: - rank.parameter( - "weight", lambda: torch.tensor(1.0), checkpoint="student" - ) - with patch.object( - _checkpoint, "_gather", wraps=_checkpoint._gather - ) as gather: - _heads.synchronize_head_buffers(trainer) - assert gather.call_count == 1 - assert gather.call_args.args[0] == {} - completed = torch.tensor(1) - dist.all_reduce(completed, group=trainer._checkpoint_group()) - assert completed.item() == 2 - - -def test_distributed_empty_buffer_sync_keeps_registration_agreement(tmp_path): - spawn_and_join( - _empty_buffer_sync_worker, - (f"file://{tmp_path / 'empty-buffer-sync'}",), - timeout=60, - failure="Empty buffer synchronization stranded a collective peer", - ) - - def _buffer_authority_worker(process_rank, init_method): from art.trainer_rank._heads import synchronize_head_buffers @@ -432,12 +402,23 @@ def test_distributed_persistent_buffers_use_dp_zero_authority(tmp_path): def _buffer_snapshot_failure_worker(process_rank, init_method): - from art.trainer_rank import _heads + from art.trainer_rank import _checkpoint, _heads with gloo_group(process_rank, init_method, timeout=15): trainer, rank = _trainer("student") trainer._checkpoint_process_group = dist.group.WORLD trainer._checkpoint_finalize_process_group = dist.group.WORLD + for parameter in (False, True): + if parameter: + rank.parameter( + "weight", lambda: torch.tensor(1.0), checkpoint="student" + ) + with mock_patch.object( + _checkpoint, "_gather", wraps=_checkpoint._gather + ) as gather: + _heads.synchronize_head_buffers(trainer) + assert gather.call_count == 1 + assert gather.call_args.args[0] == {} buffer = rank.buffer("counter", lambda: torch.tensor(1.0), checkpoint="student") if process_rank == 1: buffer.add_(1) From bfcfea5edf3da141faec22e87944d79e95dbc5d0 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 1 Oct 2026 01:00:05 +0000 Subject: [PATCH 131/150] Remove private current-Jacobian replay mode --- src/art/trainer_rank/_corrections.py | 39 +----------- src/art/trainer_rank/_graphs.py | 23 +------- src/art/trainer_rank/_impl.py | 4 +- .../lora/test_trainer_v1_graph_cache.py | 22 +++---- tests/unit/test_trainer_rank_corrections.py | 32 +--------- tests/unit/test_trainer_rank_graphs.py | 59 +++---------------- .../test_trainer_rank_memory_admission.py | 8 --- 7 files changed, 23 insertions(+), 164 deletions(-) diff --git a/src/art/trainer_rank/_corrections.py b/src/art/trainer_rank/_corrections.py index 174cb0fae..2e3e62c51 100644 --- a/src/art/trainer_rank/_corrections.py +++ b/src/art/trainer_rank/_corrections.py @@ -145,8 +145,7 @@ class ForwardCorrectionContext: Current tensors must come from the original inputs/contexts, with the same flattened output layout. Correct only stale forwards; exact original-weight - replay itself does not produce current probabilities. Validate physical - current-weight replay even when corrections are disabled. + replay itself does not produce current probabilities. """ output_count: int @@ -260,42 +259,6 @@ def correct( output = gradient = original = current = tokens = original_tokens = None sorted_tokens = order = positions = matched = logits = corrected = None - def validate_replay( - self, - gradients: Sequence[torch.Tensor | None], - current_tensors: Sequence[torch.Tensor], - ) -> None: - """Require active top-k cotangents to address the same replayed events. - - Ratio evaluation on a separate current forward can realign top-k IDs. - Physical current-weight replay cannot feed original-position cotangents - to a Jacobian whose selected token at that position has changed. - """ - try: - self.requires_current(gradients) - if len(current_tensors) != self.output_count: - raise ValueError("current tensors must match the captured output count") - for output in self.outputs: - gradient = gradients[output.index] - if output.original_tokens is None or gradient is None: - continue - active = gradient != 0 - if not bool(active.any()): - continue - assert output.token_index is not None - tokens = current_tensors[output.token_index] - original = output.original_tokens.to(tokens.device) - if tokens.shape != original.shape or bool( - ((tokens != original) & active.to(tokens.device)).any() - ): - raise RuntimeError( - "current replay changed active top-k token identities; " - "replay the original weights instead" - ) - finally: - del self, gradients, current_tensors - output = gradient = active = tokens = original = None - def capture_forward_corrections( outputs: Any, diff --git a/src/art/trainer_rank/_graphs.py b/src/art/trainer_rank/_graphs.py index 6f3657a3f..f1f2ec69c 100644 --- a/src/art/trainer_rank/_graphs.py +++ b/src/art/trainer_rank/_graphs.py @@ -229,7 +229,6 @@ class _ForwardRecord: corrections: Any = None is_stale: Callable[[], bool] | None = None current_context_factory: Callable[[], AbstractContextManager[Any]] | None = None - replay_with_current: bool = False restored: dict[tuple[torch.device, StorageWeakRef], torch.Tensor] = field( default_factory=dict ) @@ -499,16 +498,10 @@ def offload(self, handle: ForwardHandle) -> None: copies.clear() record.retention = "cpu" - def evict( - self, handle: ForwardHandle, *, replay_with_current: bool | None = None - ) -> None: + def evict(self, handle: ForwardHandle) -> None: record = self._records[handle] if not getattr(record.options, "allow_replay", True): raise RuntimeError("Graph replay is disabled for this forward") - if replay_with_current is not None: - if replay_with_current and record.current_context_factory is None: - raise ValueError("Current-weight replay requires a version context") - record.replay_with_current = replay_with_current record.release_physical() record.retention = "replay" @@ -634,7 +627,6 @@ def _prepare_correction(record, gradients, stale): stale and record.corrections is not None and record.corrections.requires_current(gradients) - and not (record.outputs is None and record.replay_with_current) ): return None # Explicit always may add a no-grad forward. Stage every correction @@ -660,14 +652,9 @@ def _prepare_backward(record, gradients, stale, prepared): gradients = gradients if prepared is None else prepared if not any(gradient is not None for gradient in gradients): return [] - current_replay = stale and record.replay_with_current if record.outputs is None: with record.rng.replay(record.rng_tracker): - physical = record.run( - context_factory=record.current_context_factory - if current_replay - else None - ) + physical = record.run() metadata = tuple( (value.shape, value.dtype, value.device, value.requires_grad) for value in physical @@ -679,11 +666,7 @@ def _prepare_backward(record, gradients, stale, prepared): ) record.replay_count += 1 if prepared is None and stale and record.corrections is not None: - if current_replay: - record.corrections.validate_replay(gradients, record.outputs) - gradients = record.corrections.correct( - gradients, record.outputs if current_replay else None - ) + gradients = record.corrections.correct(gradients) assert record.outputs is not None return [ (output, gradient.to(output.device)) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 45a153b6b..99e5b27e5 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -6820,8 +6820,8 @@ def _correction_state_bytes( else 0 ) ) - # Top-k probabilities and identities persist for replay validation even - # without corrections; explicit always also stages corrected cotangents. + # Captured top-k probabilities and identities persist even without + # corrections; explicit always also stages corrected cotangents. return topk * 8 + logprobs * ( 2 if any(c.policy == "always" for c in options.stale_gradient_corrections) diff --git a/tests/integration/megatron/lora/test_trainer_v1_graph_cache.py b/tests/integration/megatron/lora/test_trainer_v1_graph_cache.py index 12da62285..3bba8f3a8 100644 --- a/tests/integration/megatron/lora/test_trainer_v1_graph_cache.py +++ b/tests/integration/megatron/lora/test_trainer_v1_graph_cache.py @@ -135,8 +135,10 @@ def physical(items, inputs): assert trainer._forward_graph_cache().handles() == () -@pytest.mark.parametrize("mode", ["always", "current_replay"]) -def test_native_stale_logprob_correction_keeps_original_gradient_age(mode, monkeypatch): +@pytest.mark.parametrize("retention", ["gpu", "replay"]) +def test_native_stale_logprob_correction_keeps_original_gradient_age( + retention, monkeypatch +): from art.trainer_rank import ( ImportanceSamplingGradientCorrection, TrainerRankSlotStateError, @@ -166,11 +168,10 @@ def physical(items, inputs): input_tokens=tokens, target_tokens=torch.zeros_like(tokens), options=ForwardOptions( + backward_state=retention, stale_gradient_corrections=( - ImportanceSamplingGradientCorrection( - policy="always" if mode == "always" else "when_available" - ), - ) + ImportanceSamplingGradientCorrection(policy="always"), + ), ), ) group = _ForwardGroupPlan( @@ -195,13 +196,8 @@ def physical(items, inputs): current = [value.detach().clone().requires_grad_() for value in parameters] new_logprobs = (((x @ current[0]) @ current[1]) * 16).log_softmax(-1)[:, 0] weights = (new_logprobs.detach() - old_logprobs.detach()).exp().clamp(0, 5) - expected = torch.autograd.grad( - ((old_logprobs if mode == "always" else new_logprobs) * weights).sum(), - historical if mode == "always" else current, - ) + expected = torch.autograd.grad((old_logprobs * weights).sum(), historical) cache = trainer._forward_graph_cache() - if mode == "current_replay": - cache.evict(cache.handles()[0], replay_with_current=True) with trainer._gradient_transaction(): packets = trainer._forward_cotangent_collector().backward(output.sum()) cache.backward_many( @@ -209,7 +205,7 @@ def physical(items, inputs): ) for parameter, gradient in zip(parameters, expected, strict=True): torch.testing.assert_close(parameter.grad, gradient, atol=2e-4, rtol=5e-5) - assert executions == [True, mode == "current_replay"] + assert executions == [True, False] + ([True] if retention == "replay" else []) assert trainer._version_state()._origins["A"] == {(origin, 2)} trainer._checkpoint_slots["A"].revision += 2 with pytest.raises(TrainerRankSlotStateError, match="staleness 3"): diff --git a/tests/unit/test_trainer_rank_corrections.py b/tests/unit/test_trainer_rank_corrections.py index 6a6f85c6d..185b59faf 100644 --- a/tests/unit/test_trainer_rank_corrections.py +++ b/tests/unit/test_trainer_rank_corrections.py @@ -184,35 +184,7 @@ def test_bad_cotangent_shape_is_rejected_before_requesting_replay() -> None: context.requires_current((None, torch.ones(1), None, None)) -@pytest.mark.parametrize("corrections", [(), (ImportanceSamplingGradientCorrection(),)]) -def test_current_replay_rejects_changed_active_top_k_events_even_without_correction( - corrections, -) -> None: - values = torch.tensor([[-0.5, -1.0]], requires_grad=True) - tokens = torch.tensor([[2, 0]]) - output = ForwardOutput(None, TopK(values, tokens), None, None) - context = capture_forward_corrections( - output, - (values, tokens), - ResolvedForwardOptions(stale_gradient_corrections=corrections), - ) - gradients = (torch.ones_like(values), None) - context.validate_replay(gradients, (values - 1, tokens)) - for changed in (tokens.flip(-1), torch.tensor([[2, 3]])): - with pytest.raises(RuntimeError, match="changed active top-k token identities"): - context.validate_replay(gradients, (values - 1, changed)) - torch.testing.assert_close(gradients[0], torch.ones_like(values)) - - -def test_current_replay_ignores_inactive_top_k_event_changes() -> None: - context, tensors = _fixture("always") - current = (tensors[0], tensors[1], torch.tensor([[2, 3]]), tensors[3] - 1) - context.validate_replay((None, None, None, torch.tensor([[1.0, 0.0]])), current) - context.validate_replay((None, None, None, torch.zeros_like(tensors[3])), current) - context.validate_replay((None, None, None, None), current) - - -def test_current_ratio_evaluation_can_reorder_but_physical_replay_cannot() -> None: +def test_current_ratio_evaluation_can_reorder_top_k_events() -> None: context, tensors = _fixture("always", logits=True) gradients = (None, None, None, torch.ones_like(tensors[3]), None) current = ( @@ -226,5 +198,3 @@ def test_current_ratio_evaluation_can_reorder_but_physical_replay_cannot() -> No context.correct(gradients, current)[3], torch.exp(torch.full_like(tensors[3], -1)), ) - with pytest.raises(RuntimeError, match="changed active top-k token identities"): - context.validate_replay(gradients, current) diff --git a/tests/unit/test_trainer_rank_graphs.py b/tests/unit/test_trainer_rank_graphs.py index 3e48398e0..2e8ce8945 100644 --- a/tests/unit/test_trainer_rank_graphs.py +++ b/tests/unit/test_trainer_rank_graphs.py @@ -490,60 +490,15 @@ def coordinate(function): assert all(storage.expired() for storage in storages) and not cache.handles() -def test_newer_replay_opportunistically_corrects_current_jacobian(): - cache, handle, original, current, executions, _ = _corrected_cache() - cache.evict(handle, replay_with_current=True) +def test_always_correction_after_eviction_uses_original_jacobian(): + cache, handle, original, current, executions, _ = _corrected_cache(policy="always") + cache.evict(handle) cache.backward(handle, (torch.tensor(1.0),)) - assert original.grad is None - torch.testing.assert_close(current.grad, torch.tensor(0.75).exp()) - assert executions == [(False, True), (True, True)] - - -@pytest.mark.parametrize("corrections", [False, True]) -def test_current_replay_rejects_changed_selected_token_events(corrections, request): - from art.trainer_rank._corrections import capture_forward_corrections - - if gc.isenabled(): - request.addfinalizer(gc.enable) - gc.disable() - cache = GraphCache() - parameter = torch.nn.Parameter(torch.tensor([1.0, 2.0])) - tokens = [torch.tensor([0, 1])] - storages = [] - - def execute(_): - logprobs = parameter.log_softmax(-1) - output = logprobs[tokens[0]] - storages.extend( - StorageWeakRef(value.untyped_storage()) for value in (logprobs, output) - ) - return output, tokens[0] - - handle, outputs = cache.run(execute, None) - context = capture_forward_corrections( - ForwardOutput(None, TopK(outputs[0], outputs[1]), None, None), - outputs, - ResolvedForwardOptions( - stale_gradient_corrections=(ImportanceSamplingGradientCorrection(),) - if corrections - else (), - ), - ) - cache.set_corrections( - handle, context, is_stale=lambda: True, current_context_factory=nullcontext + assert current.grad is None + torch.testing.assert_close( + original.grad, torch.tensor(2.0) * torch.tensor(0.75).exp() ) - cache.evict(handle, replay_with_current=True) - tokens[0] = torch.tensor([1, 0]) - gradient = torch.ones(2) - with pytest.raises(RuntimeError, match="token identit") as failure: - cache.backward(handle, (gradient, None)) - assert failure.value.__traceback__ is not None and failure.value.__cause__ is None - assert len(storages) == 4 and all(storage.expired() for storage in storages) - torch.testing.assert_close(parameter, torch.tensor([1.0, 2.0])) - torch.testing.assert_close(gradient, torch.ones(2)) - torch.testing.assert_close(tokens[0], torch.tensor([1, 0])) - assert parameter.grad is None - assert not cache.handles() + assert executions == [(False, True), (True, False), (False, True)] @pytest.mark.parametrize("checkpointing", [False, True]) diff --git a/tests/unit/test_trainer_rank_memory_admission.py b/tests/unit/test_trainer_rank_memory_admission.py index c7e759749..48a61d08d 100644 --- a/tests/unit/test_trainer_rank_memory_admission.py +++ b/tests/unit/test_trainer_rank_memory_admission.py @@ -450,14 +450,6 @@ def test_correction_budget_covers_captured_storage_and_aliases( assert _impl._correction_state_bytes(group, options) == retained + staged if not grad_enabled: assert next(rank._graph_memory_units(plan))[2].replay_bytes == 0 - elif output.top_k is not None: - current = list(tensors) - index = next( - i for i, tensor in enumerate(tensors) if tensor is output.top_k.tokens - ) - current[index] = current[index].flip(-1) - with pytest.raises(RuntimeError, match="changed active top-k token identities"): - context.validate_replay(gradients, current) @pytest.mark.parametrize( From 9a0f6c9652960ef51490dd3cd14dc5a097764c85 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 1 Oct 2026 01:03:20 +0000 Subject: [PATCH 132/150] fix: preserve explicit correction replay context --- src/art/trainer_rank/_graphs.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/art/trainer_rank/_graphs.py b/src/art/trainer_rank/_graphs.py index f1f2ec69c..f27842548 100644 --- a/src/art/trainer_rank/_graphs.py +++ b/src/art/trainer_rank/_graphs.py @@ -666,7 +666,7 @@ def _prepare_backward(record, gradients, stale, prepared): ) record.replay_count += 1 if prepared is None and stale and record.corrections is not None: - gradients = record.corrections.correct(gradients) + gradients = record.corrections.correct(gradients, None) assert record.outputs is not None return [ (output, gradient.to(output.device)) From f9a9097a8df0f7b4d4f3d9cdbdf2cb51954a671d Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 1 Oct 2026 01:22:29 +0000 Subject: [PATCH 133/150] test: parametrize always correction retention cases --- tests/unit/test_trainer_rank_graphs.py | 22 ++++++++-------------- 1 file changed, 8 insertions(+), 14 deletions(-) diff --git a/tests/unit/test_trainer_rank_graphs.py b/tests/unit/test_trainer_rank_graphs.py index 2e8ce8945..fa01b48fa 100644 --- a/tests/unit/test_trainer_rank_graphs.py +++ b/tests/unit/test_trainer_rank_graphs.py @@ -12,7 +12,7 @@ from torch.multiprocessing.reductions import StorageWeakRef from torch.utils.checkpoint import checkpoint -from art.trainer_rank import ForwardOutput, TopK, TrainerRank +from art.trainer_rank import ForwardOutput, TrainerRank from art.trainer_rank._graphs import GraphCache from art.trainer_rank._impl import _CheckpointSlot from art.trainer_rank._options import ( @@ -340,14 +340,19 @@ def test_opportunistic_exact_replay_never_adds_correction_forward(retention): assert current.grad is None -def test_always_correction_uses_no_grad_current_evaluation_and_old_jacobian(): +@pytest.mark.parametrize("evicted", [False, True], ids=["resident", "evicted"]) +def test_always_correction_uses_no_grad_current_evaluation_and_old_jacobian(evicted): cache, handle, original, current, executions, _ = _corrected_cache(policy="always") + expected_executions = [(False, True), (True, False)] + if evicted: + cache.evict(handle) + expected_executions.append((False, True)) cache.backward(handle, (torch.tensor(1.0),)) torch.testing.assert_close( original.grad, torch.tensor(2.0) * torch.tensor(0.75).exp() ) assert current.grad is None - assert executions == [(False, True), (True, False)] + assert executions == expected_executions def _observe_backward_storage(monkeypatch, request): @@ -490,17 +495,6 @@ def coordinate(function): assert all(storage.expired() for storage in storages) and not cache.handles() -def test_always_correction_after_eviction_uses_original_jacobian(): - cache, handle, original, current, executions, _ = _corrected_cache(policy="always") - cache.evict(handle) - cache.backward(handle, (torch.tensor(1.0),)) - assert current.grad is None - torch.testing.assert_close( - original.grad, torch.tensor(2.0) * torch.tensor(0.75).exp() - ) - assert executions == [(False, True), (True, False), (False, True)] - - @pytest.mark.parametrize("checkpointing", [False, True]) def test_abandoned_release_frees_physical_record_without_autograd(checkpointing): cache = GraphCache() From 4a21f9a285d7127bf46f813740bb6c05bdca88ca Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 1 Oct 2026 01:35:37 +0000 Subject: [PATCH 134/150] test: share CUDA measurement reporting --- .../test_trainer_rank_head_memory_cuda.py | 29 +++++++------------ .../test_trainer_rank_memory_policy_cuda.py | 18 +++++------- .../test_trainer_rank_output_memory_cuda.py | 20 +++++-------- tests/unit/trainer_rank_test_support.py | 5 ++++ 4 files changed, 31 insertions(+), 41 deletions(-) diff --git a/tests/unit/test_trainer_rank_head_memory_cuda.py b/tests/unit/test_trainer_rank_head_memory_cuda.py index 9d2bed82d..d5155957a 100644 --- a/tests/unit/test_trainer_rank_head_memory_cuda.py +++ b/tests/unit/test_trainer_rank_head_memory_cuda.py @@ -1,12 +1,12 @@ """Allocator assertions; requires a validation-owned GPU reservation.""" -import json import os import pytest from test_trainer_rank_custom_tensors import _trainer import torch from torch.multiprocessing.reductions import StorageWeakRef +from trainer_rank_test_support import report_measurement from art.trainer_rank import TrainerRankMemoryError from art.trainer_rank._commands import _Executor @@ -57,16 +57,12 @@ def test_repeated_remote_head_cotangents_fit_original_gradient_reserve( torch.testing.assert_close( parameter.grad, torch.full_like(parameter, 73 if existing_gradient else 72) ) - print( - "REMOTE_HEAD_RESERVATION=" - + json.dumps( - dict( - existing_gradient=existing_gradient, - peak=peak, - reserve=reserve, - repeated_captures=8, - ) - ) + report_measurement( + "REMOTE_HEAD_RESERVATION", + existing_gradient=existing_gradient, + peak=peak, + reserve=reserve, + repeated_captures=8, ) @@ -119,11 +115,8 @@ def watch(checkpoint, name, custom): assert torch.cuda.memory_allocated() <= baseline assert not trainer._checkpoint_slots["student"].custom cache.release(handle) - print( - "LATE_HEAD_RELEASE=" - + json.dumps( - dict( - kind=kind, retained_extra_bytes=torch.cuda.memory_allocated() - baseline - ) - ) + report_measurement( + "LATE_HEAD_RELEASE", + kind=kind, + retained_extra_bytes=torch.cuda.memory_allocated() - baseline, ) diff --git a/tests/unit/test_trainer_rank_memory_policy_cuda.py b/tests/unit/test_trainer_rank_memory_policy_cuda.py index de9acbefb..86bb96ad5 100644 --- a/tests/unit/test_trainer_rank_memory_policy_cuda.py +++ b/tests/unit/test_trainer_rank_memory_policy_cuda.py @@ -2,13 +2,13 @@ from dataclasses import asdict import gc -import json import os from pathlib import Path from types import SimpleNamespace import pytest import torch +from trainer_rank_test_support import report_measurement from art.trainer_rank import TrainerRank from art.trainer_rank._graphs import GraphCache @@ -56,15 +56,11 @@ def test_sequential_gradient_publications_fit_original_admission_reserve( assert peak == parameter.numel() * parameter.element_size() * ( 2 if existing_gradient else 3 ) - print( - "GRADIENT_RESERVATION=" - + json.dumps( - { - "existing_gradient": existing_gradient, - "reserved_bytes": reserved, - "peak_bytes": peak, - } - ) + report_measurement( + "GRADIENT_RESERVATION", + existing_gradient=existing_gradient, + reserved_bytes=reserved, + peak_bytes=peak, ) torch.testing.assert_close(parameter.grad, source * 2) @@ -167,7 +163,7 @@ def _run_root(workload, state, device): "planned": asdict(planned), "host_budget_two_ranks": asdict(host_memory_budget(local_world_size=2)), } - print("MEMORY_MEASUREMENT=" + json.dumps(measurement, sort_keys=True)) + report_measurement("MEMORY_MEASUREMENT", sort_keys=True, **measurement) assert peak <= planned.gpu_required_bytes + 8 * 1024**2 if device == "cpu" and state == "replay": assert forward_retained < 1024**2 diff --git a/tests/unit/test_trainer_rank_output_memory_cuda.py b/tests/unit/test_trainer_rank_output_memory_cuda.py index 61a83e684..dbbd852d4 100644 --- a/tests/unit/test_trainer_rank_output_memory_cuda.py +++ b/tests/unit/test_trainer_rank_output_memory_cuda.py @@ -1,13 +1,13 @@ """Reserved-GPU canary for logical copies beside an existing replay backward.""" import gc -import json import os import pytest from test_trainer_rank_custom_tensors import _trainer from test_trainer_rank_output_memory import _output import torch +from trainer_rank_test_support import report_measurement from art.trainer_rank._commands import _Executor, _view from art.trainer_rank._tensors import detach_tree @@ -69,15 +69,11 @@ def test_logical_outputs_leave_existing_replay_backward_admissible(monkeypatch): atol=0, ) assert not cache.handles() - print( - "OUTPUT_BACKWARD_RESERVE=" - + json.dumps( - dict( - available_bytes=capacity, - restore_bytes=workspace, - rejected_copy_bytes=large, - admitted_copy_bytes=8 * 1024**2, - peak_bytes=peak, - ) - ) + report_measurement( + "OUTPUT_BACKWARD_RESERVE", + available_bytes=capacity, + restore_bytes=workspace, + rejected_copy_bytes=large, + admitted_copy_bytes=8 * 1024**2, + peak_bytes=peak, ) diff --git a/tests/unit/trainer_rank_test_support.py b/tests/unit/trainer_rank_test_support.py index b5cbfda55..c69707ac8 100644 --- a/tests/unit/trainer_rank_test_support.py +++ b/tests/unit/trainer_rank_test_support.py @@ -3,6 +3,7 @@ from collections.abc import Callable from contextlib import contextmanager from datetime import timedelta +import json import sys import time from types import ModuleType, SimpleNamespace @@ -211,3 +212,7 @@ def spawn_and_join(worker, args, *, timeout, failure, nprocs=2): ) except BaseException: pass + + +def report_measurement(label: str, *, sort_keys: bool = False, **values: Any) -> None: + print(label + "=" + json.dumps(values, sort_keys=sort_keys)) From 5b084333260105c404378d875166815b737ad91d Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 1 Oct 2026 02:23:34 +0000 Subject: [PATCH 135/150] Snapshot deferred releases before filtering callback handles --- src/art/trainer_rank/_commands.py | 2 +- tests/unit/test_trainer_rank_commands.py | 40 ++++++++++++++++++++++++ 2 files changed, 41 insertions(+), 1 deletion(-) diff --git a/src/art/trainer_rank/_commands.py b/src/art/trainer_rank/_commands.py index 919494115..5ca25cd31 100644 --- a/src/art/trainer_rank/_commands.py +++ b/src/art/trainer_rank/_commands.py @@ -751,7 +751,7 @@ def _invoke(self, operation: str, *args: Any, **kwargs: Any) -> Any: released = self._executor.state.released handles = tuple( handle - for handle in released + for handle in tuple(released) if handle.startswith(f"{self._executor.mode}:") ) if handles: diff --git a/tests/unit/test_trainer_rank_commands.py b/tests/unit/test_trainer_rank_commands.py index 31cd686e3..4c1f8f22e 100644 --- a/tests/unit/test_trainer_rank_commands.py +++ b/tests/unit/test_trainer_rank_commands.py @@ -574,6 +574,46 @@ def test_released_client_graph_releases_physical_bridge(): assert not rank._rank_command_state.graphs +@pytest.mark.parametrize("mode", ["rank", "zero"]) +def test_concurrent_output_release_is_deferred_to_next_command(mode): + rank: Any = _Rank() + executor = _Executor(rank, mode) + view, state = _view(executor), executor.state + first, later = view.forward(_input(2)), [view.forward(_input(3))] + old, new = state.graphs + entered, finished = threading.Event(), threading.Event() + + class PausingHandle(str): + def startswith(self, *args): + entered.set() + assert finished.wait(5) + return super().startswith(*args) + + del first + state.released = {PausingHandle(old)} + observed = [] + rank.zero_grad = lambda: observed.append(set(state.graphs)) + + def drop_later(): + if entered.wait(5): + later.clear() # Run the real autograd finalizer on this thread. + finished.set() + + thread = threading.Thread(target=drop_later, name="deferred-output-release") + thread.start() + try: + view.zero_grad() + assert state.released == {new} + assert set(state.graphs) == {new} + view.zero_grad() + assert not state.released and not state.graphs + assert observed == [{new}, set()] + finally: + entered.set() + thread.join(5) + assert not thread.is_alive() + + def test_forward_batches_captures_policy_before_iteration(): from art.trainer_rank._rng import TrainerRNG From ee99c6345c8ed537e7206d9c2a467f99f52893e0 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 1 Oct 2026 02:29:05 +0000 Subject: [PATCH 136/150] fix: skip metadata capture for disabled corrections --- src/art/trainer_rank/_corrections.py | 12 ++++---- src/art/trainer_rank/_impl.py | 11 ++------ tests/unit/test_trainer_rank_graphs.py | 28 ++++++++++++++----- .../test_trainer_rank_memory_admission.py | 10 +++++-- 4 files changed, 37 insertions(+), 24 deletions(-) diff --git a/src/art/trainer_rank/_corrections.py b/src/art/trainer_rank/_corrections.py index 2e3e62c51..fc80a44d4 100644 --- a/src/art/trainer_rank/_corrections.py +++ b/src/art/trainer_rank/_corrections.py @@ -265,7 +265,7 @@ def capture_forward_corrections( tensors: Sequence[torch.Tensor], options: ResolvedForwardOptions, ) -> ForwardCorrectionContext: - """Map a ForwardOutput tree to the caller's deduplicated flat tensor layout.""" + """Capture enabled corrections; disabled policies skip output metadata entirely.""" from ._impl import ForwardOutput def leaves(value: Any) -> Iterator[Any]: @@ -281,7 +281,9 @@ def leaves(value: Any) -> Iterator[Any]: raise TypeError("correction capture requires a ForwardOutput tree") corrections = options.stale_gradient_corrections - correction = corrections[0] if corrections else None + if not corrections: + return ForwardCorrectionContext(len(tensors), None, ()) + correction = corrections[0] indices = {id(tensor): index for index, tensor in enumerate(tensors)} if len(indices) != len(tensors): raise ValueError("correction capture requires deduplicated flat tensors") @@ -292,11 +294,7 @@ def leaves(value: Any) -> Iterator[Any]: ("target_logprobs", output.target_logprobs), ("top_k", None if output.top_k is None else output.top_k.logprobs), ): - if ( - tensor is None - or not tensor.requires_grad - or (correction is None and kind != "top_k") - ): + if tensor is None or not tensor.requires_grad: continue index = indices[id(tensor)] entry = _OutputCorrection( diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 99e5b27e5..d883aa6ca 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -6807,21 +6807,16 @@ def _resolved_request_policy(options: ForwardOptions | None) -> ResolvedForwardO def _correction_state_bytes( group: _ForwardGroupPlan, options: ResolvedForwardOptions ) -> int: - if not group.grad_enabled: + if not group.grad_enabled or not options.stale_gradient_corrections: return 0 topk = sum( item.input_ids.numel() * (item.request.top_k or 0) for item in group.items ) logprobs = 4 * ( topk - + ( - sum(item.labels.numel() for item in group.items if item.labels is not None) - if options.stale_gradient_corrections - else 0 - ) + + sum(item.labels.numel() for item in group.items if item.labels is not None) ) - # Captured top-k probabilities and identities persist even without - # corrections; explicit always also stages corrected cotangents. + # Explicit always also stages corrected cotangents. return topk * 8 + logprobs * ( 2 if any(c.policy == "always" for c in options.stale_gradient_corrections) diff --git a/tests/unit/test_trainer_rank_graphs.py b/tests/unit/test_trainer_rank_graphs.py index fa01b48fa..48143600b 100644 --- a/tests/unit/test_trainer_rank_graphs.py +++ b/tests/unit/test_trainer_rank_graphs.py @@ -216,14 +216,28 @@ def test_disabled_policy_rejects_before_execute(retention, options): ) -def test_bad_cotangent_rejects_before_replay(): +@pytest.mark.parametrize("disabled_corrections", [False, True]) +@pytest.mark.parametrize( + "gradients,match", + [ + ((), "count"), + ((torch.ones(2),), "mismatch"), + ((torch.ones((), dtype=torch.float64),), "mismatch"), + ], + ids=["count", "shape", "dtype"], +) +def test_bad_cotangent_rejects_before_replay(disabled_corrections, gradients, match): cache = GraphCache() parameter = torch.nn.Parameter(torch.tensor(2.0)) - handle, _ = cache.run( + handle, outputs = cache.run( lambda x: (x * parameter,), torch.tensor(3.0), retention="replay" ) - with pytest.raises(ValueError, match="mismatch"): - cache.backward(handle, (torch.ones(2),)) + if disabled_corrections: + cache.set_corrections( + handle, _logprob_corrections(outputs, None), is_stale=lambda: True + ) + with pytest.raises(ValueError, match=match): + cache.backward(handle, gradients) assert cache.state(handle).replay_count == 0 assert parameter.grad is None @@ -290,9 +304,9 @@ def _logprob_corrections(outputs, policy): ForwardOutput(outputs[0], None, None, None), outputs, ResolvedForwardOptions( - stale_gradient_corrections=( - ImportanceSamplingGradientCorrection(policy=policy), - ) + stale_gradient_corrections=() + if policy is None + else (ImportanceSamplingGradientCorrection(policy=policy),) ), ) diff --git a/tests/unit/test_trainer_rank_memory_admission.py b/tests/unit/test_trainer_rank_memory_admission.py index 48a61d08d..4cbb761e2 100644 --- a/tests/unit/test_trainer_rank_memory_admission.py +++ b/tests/unit/test_trainer_rank_memory_admission.py @@ -378,7 +378,7 @@ def test_correction_metadata_and_explicit_prepass_are_budgeted(rank): ForwardOptions(stale_gradient_corrections=()) ), ) - == 6 * 12 + == 0 ) @@ -439,6 +439,12 @@ def test_correction_budget_covers_captured_storage_and_aliases( gradients = tuple( torch.ones_like(tensor) if tensor.requires_grad else None for tensor in tensors ) + if corrections == (): + assert context.tensors == () + assert all( + after is before + for before, after in zip(gradients, context.correct(gradients), strict=True) + ) staged = 0 if corrections and corrections[0].policy == "always": corrected = context.correct(gradients, tensors) @@ -455,7 +461,7 @@ def test_correction_budget_covers_captured_storage_and_aliases( @pytest.mark.parametrize( "corrections,expected", [ - ((), 144), + ((), 0), (None, 144), ((ImportanceSamplingGradientCorrection(policy="always"),), 192), ], From 9e3b79b44bbb2a53a4772b3b04dbddfec3c36d40 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 1 Oct 2026 02:32:09 +0000 Subject: [PATCH 137/150] Avoid gathering logical rank outputs back to their owner --- src/art/trainer_rank/_commands.py | 4 ++-- tests/unit/test_trainer_rank_commands.py | 29 ++++++++++++++++++++---- 2 files changed, 27 insertions(+), 6 deletions(-) diff --git a/src/art/trainer_rank/_commands.py b/src/art/trainer_rank/_commands.py index 5ca25cd31..5e476295a 100644 --- a/src/art/trainer_rank/_commands.py +++ b/src/art/trainer_rank/_commands.py @@ -430,8 +430,8 @@ def _execute(self, command: _Command) -> Any: def _gather_outputs(self, value: Any) -> list[Any] | None: try: - if not self.distributed or len(self.members) == 1: - return [value] + if self.mode == "rank" or not self.distributed or len(self.members) == 1: + return [value] if self.is_leader else None def admit_serialization() -> None: from ._tensors import flatten_tensors diff --git a/tests/unit/test_trainer_rank_commands.py b/tests/unit/test_trainer_rank_commands.py index 4c1f8f22e..e2967283b 100644 --- a/tests/unit/test_trainer_rank_commands.py +++ b/tests/unit/test_trainer_rank_commands.py @@ -306,13 +306,16 @@ def backward(ctx, gradient): # ty: ignore[invalid-method-override] return gradient -def _distributed_worker(physical, rendezvous, output): +def _distributed_worker(physical, rendezvous, output, parallel_axis): with ( gloo_group(physical, f"file://{rendezvous}", world_size=4), megatron_topology(physical, dp_size=2, tp_size=2) as ps, ): dp_groups = [dist.new_group([0, 2]), dist.new_group([1, 3])] dp, tp = divmod(physical, 2) + if parallel_axis == "cp": + ps.get_context_parallel_rank = ps.get_tensor_model_parallel_rank + ps.get_tensor_model_parallel_rank = lambda: 0 rank: Any = _Rank(dp, 2) counts = [0, 0] @@ -353,9 +356,26 @@ def logical(view): if out: view.backward(_loss_tree(out)) view.optim_step() + batches = list(view.forward_batches([_input(3), _input(5)], no_grad=True)) + assert [batch.indices for batch in batches] == [ + [0] if dp == 0 else [], + [1] if dp == 1 else [], + ] + assert all( + not output.hidden_states.requires_grad + for batch in batches + for output in batch.outputs + ) + handle = view.open_forward_batches([_input(7)], no_grad=True) + batch = view.next_forward_batch(handle) + assert batch is not None and batch.indices == ([0] if dp == 0 else []) + assert view.next_forward_batch(handle) is None return dp - per_dp = asyncio.run(run_rank_callback(rank, logical)) + with patch.object( + dist, "gather_object", side_effect=AssertionError("rank output gather") + ): + per_dp = asyncio.run(run_rank_callback(rank, logical)) # Persistent streams release the command scope between waves. Every # physical participant must retain its iterator while serving other jobs. @@ -537,11 +557,12 @@ def forward(self, value): torch.save(gathered, output) -def test_gloo_dp2_tp2_participation_and_gradients(tmp_path): +@pytest.mark.parametrize("parallel_axis", ["tp", "cp"]) +def test_gloo_dp2_parallel_participation_and_gradients(tmp_path, parallel_axis): output = tmp_path / "result.pt" mp.spawn( _distributed_worker, - args=(str(tmp_path / "init"), str(output)), + args=(str(tmp_path / "init"), str(output), parallel_axis), nprocs=4, join=True, ) From 971fbab10eb02ea6dc6d17d65d8fdf651cf31054 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 1 Oct 2026 02:38:51 +0000 Subject: [PATCH 138/150] ci: follow parametrized trainer routing test selectors --- .github/workflows/prek.yml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/.github/workflows/prek.yml b/.github/workflows/prek.yml index 20a69f7af..45a638b6e 100644 --- a/.github/workflows/prek.yml +++ b/.github/workflows/prek.yml @@ -231,7 +231,7 @@ jobs: tests/unit/test_trainer_rank_weird_shapes.py \ tests/unit/test_trainer_rank_slot_graph_lifetime.py \ tests/unit/test_trainer_rank_forward_handoff.py \ - tests/unit/test_trainer_rank_commands.py::test_gloo_dp2_tp2_participation_and_gradients \ + tests/unit/test_trainer_rank_commands.py::test_gloo_dp2_parallel_participation_and_gradients \ tests/unit/test_trainer_rank_admission_inputs.py \ tests/unit/test_trainer_rank_checkpoint_memory.py \ tests/unit/test_trainer_rank_profile_warm.py \ @@ -269,7 +269,7 @@ jobs: uv run --no-sync pytest --nbval --current-env --tb=short tests/unit \ --deselect=tests/unit/test_trainer_rank_cache_recovery.py::test_dense_cp_exact_demand_fits_after_recovery \ --deselect=tests/unit/test_trainer_rank_cache_recovery.py::test_dense_cp_exact_demand_refuses_after_recovery \ - --deselect=tests/unit/test_trainer_rank_commands.py::test_gloo_dp2_tp2_participation_and_gradients \ + --deselect=tests/unit/test_trainer_rank_commands.py::test_gloo_dp2_parallel_participation_and_gradients \ --ignore=tests/unit/test_megatron_reference_logprobs.py \ --ignore=tests/unit/test_moe_routing_replay.py \ --ignore=tests/unit/test_moe_routing_real_path.py \ From 81e58a662ef6db7d379048093553026eba10640c Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 1 Oct 2026 02:56:39 +0000 Subject: [PATCH 139/150] Consolidate correction memory budget oracles in storage matrix --- .../test_trainer_rank_memory_admission.py | 62 +++++++------------ 1 file changed, 22 insertions(+), 40 deletions(-) diff --git a/tests/unit/test_trainer_rank_memory_admission.py b/tests/unit/test_trainer_rank_memory_admission.py index 4cbb761e2..8ade363a4 100644 --- a/tests/unit/test_trainer_rank_memory_admission.py +++ b/tests/unit/test_trainer_rank_memory_admission.py @@ -354,44 +354,24 @@ def test_empty_output_does_not_pin_hidden_storage_and_keeps_autograd(): assert hidden.grad is not None and hidden.grad.count_nonzero() == 0 -def test_correction_metadata_and_explicit_prepass_are_budgeted(rank): - request = ForwardInput( - input_tokens=torch.arange(3), - target_tokens=torch.arange(3), - top_k=2, - ) - group = rank._plan_flat_forward([request]).groups[0] - default = _impl._resolved_request_policy(None) - always = _impl._resolved_request_policy( - ForwardOptions( - stale_gradient_corrections=( - ImportanceSamplingGradientCorrection(policy="always"), - ) - ) - ) - assert _impl._correction_state_bytes(group, default) == 3 * 4 + 6 * 12 - assert _impl._correction_state_bytes(group, always) == 3 * 8 + 6 * 16 - assert ( - _impl._correction_state_bytes( - group, - _impl._resolved_request_policy( - ForwardOptions(stale_gradient_corrections=()) - ), - ) - == 0 - ) - - @pytest.mark.parametrize("grad_enabled", [False, True]) -@pytest.mark.parametrize("top_k", [0, 4]) +@pytest.mark.parametrize( + "top_k, hidden_states", + [(0, True), (4, True), (2, False)], + ids=["0", "4", "2-no-hidden"], +) @pytest.mark.parametrize("label_columns", [0, 1, 2]) @pytest.mark.parametrize( - "corrections", - [None, (), (ImportanceSamplingGradientCorrection(policy="always"),)], + "corrections, expected", + [ + (None, 3 * 4 + 6 * 12), + ((), 0), + ((ImportanceSamplingGradientCorrection(policy="always"),), 3 * 8 + 6 * 16), + ], ids=["default", "disabled", "always"], ) def test_correction_budget_covers_captured_storage_and_aliases( - rank, grad_enabled, top_k, label_columns, corrections + rank, grad_enabled, top_k, hidden_states, label_columns, corrections, expected ): from art.trainer_rank._corrections import capture_forward_corrections from art.trainer_rank._tensors import flatten_tensors @@ -404,18 +384,16 @@ def test_correction_budget_covers_captured_storage_and_aliases( else None ) options = _impl._resolved_request_policy( - ForwardOptions( - stale_gradient_corrections=_impl.Unset - if corrections is None - else corrections - ) + None + if corrections is None + else ForwardOptions(stale_gradient_corrections=corrections) ) request = ForwardInput( input_tokens=torch.arange(3), target_tokens=labels, top_k=top_k or None, - hidden_states=True, - no_grad=not grad_enabled, + hidden_states=hidden_states, + no_grad=None if top_k == 2 and grad_enabled else not grad_enabled, ) plan = rank._plan_flat_forward([request]) group = plan.groups[0] @@ -430,7 +408,9 @@ def test_correction_budget_covers_captured_storage_and_aliases( if top_k else None, logits=None, - hidden_states=torch.zeros(3, 1, requires_grad=grad_enabled), + hidden_states=torch.zeros(3, 1, requires_grad=grad_enabled) + if hidden_states + else None, ) tree = {"output": output, "alias": [output]} tensors, _ = flatten_tensors(tree) @@ -454,6 +434,8 @@ def test_correction_budget_covers_captured_storage_and_aliases( if after is not None and after is not before ) assert _impl._correction_state_bytes(group, options) == retained + staged + if top_k == 2 and label_columns == 1 and grad_enabled: + assert _impl._correction_state_bytes(group, options) == expected if not grad_enabled: assert next(rank._graph_memory_units(plan))[2].replay_bytes == 0 From b7db429f38452c3fee5704822c6fb74533c1c500 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 1 Oct 2026 03:27:48 +0000 Subject: [PATCH 140/150] Reuse existing LoRA helpers in trainer acceptance --- dev/trainer_v1_acceptance.py | 31 +++++++++++++------------------ 1 file changed, 13 insertions(+), 18 deletions(-) diff --git a/dev/trainer_v1_acceptance.py b/dev/trainer_v1_acceptance.py index 50080e724..7a2f83efe 100644 --- a/dev/trainer_v1_acceptance.py +++ b/dev/trainer_v1_acceptance.py @@ -70,8 +70,10 @@ def _loss_and_cotangents(first, second): def _canonical_gradients(rank, checkpoint): """Reduce once, then gather canonical LoRA shards without altering weights.""" - from art.megatron.lora import LoRA - from art.megatron.weights.lora_publish import _merge_manifest_entries + from art.megatron.weights.lora_publish import ( + iter_lora_modules, + merge_sharded_adapter_entries, + ) parameters = rank._checkpoint_slots[checkpoint].params reduced = rank._reduce_dynamic_grads(parameters, scale_grads=1.0) @@ -80,19 +82,14 @@ def _canonical_gradients(rank, checkpoint): for parameter, gradient in zip(parameters, reduced, strict=True) } local = {} - for chunk in rank.runtime.model: - for module in chunk.modules(): - if not isinstance(module, LoRA): - continue - for key, parameter, expert in module._export_items( - rank._slot_ref(checkpoint) - ): - value = by_id[id(parameter)] - value = value if expert is None else value[expert] - local[key] = ( - module._manifest_for_param(parameter), - value.T.float().cpu(), - ) + for module in iter_lora_modules(rank.runtime.model): + for key, parameter, expert in module._export_items(rank._slot_ref(checkpoint)): + value = by_id[id(parameter)] + value = value if expert is None else value[expert] + local[key] = ( + module._manifest_for_param(parameter), + value.T.float().cpu(), + ) if dist.get_rank() == 0: for name, custom in rank._checkpoint_slots[checkpoint].custom.items(): for key, parameter in custom.value.named_parameters(): @@ -108,9 +105,7 @@ def _canonical_gradients(rank, checkpoint): for shard in gathered: for key, entry in shard.items(): groups[key].append(entry) - return { - key: _merge_manifest_entries(key, entries) for key, entries in groups.items() - } + return merge_sharded_adapter_entries(groups) def _compare(actual, reference): From 465488f77fe18216e034358457072f1bf1a64489 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 1 Oct 2026 03:50:49 +0000 Subject: [PATCH 141/150] Remove redundant trainer command forwarding helpers --- src/art/trainer_rank/_commands.py | 15 ++++----------- tests/unit/test_trainer_command_transport.py | 8 +++++--- .../unit/test_trainer_rank_release_completion.py | 4 ++-- 3 files changed, 11 insertions(+), 16 deletions(-) diff --git a/src/art/trainer_rank/_commands.py b/src/art/trainer_rank/_commands.py index 5e476295a..74280b62b 100644 --- a/src/art/trainer_rank/_commands.py +++ b/src/art/trainer_rank/_commands.py @@ -61,10 +61,6 @@ class _Command: grad_enabled: bool -def _encode_command(command: _Command) -> bytes: - return _transport.encode(command) - - @dataclass(frozen=True) class _OutputPacket: packet: Any @@ -254,16 +250,13 @@ def _finish_release(self, release: _Release) -> None: {"message": self.state.release_error, "exception": error} ) - async def _join_release(self) -> asyncio.CancelledError | None: - return await join_rank_callback_release(self.rank) - async def reconcile_releases( self, *, defer_cancellation: bool = False ) -> asyncio.CancelledError | None: """Join prior cleanup before entering another all-rank boundary.""" - cancelled = await self._join_release() + cancelled = await join_rank_callback_release(self.rank) self._start_release() - cancelled = await self._join_release() or cancelled + cancelled = await join_rank_callback_release(self.rank) or cancelled if cancelled is not None and not defer_cancellation: raise cancelled return cancelled @@ -290,7 +283,7 @@ def _broadcast(self, command: _Command | None) -> _Command: payload = None if command is not None: try: - payload = _encode_command(command) + payload = _transport.encode(command) except Exception as exc: command = _Command( command.sequence, @@ -299,7 +292,7 @@ def _broadcast(self, command: _Command | None) -> _Command: {}, False, ) - payload = _encode_command(command) + payload = _transport.encode(command) objects: list[Any] = [payload] dist.broadcast_object_list(objects, src=self.leader, group=self.group) return self._decode(objects[0]) diff --git a/tests/unit/test_trainer_command_transport.py b/tests/unit/test_trainer_command_transport.py index 94a18eb05..511440570 100644 --- a/tests/unit/test_trainer_command_transport.py +++ b/tests/unit/test_trainer_command_transport.py @@ -11,8 +11,8 @@ import torch from trainer_rank_test_support import gloo_group, megatron_topology, spawn_and_join -from art.trainer_rank import ForwardInput, ForwardOutput, TrainerRank -from art.trainer_rank._commands import _Command, _encode_command, _Executor +from art.trainer_rank import ForwardInput, ForwardOutput, TrainerRank, _transport +from art.trainer_rank._commands import _Command, _Executor from art.trainer_rank._heads import HeadRegistration from art.trainer_rank._impl import _CheckpointSlot from art.trainer_rank._tensors import managed_tensor @@ -97,7 +97,9 @@ def test_command_codec_preserves_nested_types_aliases_and_closure_storage() -> N executor = _Executor(cast(Any, _Rank()), "zero") source = _payload("cpu") - result = executor._decode(_encode_command(_Command(1, "test", (source,), {}, True))) + result = executor._decode( + _transport.encode(_Command(1, "test", (source,), {}, True)) + ) assert result.operation == "test" and result.grad_enabled _check_payload(result.args[0]) assert result.args[0].base.untyped_storage() is not source.base.untyped_storage() diff --git a/tests/unit/test_trainer_rank_release_completion.py b/tests/unit/test_trainer_rank_release_completion.py index f0d6356bb..565fbf3ea 100644 --- a/tests/unit/test_trainer_rank_release_completion.py +++ b/tests/unit/test_trainer_rank_release_completion.py @@ -31,7 +31,7 @@ async def test_completed_release_is_finalized_before_queued_done_callback( ) completed.set_result(None) completed.add_done_callback(lambda _: executor._finish_release(release)) - await executor._join_release() + await join_rank_callback_release(rank) assert not state.graphs and not state.released assert state.pending_release is None # The already queued callback cannot alter the next release's ownership. @@ -58,7 +58,7 @@ async def test_background_cleanup_failure_is_reported_and_blocks_next_entry(canc else: completed.set_exception(RuntimeError("injected release transport error")) with pytest.raises(RuntimeError, match="Callback release reconciliation failed"): - await asyncio.wait_for(executor._join_release(), 1) + await asyncio.wait_for(join_rank_callback_release(rank), 1) assert len(reports) == 1 with pytest.raises(RuntimeError, match="Callback release reconciliation failed"): await executor.reconcile_releases() From 57cdfba4dfff6c9f6f99cceee3d525e18e3dcec9 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 1 Oct 2026 03:54:46 +0000 Subject: [PATCH 142/150] fix: validate gathered gradient shards --- dev/trainer_v1_acceptance.py | 1 + 1 file changed, 1 insertion(+) diff --git a/dev/trainer_v1_acceptance.py b/dev/trainer_v1_acceptance.py index 7a2f83efe..52d5b08fd 100644 --- a/dev/trainer_v1_acceptance.py +++ b/dev/trainer_v1_acceptance.py @@ -103,6 +103,7 @@ def _canonical_gradients(rank, checkpoint): return None groups = defaultdict(list) for shard in gathered: + assert shard is not None for key, entry in shard.items(): groups[key].append(entry) return merge_sharded_adapter_entries(groups) From 54ca88590660445c6d62a6c7071336b4fc68b3e3 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 1 Oct 2026 05:06:16 +0000 Subject: [PATCH 143/150] Share live module source binding --- src/art/trainer_rank/_heads.py | 33 ++++++++++++--------------------- 1 file changed, 12 insertions(+), 21 deletions(-) diff --git a/src/art/trainer_rank/_heads.py b/src/art/trainer_rank/_heads.py index aaecd5f82..de9151348 100644 --- a/src/art/trainer_rank/_heads.py +++ b/src/art/trainer_rank/_heads.py @@ -315,15 +315,20 @@ def __init__( publish: Callable[[Mapping[str, torch.Tensor]], None], ) -> None: super().__init__() - object.__setattr__(self, "_source", module) object.__setattr__(self, "_capture", capture) object.__setattr__(self, "_publish", publish) - self._parameters = module._parameters - self._buffers = module._buffers - self._modules = module._modules + self._bind_source(module) self._non_persistent_buffers_set = module._non_persistent_buffers_set self.training = module.training + def _bind_source(self, module: torch.nn.Module) -> None: + object.__setattr__(self, "_source", module) + self._parameters, self._buffers, self._modules = ( + module._parameters, + module._buffers, + module._modules, + ) + def __getattr__(self, name: str) -> Any: try: return super().__getattr__(name) @@ -358,11 +363,7 @@ def convert(value: torch.Tensor) -> torch.Tensor: ) finally: _parameter_transform.reset(token) - self._parameters, self._buffers, self._modules = ( - self._source._parameters, - self._source._buffers, - self._source._modules, - ) + self._bind_source(self._source) return self def requires_grad_(self, requires_grad: bool = True) -> ModuleHandle: @@ -409,12 +410,7 @@ def _captured_call(self, call: Callable[[], Any]) -> Any: ) ) try: - object.__setattr__(self, "_source", captured) - self._parameters, self._buffers, self._modules = ( - captured._parameters, - captured._buffers, - captured._modules, - ) + self._bind_source(captured) result = call() current_parameters = dict(captured.named_parameters()) if current_parameters.keys() != parameters.keys() or any( @@ -427,12 +423,7 @@ def _captured_call(self, call: Callable[[], Any]) -> Any: ) updated_buffers = dict(captured.named_buffers()) finally: - object.__setattr__(self, "_source", original) - self._parameters, self._buffers, self._modules = ( - original._parameters, - original._buffers, - original._modules, - ) + self._bind_source(original) _head_call.reset(token) self._publish(updated_buffers) return result From b9c12de637aa0709a3b4220a408cb3c4ba3ef550 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 1 Oct 2026 05:32:35 +0000 Subject: [PATCH 144/150] Inline parameter snapshot creation into TrainerRank --- src/art/trainer_rank/_impl.py | 11 ++++++++--- src/art/trainer_rank/_versions.py | 14 -------------- 2 files changed, 8 insertions(+), 17 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index d883aa6ca..db223d20b 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -2204,9 +2204,14 @@ def _snapshot_parameter( version: CheckpointVersion, max_gradient_staleness: int = 2, ) -> torch.nn.Parameter: - return self._version_state().snapshot( - parameter, version, max_gradient_staleness - ) + state = self._version_state() + state.validate(version, max_gradient_staleness) + with torch._C.DisableTorchFunctionSubclass(): + result = torch.nn.Parameter( + parameter.detach().clone(), requires_grad=parameter.requires_grad + ) + state.track(result, parameter, version, max_gradient_staleness) + return result def _commit_versioned_gradients( self, gradients: Sequence[VersionedGradient] diff --git a/src/art/trainer_rank/_versions.py b/src/art/trainer_rank/_versions.py index e6391b39d..8c7817015 100644 --- a/src/art/trainer_rank/_versions.py +++ b/src/art/trainer_rank/_versions.py @@ -107,20 +107,6 @@ def validate(self, version: CheckpointVersion, maximum: int = 2) -> None: f"{version.revision}, current revision {slot.revision})" ) - def snapshot( - self, - parameter: torch.nn.Parameter, - version: CheckpointVersion, - maximum: int = 2, - ) -> torch.nn.Parameter: - self.validate(version, maximum) - with torch._C.DisableTorchFunctionSubclass(): - result = torch.nn.Parameter( - parameter.detach().clone(), requires_grad=parameter.requires_grad - ) - self.track(result, parameter, version, maximum) - return result - def track( self, snapshot: torch.nn.Parameter, From 110651cb20d2271244486a99ebedd3ed18519928 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 1 Oct 2026 06:19:14 +0000 Subject: [PATCH 145/150] Use exported policy when creating live head handles --- src/art/trainer_rank/_heads.py | 13 ++----------- 1 file changed, 2 insertions(+), 11 deletions(-) diff --git a/src/art/trainer_rank/_heads.py b/src/art/trainer_rank/_heads.py index de9151348..d04d858a8 100644 --- a/src/art/trainer_rank/_heads.py +++ b/src/art/trainer_rank/_heads.py @@ -1009,16 +1009,10 @@ def __init__( state: HeadState, value: torch.nn.Module | torch.Tensor, collector: Any, - *, - max_gradient_staleness: int | None = None, ): self.state = state self.collector = collector - self.max_gradient_staleness = ( - state.max_gradient_staleness - if max_gradient_staleness is None - else max_gradient_staleness - ) + self.max_gradient_staleness = state.max_gradient_staleness self.pending = False self.invalid = False self.invalid_reason = "its checkpoint was replaced" @@ -1336,10 +1330,7 @@ def logical_register_head( if isinstance(value, torch.nn.Module) else value.to(view.device) ) - maximum = head_staleness(view._rank) - head = LiveHead( - state, value, view._executor.state.collector, max_gradient_staleness=maximum - ) + head = LiveHead(state, value, view._executor.state.collector) registry[(checkpoint, name)] = head return head.value From b5bcbd90704ec3127d094bf346d53fa663f4f461 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 1 Oct 2026 07:20:08 +0000 Subject: [PATCH 146/150] Share fixed-partition memory comparison arguments --- dev/trainer_v1_memory.py | 21 ++++----------------- 1 file changed, 4 insertions(+), 17 deletions(-) diff --git a/dev/trainer_v1_memory.py b/dev/trainer_v1_memory.py index cb7b36eaf..24f579b62 100644 --- a/dev/trainer_v1_memory.py +++ b/dev/trainer_v1_memory.py @@ -247,27 +247,14 @@ def fixed_split(requests, *, checkpoint, **_kwargs): # BF16 packed and split matmuls can differ numerically. Compare replay # against the identical admitted physical partition, with CPU reduction # on both arms; retain the packed reference as a separate diagnostic. - run( - "gpu", - "cpu", - "same_split_gpu", - chunks=accepted["telemetry"]["subforward_request_indices"], - compare="constrained_replay", - ) - run( - "cpu", - "cpu", - "same_split_cpu_offload", + split_comparison = dict( chunks=accepted["telemetry"]["subforward_request_indices"], compare="constrained_replay", ) + run("gpu", "cpu", "same_split_gpu", **split_comparison) + run("cpu", "cpu", "same_split_cpu_offload", **split_comparison) adaptive = run( - "auto", - "cpu", - "constrained_measured_auto", - cap=cap, - chunks=accepted["telemetry"]["subforward_request_indices"], - compare="constrained_replay", + "auto", "cpu", "constrained_measured_auto", cap=cap, **split_comparison ) evidence = adaptive["telemetry"]["fallback_costs"] if evidence["source"] != "measured_forward_and_transfers": From 3e608aa92dace657a251eff83f742f1939de12f0 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 1 Oct 2026 11:49:42 +0000 Subject: [PATCH 147/150] refactor(tests): consolidate callback lifecycle coverage --- .../test_trainer_rank_callback_cleanup.py | 158 ------------------ .../test_trainer_rank_callback_lifecycle.py | 149 ++++++++++++++++- 2 files changed, 147 insertions(+), 160 deletions(-) delete mode 100644 tests/unit/test_trainer_rank_callback_cleanup.py diff --git a/tests/unit/test_trainer_rank_callback_cleanup.py b/tests/unit/test_trainer_rank_callback_cleanup.py deleted file mode 100644 index 2b93a422b..000000000 --- a/tests/unit/test_trainer_rank_callback_cleanup.py +++ /dev/null @@ -1,158 +0,0 @@ -"""Independent DP2 x TP2 coverage of callback-boundary graph ownership.""" - -from __future__ import annotations - -import asyncio -import gc -from typing import Any - -import pytest -from test_trainer_rank_commands import _input, _loss_tree, _Rank -import torch -import torch.distributed as dist -import torch.multiprocessing as mp -from trainer_rank_test_support import gloo_group, megatron_topology - -from art.trainer_rank import run_rank_callback, run_rank_callback_stream - - -def _ownership(rank: Any) -> list[Any]: - gc.collect() - state = rank._rank_command_state - rows: list[Any] = [None] * dist.get_world_size() - dist.all_gather_object(rows, (tuple(state.graphs), tuple(state.released))) - return rows - - -def _empty(rank: Any, boundary: str) -> None: - rows = _ownership(rank) - assert all(not graphs and not released for graphs, released in rows), ( - boundary, - rows, - ) - - -async def _cases(rank: Any, physical: int) -> None: - inputs = [_input(2), _input(3)] - for mode, other in (("zero", "rank"), ("rank", "zero")): - leader = physical == 0 if mode == "zero" else physical % 2 == 0 - - async def call(callback, selected=mode): - return (await run_rank_callback(rank, callback, mode=selected)).value - - # The proxy dies after its creating command session has already STOPped. - returned = await call(lambda view: view.forward(inputs)) - assert any(graphs for graphs, _ in _ownership(rank)) - returned = None - gc.collect() - await call(lambda view: view.zero_grad(), other) - _empty(rank, f"{mode} -> {other} post-STOP release") - - # No later command is required to reclaim a callback-local result. - await call(lambda view: _loss_tree(view.forward(inputs)).item()) - _empty(rank, f"{mode} unused local output") - - # The common release boundary must not revoke genuinely live proxies. - returned = await call(lambda view: view.forward(inputs)) - await call(lambda view: view.zero_grad(), other) - assert any(graphs for graphs, _ in _ownership(rank)) - await call(lambda view: view.backward(_loss_tree(returned))) - expected = (2 if rank.dp == 0 else 3) if mode == "zero" else 5 - torch.testing.assert_close(rank.weight.grad, torch.tensor(float(expected))) - returned = None - gc.collect() - await call(lambda _view: None, other) - _empty(rank, f"{mode} retained output consumed in original mode") - - for ending in ("close", "error", "cancel", "live"): - - async def generate(view): - yield view.forward(inputs) - if ending == "error": - raise RuntimeError("cleanup stream failure") - if ending == "cancel": - task = asyncio.current_task() - assert task is not None - asyncio.get_running_loop().call_soon(task.cancel) - await asyncio.sleep(10) - - stream = run_rank_callback_stream(rank, generate, mode=mode) - yielded = await anext(stream) - if leader: - assert _loss_tree(yielded.value).item() == 10 - else: - assert yielded.logical_rank is None - kept = yielded.value if ending == "live" else None - del yielded - gc.collect() - if leader and ending == "error": - with pytest.raises(RuntimeError, match="cleanup stream failure"): - await anext(stream) - elif leader and ending == "cancel": - pending = asyncio.create_task(anext(stream)) - with pytest.raises(asyncio.CancelledError): - await pending - await stream.aclose() - if ending == "live": - # Closing the producing generator cannot revoke an output that - # its caller still holds, even after the other view runs. - await call(lambda view: view.zero_grad(), other) - assert any(graphs for graphs, _ in _ownership(rank)) - await call(lambda view: view.backward(_loss_tree(kept))) - torch.testing.assert_close( - rank.weight.grad, torch.tensor(float(expected)) - ) - kept = None - gc.collect() - await call(lambda _view: None, other) - if ending in ("error", "cancel"): - # Failure reaches the controller before all peers join cleanup; - # the next callback must finish that cleanup before admission. - await call(lambda _view: None, other) - _empty(rank, f"{mode} generator {ending}") - await call(lambda view: view.zero_grad(), other) - _empty(rank, f"{mode} generator {ending} followed by {other}") - - # Followers cancelled during a session must still join cleanup, and the - # next callback in the other mode must see an intact communicator. - entered = asyncio.Event() - zero_grad = rank.zero_grad - - def entered_zero_grad(): - zero_grad() - entered.set() - - rank.zero_grad = entered_zero_grad - - async def suspended(view): - unused = view.forward(inputs) - view.zero_grad() - del unused - gc.collect() - await asyncio.sleep(0.05) - - pending = asyncio.create_task(run_rank_callback(rank, suspended, mode=mode)) - await entered.wait() - if not leader: - pending.cancel() - (result,) = await asyncio.gather(pending, return_exceptions=True) - rank.zero_grad = zero_grad - if leader: - assert not isinstance(result, BaseException), result - else: - assert isinstance(result, asyncio.CancelledError), result - await call(lambda view: view.zero_grad(), other) - _empty(rank, f"{mode} cancellation followed by {other}") - - -def _worker(physical: int, rendezvous: str) -> None: - torch.set_num_threads(1) - with ( - gloo_group(physical, f"file://{rendezvous}", world_size=4, timeout=20), - megatron_topology(physical, dp_size=2, tp_size=2), - ): - asyncio.run(_cases(_Rank(physical // 2, 2), physical)) - - -def test_gloo_dp2_tp2_callback_cleanup_across_modes(tmp_path): - mp.spawn(_worker, args=(str(tmp_path / "cleanup-init"),), nprocs=4, join=True) diff --git a/tests/unit/test_trainer_rank_callback_lifecycle.py b/tests/unit/test_trainer_rank_callback_lifecycle.py index 0310e82f9..abc3ba1b4 100644 --- a/tests/unit/test_trainer_rank_callback_lifecycle.py +++ b/tests/unit/test_trainer_rank_callback_lifecycle.py @@ -1,11 +1,14 @@ -"""Checkpoint scopes and delayed cleanup stay within their callback session.""" +"""Callback sessions own checkpoint stacks, graph lifetimes, and delayed cleanup.""" + +from __future__ import annotations import asyncio import gc from typing import Any import pytest -from test_trainer_rank_commands import _input, _Rank +from test_trainer_rank_commands import _input, _loss_tree, _Rank +import torch import torch.distributed as dist import torch.multiprocessing as mp from trainer_rank_test_support import gloo_group, megatron_topology @@ -186,3 +189,145 @@ async def consume(): def test_delayed_callback_cleanup_and_checkpoint_scopes_leave_gloo_reusable(tmp_path): mp.spawn(_lifecycle_worker, args=(str(tmp_path / "lifecycle"),), nprocs=2) + + +def _ownership(rank: Any) -> list[Any]: + gc.collect() + state = rank._rank_command_state + rows: list[Any] = [None] * dist.get_world_size() + dist.all_gather_object(rows, (tuple(state.graphs), tuple(state.released))) + return rows + + +def _empty(rank: Any, boundary: str) -> None: + rows = _ownership(rank) + assert all(not graphs and not released for graphs, released in rows), ( + boundary, + rows, + ) + + +async def _cases(rank: Any, physical: int) -> None: + inputs = [_input(2), _input(3)] + for mode, other in (("zero", "rank"), ("rank", "zero")): + leader = physical == 0 if mode == "zero" else physical % 2 == 0 + + async def call(callback, selected=mode): + return (await run_rank_callback(rank, callback, mode=selected)).value + + # The proxy dies after its creating command session has already STOPped. + returned = await call(lambda view: view.forward(inputs)) + assert any(graphs for graphs, _ in _ownership(rank)) + returned = None + gc.collect() + await call(lambda view: view.zero_grad(), other) + _empty(rank, f"{mode} -> {other} post-STOP release") + + # No later command is required to reclaim a callback-local result. + await call(lambda view: _loss_tree(view.forward(inputs)).item()) + _empty(rank, f"{mode} unused local output") + + # The common release boundary must not revoke genuinely live proxies. + returned = await call(lambda view: view.forward(inputs)) + await call(lambda view: view.zero_grad(), other) + assert any(graphs for graphs, _ in _ownership(rank)) + await call(lambda view: view.backward(_loss_tree(returned))) + expected = (2 if rank.dp == 0 else 3) if mode == "zero" else 5 + torch.testing.assert_close(rank.weight.grad, torch.tensor(float(expected))) + returned = None + gc.collect() + await call(lambda _view: None, other) + _empty(rank, f"{mode} retained output consumed in original mode") + + for ending in ("close", "error", "cancel", "live"): + + async def generate(view): + yield view.forward(inputs) + if ending == "error": + raise RuntimeError("cleanup stream failure") + if ending == "cancel": + task = asyncio.current_task() + assert task is not None + asyncio.get_running_loop().call_soon(task.cancel) + await asyncio.sleep(10) + + stream = run_rank_callback_stream(rank, generate, mode=mode) + yielded = await anext(stream) + if leader: + assert _loss_tree(yielded.value).item() == 10 + else: + assert yielded.logical_rank is None + kept = yielded.value if ending == "live" else None + del yielded + gc.collect() + if leader and ending == "error": + with pytest.raises(RuntimeError, match="cleanup stream failure"): + await anext(stream) + elif leader and ending == "cancel": + pending = asyncio.create_task(anext(stream)) + with pytest.raises(asyncio.CancelledError): + await pending + await stream.aclose() + if ending == "live": + # Closing the producing generator cannot revoke an output that + # its caller still holds, even after the other view runs. + await call(lambda view: view.zero_grad(), other) + assert any(graphs for graphs, _ in _ownership(rank)) + await call(lambda view: view.backward(_loss_tree(kept))) + torch.testing.assert_close( + rank.weight.grad, torch.tensor(float(expected)) + ) + kept = None + gc.collect() + await call(lambda _view: None, other) + if ending in ("error", "cancel"): + # Failure reaches the controller before all peers join cleanup; + # the next callback must finish that cleanup before admission. + await call(lambda _view: None, other) + _empty(rank, f"{mode} generator {ending}") + await call(lambda view: view.zero_grad(), other) + _empty(rank, f"{mode} generator {ending} followed by {other}") + + # Followers cancelled during a session must still join cleanup, and the + # next callback in the other mode must see an intact communicator. + entered = asyncio.Event() + zero_grad = rank.zero_grad + + def entered_zero_grad(): + zero_grad() + entered.set() + + rank.zero_grad = entered_zero_grad + + async def suspended(view): + unused = view.forward(inputs) + view.zero_grad() + del unused + gc.collect() + await asyncio.sleep(0.05) + + pending = asyncio.create_task(run_rank_callback(rank, suspended, mode=mode)) + await entered.wait() + if not leader: + pending.cancel() + (result,) = await asyncio.gather(pending, return_exceptions=True) + rank.zero_grad = zero_grad + if leader: + assert not isinstance(result, BaseException), result + else: + assert isinstance(result, asyncio.CancelledError), result + await call(lambda view: view.zero_grad(), other) + _empty(rank, f"{mode} cancellation followed by {other}") + + +def _worker(physical: int, rendezvous: str) -> None: + torch.set_num_threads(1) + with ( + gloo_group(physical, f"file://{rendezvous}", world_size=4, timeout=20), + megatron_topology(physical, dp_size=2, tp_size=2), + ): + asyncio.run(_cases(_Rank(physical // 2, 2), physical)) + + +def test_gloo_dp2_tp2_callback_cleanup_across_modes(tmp_path): + mp.spawn(_worker, args=(str(tmp_path / "cleanup-init"),), nprocs=4, join=True) From c7b0c2ff9f0206bb90bf9d547237b018469cb53a Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 1 Oct 2026 12:09:15 +0000 Subject: [PATCH 148/150] refactor(tests): remove unused postponed annotations --- tests/integration/megatron/lora/test_trainer_v1_versions.py | 2 -- tests/unit/test_trainer_rank_graphs.py | 2 -- tests/unit/test_trainer_rank_parameter_hooks.py | 2 -- tests/unit/test_trainer_rank_parameter_no_grad.py | 2 -- tests/unit/test_trainer_rank_rng.py | 2 -- 5 files changed, 10 deletions(-) diff --git a/tests/integration/megatron/lora/test_trainer_v1_versions.py b/tests/integration/megatron/lora/test_trainer_v1_versions.py index 9e54fbb0d..9f3baef05 100644 --- a/tests/integration/megatron/lora/test_trainer_v1_versions.py +++ b/tests/integration/megatron/lora/test_trainer_v1_versions.py @@ -5,8 +5,6 @@ Full-model and distributed acceptance are additional gates, not implied here. """ -from __future__ import annotations - import json import pytest diff --git a/tests/unit/test_trainer_rank_graphs.py b/tests/unit/test_trainer_rank_graphs.py index 48143600b..41a8d6609 100644 --- a/tests/unit/test_trainer_rank_graphs.py +++ b/tests/unit/test_trainer_rank_graphs.py @@ -1,5 +1,3 @@ -from __future__ import annotations - import asyncio from contextlib import contextmanager, nullcontext from functools import partial diff --git a/tests/unit/test_trainer_rank_parameter_hooks.py b/tests/unit/test_trainer_rank_parameter_hooks.py index b7a4727dd..cd3c39817 100644 --- a/tests/unit/test_trainer_rank_parameter_hooks.py +++ b/tests/unit/test_trainer_rank_parameter_hooks.py @@ -1,5 +1,3 @@ -from __future__ import annotations - import gc import json diff --git a/tests/unit/test_trainer_rank_parameter_no_grad.py b/tests/unit/test_trainer_rank_parameter_no_grad.py index 4065f4226..4c7ccaaf4 100644 --- a/tests/unit/test_trainer_rank_parameter_no_grad.py +++ b/tests/unit/test_trainer_rank_parameter_no_grad.py @@ -1,5 +1,3 @@ -from __future__ import annotations - import asyncio import pytest diff --git a/tests/unit/test_trainer_rank_rng.py b/tests/unit/test_trainer_rank_rng.py index a3b2313fb..7275cae54 100644 --- a/tests/unit/test_trainer_rank_rng.py +++ b/tests/unit/test_trainer_rank_rng.py @@ -1,5 +1,3 @@ -from __future__ import annotations - import asyncio from contextlib import nullcontext from types import SimpleNamespace From 7a1db72c3202f1151858c56842dc70b2b21ac4ff Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 1 Oct 2026 12:44:22 +0000 Subject: [PATCH 149/150] Remove unused annotation imports from trainer controls --- tests/unit/test_trainer_batch_input_capture.py | 2 -- tests/unit/test_trainer_driver_transport.py | 2 -- tests/unit/test_trainer_rank_parameter_hooks.py | 3 +-- 3 files changed, 1 insertion(+), 6 deletions(-) diff --git a/tests/unit/test_trainer_batch_input_capture.py b/tests/unit/test_trainer_batch_input_capture.py index 349b44af4..b7876cfe4 100644 --- a/tests/unit/test_trainer_batch_input_capture.py +++ b/tests/unit/test_trainer_batch_input_capture.py @@ -1,7 +1,5 @@ """Batch iterators own submitted inputs while executing one wave at a time.""" -from __future__ import annotations - import asyncio from typing import Any, cast diff --git a/tests/unit/test_trainer_driver_transport.py b/tests/unit/test_trainer_driver_transport.py index 47f19f143..9ffa0d1fa 100644 --- a/tests/unit/test_trainer_driver_transport.py +++ b/tests/unit/test_trainer_driver_transport.py @@ -1,7 +1,5 @@ """Driver CPU transport stays separate from native output placement.""" -from __future__ import annotations - import asyncio from dataclasses import replace import gc diff --git a/tests/unit/test_trainer_rank_parameter_hooks.py b/tests/unit/test_trainer_rank_parameter_hooks.py index cd3c39817..1b0aeff28 100644 --- a/tests/unit/test_trainer_rank_parameter_hooks.py +++ b/tests/unit/test_trainer_rank_parameter_hooks.py @@ -101,8 +101,7 @@ def test_hook_removal_applies_to_retained_graph(client): @pytest.mark.parametrize("bad_result", (False, True)) def test_hook_failure_leaves_all_authoritative_gradients_unchanged(client, bad_result): trainer, native, parameter, _, collector, backward = setup(client) - rank = trainer - other = rank.parameter("q", lambda: torch.tensor(3.0), checkpoint="student") + other = trainer.parameter("q", lambda: torch.tensor(3.0), checkpoint="student") qlive = _live_head(trainer, "q", torch.tensor(3.0), collector) q = qlive.value if client else other native.grad, other.grad = torch.tensor(5.0), torch.tensor(6.0) From 75f414b2cde4c99e9a36cc01439e259d68ec116e Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 2 Oct 2026 23:47:39 +0000 Subject: [PATCH 150/150] test: account for cold recompute in retained graph admission --- tests/unit/test_trainer_rank_moe_memory.py | 17 +++++++++++++++-- 1 file changed, 15 insertions(+), 2 deletions(-) diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index 628a1ff99..7879363ea 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -746,15 +746,23 @@ def test_hybridep_high_water_needs_a_live_larger_graph( assert tuple(refs) == before and rank._pending_hybridep_graphs is refs +@pytest.mark.parametrize("profiled", [False, True], ids=["cold", "profiled"]) def test_hybridep_admission_ignores_consumed_graph_with_retained_sibling( hybrid_checkpoint_rank, + profiled, ): rank = hybrid_checkpoint_rank + signature = replace(_signature(), topology=(1, 1, 2, 1)) + if profiled: + rank._memory_profiles[signature] = _MemoryProfile( + bytes_per_token=1, packed_tokens=2 + ) + cold = 0 if profiled else COLD values = dict( packed_tokens=2, logical_tokens=2, output_bytes=8, - signature=replace(_signature(), topology=(1, 1, 2, 1)), + signature=signature, group_rows=((2, True),), ) baseline = rank._subforward_cost(**values) @@ -770,7 +778,12 @@ def test_hybridep_admission_ignores_consumed_graph_with_retained_sibling( (marker_ref,) = refs assert marker_ref() is not None and not marker_ref().item() live = rank._subforward_cost(**values) - assert live.checkpoint_workspace == 218752 * 2048 * 2 + retained, workspace = rank._checkpoint_memory_floor(values["group_rows"]) + assert workspace == 218752 * 2048 * 2 + assert live.checkpoint_adapter_gradient == 0 + # The cold allowance is separate from the live graph's dense-output extent. + assert live.checkpoint_workspace == workspace + cold + assert live.required == int((8 + 2 * retained + workspace + cold) * 1.1) assert not rank._memory_check_required(live.required).fits assert output.hidden_states is not None