Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions .github/workflows/prek.yml
Original file line number Diff line number Diff line change
Expand Up @@ -233,6 +233,7 @@ jobs:
tests/unit/test_trainer_rank_slot_memory.py \
tests/unit/test_trainer_rank_moe_memory.py \
tests/unit/test_trainer_rank_head_memory.py \
tests/unit/test_trainer_rank_eager_head_memory.py::test_optional_stats_refusal_keeps_capacity_separate_from_lower_bound \
tests/unit/test_trainer_rank_mixed_head_memory.py \
tests/unit/test_trainer_rank_ignored_mixed_head.py \
tests/unit/test_trainer_rank_pending_memory.py \
Expand All @@ -257,6 +258,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_eager_head_memory.py::test_optional_stats_refusal_keeps_capacity_separate_from_lower_bound \
--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 \
Expand Down
32 changes: 23 additions & 9 deletions src/art/trainer_rank/_impl.py
Original file line number Diff line number Diff line change
Expand Up @@ -3804,17 +3804,21 @@ def _group_head_workspace_bytes(
positions: Sequence[torch.Tensor] | None = None,
lower_bound: bool = False,
) -> int:
"""One logits buffer, or logits + both dense target-backward gradients.
"""Partial dense head component: eager statistics or no-grad logits copies.

