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" },