diff --git a/.github/workflows/prek.yml b/.github/workflows/prek.yml index 924abc0cc..20a69f7af 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 \ @@ -224,10 +224,14 @@ 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 \ 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 \ tests/unit/test_trainer_rank_profile_warm.py \ @@ -249,6 +253,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 \ @@ -264,6 +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 \ --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 \ @@ -273,10 +279,13 @@ 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 \ --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/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 7a10a4c98..477e199be 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..50080e724 --- /dev/null +++ b/dev/trainer_v1_acceptance.py @@ -0,0 +1,476 @@ +"""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 sys +import time +import traceback + +import torch +import torch.distributed as dist +from trainer_rank_diag import rank0_checked +from trainer_v1_support import ( + _cached_backward, + _leaves, + build_rank, + load_checkpoint, + nccl_group, + source_commits, +) + +from art.trainer_rank import ForwardInput + + +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 _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 _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 + with nccl_group(device): + from megatron.core import parallel_state as ps + + if args.mode in ("oracle", "retained") and dist.get_world_size() != 1: + raise ValueError( + "The mathematical global-loss oracle requires one physical rank" + ) + physical = build_rank( + args.model, + layers=args.layers or None, + print_env=dist.get_rank() == 0, + ) + if args.mode == "control" and ps.get_data_parallel_world_size() != 1: + raise ValueError("A physical topology control requires DP=1") + checkpoint = load_checkpoint(physical, args.model) + 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 = physical.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_commits(), + "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) + + +if __name__ == "__main__": + main() diff --git a/dev/trainer_v1_benchmark.py b/dev/trainer_v1_benchmark.py new file mode 100644 index 000000000..eedd3564c --- /dev/null +++ b/dev/trainer_v1_benchmark.py @@ -0,0 +1,301 @@ +"""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 time +import weakref + +from dotenv import load_dotenv +import torch +import torch.distributed as dist +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 + + +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: + 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" + with nccl_group(): + if dist.get_world_size() != 1: + raise ValueError("This paired benchmark requires one physical rank") + 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": + 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_commits(), + "torch": torch.__version__, + "device": torch.cuda.get_device_name(), + "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( + 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}" + ) + + +if __name__ == "__main__": + main() diff --git a/dev/trainer_v1_checkpoint.py b/dev/trainer_v1_checkpoint.py new file mode 100644 index 000000000..a006dbcb0 --- /dev/null +++ b/dev/trainer_v1_checkpoint.py @@ -0,0 +1,30 @@ +"""Create a native checkpoint fixture for the remote public API canary.""" + +import argparse +import os +from pathlib import Path + +from dotenv import load_dotenv +from trainer_v1_support import build_rank, load_checkpoint, nccl_group + +from art.trainer_rank import validate_checkpoint + + +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" + 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) + + +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..9b1629253 --- /dev/null +++ b/dev/trainer_v1_facade_benchmark.py @@ -0,0 +1,219 @@ +"""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_v1_support import build_rank, load_checkpoint, nccl_group + +from art.trainer_rank import ForwardInput, ForwardOptions, _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") + 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( + world if axis == "DATA" else 1 + ) + os.environ["ART_TRAINER_RANK_TEST_HOOKS"] = "1" + os.environ["ART_TRAINER_RANK_TEST_ANCHOR"] = "no_sharing" + 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=() + ) + 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) + + +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..7cbdefdce --- /dev/null +++ b/dev/trainer_v1_head_memory.py @@ -0,0 +1,170 @@ +"""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 +from torch.multiprocessing.reductions import StorageWeakRef +from trainer_v1_support import build_rank, load_checkpoint, nccl_group + +from art.trainer_rank import ( + ForwardInput, + ForwardOptions, + 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) + with nccl_group(int(os.environ.get("LOCAL_RANK", "0"))): + os.environ["ART_TRAINER_RANK_TEST_HOOKS"] = "1" + rank = build_rank( + "Qwen/Qwen3-0.6B", + seed=716, + layers=2, + eval_mode=False, + model_initialization="random", + ) + checkpoint = load_checkpoint(rank, "Qwen/Qwen3-0.6B") + 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) + + +if __name__ == "__main__": + main() diff --git a/dev/trainer_v1_memory.py b/dev/trainer_v1_memory.py new file mode 100644 index 000000000..cb7b36eaf --- /dev/null +++ b/dev/trainer_v1_memory.py @@ -0,0 +1,361 @@ +"""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_v1_support import ( + _cached_backward, + _leaves, + build_rank, + deterministic_kernels, + load_checkpoint, + nccl_group, +) + +from art.trainer_rank import ( + ForwardInput, + ForwardOptions, + TrainerRankMemoryError, +) +from art.trainer_rank._impl import _SplitForwardPlan +from art.trainer_rank._memory_policy import host_memory_budget, placement_cost + + +def _backward(rank, outputs): + loss = sum( + output.hidden_states.float().mean() for output in _leaves(outputs) + ).square() + _cached_backward(rank, loss) + 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: + 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" + + def configure(provider): + provider.num_layers = args.layers + provider.recompute_granularity = None + provider.recompute_method = None + provider.recompute_num_layers = None + provider.recompute_modules = [] + + rank = build_rank( + args.model, + seed=913, + model_initialization="random", + provider_configure=configure, + ) + 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( + 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, + ) + ) + + +if __name__ == "__main__": + 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/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..939ffcd93 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)' )" @@ -10,7 +11,21 @@ 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 \ + 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 +36,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/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: 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 f3baead14..f02483b95 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 @@ -69,15 +71,41 @@ 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) _logger = logging.getLogger(__name__) @@ -130,15 +158,13 @@ def _collect_compile_garbage() -> None: 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: _collect_compile_garbage() - token = _CURRENT_LORA_SLOT.set(context) - try: + with use_lora_slot(context, version=version): result = function(*args, **kwargs) - finally: - _CURRENT_LORA_SLOT.reset(token) # A compile inside this call has unwound; collect before its backward. # Failed calls leave the mark for the next call, keeping their error. _collect_compile_garbage() @@ -157,7 +183,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) @@ -967,7 +993,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/model_support/handlers/gpt_oss.py b/src/art/megatron/model_support/handlers/gpt_oss.py index 320d223f7..d03823d0e 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, *, @@ -324,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") @@ -1272,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: @@ -1299,11 +1269,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"): @@ -1318,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: @@ -1337,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/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/megatron/weights/lora_publish.py b/src/art/megatron/weights/lora_publish.py index 02f5cb22f..235eea0dd 100644 --- a/src/art/megatron/weights/lora_publish.py +++ b/src/art/megatron/weights/lora_publish.py @@ -286,6 +286,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]]: @@ -647,21 +668,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/tinker/__init__.py b/src/art/tinker/__init__.py index a74a3d9a8..c1b959e2e 100644 --- a/src/art/tinker/__init__.py +++ b/src/art/tinker/__init__.py @@ -22,3 +22,7 @@ def __getattr__(name: str) -> Any: value = getattr(import_module(_EXPORTS[name], __name__), name) globals()[name] = value return value + + +def __dir__() -> list[str]: + return sorted(set(globals()) | set(__all__)) diff --git a/src/art/trainer_rank/__init__.py b/src/art/trainer_rank/__init__.py index 04b78329e..825025223 100644 --- a/src/art/trainer_rank/__init__.py +++ b/src/art/trainer_rank/__init__.py @@ -1,14 +1,21 @@ 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, Sequence +from typing import Literal, TypeVar import torch import torch.distributed as dist 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 @@ -33,16 +40,16 @@ MaterializedCheckpoint = _impl.MaterializedCheckpoint PushedCheckpoint = _impl.PushedCheckpoint -if TYPE_CHECKING: - from art.megatron.train import TrainingRuntime - ModuleT = TypeVar("ModuleT", bound=torch.nn.Module) for _public_type in ( AdamParams, ForwardInput, + ForwardOptions, ForwardOutput, + ImportanceSamplingGradientCorrection, + ResolvedForwardOptions, MicroBatch, MicroBatchStats, TopK, @@ -60,8 +67,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,346 +76,52 @@ class TrainerRank(_impl.TrainerRank): profile is keyed by topology and calibrates itself online. """ - def __init__(self, runtime: TrainingRuntime) -> None: - super().__init__(runtime) - @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, 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( - 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_micro_batches( - self, - inputs: Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]], - *, - 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_micro_batches( - self, - inputs: Iterable[ - Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]] - ], - *, - 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_micro_batches( - self, - inputs: Iterable[ - Iterable[Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]]] - ], - *, - 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_micro_batches( - self, - inputs: Iterable[ - Iterable[ - Iterable[ - Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]] - ] - ] - ], - *, - 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_micro_batches( - self, - inputs: Iterable[ForwardInputs], - *, - 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. 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. - - 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. - - Admission learns each wave's memory peak, including the caller's loss - and backward when they run inside the yield; a backward deferred past the - yield is not learned. A wave with another TrainerRank forward inside its - yield cannot lower later estimates. - """ - forward = cast( - Callable[..., Iterator[MicroBatch[ForwardInputs, ForwardOutputs]]], - super().forward_micro_batches, - ) - return forward( - inputs, checkpoint=checkpoint, no_grad=no_grad, yield_empty=yield_empty - ) - - @overload - def dp_rank_forward( - self, - inputs: Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]], - *, - checkpoint: AdapterSelection = Unset, - no_grad: bool | None = None, - ) -> Sequence[ForwardOutput[LogprobsT, TopKT, LogitsT, HiddenStatesT]]: ... - - @overload - def dp_rank_forward( - self, - inputs: Iterable[ - Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]] - ], - *, - checkpoint: AdapterSelection = Unset, - no_grad: bool | None = None, - ) -> Sequence[ - Sequence[ForwardOutput[LogprobsT, TopKT, LogitsT, HiddenStatesT]] - ]: ... - - @overload - def dp_rank_forward( - self, - inputs: Iterable[ - Iterable[Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]]] - ], - *, - checkpoint: AdapterSelection = Unset, - no_grad: bool | None = None, - ) -> Sequence[ - Sequence[Sequence[ForwardOutput[LogprobsT, TopKT, LogitsT, HiddenStatesT]]] - ]: ... - - @overload - def dp_rank_forward( - self, - inputs: Iterable[ - Iterable[ - Iterable[ - Iterable[ForwardInput[LogprobsT, TopKT, LogitsT, HiddenStatesT]] - ] - ] - ], - *, - checkpoint: AdapterSelection = Unset, - no_grad: bool | None = None, - ) -> Sequence[ - Sequence[ - Sequence[Sequence[ForwardOutput[LogprobsT, TopKT, LogitsT, HiddenStatesT]]] - ] - ]: ... - - def dp_rank_forward( - self, - inputs: ForwardInputs, - *, - 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`. - - 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().dp_rank_forward, - ) - return forward(inputs, checkpoint=checkpoint, no_grad=no_grad) - - def dp_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) - - 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. - - 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. - """ - return super().optim_step( - params=params, - scale_grads=scale_grads, - checkpoints=checkpoints, - on_live_graphs=on_live_graphs, - ) +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 f15ad626c..977db170b 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,87 @@ 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. + + Unregistered allocations, unusually large serialization metadata, + payload-expanding hooks and concurrent 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 + ] + 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) + 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: + (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()) + 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) + 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( + 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: @@ -970,6 +1067,14 @@ 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 + + _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)): raise trainer._slot_state_error( @@ -1010,6 +1115,14 @@ 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: + 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: raise_distributed(error, "prepare checkpoint", group) if any(value != optimizer for value in _gather(optimizer, group)): @@ -1076,7 +1189,6 @@ def start_writer() -> Future[None] | BaseException: 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( @@ -1210,13 +1322,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: @@ -1410,7 +1516,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: @@ -1425,38 +1530,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, @@ -1492,18 +1565,26 @@ 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 = ( - _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: @@ -1591,7 +1672,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) @@ -1599,7 +1679,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: @@ -1630,13 +1709,13 @@ 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: + 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, @@ -1644,9 +1723,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) ) @@ -1669,6 +1746,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 ( @@ -1703,6 +1796,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, ) @@ -1717,6 +1811,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(): @@ -1757,6 +1852,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 @@ -1889,7 +1985,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}") @@ -1944,27 +2042,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, @@ -2020,6 +2097,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( @@ -2038,26 +2116,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 @@ -2079,12 +2157,24 @@ 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 _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 new file mode 100644 index 000000000..919494115 --- /dev/null +++ b/src/art/trainer_rank/_commands.py @@ -0,0 +1,1264 @@ +"""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, contextmanager, nullcontext +from dataclasses import dataclass, field, replace +from functools import partial +import inspect +from typing import TYPE_CHECKING, Any, Literal, TypeVar, cast +import weakref + +import cloudpickle +import torch +import torch.distributed as dist + +from . import _impl, _transport + +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: + 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) +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: + return _transport.encode(command) + + +@dataclass(frozen=True) +class _OutputPacket: + packet: Any + cpu: tuple[bool, ...] + managed: bool + + +@dataclass +class _Release: + completed: asyncio.Future[None] + gathered: list[Any] + finish: Callable[[_Release], None] + + +@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 + ) + + +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 + + 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") + 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: + 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, 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)) + + 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) + 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}" + ) + release.completed.get_loop().call_exception_handler( + {"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() + 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: + decoded = _transport.decode(payload) + 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: + deferred: BaseException | None = 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. + 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 deferred is not None: + raise deferred + return + try: + self._execute(command) + except BaseException as error: + # The leader receives the same coordinated error and chooses + # whether to catch it, continue, or stop the callback. + # 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() + + 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: + 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): + 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"): + 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: + try: + 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. + 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) + 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) + finally: + value = payload = values = None + + 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 + + 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 + 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", "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) + 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) + 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, + local_lookup=self.mode == "rank", + **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") + ) + 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] = [] + 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) + 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_failed_iterator("batches_close", 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 _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, + kwargs: dict[str, Any], + handles: list[str] | None = None, + ) -> Iterator[_impl.MicroBatch]: + identifier = self._invoke("batches", items, **kwargs) + delivery_failed = False + try: + while True: + try: + batch = self._combine_wave( + self._invoke("next", identifier), items, handles + ) + except BaseException: + delivery_failed = True + raise + 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. + if not self._executor.stopped: + 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 + ) -> _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): + 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()) + placements = iter( + choose_output_placements(costs, gpu_available_bytes=available) + ) + 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) + ) + + +@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: + 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) + primary: BaseException | None = None + 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) + except BaseException as error: + primary = error + raise + finally: + with _preserve_callback_error(primary): + 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 + primary: BaseException | None = 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 + exhausted = ( + StopAsyncIteration if inspect.isasyncgen(iterator) else StopIteration + ) + while True: + try: + value = ( + await iterator.asend(sent) + if inspect.isasyncgen(iterator) + else iterator.send(sent) + ) + except exhausted: + return + sent = yield RankCallbackResult( + 0 if mode == "zero" else executor.dp_rank, value + ) + except GeneratorExit: + raise + except BaseException as error: + primary = error + raise + finally: + 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/_corrections.py b/src/art/trainer_rank/_corrections.py new file mode 100644 index 000000000..174cb0fae --- /dev/null +++ b/src/art/trainer_rank/_corrections.py @@ -0,0 +1,364 @@ +"""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. + """ + 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) + ) + finally: + del original_logprobs, current_logprobs + log_ratio = None + + +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. + """ + try: + 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 + finally: + del cotangent, original_logprobs, current_logprobs + del original_tokens, current_tokens + active = selected = weights = corrected = None + + +@dataclass(frozen=True) +class _OutputCorrection: + index: int + 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.""" + 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 + ) + 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) + 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, + 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, + 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] = {} + 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), + ): + 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, + ) + 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/src/art/trainer_rank/_graphs.py b/src/art/trainer_rank/_graphs.py new file mode 100644 index 000000000..6f3657a3f --- /dev/null +++ b/src/art/trainer_rank/_graphs.py @@ -0,0 +1,695 @@ +"""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 +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 + +from ._rng import RNGState as _RNGState +from ._tensors import _map_tensor_arguments + +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 + }, + ) + return _map_tensor_arguments(_restore_inputs, 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 _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, 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 = storage.nbytes() + 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": + # Weak storage identity prevents allocator address reuse from + # confusing distinct activations without pinning CUDA storage. + key = self.source + if key not in copies: + copies[key] = self.transfer_stats.copy(self, "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: + self.restored[key] = self.transfer_stats.copy(self, 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 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 (): + if (cell := reference()) is not 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 + self.validate_backward = self.current_context_factory = None + self.is_stale = self.keep_on_device = None + + 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, + ) + saved.append(weakref.ref(cell)) + del tensor + if retention == "cpu": + cell.offload(copies) + 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) + # These frame aliases otherwise outlive record.release(). + del self + context_factory = keep_on_device = None + cell = None + 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 + 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() + .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": + 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.release_physical() + record.retention = "replay" + + def release(self, handle: ForwardHandle) -> None: + if (record := self._records.pop(handle, None)) is not None: + record.release() + + 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: + 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. + for handle, _ in packets: + self.release(handle) + raise + + @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. + 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) + ) + finally: + current = corrected = None + + @staticmethod + def _prepare_backward(record, gradients, stale, prepared): + 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 + ) + 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 + ) + 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/src/art/trainer_rank/_heads.py b/src/art/trainer_rank/_heads.py new file mode 100644 index 000000000..af3eea98c --- /dev/null +++ b/src/art/trainer_rank/_heads.py @@ -0,0 +1,1385 @@ +"""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 torch + +from ._tensors import _map_tensor_arguments, _map_tensors + +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 + 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", + } +) +_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: + """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 _TENSOR_MUTATION_DUNDERS + 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: + 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 + + 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, tracker: _CustomTensorTracker, module: torch.nn.Module): + 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( + custom: _CustomObject, tracker: _CustomTensorTracker +) -> ModuleHandle: + state = _NativeModuleState(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 | None + factory_error: str | None = None + + +@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 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 + 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 _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 = ( + _head_tensors(custom.kind, custom.value, parameters=True), + _head_tensors(custom.kind, custom.value, parameters=False), + ) + 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, + 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 + 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 + 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: + _validate_registration(trainer, existing, registration) + trainer._custom_object( + registration.name, + registration.kind, + lambda: deepcopy(registration.value), + checkpoint=checkpoint, + ) + return _observe_head( + trainer, export_head(trainer, checkpoint, registration.name) + ) + if kind == "head_export": + return tuple( + _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": + 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 = _head_tensors(custom.kind, custom.value, parameters=False) + 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 = _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): + 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"], + cast(torch.nn.Parameter, 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)] + return _map_tensor_arguments(replace, 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 _CLIENT_BUFFER_MUTATIONS + ): + 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)] + return _map_tensor_arguments(replace, 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]: + return _head_tensors(self.state.kind, self.source, parameters=True) + + def buffers(self) -> dict[str, torch.Tensor]: + return _head_tensors(self.state.kind, self.source, parameters=False) + + 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 = _head_tensors(custom.kind, custom.value, parameters=False) + 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 _head_tensors( + custom.kind, custom.value, parameters=False + ).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 = _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 + 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)) + 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: + return current.value + if state is None: + 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) + 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, observation in zip(keys, states, strict=True): + state = observation.state + 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 a25505c53..03a0f7e4e 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -6,17 +6,19 @@ from collections import OrderedDict from collections.abc import ( Callable, + Generator, Iterable, Iterator, Mapping, Sequence, ) from concurrent.futures import Future, ThreadPoolExecutor -from contextlib import contextmanager +from contextlib import contextmanager, nullcontext from contextvars import ContextVar -from dataclasses import dataclass, replace +from copy import deepcopy +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 logging import math import os @@ -25,7 +27,7 @@ import threading import time import traceback -from types import TracebackType +from types import MethodType, TracebackType from typing import ( TYPE_CHECKING, Any, @@ -50,28 +52,42 @@ from art.megatron.prefix_tree_packing import ( PrefixTreePack, _local_position_pairs, - estimate_prefix_tree_packed_tokens, # noqa: F401 # read as ``_impl.X`` by _memory / _micro_batch_planner -) -from art.trainer_rank import ( # noqa: F401 # ``_gdn_memory`` is read as ``_impl.X`` by _memory / _micro_batch_planner - _gdn_memory, - _planner_evidence, - _planner_misses, + estimate_prefix_tree_packed_tokens, ) +from art.trainer_rank import _gdn_memory, _planner_evidence, _planner_misses from art.trainer_rank._backward_work import BackwardWork from art.trainer_rank._backward_work import region as _backward_region +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 ( ModelGeometry, ParallelShape, select_scoring, ) -from art.trainer_rank._prefix_tree_materializer import ( # noqa: F401 # read as ``_impl.X`` by _memory / _micro_batch_planner - materialize_prefix_tree_layout, -) +from art.trainer_rank._prefix_tree_materializer import materialize_prefix_tree_layout from art.trainer_rank._prefix_tree_planner import ( CanonicalPrefixTree, PrefixTreeLayout, ) +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, + CheckpointVersions, + VersionedGradient, +) if TYPE_CHECKING: from megatron.core.models.gpt.gpt_model import GPTModel @@ -81,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 ( @@ -92,7 +108,9 @@ _PreparedSave, _SnapshotSpill, ) - from art.trainer_rank._lora_export import _PreparedLoraExport + from art.trainer_rank._lora_export import _VllmLoraPublishInputs + + from ._heads import ModuleHandle @dataclass(frozen=True) @@ -173,11 +191,6 @@ class _AdapterConfig(TypedDict): hidden_size: NotRequired[int] -class _Unset: - pass - - -Unset = _Unset() type AdapterSelection = str | None | _Unset @@ -207,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__( @@ -219,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 @@ -232,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 @@ -245,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 @@ -258,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 @@ -271,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 @@ -284,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 @@ -297,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 @@ -310,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 @@ -323,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 @@ -336,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 @@ -349,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 @@ -362,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 @@ -375,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 @@ -388,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 @@ -401,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 @@ -414,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 @@ -427,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__( @@ -439,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) @@ -452,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 @@ -460,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") @@ -592,6 +629,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 sample: _planner_evidence.MemorySample | None = dataclass_field( default=None, compare=False, repr=False ) @@ -637,23 +678,6 @@ class _CandidateMicroBatch(Generic[ForwardInputsT]): fallback: _CandidateMicroBatch[ForwardInputsT] | None = None -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( @@ -709,6 +733,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] @@ -716,6 +751,7 @@ class _CustomTensorTracker: name: str generation: object active: bool = False + buffer_revision: int = 0 def validate(self) -> TrainerRank: trainer = self.trainer() @@ -803,6 +839,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(): @@ -852,6 +900,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 @@ -863,6 +912,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) @@ -949,6 +999,7 @@ class _MemorySignature: request_mix: tuple[str, ...] grad_enabled: bool grad_modes: tuple[bool, ...] + memory_placement: tuple[tuple[str, str], ...] = () slot_shapes: tuple[tuple[bool, tuple[tuple[int, ...], ...]], ...] = () # Short single-target requests keep the logical extrapolation, but share # the profile learned from longer requests of the same signature. @@ -962,6 +1013,7 @@ class _ForwardGroupPlan: request_indices: tuple[int, ...] items: tuple[_ForwardItem, ...] packed: PrefixTreePack + memory_placement: MemoryPlacement | None = None layout: PrefixTreeLayout | None = None @@ -1068,7 +1120,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." ) @@ -1086,7 +1138,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, @@ -1850,10 +1908,12 @@ def _moe_output_bytes_per_token( class TrainerRank: - def __init__(self, runtime: TrainingRuntime) -> None: - options = _planner_misses.parse_options(os.environ) - self._allow_oversized_batches = options.allow_oversized_batches - self._planner_reporter = _planner_misses.Reporter(options.threshold_pct) + def __init__( + self, runtime: TrainingRuntime, *, options: ForwardOptions | None = None + ) -> None: + planner_options = _planner_misses.parse_options(os.environ) + self._allow_oversized_batches = planner_options.allow_oversized_batches + self._planner_reporter = _planner_misses.Reporter(planner_options.threshold_pct) self._planner_observation_context: ContextVar[dict[str, Any] | None] = ( ContextVar("trainer_rank_planner_observation", default=None) ) @@ -1877,8 +1937,11 @@ 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._rng = TrainerRNG(self.device) self._param_dtype_size = _dtype_size(next(runtime.model[0].parameters()).dtype) try: metadata_model = _language_model(runtime.model[0]) @@ -2019,7 +2082,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() @@ -2035,7 +2098,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] = {} @@ -2048,6 +2110,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() + ) # Tracked peak-counter resets, and the latest (resets, peak) reading. self._peak_resets = 0 self._peak_reading: tuple[int, int] | None = None @@ -2083,22 +2148,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, @@ -2107,7 +2328,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) @@ -2118,7 +2342,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) @@ -2156,6 +2383,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 @@ -2170,7 +2398,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() @@ -2178,7 +2406,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": @@ -2222,18 +2452,45 @@ 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 _initialize_custom_object( self, @@ -2393,10 +2650,11 @@ def _slot_state_error(message: str) -> TrainerRankSlotStateError: return TrainerRankSlotStateError(message) @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, @@ -2408,12 +2666,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, @@ -2425,12 +2684,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, @@ -2442,7 +2702,7 @@ def forward_micro_batches( ]: ... @overload - def forward_micro_batches( + def forward_batches( self, inputs: Iterable[ Iterable[ @@ -2452,6 +2712,7 @@ def forward_micro_batches( ] ], *, + options: ForwardOptions | None = None, checkpoint: AdapterSelection = Unset, no_grad: bool | None = None, yield_empty: bool = False, @@ -2470,15 +2731,49 @@ 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, ) -> Iterator[MicroBatch[ForwardInputs, ForwardOutputs]]: - """Yield admitted micro-batches; the caller runs its loss and backward. + """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. + + 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 + 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. Admission learns each call's whole peak, including the caller's loss and backward. For grad-enabled single-target requests of at least 64 @@ -2494,20 +2789,35 @@ 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") - with torch.set_grad_enabled(enabled): + self._guard_forward_collective("forward_batches") + 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] @@ -2525,11 +2835,40 @@ 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." @@ -2571,21 +2910,33 @@ def _release_cached_memory_for_backward( ) @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[ @@ -2593,12 +2944,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[ @@ -2606,7 +2958,7 @@ def dp_rank_forward( ]: ... @overload - def dp_rank_forward( + def forward( self, inputs: Iterable[ Iterable[ @@ -2616,6 +2968,7 @@ def dp_rank_forward( ] ], *, + options: ForwardOptions | None = None, checkpoint: AdapterSelection = Unset, no_grad: bool | None = None, ) -> Sequence[ @@ -2624,31 +2977,63 @@ 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") + """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") backward = self._backward_work() if backward is not None: 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() - materialized = _materialize(inputs) + # 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)) - plan, check = self._plan_admissible_forward( - requests, checkpoint=checkpoint, context="dp_rank_forward" - ) - tracked_outputs = self._execute_admitted_plan( - plan, check=check, context="dp_rank_forward" + error: BaseException | None = None + 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)) + except BaseException as exc: + error = exc + # 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 backward is not None: - backward.attach(tracked_outputs) - return _unflatten(materialized, iter(tracked_outputs)) + if error is not None: + raise error + return outputs def _execute_admitted_plan( self, plan: _AnyForwardPlan, *, check: _MemoryCheck, context: str @@ -2665,6 +3050,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 @@ -2672,40 +3100,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: @@ -2971,7 +3410,7 @@ 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 @@ -2986,13 +3425,14 @@ 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") + """Reduce in place over data-parallel batches, excluding TP/CP replicas.""" + self._guard_forward_collective("reduce") from megatron.core import parallel_state as ps # Public outputs are CP-replicated; internal shard reductions still include CP. @@ -3002,6 +3442,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 _compact_lora_slot_keys(self) -> None: from art.megatron.lora import LoRA @@ -3175,9 +3645,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}." ) @@ -3205,7 +3675,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 @@ -3221,10 +3693,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() + ) @_backward_region def _run_flat_plan_with_memory_tracking( @@ -3279,6 +3755,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 @@ -3390,38 +3867,240 @@ 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): - 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 - ) - 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]) 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 + # 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: + 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 = 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: + 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) + except BaseException: + # No caller owns a failed handoff; a partial bridge may also release. + 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( self, request: AnyForwardInput, @@ -3435,7 +4114,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 ) @@ -3450,7 +4129,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 _topology_key(self) -> tuple[int, int, int, int]: @@ -3477,6 +4156,459 @@ 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), + slot_groups=tuple( + (g.slot_ref, 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 _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] + ) + @contextmanager def _cache_recovery_episode( self, *, error: BaseException | None = None @@ -3568,7 +4700,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: @@ -3642,7 +4774,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. @@ -3650,7 +4785,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: @@ -4690,9 +5825,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 @@ -4722,11 +5854,7 @@ 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_micro_batches = _micro_batch_planner._forward_micro_batches + _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 _admit_split_rung = _micro_batch_planner._admit_split_rung @@ -4763,9 +5891,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 @@ -4798,9 +5923,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 @@ -5006,6 +6128,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() @@ -5030,6 +6160,18 @@ 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, + ) + from ._tensors import _map_tensor_arguments + + if (captured := head_call_arguments(args, kwargs)) is not None: + return func(*captured[0], **captured[1]) del types tracked = tuple( value @@ -5039,9 +6181,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__", @@ -5049,10 +6191,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)) @@ -5062,42 +6226,70 @@ 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 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)), **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( @@ -5131,36 +6323,25 @@ 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 - return custom + getattr(child, f"_{kind}s")[key] = replacements[identity] + from ._heads import native_module_handle + + tracker.validate() + return replace(custom, handle=native_module_handle(custom, tracker)) def _custom_layout( @@ -5316,11 +6497,11 @@ 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 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 " - f"finite step (invalid={invalid}, step={state.step})." + f"finite integer step (invalid={invalid}, step={state.step})." ) @@ -5559,10 +6740,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( @@ -5580,7 +6808,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]: @@ -5600,7 +6840,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/_lora_export.py b/src/art/trainer_rank/_lora_export.py index 0bb01707e..9e97fb7e9 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 @@ -106,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( @@ -131,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, @@ -161,14 +134,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), @@ -215,6 +186,25 @@ 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, + ) + + sizes = _gpt_oss_padding_sizes_from_adapter_config(inputs.adapter_config) + _, _, 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") return inputs.handler.to_vllm_lora_tensors( merged_tensors, adapter_config=inputs.adapter_config, @@ -226,7 +216,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] = {} @@ -247,14 +237,29 @@ 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, ), "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() @@ -263,7 +268,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( @@ -276,7 +281,7 @@ def stage() -> _PreparedLoraExport | None: ), ) stager.finish() - return _PreparedLoraExport(staged) + return staged return None prepared = _checkpoint._phase(stage, "stage LoRA publish tensors", group) @@ -285,14 +290,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) @@ -311,7 +314,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 = ( @@ -360,7 +363,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: @@ -374,7 +377,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/src/art/trainer_rank/_memory.py b/src/art/trainer_rank/_memory.py index 70e3c9250..60b6fba1c 100644 --- a/src/art/trainer_rank/_memory.py +++ b/src/art/trainer_rank/_memory.py @@ -110,6 +110,7 @@ def feed(value: Any) -> None: signature.request_mix, signature.grad_enabled, signature.grad_modes, + signature.memory_placement, signature.slot_shapes, signature.short_requests, p.packed_tokens, @@ -583,7 +584,9 @@ def _checkpoint_memory_floor( and self._moe_memory_supported ): hybrid_rows = max(rows for rows, _ in group_rows) - if any(ref() is not None for ref in self._pending_hybridep_graphs): + if any( + _impl._graph_marker_is_live(ref) for ref in self._pending_hybridep_graphs + ): hybrid_rows = max(hybrid_rows, self._hybridep_rows_high_water) if hybrid_rows is not None: workspace = max(workspace, -(-hybrid_rows // 4) * 4 * self._hidden_size * 2) @@ -967,9 +970,16 @@ def _refresh_memory_check( # The existing admission operand can already be a cross-rank maximum. # Preserve its local producer separately; never price the plan again. with decision.refresh_of(check.sample) if decision is not None else nullcontext(): - return self._memory_check_required( + refreshed = self._memory_check_required( check.estimated_required_bytes, sync_across_dp=sync_across_dp ) + return _impl.replace( + check, + estimated_required_bytes=refreshed.estimated_required_bytes, + available_bytes=refreshed.available_bytes, + fits=refreshed.fits and check.cpu_fits, + sample=refreshed.sample, + ) def _memory_check_required( diff --git a/src/art/trainer_rank/_memory_policy.py b/src/art/trainer_rank/_memory_policy.py new file mode 100644 index 000000000..e7809f22b --- /dev/null +++ b/src/art/trainer_rank/_memory_policy.py @@ -0,0 +1,304 @@ +"""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), + ) diff --git a/src/art/trainer_rank/_micro_batch_planner.py b/src/art/trainer_rank/_micro_batch_planner.py index 2a077b57d..c84f0b3f0 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 @@ -72,7 +60,7 @@ ) -def _forward_micro_batches( +def _forward_batches( self: TrainerRank, inputs: Iterable[ForwardInputs], *, @@ -112,7 +100,7 @@ def _forward_micro_batches( self._run_flat_plan_with_memory_tracking( candidate.plan, check=candidate.check, - context="forward_micro_batches", + context="forward_batches", ) ) # This wave's peak interval, which its caller phase continues. @@ -122,7 +110,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) @@ -137,8 +125,6 @@ def _forward_micro_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( @@ -252,7 +238,10 @@ def _find_admissible_forward( if ensure_slots: self._ensure_checkpoint_slots_for(requests, checkpoint=checkpoint) 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 = (plan, check) @@ -261,7 +250,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 if check.estimated_required_bytes < best[1].estimated_required_bytes: @@ -361,20 +353,22 @@ 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) + managed = self._graph_memory_policy_enabled() if keep_rejected is None: keep_rejected = getattr(self, "_allow_oversized_batches", False) - if not check.fits and not keep_rejected: - return None, check - best: tuple[_SplitForwardPlan | None, _MemoryCheck] = (None, check) + 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 and not keep_rejected: + return None, check + best: tuple[_impl._SplitForwardPlan, _MemoryCheck] | None = None for memory_minimal in (False, True): plans = [ self._plan_flat_forward( @@ -399,14 +393,18 @@ 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 if ( - best[0] is None + best is None or check.estimated_required_bytes < best[1].estimated_required_bytes ): best = (split, check) + assert best is not None return best if keep_rejected else (None, check) @@ -610,6 +608,8 @@ 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 _select_next_micro_batch( @@ -619,7 +619,9 @@ def _select_next_micro_batch( *, checkpoint: AdapterSelection = _impl.Unset, ) -> _CandidateMicroBatch[ForwardInputsT]: - def admit(refusal: _ForwardRefusal) -> _CandidateMicroBatch[ForwardInputsT]: + def admit( + refusal: _impl._ForwardRefusal, + ) -> _impl._CandidateMicroBatch[ForwardInputsT]: dp_rank, dp_size = self._dp_rank_and_size() width = min(len(items) - start, dp_size) indices = _impl._local_wave_indices(start, width, dp_rank, dp_size) @@ -637,7 +639,7 @@ def admit(refusal: _ForwardRefusal) -> _CandidateMicroBatch[ForwardInputsT]: 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, admit_refusal=admit, ) @@ -663,7 +665,7 @@ def local_slice(width: int) -> tuple[tuple[int, ...], list[ForwardInputsT]]: return indices, [items[index] for index in indices] estimates: dict[int, tuple[_MemoryCheck, bool, bool] | None] = {} - plans: dict[int, _FlatForwardPlan] = {} + plans: dict[int, _impl._FlatForwardPlan] = {} checked_plans: dict[int, _MemoryCheck] = {} # Per-width layout mode chosen by admission: False = cost-optimal, # True = memory-minimal (full sharing). Materialization must build the @@ -803,7 +805,7 @@ def fits(width: int) -> tuple[bool, bool]: # admit on the materialized plan, trying the cost-optimal # layouts first and the memory-minimal layouts if those do not # fit or fall outside the profile's trust window. - def price(plan: _FlatForwardPlan) -> tuple[_MemoryCheck, bool, bool]: + def price(plan: _impl._FlatForwardPlan) -> tuple[_MemoryCheck, bool, bool]: check = self._memory_check( plan, sync_across_dp=True, sync_planning_errors=True ) @@ -831,7 +833,7 @@ def price(plan: _FlatForwardPlan) -> tuple[_MemoryCheck, bool, bool]: rejected_widths.add(width) return check.fits and (trusted or not profiled), trusted - def materialize(width: int) -> _FlatForwardPlan: + def materialize(width: int) -> _impl._FlatForwardPlan: width = normalize(width) plan = plans.get(width) if plan is None: @@ -845,7 +847,7 @@ def materialize(width: int) -> _FlatForwardPlan: plans[width] = plan return plan - def candidate(width: int) -> _CandidateMicroBatch[ForwardInputsT]: + def candidate(width: int) -> _impl._CandidateMicroBatch[ForwardInputsT]: width = normalize(width) indices, local_inputs = local_slice(width) plan = materialize(width) @@ -857,6 +859,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, @@ -1402,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), @@ -1434,6 +1443,12 @@ def _fill_planner_snapshot( self._split_required_memory(costs), int(floor * _impl._MEMORY_SAFETY_FACTOR), ) + if any(group.memory_placement is not None for group in plan.groups): + # Placement also prices version snapshots and outstanding graphs; + # the unplaced model estimate cannot recreate that admission. + if check.sample is None or check.sample.local_required_bytes is None: + raise ValueError("graph placement admission sample unavailable") + local_required = check.sample.local_required_bytes rank_fields = { name: getattr(self, "_" + name) for name in ( @@ -1455,7 +1470,7 @@ def _fill_planner_snapshot( rank_fields["geometry"] = asdict(self._geometry) rank_fields["topology"] = list(plan.signature.topology) - def version(tensor: torch.Tensor | None) -> int | None: + def version(tensor: _impl.torch.Tensor | None) -> int | None: try: return None if tensor is None else tensor._version except RuntimeError: @@ -1498,6 +1513,9 @@ def capture(tensor: torch.Tensor | None) -> InputSnapshot: "hidden_states": item.request.hidden_states, "no_grad": item.request.no_grad, "checkpoint": str(item.request.checkpoint), + "options": asdict( + _impl._resolved_request_policy(item.request.options) + ), }, ) for group in plan.groups @@ -1710,11 +1728,11 @@ 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: - dp_rank_forward's return or forward_micro_batches' iterator resume. + forward's return or forward_batches' iterator resume. Caladan calls this at execution end. Direct callers can report a caught caller OOM before cleanup, then call this method to release the context. An abandoned microbatch iterator cannot mint a completed comparison. @@ -1725,7 +1743,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 @@ -1817,7 +1835,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: @@ -1970,19 +1988,20 @@ def _recover_admission_impl( ) -> Any: """Pure search, at most one smaller-plan refresh, then one recovery.""" original: TrainerRankMemoryError | None = None - refused: _ForwardRefusal | None = None - best: _ForwardRefusal | None = None + refused: _impl._ForwardRefusal | None = None + best: _impl._ForwardRefusal | None = None def reject() -> Any: assert refused is not None if admit_refusal is not None: # Only the exhausted memory-refusal path changes. Every peer # must have a supported candidate; never override an EP or - # failed planning/runtime capability guard. + # failed planning/runtime capability guard or host placement budget. allowed = ( getattr(self, "_allow_oversized_batches", False) and refused.overridable and best is not None + and best.check.cpu_fits ) selected = None if allowed: @@ -2084,12 +2103,16 @@ def finish(value: Any) -> Any: self._snapshot_planning_telemetry(refused.plan, refused.check) return reject() 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: diff --git a/src/art/trainer_rank/_operations.py b/src/art/trainer_rank/_operations.py new file mode 100644 index 000000000..401790002 --- /dev/null +++ b/src/art/trainer_rank/_operations.py @@ -0,0 +1,223 @@ +"""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, 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: OperationId + kind: str + payload: bytes + + @classmethod + def capture(cls, id: OperationId, kind: str, payload: Any) -> TrainerOperation: + """Freeze arguments before asynchronous submission can observe mutations.""" + 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] + + +@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) +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 client retired this operation; its result is no longer available.""" + + +def _ledger(rank_zero: Any) -> _Ledger: + owner = rank_zero._rank + if not hasattr(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. + + 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) + if operation.kind == "acknowledge": + 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.outcomes.get(operation.id)) is not None: + if outcome.fingerprint != fingerprint: + raise ValueError( + "trainer operation identity was reused with different arguments" + ) + 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.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. + transport = copy(rank_zero) + transport._transport_handles = [] + 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) + 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"]) + 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 + 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/_optimizer.py b/src/art/trainer_rank/_optimizer.py index 0285582af..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 @@ -91,13 +81,27 @@ 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", ) -> 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( @@ -169,11 +173,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( @@ -235,6 +243,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) selected = [] for name in checkpoint_names: @@ -271,6 +289,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: ( @@ -323,6 +342,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 diff --git a/src/art/trainer_rank/_options.py b/src/art/trainer_rank/_options.py new file mode 100644 index 000000000..201e5a47f --- /dev/null +++ b/src/art/trainer_rank/_options.py @@ -0,0 +1,156 @@ +"""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") + for name, choices in ( + ("backward_state", ("auto", "gpu", "cpu", "replay")), + ("output_device", ("auto", "model", "cpu")), + ): + 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( + 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/_planner_evidence.py b/src/art/trainer_rank/_planner_evidence.py index 14d20e1be..58316b2bd 100644 --- a/src/art/trainer_rank/_planner_evidence.py +++ b/src/art/trainer_rank/_planner_evidence.py @@ -226,7 +226,7 @@ def validate(decision: Any, failed: Any) -> None: } or not isinstance(decision["attempt_id"], str) or re.fullmatch("[0-9a-f]{32}", decision["attempt_id"]) is None - or decision["operation"] not in {"forward_micro_batches", "dp_rank_forward"} + or decision["operation"] not in {"forward_batches", "forward"} or decision["reduction_scope"] not in {"world", "tp_cp_or_local"} or decision["outcome"] not in {"admitted", "admitted_oversized", "refused", "planning_error"} diff --git a/src/art/trainer_rank/_planner_misses.py b/src/art/trainer_rank/_planner_misses.py index bb7e8297c..1d79e8976 100644 --- a/src/art/trainer_rank/_planner_misses.py +++ b/src/art/trainer_rank/_planner_misses.py @@ -52,6 +52,8 @@ "_prefix_tree_performance_search.py", "_planner_misses.py", "_gdn_memory.py", + "_memory_policy.py", + "_options.py", "_planner_replay.py", "_planner_evidence.py", "_planner_retention.py", @@ -541,6 +543,9 @@ def _signature_values(values: dict[str, Any]) -> dict[str, Any]: values = dict(values) for name in ("topology", "planner_coefficients", "request_mix", "grad_modes"): values[name] = tuple(values[name]) + values["memory_placement"] = tuple( + tuple(placement) for placement in values.get("memory_placement", ()) + ) slots = [] raw_slots = values.get("slot_shapes", ()) if not isinstance(raw_slots, (list, tuple)): 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/src/art/trainer_rank/_slots.py b/src/art/trainer_rank/_slots.py index 7e867632d..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 @@ -44,7 +34,7 @@ def _resolve_custom_checkpoint(self: TrainerRank, checkpoint: AdapterSelection) 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 _impl.TrainerRankSlotStateError( "Custom checkpoint objects require a loaded named checkpoint" @@ -56,7 +46,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 +172,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 +222,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") @@ -411,7 +401,7 @@ def _resolve_slot_ref( request.checkpoint if request.checkpoint is not _impl.Unset else checkpoint ) if selection is not _impl.Unset: - name = cast(str | None, selection) + name = selection if name is not None and name not in self._checkpoint_slots: raise _impl.TrainerRankSlotStateError( f"Forward selects unloaded checkpoint {name!r}" @@ -506,7 +496,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 @@ -536,15 +526,15 @@ def _track_slot_graph_outputs( if not track_slot and not track_hybridep: return list(outputs) - marker: torch.Tensor | None = None + marker: _impl.torch.Tensor | None = None - def track(tensor: torch.Tensor | None) -> torch.Tensor | None: + def track(tensor: _impl.torch.Tensor | None) -> _impl.torch.Tensor | None: nonlocal marker if tensor is None or not tensor.requires_grad: return tensor if marker is None: - marker = tensor.new_empty(0) - return cast(_impl.torch.Tensor, _impl._SlotGraphSentinel.apply(tensor, marker)) + marker = _impl.torch.zeros((), dtype=_impl.torch.bool, device="cpu") + return _impl._track_slot_graph_tensor(tensor, marker) tracked_outputs = [ _impl.ForwardOutput( diff --git a/src/art/trainer_rank/_tensors.py b/src/art/trainer_rank/_tensors.py new file mode 100644 index 000000000..47ff44097 --- /dev/null +++ b/src/art/trainer_rank/_tensors.py @@ -0,0 +1,508 @@ +"""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 _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(): + 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.""" + 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): + @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/_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/src/art/trainer_rank/_versions.py b/src/art/trainer_rank/_versions.py new file mode 100644 index 000000000..e6391b39d --- /dev/null +++ b/src/art/trainer_rank/_versions.py @@ -0,0 +1,330 @@ +"""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() + prepared = None + try: + self.validate_gradients(gradients) + for entry in gradients: + batch.add(*entry) + 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 + + 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 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/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index e1d2c4438..3751285c7 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -5062,58 +5062,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"] = "" - - def _recorded_boundary_evidence( history: ChatCompletionsHistory, messages: list[dict[str, Any]], @@ -5378,7 +5326,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 e6a318c57..0e75dce1b 100644 --- a/src/art_inference/chat_template.py +++ b/src/art_inference/chat_template.py @@ -340,7 +340,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) @@ -385,18 +385,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,43 +447,18 @@ def visit( # Publishing a new binding can release an old # object whose destructor observes the new value. shared.update(bindings) - else: - # Macro/loop/with/block bodies have independent bindings. - for _, value in node.iter_fields(): - if isinstance(value, list) and all( - isinstance(n, nodes.Node) for n in value - ): - # Jinja initializes loop locals before each body; - # parameters/targets already have arbitrary values. - local_names = ( - { - n.name - for n in ( - node.target, - *node.target.find_all(nodes.Name), - ) - if isinstance(n, nodes.Name) - } - if isinstance(node, nodes.For) and value is node.body - else None - ) - 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() + elif isinstance(node, nodes.Macro): + # Macro parameters already have arbitrary values. + visit(node.body, set(), {arg.name for arg in node.args}) + elif isinstance(node, nodes.For): + # Only the loop body has fresh loop-local bindings. + local_names = { + n.name + for n in (node.target, *node.target.find_all(nodes.Name)) + if isinstance(n, nodes.Name) + } + visit(node.body, set(), local_names, True) + visit(node.else_, set()) return bindings visit(tree.body, set()) 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..c80657e20 --- /dev/null +++ b/tests/integration/megatron/cp_attn/test_cpu_offload_residency.py @@ -0,0 +1,196 @@ +"""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 +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 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( + 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, 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) + 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 + replay_plan = placement_cost( + [cost] * count, backward_state="replay", output_device="cpu" + ) + assert replay_plan.gpu_required_bytes <= cap + assert replay_plan.cpu_required_bytes <= 1 << 40 + replay = run("replay", count) + assert replay["peak"] <= cap + 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( + "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..870d1adc5 --- /dev/null +++ b/tests/integration/megatron/cp_attn/test_retained_backward.py @@ -0,0 +1,169 @@ +"""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 +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.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 + + +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, state, plan, indices = prepare_cp2_attention(rank, device, length) + 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..d902b6d2a 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)) @@ -594,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], @@ -601,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/integration/megatron/lora/test_lora_disk_codecs.py b/tests/integration/megatron/lora/test_lora_disk_codecs.py index 800143467..352da5f00 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,12 +1589,133 @@ 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( +def _export_preparation_failure_worker( + 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, + ): + 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 == failing_rank + and failure_site == "metadata" + and len(metadata_calls) == failure_call + ): + 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 == failing_rank: + 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 == 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) == ( + failure_call 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_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, -): - 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()) + 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_case), + timeout=90, + failure=f"collective export {failure_case[0]} failure test hung", + ) + + +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 +1741,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 +1767,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 +1924,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 +1948,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 +1957,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 +1971,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 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..fe419ea69 --- /dev/null +++ b/tests/integration/megatron/lora/test_trainer_v1_graph_cache.py @@ -0,0 +1,220 @@ +"""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.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 + 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: ParallelTopology()) + 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: ParallelTopology()) + 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/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 diff --git a/tests/unit/test_checkpoint_moment_capture.py b/tests/unit/test_checkpoint_moment_capture.py index bafd81a3b..e640a5ac7 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 +from test_trainer_rank_validation import _adapter_config, _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 @@ -173,3 +176,70 @@ def pack_owned(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 = _adapter_config( + rank=2, 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) diff --git a/tests/unit/test_checkpoint_snapshot_spill.py b/tests/unit/test_checkpoint_snapshot_spill.py index 8cb551eac..c1904162d 100644 --- a/tests/unit/test_checkpoint_snapshot_spill.py +++ b/tests/unit/test_checkpoint_snapshot_spill.py @@ -3,20 +3,49 @@ from concurrent.futures import Future from dataclasses import replace from pathlib import Path +import pickle +import sys import threading +from types import SimpleNamespace import weakref 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 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() + if parameter is None: + parameter = torch.nn.Parameter(torch.tensor([1.0])) + trainer._checkpoint_slots["a"] = _CheckpointSlot( + params=(parameter,), + config=_adapter_config(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: ({}, [])), + ) + 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()) @@ -26,20 +55,14 @@ def test_captured_cpu_custom_optimizer_is_independent() -> None: "exp_avg": torch.ones(2), "exp_avg_sq": torch.full((2,), 2.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"], - }, - optimizer=_DynamicOptimizer(optimizer, (master,)), - custom={ - "p": _CustomObject("parameter", parameter, object()), - "b": _CustomObject("buffer", buffer, object()), - }, - ) + 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() @@ -249,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 = [] @@ -441,3 +438,144 @@ 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] + + +@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") 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() 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_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 4c0aa0a01..9886efc34 100644 --- a/tests/unit/test_reasoning_parser_joined_branches.py +++ b/tests/unit/test_reasoning_parser_joined_branches.py @@ -1,14 +1,19 @@ +from pathlib import Path + 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, ) +_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/test_tinker_import_boundary.py b/tests/unit/test_tinker_import_boundary.py index 5d3940ab4..5305e553b 100644 --- a/tests/unit/test_tinker_import_boundary.py +++ b/tests/unit/test_tinker_import_boundary.py @@ -1,13 +1,18 @@ -"""Exercise the package initializer without importing optional training extras.""" +"""Keep inference imports isolated and public Tinker exports lazy.""" import builtins from importlib.util import module_from_spec, spec_from_file_location +import os from pathlib import Path +import subprocess import sys +import textwrap from types import SimpleNamespace import unittest from unittest.mock import Mock, patch +import pytest + class TinkerImportBoundaryTests(unittest.TestCase): def load_package(self): @@ -64,3 +69,98 @@ def test_requested_export_preserves_missing_dependency_error(self): if __name__ == "__main__": unittest.main() + + +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"] + """ + ) 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..349b44af4 --- /dev/null +++ b/tests/unit/test_trainer_batch_input_capture.py @@ -0,0 +1,170 @@ +"""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, +) +from art.trainer_rank._rng import TrainerRNG + + +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._rng = TrainerRNG(torch.device("cpu")) + 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..94a18eb05 --- /dev/null +++ b/tests/unit/test_trainer_command_transport.py @@ -0,0 +1,226 @@ +"""Commands cross ranks through CPU storage, independent of CUDA ordinals.""" + +from __future__ import annotations + +from dataclasses import dataclass +import gc +from typing import Any, Literal, cast +import weakref + +import pytest +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._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) + with ( + gloo_group(physical, f"file://{rendezvous}"), + megatron_topology(physical, dp_size=1, tp_size=2), + ): + 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) + 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 * model.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 + + +@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") + spawn_and_join( + _transport_worker, + (str(tmp_path / "init"), cuda), + timeout=90, + failure="command transport did not complete", + ) diff --git a/tests/unit/test_trainer_driver_transport.py b/tests/unit/test_trainer_driver_transport.py new file mode 100644 index 000000000..47f19f143 --- /dev/null +++ b/tests/unit/test_trainer_driver_transport.py @@ -0,0 +1,340 @@ +"""Driver CPU transport stays separate from native output placement.""" + +from __future__ import annotations + +import asyncio +from dataclasses import replace +import gc +import sys +from typing import Any, cast +import weakref + +import pytest +from test_trainer_rank_commands import _input, _Rank +import torch + +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 +from art.trainer_rank._tensors import CotangentCollector + + +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)), 1), kind, payload) + ) + + +@pytest.mark.parametrize("device", ["cpu", "cuda"]) +@pytest.mark.parametrize("batches", [False, True]) +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") + + 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 + + +@pytest.mark.parametrize("policy", ["model", "cpu", "auto"]) +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", "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, batches +): + 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) + 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 + + def constrained(): + nonlocal calls + 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: + + def reject_clone(handle, *args, **kwargs): + if handle.startswith("client:"): + raise error + return detach(handle, *args, **kwargs) + + monkeypatch.setattr(_tensors, "detach_tree", reject_clone) + 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 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(): + 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() + + +@pytest.mark.parametrize("kind", ["forward", "batches_open"]) +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_live_parameter_roots.py b/tests/unit/test_trainer_live_parameter_roots.py new file mode 100644 index 000000000..dcea359d5 --- /dev/null +++ b/tests/unit/test_trainer_live_parameter_roots.py @@ -0,0 +1,125 @@ +"""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_live_heads import _live_head, _native_head +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, parameter = _native_head("parameter", "weight", factory) + collector = CotangentCollector() + live = _live_head(trainer, "weight", parameter.detach(), collector) + return trainer, parameter, collector, live + + +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, parameter = _native_head("parameter", "weight", lambda: torch.tensor(2.0)) + 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, parameter, collector, live = _live_parameter( + lambda: torch.tensor([2.0, 3.0], dtype=dtype) + ) + 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, parameter, collector, live = _live_parameter() + 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, 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"): + 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..6a5dfa2c3 --- /dev/null +++ b/tests/unit/test_trainer_operations.py @@ -0,0 +1,400 @@ +import asyncio +from collections.abc import Coroutine +from contextlib import nullcontext +import gc +from types import SimpleNamespace +from typing import Any, cast +import weakref + +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( + 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)) + + +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(): + 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(("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): + 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, + _release_on_error=lambda handles: nullcontext(), + ) + assert torch.equal(await execute_operation(rank, operation), torch.tensor([6.0])) + + +@pytest.mark.parametrize("cuda_tag", [False, True]) +async def test_operation_codec_remaps_storages_and_preserves_aliases( + monkeypatch, cuda_tag +): + from test_trainer_command_transport import _check_payload, _payload + + 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() + + +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(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(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), + _release_on_error=lambda handles: nullcontext(), + 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]) +async def test_failed_outcome_releases_traceback_activations_and_replays_error( + unserializable, +): + class UnserializableError(RuntimeError): + def __reduce__(self): + raise TypeError("cannot serialize this error") + + 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(("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 + ) + + +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"]) +async def test_failed_exported_backward_replays_once_and_releases_only_consumed_graphs( + retain_graph, + failure_stage, +): + 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}, + ) + 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(): + 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),)}, + ) + assert "unrelated" in state.exports + assert ("known" in state.exports) is (handle == "missing") diff --git a/tests/unit/test_trainer_rank_active_memory.py b/tests/unit/test_trainer_rank_active_memory.py index a2033710d..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, ) @@ -128,7 +120,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,) @@ -264,7 +256,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_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_cache_recovery.py b/tests/unit/test_trainer_rank_cache_recovery.py index ef34113c3..a17091f59 100644 --- a/tests/unit/test_trainer_rank_cache_recovery.py +++ b/tests/unit/test_trainer_rank_cache_recovery.py @@ -1,6 +1,7 @@ """Actual admission methods with scalar allocator/search/clock facades, no CUDA.""" from contextlib import nullcontext +from dataclasses import replace import math import os import types @@ -8,7 +9,7 @@ import unittest from unittest.mock import patch -from art.trainer_rank import _backward_work, _impl +from art.trainer_rank import ForwardOptions, _backward_work, _impl Refusal = _impl.TrainerRankMemoryError Partial = _impl.TrainerRankPartialExecutionError @@ -140,7 +141,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: @@ -172,7 +173,9 @@ def make(self): self.addCleanup(observer_clock.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, **k: 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: {} @@ -280,7 +283,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) @@ -362,7 +365,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" @@ -393,8 +396,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) @@ -426,26 +429,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() @@ -459,9 +465,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) @@ -497,7 +501,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) @@ -520,14 +524,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) @@ -543,9 +547,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) @@ -564,7 +566,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) @@ -928,8 +930,17 @@ def _check_component_demand_recovery( import pytest assert not _impl.dist.is_initialized() and not _impl.torch.cuda.is_initialized() + # Hold placement fixed: automatic offload can legitimately lower demand. + requests = [ + replace( + request, options=ForwardOptions(backward_state="gpu", output_device="model") + ) + for request in requests + ] plan = rank._plan_flat_forward(requests) - required = rank._memory_check(plan).estimated_required_bytes + model_required = rank._memory_check(plan).estimated_required_bytes + required = rank._admit_graph_memory(plan)[1].estimated_required_bytes + assert required >= model_required # Includes v1's detached caller outputs. profiles = dict(rank._memory_profiles) groups = rank._plan_group_rows(plan) head = rank._plan_head_workspace_bytes(plan) @@ -1018,7 +1029,7 @@ def observed_error(refused, context): assert captured.value.__cause__ is errors[0] assert len(searches) == 2 and isinstance(searches[0], _impl._ForwardRefusal) assert releases == [1] and outcomes[0] == (1, 1) - assert (1, required) in demands and (2, required) in demands + assert (1, model_required) in demands and (2, model_required) in demands assert rank._recovery_state().first_consumed assert rank._recovery_state().owner is None assert rank._memory_profiles == profiles 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..2b93a422b --- /dev/null +++ b/tests/unit/test_trainer_rank_callback_cleanup.py @@ -0,0 +1,158 @@ +"""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_failures.py b/tests/unit/test_trainer_rank_callback_failures.py new file mode 100644 index 000000000..70af8ab06 --- /dev/null +++ b/tests/unit/test_trainer_rank_callback_failures.py @@ -0,0 +1,214 @@ +"""Controller cancellation must remain possible while global cleanup is pending.""" + +from __future__ import annotations + +import asyncio +from contextlib import closing +import gc +from multiprocessing.connection import Connection +from typing import Any + +import pytest +from test_trainer_rank_commands import _input, _loss_tree, _Rank +import torch +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 _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) + 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)) + + +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..0310e82f9 --- /dev/null +++ b/tests/unit/test_trainer_rank_callback_lifecycle.py @@ -0,0 +1,188 @@ +"""Checkpoint scopes and delayed cleanup stay within their callback session.""" + +import asyncio +import gc +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, run_rank_callback_stream + + +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): + 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 = [] + + 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() + + 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_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_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 c6e33ee9d..59f5fb34d 100644 --- a/tests/unit/test_trainer_rank_checkpoint_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_memory.py @@ -2,10 +2,10 @@ 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, recompute_model from art.trainer_rank import ForwardInput, TrainerRank from art.trainer_rank._impl import Unset, _ForwardRefusal, _MemoryProfile @@ -14,43 +14,8 @@ def rank(): from megatron.core.transformer.transformer_block import TransformerBlock - 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.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 - 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), - ), - ) - ) + 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 return result @@ -157,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]) @@ -340,39 +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( - "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() @@ -480,20 +423,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/test_trainer_rank_commands.py b/tests/unit/test_trainer_rank_commands.py new file mode 100644 index 000000000..31cd686e3 --- /dev/null +++ b/tests/unit/test_trainer_rank_commands.py @@ -0,0 +1,872 @@ +"""Command participation tests; native model numerics are a separate GPU gate.""" + +from __future__ import annotations + +import asyncio +from dataclasses import replace +from datetime import timedelta +import gc +import sys +import threading +from types import SimpleNamespace +from typing import Any, cast +from unittest.mock import patch +import weakref + +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, + ForwardOptions, + ForwardOutput, + MicroBatch, + MicroBatchStats, + TrainerRank, + TrainerRankZero, + 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: + 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) + 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 + + 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) + ) + 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)) + + +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() + 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} + + +@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, 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] + + +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): # ty: ignore[invalid-method-override] + if dist.get_rank() == 1: + raise RuntimeError("intentional physical backward failure") + return gradient + + +def _distributed_worker(physical, rendezvous, output): + 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) + 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 + + with patch.object(_commands.cloudpickle, "loads", track_loads): + 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, + ) + ) + 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)]) + + with patch.object(_commands.cloudpickle, "loads", fail_result_decode): + run(decode_refusal) + 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 + try: + with patch.object(_commands.cloudpickle, "dumps", track_dumps): + run(host_refusal) + finally: + 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): + 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)]))) + + with patch.object(rank, "forward", failing_forward): + run(backward_refusal) + 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 = 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() + ) + 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) + + +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._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) + 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 + + 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 + + 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): + count: torch.Tensor + + 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 = 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 + + +@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")) + + +@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 + + 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(): + 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]) + + +@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_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..e8f180c8f 100644 --- a/tests/unit/test_trainer_rank_custom_tensors.py +++ b/tests/unit/test_trainer_rank_custom_tensors.py @@ -4,26 +4,27 @@ 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 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 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 ( AdamParams, MaterializedCheckpoint, + ModuleHandle, TrainerRank, TrainerRankSlotStateError, - Unset, ) from art.trainer_rank._checkpoint import ( PreparedCustomPayload, @@ -34,34 +35,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, - ) -> ModuleT: ... - - 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: @@ -124,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", @@ -159,11 +106,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: @@ -180,14 +140,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": @@ -231,8 +184,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( @@ -240,24 +191,16 @@ 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") 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) - finally: - dist.destroy_process_group() def test_module_accepts_class_and_lambda_factories_and_is_idempotent() -> None: @@ -279,8 +222,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 +343,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 +367,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 +384,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 +400,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 +428,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 +445,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 +457,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"): @@ -641,17 +592,9 @@ 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 - ), - ) - (a * 3).backward() + _use_local_gradients(trainer, monkeypatch) + 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), @@ -667,17 +610,9 @@ 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 - ), - ) - (used * 2).backward() + _use_local_gradients(trainer, monkeypatch) + with trainer._gradient_transaction(): + (used * 2).backward() trainer.optim_step( params=AdamParams(learning_rate=0.1, weight_decay=0.5), @@ -949,17 +884,9 @@ 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 - ), - ) - head.q_proj.lora_A.weight.sum().backward() + _use_local_gradients(trainer, monkeypatch) + 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"], @@ -1103,10 +1030,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 @@ -1119,6 +1047,17 @@ 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) + 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() async def prefetch() -> None: @@ -1146,6 +1085,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" ) @@ -1247,7 +1207,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]]), @@ -1261,17 +1221,17 @@ 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, -) -> tuple[_ValueHead, torch.nn.Parameter, torch.Tensor]: + rank: TrainerRank, +) -> 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,24 +1244,16 @@ def _register_custom_tensors( def _step_custom_tensors( trainer: TrainerRank, - head: _ValueHead, + head: ModuleHandle, temperature: torch.nn.Parameter, monkeypatch: pytest.MonkeyPatch, *, 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]]) - (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_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_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 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..f2863224e --- /dev/null +++ b/tests/unit/test_trainer_rank_graph_order.py @@ -0,0 +1,99 @@ +"""Physical peers must enter cached/replayed backward collectives identically.""" + +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._commands import _coordinate_call +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): + 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() + 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): + return _coordinate_call(function, group=None) + + 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 + + +@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..3e48398e0 --- /dev/null +++ b/tests/unit/test_trainer_rank_graphs.py @@ -0,0 +1,831 @@ +from __future__ import annotations + +import asyncio +from contextlib import contextmanager, nullcontext +from functools import partial +import gc +from types import SimpleNamespace +import weakref + +import pytest +import torch +from torch.multiprocessing.reductions import StorageWeakRef +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(): + 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 _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 +): + cache = cache or 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 = _logprob_corrections(outputs, 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 _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() + 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: + watch(args[:2]) + result = function(*args, **kwargs) + watch(result) + return result + except BaseException as error: + errors.append(error) + raise + finally: + args = kwargs = result = 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), + ) + 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", cache=cache + ) + 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 + 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 = commands._coordinate_call(function, group=None) + if phase == "coordinator" and calls == 3: + raise primary from cause + return result + + if phase == "always": + with torch.no_grad(): + current.fill_(float("nan")) + with pytest.raises(ValueError if phase == "always" else RuntimeError) as failure: + 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 "current logprobs" in str(failure.value) + assert original.grad is None + 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)] + 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_many(((retry, (gradient,)),), coordinate=coordinate) + 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) + 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 + ) + 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() + + +@pytest.mark.parametrize("checkpointing", [False, True]) +def test_abandoned_release_frees_physical_record_without_autograd(checkpointing): + 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(): + 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)) + + +@pytest.mark.parametrize("always_prepass", [False, True], ids=["uncorrected", "always"]) +def test_replay_failure_discards_transaction_and_releases_participating_records( + monkeypatch, request, always_prepass +): + 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() + older_weight = torch.nn.Parameter(torch.tensor(3.0)) + older, (older_output,) = cache.run(lambda _: (older_weight.square(),), ()) + 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] + + def failing(x): + executions[0] += 1 + activation = snapshot * x + 1 + result = activation.square() + storages.extend( + StorageWeakRef(value.untyped_storage()) for value in (activation, result) + ) + 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, 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, (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 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),)) + 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() == () + + +@pytest.mark.parametrize("failure_type", [MemoryError, asyncio.CancelledError]) +@pytest.mark.parametrize("retention", ["gpu", "cpu", "replay"]) +@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 +): + 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 + 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") + snapshot = trainer._snapshot_parameter(parameter, version) + snapshots.append(weakref.ref(snapshot)) + versions.append(weakref.ref(version)) + return dict( + execute=partial(execute, snapshot), + inputs=torch.tensor(3.0), + context_factory=lambda: nullcontext(snapshot), + 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): + records.append(weakref.ref(record)) + 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 + 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 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, + *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(): + cache.backward(handle, (torch.tensor(1.0), torch.tensor(1.0))) + torch.testing.assert_close(parameter.grad, torch.tensor(7.0)) + 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_handoff_budget.py b/tests/unit/test_trainer_rank_handoff_budget.py index 4f5f4f8ee..72ae2dd68 100644 --- a/tests/unit/test_trainer_rank_handoff_budget.py +++ b/tests/unit/test_trainer_rank_handoff_budget.py @@ -43,7 +43,7 @@ def test_admission_consumes_the_only_first_release(rig): rank._release_cached_memory_for_backward(plan(True)) assert cuda.events.count("release") == 1 assert rank._recovery_state().cost > before - rank._record_recovery_work("forward_micro_batches", 100.0) + rank._record_recovery_work("forward_batches", 100.0) rank._release_cached_memory_for_backward(plan(True)) assert cuda.events.count("release") == 2 @@ -85,7 +85,7 @@ def run(): lambda: next(searches), lambda value: value, lambda value, check: (value[0], check), - context="forward_micro_batches", + context="forward_batches", sync_across_dp=True, admit_refusal=lambda refused: (refused.plan, refused.check), ) @@ -142,7 +142,7 @@ def run(): lambda: next(searches), lambda value: value, lambda value, check: (value[0], check), - context="forward_micro_batches", + context="forward_batches", sync_across_dp=True, admit_refusal=lambda refused: (refused.plan, refused.check), ) @@ -301,7 +301,7 @@ def forward(plan, **kwargs): primary = RuntimeError("post-release sample") state["failure"] = primary with torch.no_grad(): - iterator = rank.forward_micro_batches([_target_request(1)], no_grad=False) + iterator = rank.forward_batches([_target_request(1)], no_grad=False) with pytest.raises(RuntimeError) as caught: next(iterator) assert caught.value is primary and not torch.is_grad_enabled() 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_head_memory.py b/tests/unit/test_trainer_rank_head_memory.py index 21b7ad071..a20bea90c 100644 --- a/tests/unit/test_trainer_rank_head_memory.py +++ b/tests/unit/test_trainer_rank_head_memory.py @@ -128,7 +128,7 @@ def test_real_head_split_fits_before_cache_recovery(monkeypatch): r, "_try_cache_recovery", lambda *a, **kw: pytest.fail("Split already fits") ) executed = _recording_executor(monkeypatch, r) - batches = list(r.forward_micro_batches([requests])) + batches = list(r.forward_batches([requests])) assert len(batches) == 1 and batches[0].stats.subforward_count == 2 assert batches[0].stats.global_count == 1 and len(executed) == 2 assert [ 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 3e40d2fa8..77c0b1416 100644 --- a/tests/unit/test_trainer_rank_head_recompute.py +++ b/tests/unit/test_trainer_rank_head_recompute.py @@ -2,7 +2,6 @@ from contextlib import nullcontext from dataclasses import replace -from datetime import timedelta import sys from types import SimpleNamespace from unittest.mock import patch @@ -12,6 +11,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 @@ -177,14 +177,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) @@ -218,8 +217,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( @@ -329,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_head_memory.py b/tests/unit/test_trainer_rank_live_head_memory.py new file mode 100644 index 000000000..3dce0d426 --- /dev/null +++ b/tests/unit/test_trainer_rank_live_head_memory.py @@ -0,0 +1,192 @@ +"""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, detach_tree + + +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", + ) + 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") + 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 == () + 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 new file mode 100644 index 000000000..a0f4d1b06 --- /dev/null +++ b/tests/unit/test_trainer_rank_live_heads.py @@ -0,0 +1,1487 @@ +from __future__ import annotations + +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, megatron_topology, spawn_and_join + +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, + LiveHead, + 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 + + +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 _native_head(kind="module", name="head", factory=TiedHead): + trainer, rank = _trainer("student") + native = getattr(rank, kind)(name, factory, checkpoint="student") + 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: + 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 + + +def _tensor(live: LiveHead) -> torch.Tensor: + assert isinstance(live.value, torch.Tensor) + return live.value + + +def _step(trainer, monkeypatch): + _use_local_gradients(trainer, monkeypatch) + 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, 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(): + 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, head = _native_head( + "module", "bn", lambda: torch.nn.BatchNorm1d(2, dtype=torch.float64) + ) + 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, 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 + torch.testing.assert_close(head.running_mean, torch.zeros(2)) + + +def test_client_tied_head_old_backward_after_native_refresh(monkeypatch): + trainer, native = _native_head() + collector = CotangentCollector() + 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(): + 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, 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) + 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, 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() + 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) + ).state + 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, head = _native_head() + 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, head = _native_head() + 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(): + 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 art.trainer_rank._heads import synchronize_head_buffers + + 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): + 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 + + +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 _buffer_snapshot_failure_worker(process_rank, init_method): + 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") + 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 = _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) + 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, 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) + 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, native = _native_head() + 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 + 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, native = _native_head() + collector = CotangentCollector() + 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) + 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, 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 + 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 = _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])) + + +def test_head_export_preserves_strict_constructor_policy_for_client(): + 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 = _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"): + 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, native = _native_head() + collector = CotangentCollector() + with torch.no_grad(): + 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])) + 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, 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) + 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 = _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) + 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) + ).state + 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 = _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 + 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): + trainer, native = _native_head( + "parameter", "weight", lambda: torch.tensor([2.0, 4.0]) + ) + collector = CotangentCollector() + 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 + ) + 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, 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 + 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, 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 + 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, 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) + + 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, 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 + 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, 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) + 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, native = _native_head() + live = _live_head(trainer, "head", TiedHead()) 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: _live_head(trainer, 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, 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) + 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 art.trainer_rank._heads import synchronize_head_buffers + + with gloo_group(process_rank, init_method): + 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 = _live_head(trainer, "head", torch.nn.BatchNorm1d(2)) + 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 + + +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 = _live_head(trainer, "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(): + 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, 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 + 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, 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 + 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, 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) + + 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, 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) + 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, 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"): + 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, 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( + 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, 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 ( + 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 = ( + _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 + + 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 + + +@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()) 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..c7e759749 --- /dev/null +++ b/tests/unit/test_trainer_rank_memory_admission.py @@ -0,0 +1,544 @@ +from dataclasses import replace +from types import SimpleNamespace + +import pytest +from test_trainer_rank_active_memory import _rank +import torch + +from art.trainer_rank import ( + ForwardInput, + ForwardOptions, + ImportanceSamplingGradientCorrection, + _impl, +) + + +@pytest.fixture +def rank(monkeypatch): + 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) + 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 + + +@pytest.mark.parametrize("allow_oversized", [False, True]) +def test_cpu_shortage_is_not_overwritten_by_fresh_gpu_check( + rank, monkeypatch, allow_oversized +): + rank._allow_oversized_batches = allow_oversized + 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, + admit_refusal=lambda refused: (refused.plan, refused.check), + ) + 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..6ae2d7fce --- /dev/null +++ b/tests/unit/test_trainer_rank_memory_policy.py @@ -0,0 +1,219 @@ +from pathlib import Path +from typing import Literal + +import pytest + +from art.trainer_rank._memory_policy import ( + ForwardMemoryCost, + HostMemoryBudget, + MemoryScope, + 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 + + +def test_no_grad_cpu_outputs_release_gpu_storage_without_replay(): + costs = (ForwardMemoryCost(100, 80, 80, backward_required=False),) * 3 + 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 + + +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..de9acbefb --- /dev/null +++ b/tests/unit/test_trainer_rank_memory_policy_cuda.py @@ -0,0 +1,198 @@ +"""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, + 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_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 + 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 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..33fef6e29 --- /dev/null +++ b/tests/unit/test_trainer_rank_memory_recovery_distributed.py @@ -0,0 +1,165 @@ +"""Real collectives around injected transfer failure; no CUDA memory claims.""" + +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 trainer_rank_test_support import gloo_group + +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: + 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 + 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)) + + +@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: + 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 + # 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" + + +@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 + + 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) + 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) + 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) + torch.testing.assert_close(parameter.grad, torch.ones_like(parameter)) + assert trainer._version_state()._transaction is None + dist.barrier() + + +@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_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index 3fbdcee89..541f4c2e4 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -2,13 +2,14 @@ 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, recompute_model -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, @@ -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, ) @@ -636,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) @@ -651,51 +643,18 @@ 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 = 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.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)) - 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 @@ -780,6 +739,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"] ) 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..abade529e --- /dev/null +++ b/tests/unit/test_trainer_rank_offload_lifetime.py @@ -0,0 +1,148 @@ +"""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, "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) + + @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(bool(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 + 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 + 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() 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_delivery.py b/tests/unit/test_trainer_rank_output_delivery.py new file mode 100644 index 000000000..0170e14bb --- /dev/null +++ b/tests/unit/test_trainer_rank_output_delivery.py @@ -0,0 +1,472 @@ +"""Undelivered logical outputs must not retain physical forward graphs.""" + +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, + ForwardOutput, + MicroBatch, + MicroBatchStats, + _memory_policy, + _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 placement, copy, or collection 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) + 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 + + 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", "placement"]) +def test_failed_delivery_releases_registered_graph_and_native_cache( + monkeypatch, mode, kind +): + rank: Any = _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( + "phase", ["peer", "snapshot", "copy", "serialization", "transfer", "exchange"] +) +@pytest.mark.parametrize("operation", ["forward", "next", "batches_next"]) +def test_peer_rejection_releases_undelivered_packet_storage( + monkeypatch, request, operation, phase +): + 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 + inputs = [_input(3), _input(5)] if phase == "copy" else _input(3) + argument = ( + inputs + if operation == "forward" + else executor.invoke( + "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 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( + StorageWeakRef(t.untyped_storage()) for t in value.packet.tensors + ) + return value + 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): + 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) + 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"): + 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 = ( + "required rank 4" + if phase == "copy" + else "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 == "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 + # 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[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 + 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()) + 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 +): + rank: Any = _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) + 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) + 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: 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 + 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 + + +@pytest.mark.parametrize("kind", ["copy", "placement"]) +@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, delivery, close_error, kind +): + rank: Any = _CachedRank() + executor = _Executor(rank, "rank" if delivery == "rank" else "zero") + view = _view(executor) + invoke = executor.invoke + release_calls = 0 + flushes = [] + + 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", 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, kind) + 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: + advance() + assert failure.value is error and error.__traceback__ is not None + 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: Any = _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 isinstance(iterator, Generator) + 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() 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..4422564cc --- /dev/null +++ b/tests/unit/test_trainer_rank_output_memory.py @@ -0,0 +1,100 @@ +"""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")) + if policy == "model": + with pytest.raises(MemoryError, match="only 20 bytes"): + view._place_outputs([_output(policy=policy)]) + else: + assert view._place_outputs([_output(policy=policy)])[0].cpu == (True,) + + +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..61a83e684 --- /dev/null +++ b/tests/unit/test_trainer_rank_output_memory_cuda.py @@ -0,0 +1,83 @@ +"""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")) + 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 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..b7a4727dd --- /dev/null +++ b/tests/unit/test_trainer_rank_parameter_hooks.py @@ -0,0 +1,271 @@ +from __future__ import annotations + +import gc +import json + +import pytest +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 export_head, head_gradient_targets +from art.trainer_rank._tensors import CotangentCollector + + +def setup(client, values=None): + factory = lambda: torch.tensor(2.0) if values is None else values.clone() + trainer, native = _native_head("parameter", "p", factory) + collector = CotangentCollector() + 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) + + 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 = _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) + + 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 +): + values = torch.arange(1.0, 7).reshape(3, 2) + _, native, parameter, _, _, backward = setup(client, values) + 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) + + if not mixed: + with pytest.raises(ValueError, match="layout"): + backward(loss) + torch.testing.assert_close(native.grad, torch.ones_like(values)) + if client: + assert seen == [torch.sparse_coo] + else: + 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) + + +@pytest.mark.parametrize("client", (False, True)) +def test_tied_module_parameter_hook_sums_all_calls(client): + trainer, native = _native_head(factory=TiedHead) + collector = CotangentCollector() + live = _live_head(trainer, "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..4065f4226 --- /dev/null +++ b/tests/unit/test_trainer_rank_parameter_no_grad.py @@ -0,0 +1,99 @@ +from __future__ import annotations + +import asyncio + +import pytest +from test_trainer_rank_custom_tensors import _trainer +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 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 +): + initial = torch.tensor([[2.0, 3.0], [4.0, 5.0]]) + trainer, native = _native_head("parameter", "p", lambda: initial.clone()) + collector = CotangentCollector() + live = _live_head(trainer, "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_pending_memory.py b/tests/unit/test_trainer_rank_pending_memory.py index 359de8daa..c41a97229 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, recompute_model from art.megatron.prefix_tree_packing import prefix_tree_pack from art.trainer_rank import ForwardInput, TrainerRank @@ -29,27 +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 = 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.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) @@ -75,26 +56,12 @@ 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 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_physical_reserve.py b/tests/unit/test_trainer_rank_physical_reserve.py index b41a6e615..e15b91927 100644 --- a/tests/unit/test_trainer_rank_physical_reserve.py +++ b/tests/unit/test_trainer_rank_physical_reserve.py @@ -172,12 +172,23 @@ def phase(name, evidence, **kwargs): ] ] with torch.no_grad(): - iterator = rank.forward_micro_batches(inputs, no_grad=True) + iterator = rank.forward_batches(inputs, no_grad=True) batch = next(iterator) assert not torch.is_grad_enabled() assert [g.grad_enabled for g in executed[0].groups] == [False, True] assert state["releases"] == 1 and state["current"] == torch.device("cuda:7") - assert batch.inputs == inputs and batch.indices == (0,) + assert len(batch.inputs) == 1 and batch.indices == (0,) + for actual, expected in zip(batch.inputs[0], inputs[0], strict=True): + torch.testing.assert_close(actual.input_tokens, expected.input_tokens) + torch.testing.assert_close(actual.target_tokens, expected.target_tokens) + assert ( + replace( + actual, + input_tokens=expected.input_tokens, + target_tokens=expected.target_tokens, + ) + == expected + ) assert ( batch.stats.estimated_required_bytes == checks[0].estimated_required_bytes ) @@ -222,7 +233,7 @@ def fail(): else: state["failure"] = original with torch.no_grad(): - iterator = rank.forward_micro_batches([_target_request(1)], no_grad=False) + iterator = rank.forward_batches([_target_request(1)], no_grad=False) with pytest.raises(RuntimeError) as caught: next(iterator) assert caught.value is original @@ -312,6 +323,6 @@ def forward(plan, **kwargs): "_release_cached_memory_for_backward", lambda plan: pytest.fail("direct forward is outside the iterator handoff"), ) - output = rank.dp_rank_forward([_target_request(1)])[0] + output = rank.forward([_target_request(1)])[0] output.target_logprobs.sum().backward() assert len(executed) == 1 diff --git a/tests/unit/test_trainer_rank_planner_evidence.py b/tests/unit/test_trainer_rank_planner_evidence.py index 8fe601304..6cf41dd45 100644 --- a/tests/unit/test_trainer_rank_planner_evidence.py +++ b/tests/unit/test_trainer_rank_planner_evidence.py @@ -41,7 +41,7 @@ def scalar(request): def test_first_and_selected_are_not_latest_or_evicted(): - decision = evidence.Decision("dp_rank_forward", sync_across_dp=False) + decision = evidence.Decision("forward", sync_across_dp=False) first = sample(decision, 100) selected = sample(decision, 80) for i in range(100): @@ -58,8 +58,8 @@ def test_first_and_selected_are_not_latest_or_evicted(): def test_nested_rank_context_isolated_and_restored(): a, b = object(), object() - outer = evidence.Decision("dp_rank_forward", sync_across_dp=False, owner=a) - inner = evidence.Decision("dp_rank_forward", sync_across_dp=False, owner=b) + outer = evidence.Decision("forward", sync_across_dp=False, owner=a) + inner = evidence.Decision("forward", sync_across_dp=False, owner=b) with evidence.scope(outer): assert evidence.current(a) is outer and evidence.current(b) is None with evidence.scope(inner): @@ -93,7 +93,7 @@ def stats(device): ordinary_calls, ordinary_events = calls[:], cuda.events[:] calls.clear() cuda.events.clear() - decision = evidence.Decision("dp_rank_forward", sync_across_dp=False, owner=rank) + decision = evidence.Decision("forward", sync_across_dp=False, owner=rank) with evidence.scope(decision): observed = rank._memory_check_required(80) assert observed == ordinary @@ -121,9 +121,7 @@ def reduce(value, *, op, group): dist.all_reduce = reduce monkeypatch.setattr(tr, "dist", dist) - decision = evidence.Decision( - "forward_micro_batches", sync_across_dp=True, owner=rank - ) + decision = evidence.Decision("forward_batches", sync_across_dp=True, owner=rank) with evidence.scope(decision): check = rank._memory_check_required(80, sync_across_dp=True) assert calls == [("MAX", None), ("MIN", None)] @@ -141,7 +139,7 @@ def test_final_refusal_emits_without_forward_or_memory_window(monkeypatch, tmp_p rank._planner_reporter = reports.Reporter(1e9, spool_dir=tmp_path / "reports") rank._planner_device_identity = {} with pytest.raises(tr.TrainerRankMemoryError): - rank.dp_rank_forward([_request(i) for i in range(4)]) + rank.forward([_request(i) for i in range(4)]) raw = next(rank._planner_reporter.spool_dir.glob("*.json")).read_bytes() record = reports.validate_report(raw) assert record["event"] == "admission_refused" @@ -192,7 +190,7 @@ def snapshot(plan, check, observation): if microbatch: rank._select_next_micro_batch([requests], 0) else: - rank.dp_rank_forward(requests) + rank.forward(requests) record = reports.validate_report( next(rank._planner_reporter.spool_dir.glob("*.json")).read_bytes() ) @@ -231,7 +229,7 @@ def search(): search, lambda x: x, lambda x, c: x, - context="dp_rank_forward", + context="forward", sync_across_dp=False, ) assert not rank._planner_reporter.spool_dir.exists() @@ -267,7 +265,7 @@ def test_actual_cache_release_trace_and_nonfinite_budget(scalar, tmp_path): if item["kind"] == "recovery" and item["status"] == "completed" ) assert completed["values"]["observed_available_delta_bytes"] == 160 - decision = evidence.Decision("dp_rank_forward", sync_across_dp=False) + decision = evidence.Decision("forward", sync_across_dp=False) decision.outcome = "planning_error" decision.record("recovery", "budget_observed", local_high_seconds=math.inf) value = decision.snapshot() @@ -286,9 +284,7 @@ def test_budget_trace_uses_completed_backward_and_observer_operands(scalar): backward.cost_ns = 500_000_000 ticks = iter((10.0, 11.0, 13.0)) rank._recovery_clock = lambda: next(ticks) - decision = evidence.Decision( - "forward_micro_batches", sync_across_dp=True, owner=rank - ) + decision = evidence.Decision("forward_batches", sync_across_dp=True, owner=rank) with evidence.scope(decision): _, error, _ = run(rank, [fail(names), recovery.success(names)]) assert error is None and cuda.events.count("release") == 1 @@ -307,9 +303,7 @@ def test_handoff_sentinel_is_not_reported_as_available_memory(scalar): cuda.free = 1 cuda.memory_reserved = lambda device: cuda.allocated + 100 cuda.device = lambda device: nullcontext() - decision = evidence.Decision( - "forward_micro_batches", sync_across_dp=True, owner=rank - ) + decision = evidence.Decision("forward_batches", sync_across_dp=True, owner=rank) with evidence.scope(decision): rank._release_cached_memory_for_backward( SimpleNamespace(groups=[SimpleNamespace(grad_enabled=True)]) @@ -530,7 +524,7 @@ def encode(*args, **kwargs): def test_summary_trimming_preserves_first_selected(monkeypatch): - decision = evidence.Decision("dp_rank_forward", sync_across_dp=False) + decision = evidence.Decision("forward", sync_across_dp=False) decision.first = decision.selected = sample(decision) decision.outcome = "refused" for i in range(60): @@ -562,7 +556,7 @@ def test_legacy_report_reader_remains_supported(tmp_path): def test_planning_oom_does_not_borrow_previous_forward(monkeypatch, tmp_path): rank, plan, _ = _reporting_rank(monkeypatch, tmp_path) rank._run_flat_plan_with_memory_tracking( - plan, check=tr._MemoryCheck(220, 10000, True), context="dp_rank_forward" + plan, check=tr._MemoryCheck(220, 10000, True), context="forward" ) original = tr.torch.cuda.OutOfMemoryError("new planning allocation") _, caught, _ = run(rank, [original]) @@ -595,9 +589,7 @@ def reduce(value, *, op, group): dist.all_reduce = reduce monkeypatch.setattr(tr, "dist", dist) - decision = evidence.Decision( - "forward_micro_batches", sync_across_dp=True, owner=rank - ) + decision = evidence.Decision("forward_batches", sync_across_dp=True, owner=rank) with evidence.scope(decision): first = rank._memory_check_required(80, sync_across_dp=True) selected = rank._refresh_memory_check(first, sync_across_dp=True) @@ -620,7 +612,7 @@ def reduce(value, *, op, group): def test_refresh_without_original_sample_leaves_local_requirement_unknown(): - decision = evidence.Decision("dp_rank_forward", sync_across_dp=False) + decision = evidence.Decision("forward", sync_across_dp=False) with decision.refresh_of(None): value = sample(decision, 150) assert value.local_required_bytes is value.required_from_ordinal is None diff --git a/tests/unit/test_trainer_rank_planner_options.py b/tests/unit/test_trainer_rank_planner_options.py index 0564cd7cb..4d77076aa 100644 --- a/tests/unit/test_trainer_rank_planner_options.py +++ b/tests/unit/test_trainer_rank_planner_options.py @@ -26,7 +26,7 @@ def test_default_refusal_does_not_execute(monkeypatch): rank._allow_oversized_batches = False executed = _recording_executor(monkeypatch, rank) with pytest.raises(tr.TrainerRankMemoryError): - rank.dp_rank_forward([_request(i) for i in range(4)]) + rank.forward([_request(i) for i in range(4)]) assert executed == [] @@ -37,7 +37,7 @@ def test_override_exhausts_recovery_then_runs_lowest_exact_split(monkeypatch): rank, "_try_cache_recovery", lambda *a, **k: events.append("recovery") or False ) executed = _recording_executor(monkeypatch, rank) - rank.dp_rank_forward([_request(i) for i in range(4)]) + rank.forward([_request(i) for i in range(4)]) assert events == ["recovery"] assert len(executed) == 4 assert all(plan.packed_tokens == 10 for plan in executed) @@ -47,7 +47,7 @@ def test_override_exhausts_recovery_then_runs_lowest_exact_split(monkeypatch): def test_fitting_split_is_unchanged_when_override_enabled(monkeypatch): rank = _oversized(monkeypatch, limit=20) executed = _recording_executor(monkeypatch, rank) - rank.dp_rank_forward([_request(i) for i in range(4)]) + rank.forward([_request(i) for i in range(4)]) assert [plan.packed_tokens for plan in executed] == [20, 20] @@ -60,7 +60,7 @@ def recover(*args, **kwargs): monkeypatch.setattr(rank, "_try_cache_recovery", recover) executed = _recording_executor(monkeypatch, rank) - rank.dp_rank_forward([_request(i) for i in range(4)]) + rank.forward([_request(i) for i in range(4)]) assert [plan.packed_tokens for plan in executed] == [20, 20] @@ -72,7 +72,7 @@ def test_override_selects_best_rung_not_last(monkeypatch): lambda *, packed_tokens, **k: {40: 400, 20: 20, 10: 30}[packed_tokens], ) executed = _recording_executor(monkeypatch, rank) - rank.dp_rank_forward([_request(i) for i in range(4)]) + rank.forward([_request(i) for i in range(4)]) assert [plan.packed_tokens for plan in executed] == [20, 20] @@ -81,7 +81,7 @@ def test_ep_unsupported_split_is_never_overridden(monkeypatch): monkeypatch.setattr(rank, "_expert_parallel_active", lambda: True) executed = _recording_executor(monkeypatch, rank) with pytest.raises(tr.TrainerRankMemoryError, match="expert parallelism"): - rank.dp_rank_forward([_request(i) for i in range(2)]) + rank.forward([_request(i) for i in range(2)]) assert not executed @@ -96,7 +96,7 @@ def reduce(values, *, op, sync_across_dp): monkeypatch.setattr(rank, "_recovery_reduce", reduce) with pytest.raises(tr.TrainerRankMemoryError): - rank.dp_rank_forward([_request(0)]) + rank.forward([_request(0)]) assert seen == [([1.0, 0.0, 0.0], "MIN", False)] @@ -114,9 +114,7 @@ def test_microbatch_override_keeps_minimum_wave_inputs(monkeypatch): def test_nonmemory_validation_is_not_overridden(monkeypatch): rank = _oversized(monkeypatch) with pytest.raises(ValueError, match="exceeds vocabulary"): - rank.dp_rank_forward( - [tr.ForwardInput(input_tokens=torch.tensor([1]), top_k=1000)] - ) + rank.forward([tr.ForwardInput(input_tokens=torch.tensor([1]), top_k=1000)]) def _reporting_rank(monkeypatch, tmp_path): @@ -164,10 +162,43 @@ def _records(root: Path): return [json.loads(path.read_text()) for path in root.rglob("*.json")] +def test_graph_placement_report_preserves_selected_local_admission( + monkeypatch, tmp_path +): + from art.trainer_rank import ForwardOptions, _planner_evidence + + rank, _, _ = _reporting_rank(monkeypatch, tmp_path) + monkeypatch.setattr(rank, "_available_memory_bytes", lambda *args: 10_000) + monkeypatch.setattr(rank, "_available_cpu_memory_bytes", lambda: 1_000_000) + request = tr.replace( + _request(0), + options=ForwardOptions(backward_state="replay", output_device="cpu"), + ) + decision = _planner_evidence.Decision("forward", sync_across_dp=False, owner=rank) + with _planner_evidence.scope(decision): + plan, check = rank._admit_graph_memory(rank._plan_flat_forward([request])) + assert check.fits and check.sample is not None + local = check.sample.local_required_bytes + assert local is not None + rank._begin_planner_observation( + plan, + tr.replace( + check, estimated_required_bytes=check.estimated_required_bytes + 123 + ), + ) + observation = rank._planner_observation + assert observation is not None and observation["comparable"] + snapshot = observation["replay"]() + assert observation["predicted"] == round(local / tr._MEMORY_SAFETY_FACTOR) + assert snapshot["local_admission_peak_bytes"] == local + assert snapshot["reduced_admission_peak_bytes"] == local + 123 + assert snapshot["requests"][0]["options"]["backward_state"] == "replay" + + def test_report_uses_local_prediction_before_profile_update(monkeypatch, tmp_path): rank, plan, counters = _reporting_rank(monkeypatch, tmp_path) rank._run_flat_plan_with_memory_tracking( - plan, check=tr._MemoryCheck(9999, 10000, True), context="dp_rank_forward" + plan, check=tr._MemoryCheck(9999, 10000, True), context="forward" ) counters["peak"] = 500 # caller backward exceeds the earlier forward peak rank._complete_planner_observation() @@ -190,7 +221,7 @@ def test_caught_backward_oom_is_reported_once_without_completed_peak( ): rank, plan, counters = _reporting_rank(monkeypatch, tmp_path) rank._run_flat_plan_with_memory_tracking( - plan, check=tr._MemoryCheck(220, 10000, True), context="dp_rank_forward" + plan, check=tr._MemoryCheck(220, 10000, True), context="forward" ) error = torch.cuda.OutOfMemoryError("caller backward allocation") rank.report_planner_oom(error) @@ -217,7 +248,7 @@ def fail(plan): monkeypatch.setattr(rank, "_execute_flat_plan", fail) with pytest.raises(tr.TrainerRankMemoryError) as raised: rank._run_flat_plan_with_memory_tracking( - plan, check=tr._MemoryCheck(220, 10000, True), context="dp_rank_forward" + plan, check=tr._MemoryCheck(220, 10000, True), context="forward" ) assert raised.value.__cause__ is error [record] = _records(tmp_path) @@ -230,7 +261,7 @@ def test_snapshot_failure_still_persists_minimal_oom(monkeypatch, tmp_path): rank, plan, _ = _reporting_rank(monkeypatch, tmp_path) monkeypatch.setattr(rank, "_plan_cost", lambda plan: 1 / 0) rank._run_flat_plan_with_memory_tracking( - plan, check=tr._MemoryCheck(220, 10000, True), context="dp_rank_forward" + plan, check=tr._MemoryCheck(220, 10000, True), context="forward" ) rank.report_planner_oom(torch.cuda.OutOfMemoryError("backward")) [record] = _records(tmp_path) @@ -241,7 +272,7 @@ def test_snapshot_failure_still_persists_minimal_oom(monkeypatch, tmp_path): def test_nonoom_does_not_mint_oom_report(monkeypatch, tmp_path): rank, plan, _ = _reporting_rank(monkeypatch, tmp_path) rank._run_flat_plan_with_memory_tracking( - plan, check=tr._MemoryCheck(220, 10000, True), context="dp_rank_forward" + plan, check=tr._MemoryCheck(220, 10000, True), context="forward" ) rank.report_planner_oom(RuntimeError("unrelated failure")) rank.discard_planner_observation() @@ -261,7 +292,7 @@ def test_delayed_report_keeps_original_request_options(monkeypatch, tmp_path, oo } output_bytes = plan.output_bytes rank._run_flat_plan_with_memory_tracking( - plan, check=tr._MemoryCheck(220, 10000, True), context="dp_rank_forward" + plan, check=tr._MemoryCheck(220, 10000, True), context="forward" ) # Caller/backward work can reuse the mutable ForwardInput before emission. request.top_k = 9 @@ -290,7 +321,7 @@ def test_delayed_report_keeps_original_request_options(monkeypatch, tmp_path, oo def test_changed_tokens_mark_replay_incomplete(monkeypatch, tmp_path): rank, plan, _ = _reporting_rank(monkeypatch, tmp_path) rank._run_flat_plan_with_memory_tracking( - plan, check=tr._MemoryCheck(220, 10000, True), context="dp_rank_forward" + plan, check=tr._MemoryCheck(220, 10000, True), context="forward" ) plan.groups[0].items[0].input_ids[0] = 123 rank._complete_planner_observation() @@ -308,7 +339,7 @@ def test_disabled_reports_do_not_sample_extra_counters(monkeypatch, tmp_path): lambda plan: pytest.fail("disabled observation built snapshot"), ) rank._run_flat_plan_with_memory_tracking( - plan, check=tr._MemoryCheck(220, 10000, True), context="dp_rank_forward" + plan, check=tr._MemoryCheck(220, 10000, True), context="forward" ) rank.finish_planner_observation() assert not _records(tmp_path) @@ -333,7 +364,7 @@ def plan(*args, **kwargs): monkeypatch.setattr(rank, "_plan_flat_forward", plan) with pytest.raises(tr.TrainerRankMemoryError): - rank.dp_rank_forward([_request(i) for i in range(4)]) + rank.forward([_request(i) for i in range(4)]) assert planned == [4, 4] # disabled peer keeps lower-bound pruning everywhere assert all(entry[1]["sync_across_dp"] is False for entry in seen) @@ -346,13 +377,13 @@ def test_fitting_path_adds_no_option_agreement_collective(monkeypatch): lambda *a, **k: pytest.fail("new normal-path collective"), ) _recording_executor(monkeypatch, rank) - rank.dp_rank_forward([_request(0)]) + rank.forward([_request(0)]) def test_iterator_close_preserves_pending_backward_oom_context(monkeypatch, tmp_path): rank, plan, _ = _reporting_rank(monkeypatch, tmp_path) monkeypatch.setattr(rank, "_available_memory_bytes", lambda sample=None: 10000) - iterator = rank.forward_micro_batches([_request(0)]) + iterator = rank.forward_batches([_request(0)]) next(iterator) iterator.close() assert rank._planner_observation is not None @@ -377,7 +408,7 @@ def execute(plan): monkeypatch.setattr(rank, "_execute_flat_plan", execute) rank._execute_split_plan_with_memory_tracking( - split, check=tr._MemoryCheck(440, 10000, True), context="dp_rank_forward" + split, check=tr._MemoryCheck(440, 10000, True), context="forward" ) counters["peak"] = 350 rank._complete_planner_observation() @@ -409,7 +440,7 @@ def test_volatile_budget_uses_smaller_materialized_candidate(monkeypatch): lambda: selected, lambda value: (value.plan, value.check), lambda value, check: tr.replace(value, check=check), - context="forward_micro_batches", + context="forward_batches", sync_across_dp=True, admit_refusal=lambda refusal: pytest.fail("lost original candidate mapping"), ) @@ -424,14 +455,10 @@ def test_interleaved_execution_scopes_preserve_both_oom_inputs(monkeypatch, tmp_ one, two = {}, {} check = tr._MemoryCheck(220, 10000, True) with rank.planner_observation_scope(one): - rank._run_flat_plan_with_memory_tracking( - first, check=check, context="dp_rank_forward" - ) + rank._run_flat_plan_with_memory_tracking(first, check=check, context="forward") with rank.planner_observation_scope(two): rank.discard_planner_observation() # new actor execution cannot erase one - rank._run_flat_plan_with_memory_tracking( - second, check=check, context="dp_rank_forward" - ) + rank._run_flat_plan_with_memory_tracking(second, check=check, context="forward") with rank.planner_observation_scope(one): rank.report_planner_oom(torch.cuda.OutOfMemoryError("first backward")) with rank.planner_observation_scope(two): @@ -461,7 +488,7 @@ def test_overlapping_scope_does_not_claim_completed_comparison(monkeypatch, tmp_ for context in (one, two): with rank.planner_observation_scope(context): rank._run_flat_plan_with_memory_tracking( - plan, check=check, context="dp_rank_forward" + plan, check=check, context="forward" ) for context in (one, two): with rank.planner_observation_scope(context): @@ -477,7 +504,7 @@ def test_sequential_scopes_keep_independent_completed_peaks(monkeypatch, tmp_pat counters.update(allocated=100, peak=100) with rank.planner_observation_scope({}): rank._run_flat_plan_with_memory_tracking( - plan, check=tr._MemoryCheck(220, 10000, True), context="dp_rank_forward" + plan, check=tr._MemoryCheck(220, 10000, True), context="forward" ) rank._complete_planner_observation() assert len(_records(tmp_path)) == 2 @@ -511,7 +538,7 @@ def world_price(required, *, sync_across_dp): lambda: next(searches), lambda value: (value.plan, value.check), lambda value, check: tr.replace(value, check=check), - context="forward_micro_batches", + context="forward_batches", sync_across_dp=True, admit_refusal=lambda refused: tr.replace( candidate, plan=refused.plan, check=refused.check, stats_global_count=1 @@ -534,7 +561,7 @@ def test_dp_disagreeing_fallback_widths_refuse_before_execution(monkeypatch): lambda: refusal, lambda value: (value.plan, value.check), lambda value, check: value, - context="forward_micro_batches", + context="forward_batches", sync_across_dp=True, admit_refusal=lambda r: tr._CandidateMicroBatch( items, (0,), r.plan, r.check, 1, 0, True @@ -551,7 +578,7 @@ def test_dp_forward_reports_at_forward_boundary_not_later_optimizer( "_plan_admissible_forward", lambda *a, **k: (plan, tr._MemoryCheck(220, 10000, True)), ) - rank.dp_rank_forward([_request(0)]) + rank.forward([_request(0)]) [record] = _records(tmp_path) assert record["observed_peak_bytes"] == 250 assert record["phase"] == "forward" @@ -568,7 +595,7 @@ def test_closed_forward_keeps_backward_oom_without_false_partial_peak( ): rank, plan, counters = _reporting_rank(monkeypatch, tmp_path) rank._execute_admitted_plan( - plan, check=tr._MemoryCheck(220, 10000, True), context="dp_rank_forward" + plan, check=tr._MemoryCheck(220, 10000, True), context="forward" ) counters["peak"] = 10000 rank.report_planner_oom(torch.cuda.OutOfMemoryError("later backward")) @@ -592,7 +619,7 @@ def test_execution_finish_does_not_complete_abandoned_microbatch_window( ): rank, plan, counters = _reporting_rank(monkeypatch, tmp_path) rank._run_flat_plan_with_memory_tracking( - plan, check=tr._MemoryCheck(220, 10000, True), context="forward_micro_batches" + plan, check=tr._MemoryCheck(220, 10000, True), context="forward_batches" ) counters["peak"] = 10000 rank.finish_planner_observation() @@ -607,7 +634,7 @@ def test_abandoned_execution_does_not_poison_future_measurements(monkeypatch, tm rank._run_flat_plan_with_memory_tracking( plan, check=tr._MemoryCheck(220, 10000, True), - context="forward_micro_batches", + context="forward_batches", ) assert rank._planner_active_observations del abandoned @@ -616,7 +643,7 @@ def test_abandoned_execution_does_not_poison_future_measurements(monkeypatch, tm counters.update(allocated=100, peak=100) with rank.planner_observation_scope({}): rank._execute_admitted_plan( - plan, check=tr._MemoryCheck(220, 10000, True), context="dp_rank_forward" + plan, check=tr._MemoryCheck(220, 10000, True), context="forward" ) [record] = _records(tmp_path) assert record["observed_peak_bytes"] == 250 @@ -635,7 +662,7 @@ def execute(plan): monkeypatch.setattr(rank, "_execute_flat_plan", execute) with pytest.raises(tr.TrainerRankMemoryError) as raised: rank._run_flat_plan_with_memory_tracking( - plan, check=tr._MemoryCheck(220, 10000, True), context="dp_rank_forward" + plan, check=tr._MemoryCheck(220, 10000, True), context="forward" ) assert raised.value.__cause__ is error rank.finish_planner_observation() # diagnostic cleanup must also be auxiliary @@ -644,7 +671,7 @@ def execute(plan): def test_complete_diagnostic_failure_and_cancellation_semantics(monkeypatch, tmp_path): rank, plan, _ = _reporting_rank(monkeypatch, tmp_path) rank._run_flat_plan_with_memory_tracking( - plan, check=tr._MemoryCheck(220, 10000, True), context="dp_rank_forward" + plan, check=tr._MemoryCheck(220, 10000, True), context="forward" ) del rank._planner_overlap_generation rank._complete_planner_observation() # ordinary diagnostic AttributeError is contained @@ -666,13 +693,13 @@ def test_closed_dp_context_invalidates_other_executions_open_caller_window( one, two = {}, {} with rank.planner_observation_scope(one): rank._execute_admitted_plan( - plan, check=tr._MemoryCheck(220, 10000, True), context="dp_rank_forward" + plan, check=tr._MemoryCheck(220, 10000, True), context="forward" ) [forward_report] = _records(tmp_path) assert forward_report["observed_peak_bytes"] == 250 assert one["observation"]["window_open"] is False - iterator = rank.forward_micro_batches([_request(1)]) + iterator = rank.forward_batches([_request(1)]) with rank.planner_observation_scope(two): next(iterator) # A resumes caller backward while B's yielded microbatch interval is open. diff --git a/tests/unit/test_trainer_rank_planner_reports.py b/tests/unit/test_trainer_rank_planner_reports.py index a58adfbc1..106e8d722 100644 --- a/tests/unit/test_trainer_rank_planner_reports.py +++ b/tests/unit/test_trainer_rank_planner_reports.py @@ -357,11 +357,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() assert "_memory.py" in reports._source_files() assert "_micro_batch_planner.py" in reports._source_files() @@ -491,7 +494,8 @@ def test_actual_emitted_split_recomputes_frozen_runtime_facts( @pytest.mark.parametrize("slots", [[], [[True, [[3, 8], []]], [False, []]]]) -def test_signature_json_roundtrip_is_immutable(slots): +@pytest.mark.parametrize("placements", [[], [["gpu", "model"], ["replay", "cpu"]]]) +def test_signature_json_roundtrip_is_immutable(slots, placements): from art.trainer_rank._impl import _MemorySignature values: dict[str, Any] = dict( @@ -502,6 +506,7 @@ def test_signature_json_roundtrip_is_immutable(slots): grad_enabled=True, grad_modes=[True], slot_shapes=slots, + memory_placement=placements, ) old = dict(values) for name in ("topology", "planner_coefficients", "request_mix", "grad_modes"): @@ -513,6 +518,7 @@ def test_signature_json_roundtrip_is_immutable(slots): assert key.slot_shapes == tuple( (enabled, tuple(map(tuple, shapes))) for enabled, shapes in slots ) + assert key.memory_placement == tuple(map(tuple, placements)) @pytest.mark.parametrize( diff --git a/tests/unit/test_trainer_rank_planning_status.py b/tests/unit/test_trainer_rank_planning_status.py index b684aadce..7e43eb3e0 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", @@ -214,8 +207,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_profile_warm.py b/tests/unit/test_trainer_rank_profile_warm.py index f1fa1c38e..33612af68 100644 --- a/tests/unit/test_trainer_rank_profile_warm.py +++ b/tests/unit/test_trainer_rank_profile_warm.py @@ -78,7 +78,7 @@ def test_the_first_plans_one_time_costs_do_not_price_later_waves(): def test_forward_only_observations_never_set_the_warm_fit(): - """Split children and dp_rank_forward observe forward only; a flat wave's + """Split children and forward observe forward only; a flat wave's caller phase also includes its backward.""" r = rank() first, larger = _plans(r) @@ -261,7 +261,7 @@ def run(plan, **_kwargs): monkeypatch.setattr(r, "_run_flat_plan_with_memory_tracking", run) _packed_budget(monkeypatch, r, 20) items = [[_request(0)], [_request(m) for m in range(1, 5)], [_request(5)]] - batches = r.forward_micro_batches(items) + batches = r.forward_batches(items) next(batches) assert next(batches).stats.subforward_count > 1 (signature,) = r._memory_profiles @@ -319,7 +319,7 @@ def test_a_width_accepted_on_the_no_sharing_bound_never_executes_above_it( None, ), ) - (batch,) = list(r.forward_micro_batches(items)) + (batch,) = list(r.forward_batches(items)) # Accepted on the no-sharing bound, which priced more tokens than ran. assert batch.stats.packed_tokens == plan.packed_tokens assert batch.stats.estimated_required_bytes > own @@ -413,7 +413,7 @@ def run(plan, **_kwargs): monkeypatch.setattr(r, "_run_flat_plan_with_memory_tracking", run) _packed_budget(monkeypatch, r, 10) items = [[_request(m)] for m in range(3)] - batches = r.forward_micro_batches(items) + batches = r.forward_batches(items) for index, batch in enumerate(batches): (signature,) = r._memory_profiles state["peak"] += 200 * batch.stats.packed_tokens # the caller's backward diff --git a/tests/unit/test_trainer_rank_recompute_memory.py b/tests/unit/test_trainer_rank_recompute_memory.py index 1d6eca3cb..a414586d8 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 42ad7b4b9..9e0326c19 100644 --- a/tests/unit/test_trainer_rank_recovery_slots_distributed.py +++ b/tests/unit/test_trainer_rank_recovery_slots_distributed.py @@ -65,142 +65,134 @@ 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") - if mode.startswith("handoff"): - # The second real peer has no local groups, but must join every phase. + with gloo_group(index, f"file://{directory}/rendezvous", timeout=3): + rank = TrainerRank.__new__(TrainerRank) + rank.device = torch.device("cpu") + if mode.startswith("handoff"): + # The second real peer has no local groups, but must join every phase. + plan = SimpleNamespace( + groups=[SimpleNamespace(grad_enabled=True)] if index == 0 else [] + ) + original = ( + KeyboardInterrupt("original forward cancellation") + if mode == "handoff-cancel" + else RuntimeError("original forward failure") + ) + error = barrier_error = None + caught = None + try: + rank._release_cached_memory_for_backward( + plan, error=original if index == 0 and mode != "handoff" else None + ) + except BaseException as exc: + caught = exc + if mode == "handoff": + valid = caught is None + else: + valid = ( + caught is original + if index == 0 + else ( + isinstance(caught, RuntimeError) + and "another rank" in str(caught) + ) + ) + if not valid or rank._recovery_state().owner is not None: + error = "handoff disposition or owner differs" + 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(dict(error=error, barrier_error=barrier_error, ensures=0)) + ) + return + 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( - groups=[SimpleNamespace(grad_enabled=True)] if index == 0 else [] + groups=(), # This slot-only fixture has no retained/head/GDN groups. + 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,), + ), ) - original = ( - KeyboardInterrupt("original forward cancellation") - if mode == "handoff-cancel" - else RuntimeError("original forward failure") + 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 - caught = None try: - rank._release_cached_memory_for_backward( - plan, error=original if index == 0 and mode != "handoff" else None - ) + rank._plan_admissible_forward([request], checkpoint=None, context="forward") except BaseException as exc: - caught = exc - if mode == "handoff": - valid = caught is None - else: - valid = ( - caught is original - if index == 0 - else ( - isinstance(caught, RuntimeError) and "another rank" in str(caught) - ) - ) - if not valid or rank._recovery_state().owner is not None: - error = "handoff disposition or owner differs" + 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(dict(error=error, barrier_error=barrier_error, ensures=0)) - ) - dist.destroy_process_group() - return - 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( - groups=(), # This slot-only fixture has no retained/head/GDN groups. - 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="dp_rank_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, + 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/test_trainer_rank_release_completion.py b/tests/unit/test_trainer_rank_release_completion.py new file mode 100644 index 000000000..f0d6356bb --- /dev/null +++ b/tests/unit/test_trainer_rank_release_completion.py @@ -0,0 +1,107 @@ +"""Pending callback cleanup remains observable and ordered after failure.""" + +import asyncio +from typing import Any, cast + +import pytest +from test_trainer_rank_commands import _Rank + +from art.trainer_rank._commands import _Executor, _Release, join_rank_callback_release + + +@pytest.mark.parametrize("peer_buffers", [False, True]) +async 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 + ) + + 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]) +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 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_rng.py b/tests/unit/test_trainer_rank_rng.py new file mode 100644 index 000000000..161144c74 --- /dev/null +++ b/tests/unit/test_trainer_rank_rng.py @@ -0,0 +1,658 @@ +from __future__ import annotations + +import asyncio +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 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 _coordinate_call, join_rank_callback_release +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( + "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 "model failure" not 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, + (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), + 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), + ): + 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)) +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]) 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..9ff8e850a --- /dev/null +++ b/tests/unit/test_trainer_rank_slot_graph_lifetime.py @@ -0,0 +1,365 @@ +from __future__ import annotations + +import asyncio +from contextlib import nullcontext +import gc +from typing import Literal +import weakref + +import pytest +from test_trainer_rank_custom_tensors import _real_lora_trainer, _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, TopK +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 _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)) + 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), + ) + ], + ) + 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(), + 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, detached, copies = [], [], [], [], [] + to = torch.Tensor.to + + def fail_snapshot(value, *args, **kwargs): + 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 ()) + 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"), + ) + del value + raise primary + 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 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) + 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 + trainer.backward(output) + torch.testing.assert_close(weight.grad, torch.tensor(4.0)) + 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 +): + 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 c9d5b3c9f..1a8f6f74a 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. @@ -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,6 +40,7 @@ import pytest import torch +from trainer_rank_test_support import _FakeGPT, _packed_budget from art.trainer_rank import ( ForwardInput, @@ -55,7 +55,6 @@ _PACKED_PRICED_LOGICAL_ROW_BYTES, Unset, _FlatForwardPlan, - _MemoryCheck, _MemoryProfile, _SplitForwardPlan, ) @@ -66,21 +65,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()], @@ -107,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]: @@ -157,7 +121,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) @@ -169,7 +133,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 @@ -191,7 +155,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 @@ -211,7 +175,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 @@ -242,7 +206,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 @@ -273,7 +237,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 == [] @@ -297,13 +261,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) @@ -314,7 +278,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 @@ -340,10 +304,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 @@ -371,7 +335,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 @@ -405,11 +369,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 == [] @@ -877,7 +841,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 @@ -905,7 +869,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 @@ -922,9 +886,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 @@ -948,10 +912,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 @@ -966,9 +930,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) @@ -977,7 +939,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 = 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) monkeypatch.setattr(rank, "_prepare_packed_forward", lambda _packed: None) @@ -994,23 +958,35 @@ 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] + ) + + 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) 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 4ed9aff00..40a429ca2 100644 --- a/tests/unit/test_trainer_rank_split_peak.py +++ b/tests/unit/test_trainer_rank_split_peak.py @@ -4,40 +4,21 @@ 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 as _active_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), - ), - ) - ) + 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): @@ -157,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, @@ -170,6 +155,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() @@ -218,7 +205,7 @@ def test_completed_iterator_preserves_caller_peak_for_next_admission(monkeypatch rank, requests, counters = _counter_split(monkeypatch) 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] @@ -229,7 +216,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 == [] @@ -238,7 +225,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) @@ -260,7 +247,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..780bfc86a --- /dev/null +++ b/tests/unit/test_trainer_rank_tensors.py @@ -0,0 +1,996 @@ +from __future__ import annotations + +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 + +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): + 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(): + 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(): + 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): + 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): + 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(): + 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): + @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 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_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_tp_floor.py b/tests/unit/test_trainer_rank_tp_floor.py index 6d5f4d00d..b79d6acaa 100644 --- a/tests/unit/test_trainer_rank_tp_floor.py +++ b/tests/unit/test_trainer_rank_tp_floor.py @@ -7,11 +7,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, recompute_model from art.trainer_rank import TrainerRank from art.trainer_rank._impl import _MemorySignature @@ -28,44 +28,8 @@ 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 = 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.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 - 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), - ), - ) - ) + 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( r._geometry, @@ -138,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": diff --git a/tests/unit/test_trainer_rank_validation.py b/tests/unit/test_trainer_rank_validation.py index 0d85e276a..e6a6d17a3 100644 --- a/tests/unit/test_trainer_rank_validation.py +++ b/tests/unit/test_trainer_rank_validation.py @@ -3,8 +3,7 @@ import asyncio from collections.abc import Iterable from concurrent.futures import ThreadPoolExecutor -from dataclasses import dataclass, replace -from datetime import timedelta +from dataclasses import dataclass, fields, replace import gc from importlib.util import find_spec import inspect @@ -20,7 +19,8 @@ import pytest 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, spawn_and_join from art.megatron.prefix_tree_packing import prefix_tree_pack from art.trainer_rank import ( @@ -67,7 +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: @@ -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 @@ -213,6 +181,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, @@ -220,11 +198,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 @@ -263,7 +237,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 +249,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 +266,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 +299,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 +311,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 +323,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 +338,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 +364,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 +398,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 +417,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 +428,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 +439,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 +453,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 +475,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 +535,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 +769,7 @@ def test_snapshot_disposal_is_not_public() -> None: "parameter", "buffer", "forward", - "forward_micro_batches", + "forward_batches", "optim_step", "save", "export_lora", @@ -833,10 +799,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 +1389,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( @@ -1455,6 +1428,21 @@ def load( snapshot_prepared_checkpoint(trainer, source, "loaded") +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": rank, + "lora_alpha": alpha, + "target_modules": list(target_modules), + } + + def test_checkpoint_export_requires_retained_adapter_config() -> None: trainer = TrainerRank(_runtime()) with pytest.raises(TrainerRankSlotStateError, match="unloaded checkpoint"): @@ -1486,12 +1474,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, ) @@ -1529,12 +1512,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") @@ -1652,11 +1630,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), @@ -2044,11 +2018,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 @@ -2060,28 +2039,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( @@ -2108,13 +2136,14 @@ 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 +@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 @@ -2126,6 +2155,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 @@ -2135,11 +2165,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( @@ -2177,12 +2213,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 +2234,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 @@ -2234,24 +2260,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=[], @@ -2268,17 +2283,12 @@ 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( - 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") @@ -2287,6 +2297,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, @@ -2313,96 +2329,110 @@ 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", ) - - setattr( - checkpoint_module, - "_load_adapter", - ( - lambda *_args: ( - (_ for _ in ()).throw(RuntimeError("injected snapshot read")) - if phase == "read" and rank == 1 - else {} - ) - ), - ) - 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) - ) - ) - ), - ) - setattr( - checkpoint_module, - "_commit_slot", - ( - lambda *_args: ( - (_ for _ in ()).throw(RuntimeError("injected rank-zero commit")) - if phase == "commit" and rank == 0 - else None - ) - ), - ) - - with pytest.raises(RuntimeError, match="injected|Another rank failed"): - checkpoint_module.load_checkpoint(trainer, source, "student") + 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) + + 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" + ) 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 ) - 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() + 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) + 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("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: - 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") @@ -2422,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( @@ -2453,11 +2475,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() @@ -2470,6 +2488,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,11 +2496,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 @@ -2509,6 +2524,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()) @@ -2557,11 +2609,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( @@ -2581,11 +2629,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() @@ -2674,11 +2718,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] ) @@ -2720,11 +2760,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", @@ -2783,11 +2819,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( @@ -2807,11 +2839,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: @@ -2877,15 +2905,9 @@ 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 - ) - ) - monkeypatch.setattr( - trainer, - "_reduce_dynamic_grads", - lambda params, **_kwargs: tuple(param.grad.float() for param in params), + trainer.forward_batches([_target_request(1)], checkpoint="saved", no_grad=True) ) + _use_strict_local_gradients(trainer, monkeypatch) before_student = student.detach().clone() before_snapshot = snapshot.detach().clone() @@ -2931,11 +2953,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( @@ -3134,7 +3152,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"): @@ -3234,15 +3254,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])) @@ -3267,33 +3280,24 @@ 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_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 @@ -3303,11 +3307,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()) @@ -3317,7 +3321,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"]]]]], @@ -3327,7 +3331,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: @@ -3335,7 +3339,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 ) ) @@ -3351,7 +3355,7 @@ def test_forward_micro_batches_uses_deterministic_dp_windows( @pytest.mark.parametrize( "operation", [ - "dp_reduce", + "reduce", "optim_step", "parameter", "module", @@ -3364,8 +3368,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( @@ -3373,7 +3377,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 @@ -3383,7 +3387,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), @@ -3396,9 +3400,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"): @@ -3413,7 +3417,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() @@ -3440,9 +3444,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) @@ -3453,14 +3457,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 @@ -3496,9 +3493,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) @@ -3506,11 +3503,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() @@ -3526,40 +3523,32 @@ 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( + 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_micro_batches_syncs_fit_decision_across_dp( +def test_forward_batches_syncs_fit_decision_across_dp( monkeypatch: pytest.MonkeyPatch, ) -> None: trainer = TrainerRank(_runtime()) @@ -3575,13 +3564,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()) @@ -3592,9 +3581,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"]]]]], @@ -3602,7 +3591,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()) @@ -3618,9 +3607,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 @@ -3628,7 +3615,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()) @@ -3650,7 +3637,7 @@ def run(*_args, **_kwargs): ), ) - batches = trainer.forward_micro_batches([_target_request(1)]) + batches = trainer.forward_batches([_target_request(1)]) next(batches) assert profiles == [] @@ -3663,7 +3650,7 @@ def run(*_args, **_kwargs): @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()) @@ -3693,7 +3680,7 @@ def profile(*_args, **_kwargs): 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) @@ -3727,7 +3714,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()) @@ -3745,7 +3732,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()) @@ -3765,15 +3752,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()) @@ -3793,10 +3778,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()) @@ -3813,7 +3798,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)) @@ -3821,7 +3806,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: @@ -3887,6 +3872,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..d93455b36 --- /dev/null +++ b/tests/unit/test_trainer_rank_versions.py @@ -0,0 +1,471 @@ +from __future__ import annotations + +from contextlib import nullcontext +import gc +from pathlib import Path +from types import SimpleNamespace +import weakref + +import pytest +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 art.trainer_rank import TrainerRank, TrainerRankSlotStateError +from art.trainer_rank._commands import _coordinate_call +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 + + +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, +) -> None: + 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(): + 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, old = _snapshot_trainer() + 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, old = _snapshot_trainer() + + 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, 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"): + 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, 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"): + 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, snapshot = _snapshot_trainer() + 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, old = _snapshot_trainer() + + 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, old = _snapshot_trainer() + 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, old = _snapshot_trainer() + 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: + with gloo_group(rank, rendezvous): + 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 + + +def test_divergent_optimizer_provenance_rejects_collectively_before_mutation( + tmp_path: Path, +) -> None: + spawn_and_join( + _divergent_version_worker, + args=(f"file://{tmp_path / 'versions'}",), + timeout=90, + failure="Divergent version preflight did not complete collectively", + ) + + +def test_transaction_coalesces_many_children_but_keeps_each_origin() -> None: + trainer, current, old = _snapshot_trainer() + 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, old = _snapshot_trainer() + 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: + with gloo_group(rank, rendezvous, timeout=15): + 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) + _coordinate_call(validate, group=None) + + 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 + + +@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: + spawn_and_join( + _transaction_exit_failure_worker, + args=(f"file://{tmp_path / 'exit'}", nested), + timeout=60, + failure="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 51b0c313a..6bdff3e52 100644 --- a/tests/unit/test_trainer_rank_weird_shapes.py +++ b/tests/unit/test_trainer_rank_weird_shapes.py @@ -1,11 +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, @@ -32,21 +34,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. @@ -58,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) @@ -89,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. @@ -238,7 +218,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()) @@ -251,17 +231,10 @@ def test_forward_micro_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_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 @@ -272,7 +245,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()) @@ -282,16 +255,9 @@ def test_forward_micro_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_micro_batches(inputs) + generator = rank.forward_batches(inputs) first = next(generator) assert first.stats.global_count == 4 @@ -336,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 @@ -364,7 +323,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]: @@ -389,7 +348,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 @@ -419,21 +378,14 @@ 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 # (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] @@ -473,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 @@ -501,13 +443,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)) @@ -529,7 +471,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] @@ -651,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. @@ -667,24 +602,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, @@ -692,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) @@ -711,10 +639,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 @@ -914,7 +842,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()) @@ -948,7 +876,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] == [ @@ -1010,13 +938,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", @@ -1038,9 +966,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/tests/unit/trainer_rank_test_support.py b/tests/unit/trainer_rank_test_support.py new file mode 100644 index 000000000..b5cbfda55 --- /dev/null +++ b/tests/unit/trainer_rank_test_support.py @@ -0,0 +1,213 @@ +"""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 +import time +from types import ModuleType, SimpleNamespace +from typing import TYPE_CHECKING, Any + +import pytest +import torch +import torch.distributed as dist +import torch.multiprocessing as mp + +if TYPE_CHECKING: + from art.megatron.train import TrainingRuntime + from art.trainer_rank import TrainerRank + + +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 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 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, + 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.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, + *, + 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 _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( + backend, + init_method=rendezvous, + rank=rank, + world_size=world_size, + timeout=None if timeout is None else timedelta(seconds=timeout), + ) + try: + yield + finally: + 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.""" + 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 pytest.MonkeyPatch.context() as modules: + modules.setitem(sys.modules, "megatron", megatron) + modules.setitem(sys.modules, "megatron.core", core) + yield getattr(core, "parallel_state") + + +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: + 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 diff --git a/tests/unit/trajectories/test_literal_thinking_off.py b/tests/unit/trajectories/test_literal_thinking_off.py index 66c434856..959ddce74 100644 --- a/tests/unit/trajectories/test_literal_thinking_off.py +++ b/tests/unit/trajectories/test_literal_thinking_off.py @@ -234,7 +234,56 @@ 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") - _outcome(history, 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 = ( + "<|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 + 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, + # 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], + ) + ) # 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 @@ -360,6 +409,18 @@ def apply_chat_template( def test_literal_next_turn_preserves_preceding_length_stop_boundary( monkeypatch: pytest.MonkeyPatch, +) -> None: + _check_literal_next_turn_boundary(monkeypatch, recorded=True) + + +def test_literal_next_turn_renderer_preserves_preceding_length_stop_boundary( + monkeypatch: pytest.MonkeyPatch, +) -> None: + _check_literal_next_turn_boundary(monkeypatch, recorded=False) + + +def _check_literal_next_turn_boundary( + monkeypatch: pytest.MonkeyPatch, *, recorded: bool ) -> None: tokenizer = _NewlineRunTokenizer() messages = [ @@ -423,24 +484,50 @@ def observe(*args: Any, **kwargs: Any) -> tr.TokenizedHistory | None: observed.append((boundary, result)) return result + def tokenize() -> tr.TokenizedHistory: + if recorded: + return history.tokenize(tokenizer=tokenizer) + # Isolate the existing renderer stage from the new default native policy. + return _tokenize._tokenize_chat_view( + history, + base_model=None, + tokenizer=tokenizer, + chat_template=None, + chat_template_kwargs=None, + _projection_matches=True, + _recorded_boundaries=False, + ) + monkeypatch.setattr(_tokenize, "_tokenize_exact_projected_chat_history", observe) with monkeypatch.context() as patch: patch.setattr( _tokenize, "chat_template_with_preserved_thinking", lambda value: value ) - _outcome(history, tokenizer) + try: + tokenize() + except ValueError: + pass boundary, old_exact = observed[0] - # The final recorded body no longer needs a reconstructed terminal tail. - # Disabling literal normalization cannot invalidate the proved earlier gap. - assert old_exact is not None - assert list(boundary.tail + boundary.following) == native_boundary + if recorded: + # The final recorded body no longer needs a reconstructed terminal tail. + # Disabling literal normalization cannot invalidate the proved earlier gap. + assert old_exact is not None + assert list(boundary.tail + boundary.following) == native_boundary + else: + stored = list(boundary.tail + boundary.following) + assert old_exact is None + assert len(native_boundary) - len(stored) == 2 + assert stored[:-1] == native_boundary[:-3] + assert tokenizer.decode(stored[-1:]) == "\n" + assert tokenizer.decode(native_boundary[-3:]) == "\n\n\n\n" observed.clear() - value = history.tokenize(tokenizer=tokenizer) + value = tokenize() fixed_boundary, fixed_exact = observed[0] assert fixed_exact is value - assert value.tokens == old_exact.tokens - assert value.flags == old_exact.flags + if old_exact is not None: + assert value.tokens == old_exact.tokens + assert value.flags == old_exact.flags assert list(fixed_boundary.tail + fixed_boundary.following) == native_boundary assert ( value.tokens[: len(last["prompt_token_ids"]) + len(last["token_ids"])] 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: 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" },