The supported head path overlaps indexing and statistics gradients
with recomputed logits; cold library workspaces remain outside this
component. Pair each group's mode with its own projected rows.
Capacity reserves the eager path even when optional Triton may succeed.
Its BF16 logits, FP32 conversion, subtraction and exp overlap. This is
not a bound for row vectors, inter-chunk liveness or library workspaces.
Rejection lower bounds retain only the unconditional dense components.
"""
dense = self._head_workspace_bytes(rows)
if (
not dense
or not grad_enabled
or not any(request.target_tokens is not None for request in requests)
needs_statistics = any(
request.target_tokens is not None or request.top_k is not None
for request in requests
)
if not dense or (
not needs_statistics
and (grad_enabled or lower_bound or not any(r.logits for r in requests))
):
return dense
from megatron.core.models.common.language_module.language_module import (
Expand All @@ -3829,8 +3833,18 @@ def _group_head_workspace_bytes(
and scale.__func__ is LanguageModule._scale_logits
and getattr(model.config, "use_mup", None) is False
):
if not lower_bound:
# need_log_z is group-wide, including logits-only chunks and
# chunks overlapping ignored labels. A short final chunk can
# also take the eager path; optional success is not guaranteed.
# Without statistics, local logits and both indexed copies
# overlap. Requested output storage is charged separately.
return (7 if needs_statistics else 3) * dense
if not grad_enabled or not any(
request.target_tokens is not None for request in requests
):
return dense
# IndexBackward's dense result overlaps saved logits and grad_logits.
# The FP32 fallback already exceeds this three-buffer component.
target_dense = (
self._head_workspace_bytes(
self._head_target_chunk_rows(
Expand Down
4 changes: 2 additions & 2 deletions tests/unit/test_trainer_rank_admission_inputs.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,7 +49,7 @@ def test_shared_head_lower_and_plan_keep_all_keywords(monkeypatch):
0,
2 * (188416 + 4 * 2048 * 2),
)
assert rank._plan_head_workspace_bytes(plan) == 2 * 248320 * 2
assert rank._plan_head_workspace_bytes(plan) == 7 * 2 * 248320 * 2
calls = record_prices(monkeypatch, rank)
lower = rank._split_chunk_lower_cost(
requests, tuple(item.input_tokens for item in requests), checkpoint=Unset
Expand All @@ -69,7 +69,7 @@ def test_shared_head_lower_and_plan_keep_all_keywords(monkeypatch):
assert_plan_values(rank, plan, values)
# All shape/floor inputs remain available together, even though the dense
# head dominates this two-row source floor and sharing changes logical rows.
assert cost.required == int((plan.output_bytes + 2 * 248320 * 2) * 1.1)
assert cost.required == int((plan.output_bytes + 7 * 2 * 248320 * 2) * 1.1)


def test_cp_gdn_segments_groups_and_retained_tokens_reach_exact_search(monkeypatch):
Expand Down
215 changes: 215 additions & 0 deletions tests/unit/test_trainer_rank_eager_head_memory.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,215 @@
"""Eager source allocations and admission; no CUDA lifetime/whole-plan bound."""

from types import SimpleNamespace
from typing import cast
import weakref

import pytest
from test_trainer_rank_head_memory import rank, request
from test_trainer_rank_head_recompute import _Head, _patch_local_head
import torch

from art import trainer_rank
from art.trainer_rank import ForwardInput, TrainerRank, _impl


@pytest.mark.parametrize("rows", [4, 5])
def test_no_grad_logits_assignment_has_output_and_three_distinct_cpu_storages(
monkeypatch, rows
):
_patch_local_head(monkeypatch)
monkeypatch.setattr(_impl, "_HEAD_CHUNK_TOKENS", 4)
weight = torch.arange(85, dtype=torch.bfloat16).reshape(17, 5) / 100
hidden = torch.arange(rows * 5, dtype=torch.bfloat16).reshape(rows, 5) / 100
model = SimpleNamespace(
output_layer=_Head(weight),
vocab_size=17,
share_embeddings_and_output_weights=False,
_scale_logits=lambda value: value,
)
r = object.__new__(TrainerRank)
r.runtime = SimpleNamespace(model=[model])
original_select = torch.Tensor.index_select
original_setitem = torch.Tensor.__setitem__
selections, observed = [], []

def select(self, *args, **kwargs):
result = original_select(self, *args, **kwargs)
if self.ndim == 2 and self.shape[1] == 17:
selections.append((weakref.ref(self), weakref.ref(result)))
return result

def assign(self, key, value):
if self.ndim == 2 and self.shape[1] == 17:
# Weak references do not extend any allocation across callbacks.
assert len(selections) == 2
local, selected = (ref() for ref in selections[0])
chunk, rhs = (ref() for ref in selections[1])
assert rhs is value and selected is not None
assert (
selected.untyped_storage().data_ptr()
== chunk.untyped_storage().data_ptr()
)
tensors = (self, local, chunk, rhs)
assert all(t is not None and t.dtype == torch.bfloat16 for t in tensors)
assert len({t.untyped_storage().data_ptr() for t in tensors}) == 4
observed.append(sum(t.untyped_storage().nbytes() for t in tensors))
selections.clear()
return original_setitem(self, key, value)

monkeypatch.setattr(torch.Tensor, "index_select", select)
monkeypatch.setattr(torch.Tensor, "__setitem__", assign)
req = ForwardInput(input_tokens=torch.arange(rows), logits=True, no_grad=True)
positions = (torch.arange(rows),)
with torch.no_grad():
outputs = r._project_head(
[r._forward_item(req)],
SimpleNamespace(
positions_by_item=positions, source_positions_by_item=positions
),
hidden,
)
assert observed == [
(rows + 3 * min(4, rows - start)) * 17 * 2 for start in range(0, rows, 4)
]
torch.testing.assert_close(outputs[0].logits, hidden @ weight.T)


@pytest.mark.parametrize("mode", ["short", "disabled", "error", "strict"])
def test_optional_stats_refusal_keeps_capacity_separate_from_lower_bound(
monkeypatch, mode
):
calls = []

def kernel(*args, **kwargs):
calls.append(True)
raise RuntimeError("optional kernel failed")

# Exercise the real dispatch without a CUDA tensor or a Triton import.
monkeypatch.setattr(
trainer_rank,
"topk",
SimpleNamespace(local_logsumexp_stats=kernel),
raising=False,
)
monkeypatch.setenv("ART_TRAINER_RANK_TRITON_MIN_ROWS", "64")
monkeypatch.setenv(
"ART_TRAINER_RANK_TRITON_TOPK",
"0" if mode == "disabled" else "strict" if mode == "strict" else "1",
)
rows = 63 if mode == "short" else 64
probe = cast(torch.Tensor, SimpleNamespace(is_cuda=True, shape=(rows, 17)))
if mode == "strict":
with pytest.raises(RuntimeError, match="optional kernel failed"):
_impl._try_triton_stats("local_logsumexp_stats", probe)
else:
assert _impl._try_triton_stats("local_logsumexp_stats", probe) is None
assert len(calls) == int(mode in {"error", "strict"})
r = rank()
dense = rows * 248320 * 2
for grad in (False, True):
req = [request(rows, grad=grad)]
assert r._group_head_workspace_bytes(rows, req, grad_enabled=grad) == 7 * dense
assert (
r._group_head_workspace_bytes(
rows, req, grad_enabled=grad, lower_bound=True
)
== (3 if grad else 1) * dense
)


@pytest.mark.parametrize("grad", [False, True])
@pytest.mark.parametrize("rows", [1, 63])
def test_eager_exp_boundary_has_four_distinct_cpu_dense_storages(
monkeypatch, grad, rows
):
monkeypatch.setattr(_impl, "_all_reduce_tensor_parallel_max", lambda x: x)
monkeypatch.setattr(_impl, "_all_reduce_tensor_parallel_sum", lambda x: x)
logits = torch.linspace(-2, 2, rows * 17, dtype=torch.bfloat16).reshape(rows, 17)
logits.requires_grad_(grad)
original_float, original_exp = torch.Tensor.float, torch.exp
converted = []
observed = []

def as_float(self, *args, **kwargs):
result = original_float(self, *args, **kwargs)
if self is logits:
converted.append(weakref.ref(result))
return result

def exp(subtraction):
result = original_exp(subtraction)
# Only weak references survive between callbacks; do not manufacture
# overlap by retaining the converted tensor ourselves.
tensors = (logits, converted[0](), subtraction, result)
assert all(tensor is not None for tensor in tensors)
assert [tensor.dtype for tensor in tensors] == [
torch.bfloat16,
torch.float32,
torch.float32,
torch.float32,
]
assert len({tensor.untyped_storage().data_ptr() for tensor in tensors}) == 4
observed.append(sum(tensor.untyped_storage().nbytes() for tensor in tensors))
return result

monkeypatch.setattr(torch.Tensor, "float", as_float)
monkeypatch.setattr(torch, "exp", exp)
with torch.set_grad_enabled(grad):
actual = _impl._vocab_parallel_log_z(logits)
assert observed == [7 * rows * 17 * 2]
torch.testing.assert_close(actual, torch.logsumexp(original_float(logits), dim=-1))
if grad:
actual.sum().backward()
assert logits.grad is not None and logits.grad.isfinite().all()


@pytest.mark.parametrize("grad", [False, True])
@pytest.mark.parametrize("label", [1, -100, None])
def test_group_stats_cover_logits_chunks_before_short_target_tail(
monkeypatch, grad, label
):
_patch_local_head(monkeypatch)
monkeypatch.setattr(_impl, "_HEAD_CHUNK_TOKENS", 4)
original = _impl._vocab_parallel_log_z
calls = []

def log_z(logits):
calls.append(tuple(logits.shape))
return original(logits)

monkeypatch.setattr(_impl, "_vocab_parallel_log_z", log_z)
model = SimpleNamespace(
output_layer=_Head(torch.arange(85, dtype=torch.float32).reshape(17, 5) / 100),
vocab_size=17,
share_embeddings_and_output_weights=False,
_scale_logits=lambda value: value,
)
r = object.__new__(TrainerRank)
r.runtime = SimpleNamespace(model=[model])
requests = [
ForwardInput(input_tokens=torch.arange(13), logits=True, no_grad=not grad)
]
positions = (torch.arange(13),)
if label is not None:
requests.append(
ForwardInput(
input_tokens=torch.tensor([12]),
target_tokens=torch.tensor([label]),
no_grad=not grad,
)
)
positions += (torch.tensor([12]),)
with torch.set_grad_enabled(grad):
outputs = r._project_head(
[r._forward_item(item) for item in requests],
SimpleNamespace(
positions_by_item=positions,
source_positions_by_item=tuple(torch.arange(len(p)) for p in positions),
),
torch.arange(65, dtype=torch.float32).reshape(13, 5) / 100,
)
assert calls == ([] if label is None else [(4, 17)] * 3 + [(1, 17)])
assert outputs[0].logits.shape == (13, 17)
if label == -100:
assert outputs[1].target_logprobs.item() == 0
Loading
Loading