diff --git a/.github/workflows/prek.yml b/.github/workflows/prek.yml index 453b7f8b2..7697fcc5b 100644 --- a/.github/workflows/prek.yml +++ b/.github/workflows/prek.yml @@ -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 \ @@ -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 \ diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index d4718cab0..798c900a9 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -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 ( @@ -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( diff --git a/tests/unit/test_trainer_rank_admission_inputs.py b/tests/unit/test_trainer_rank_admission_inputs.py index a79d198e1..0b574cf0f 100644 --- a/tests/unit/test_trainer_rank_admission_inputs.py +++ b/tests/unit/test_trainer_rank_admission_inputs.py @@ -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 @@ -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): diff --git a/tests/unit/test_trainer_rank_eager_head_memory.py b/tests/unit/test_trainer_rank_eager_head_memory.py new file mode 100644 index 000000000..5a548d134 --- /dev/null +++ b/tests/unit/test_trainer_rank_eager_head_memory.py @@ -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 diff --git a/tests/unit/test_trainer_rank_head_memory.py b/tests/unit/test_trainer_rank_head_memory.py index 21b7ad071..fbbed07f1 100644 --- a/tests/unit/test_trainer_rank_head_memory.py +++ b/tests/unit/test_trainer_rank_head_memory.py @@ -75,7 +75,7 @@ def test_mixed_checkpoint_head_demand_survives_recovery(monkeypatch, fits_after) values = r._estimate_flat_forward(requests, exact=True) assert values[3] == ((8, True), (16, False)) assert r._checkpoint_memory_floor(values[3])[0] == 8 * 40 * 2048 * 2 - assert values[4] == 3 * 8 * 248320 * 2 + assert values[4] == 7 * 16 * 248320 * 2 _check_component_demand_recovery(monkeypatch, r, requests, fits_after=fits_after) @@ -144,12 +144,12 @@ def test_outputs_retention_and_empirical_peak_are_counted_once(): plan = r._plan_flat_forward([request(grad=True)]) retained = 512 * 40 * 2048 * 2 gradient = 512 * 40 * 2048 * 2 - head = 3 * 512 * 248320 * 2 + head = 7 * 512 * 248320 * 2 cost = r._plan_cost(plan) assert cost.retained == int((plan.output_bytes + retained + head) * 1.1) assert cost.required == int((plan.output_bytes + retained + gradient + head) * 1.1) r._memory_profiles[plan.signature] = _MemoryProfile( - bytes_per_token=2_000_000, + bytes_per_token=10_000_000, packed_tokens=512, logical_per_packed=1, retained_compute_bytes_per_token=1, @@ -157,7 +157,7 @@ def test_outputs_retention_and_empirical_peak_are_counted_once(): cost = r._plan_cost(plan) # Packed pricing adds head and caller memory for every logical row. rows = _PACKED_PRICED_LOGICAL_ROW_BYTES * 512 - assert cost.required == int((plan.output_bytes + 512 * 2_000_000 + rows) * 1.1) + assert cost.required == int((plan.output_bytes + 512 * 10_000_000 + rows) * 1.1) assert cost.retained == int((plan.output_bytes + retained) * 1.1) @@ -171,9 +171,9 @@ def test_ignored_targets_hidden_only_and_tiny_target_group(): assert r._head_projection_rows([ignored, hidden]) == 0 assert r._plan_head_workspace_bytes(r._plan_flat_forward([ignored, hidden])) == 0 plan = r._plan_flat_forward([hidden, target]) - assert r._plan_head_workspace_bytes(plan) == 3 * 248320 * 2 - assert r._estimate_flat_forward([hidden, target])[-1] == 3 * 248320 * 2 - assert r._estimate_flat_forward([hidden, target], exact=True)[-1] == 3 * 248320 * 2 + assert r._plan_head_workspace_bytes(plan) == 7 * 248320 * 2 + assert r._estimate_flat_forward([hidden, target])[-1] == 7 * 248320 * 2 + assert r._estimate_flat_forward([hidden, target], exact=True)[-1] == 7 * 248320 * 2 def test_multilabel_row_validity_matches_projection(): @@ -183,7 +183,9 @@ def test_multilabel_row_validity_matches_projection(): target_tokens=torch.tensor([[-100, -100], [-100, 2], [3, -100], [-100, -100]]), ) assert r._head_projection_rows([item]) == 2 - assert r._plan_head_workspace_bytes(r._plan_flat_forward([item])) == 2 * 248320 * 2 + assert ( + r._plan_head_workspace_bytes(r._plan_flat_forward([item])) == 7 * 2 * 248320 * 2 + ) def test_shared_rows_use_lower_upper_and_exact_layout_union(): @@ -195,7 +197,7 @@ def test_shared_rows_use_lower_upper_and_exact_layout_union(): assert r._head_projection_rows(req) == 4 exact = r._estimate_flat_forward(req, exact=True, memory_minimal=True) plan = r._plan_flat_forward(req, memory_minimal=True) - assert exact[-1] == r._plan_head_workspace_bytes(plan) == 2 * 248320 * 2 + assert exact[-1] == r._plan_head_workspace_bytes(plan) == 7 * 2 * 248320 * 2 lower = r._split_chunk_lower_cost( req, tuple(x.input_tokens for x in req), checkpoint=Unset ) @@ -306,14 +308,14 @@ def test_tied_standard_head_weight_uses_the_same_capacity(): @pytest.mark.parametrize("rows", [128, 512]) -def test_target_backward_refuses_budget_below_logits_and_both_gradients(rows): +def test_eager_statistics_refuses_budget_between_old_and_new_components(rows): r = rank() plan = r._plan_flat_forward([request(rows, grad=True)]) retained, _ = r._checkpoint_memory_floor(r._plan_group_rows(plan)) gradient = rows * 40 * 2048 * 2 dense = min(rows, 512) * 248320 * 2 - before = int((plan.output_bytes + retained + gradient + 2 * dense) * 1.1) - expected = int((plan.output_bytes + retained + gradient + 3 * dense) * 1.1) + before = int((plan.output_bytes + retained + gradient + 3 * dense) * 1.1) + expected = int((plan.output_bytes + retained + gradient + 7 * dense) * 1.1) r._available_memory_bytes = lambda: (before + expected) // 2 check = r._memory_check(plan) assert check.estimated_required_bytes == expected @@ -326,16 +328,17 @@ def test_group_head_workspace_keeps_gradient_mode_with_its_rows( ): r = rank() requests = [request(gradient_rows, grad=True), request(reference_rows)] - expected = max(3 * gradient_rows, reference_rows) * 248320 * 2 + expected = 7 * max(gradient_rows, reference_rows) * 248320 * 2 plan = r._plan_flat_forward(requests) assert r._plan_head_workspace_bytes(plan) == expected for exact in (False, True): for minimal in (False, True): - assert ( - r._estimate_flat_forward(requests, exact=exact, memory_minimal=minimal)[ - -1 - ] - == expected + assert r._estimate_flat_forward( + requests, exact=exact, memory_minimal=minimal + )[-1] == ( + max(3 * gradient_rows, reference_rows) * 248320 * 2 + if minimal and not exact + else expected ) @@ -352,8 +355,13 @@ def test_gradient_statistics_floor_requires_exact_effective_scaling(mutation): r = rank() model = r.runtime.model[0] req = [request(128, grad=True)] + logits = [replace(request(128), target_tokens=None, logits=True)] + assert ( + r._plan_head_workspace_bytes(r._plan_flat_forward(logits)) + == 3 * 128 * 248320 * 2 + ) assert ( - r._plan_head_workspace_bytes(r._plan_flat_forward(req)) == 3 * 128 * 248320 * 2 + r._plan_head_workspace_bytes(r._plan_flat_forward(req)) == 7 * 128 * 248320 * 2 ) if mutation == "custom": model._scale_logits = lambda logits: logits @@ -378,6 +386,57 @@ def __call__(self, logits): del model._scale_logits # The standard head still allocates its original one-buffer component. assert r._plan_head_workspace_bytes(r._plan_flat_forward(req)) == 128 * 248320 * 2 + assert ( + r._plan_head_workspace_bytes(r._plan_flat_forward(logits)) == 128 * 248320 * 2 + ) + + +@pytest.mark.parametrize("rows", [1, 513]) +@pytest.mark.parametrize("hidden", [False, True]) +def test_no_grad_logits_copies_refuse_budget_between_old_and_new_components( + rows, hidden +): + r = rank() + req = [replace(request(rows, hidden=hidden), target_tokens=None, logits=True)] + plan = r._plan_flat_forward(req) + dense = min(rows, 512) * 248320 * 2 + output = rows * (248320 + (2048 if hidden else 0)) * 2 + assert plan.output_bytes == output + assert r._plan_head_workspace_bytes(plan) == 3 * dense + for exact in (False, True): + assert r._estimate_flat_forward(req, exact=exact)[-1] == 3 * dense + # Capacity is not a universal rejection floor or a grad-enabled bound. + assert ( + r._group_head_workspace_bytes(rows, req, grad_enabled=False, lower_bound=True) + == dense + ) + grad_req = [replace(req[0], no_grad=False)] + assert r._plan_head_workspace_bytes(r._plan_flat_forward(grad_req)) == dense + before = int((output + dense) * 1.1) + expected = int((output + 3 * dense) * 1.1) + r._available_memory_bytes = lambda: (before + expected) // 2 + check = r._memory_check(plan) + assert check.estimated_required_bytes == expected + assert not check.fits + + +def test_no_grad_logits_shared_rows_keep_outputs_and_statistics_separate(): + r = rank() + item = replace(request(4), target_tokens=None, logits=True) + req = [item, item] + plan = r._plan_flat_forward(req, memory_minimal=True) + dense = 4 * 248320 * 2 + assert plan.output_bytes == 2 * dense + assert r._plan_head_workspace_bytes(plan) == 3 * dense + assert ( + r._estimate_flat_forward(req, exact=True, memory_minimal=True)[-1] == 3 * dense + ) + # Any labels enable statistics group-wide, including ignored labels. + req.append(request(4, ignored=True)) + assert ( + r._plan_head_workspace_bytes(r._plan_flat_forward(req, memory_minimal=True)) + == 7 * dense + ) @pytest.mark.parametrize("extra", [{"logits": True}, {"top_k": 2}]) @@ -385,7 +444,7 @@ def test_gradient_statistics_floor_survives_additional_output_modes(extra): r = rank() req = [replace(request(128, grad=True), **extra)] assert ( - r._plan_head_workspace_bytes(r._plan_flat_forward(req)) == 3 * 128 * 248320 * 2 + r._plan_head_workspace_bytes(r._plan_flat_forward(req)) == 7 * 128 * 248320 * 2 ) @@ -395,10 +454,10 @@ def test_gradient_shared_rows_price_same_union_in_exact_and_split_lower_cost(): b = replace(a, target_tokens=torch.tensor([-100, 1])) requests = [a, a, b, b] plan = r._plan_flat_forward(requests, memory_minimal=True) - expected = 3 * 2 * 248320 * 2 + expected = 7 * 2 * 248320 * 2 exact = r._estimate_flat_forward(requests, exact=True, memory_minimal=True) assert exact[-1] == r._plan_head_workspace_bytes(plan) == expected - assert r._estimate_flat_forward(requests, memory_minimal=True)[-1] == expected // 2 + assert r._estimate_flat_forward(requests, memory_minimal=True)[-1] == 3 * 248320 * 2 lower = r._split_chunk_lower_cost( requests, tuple(x.input_tokens for x in requests), checkpoint=Unset ) @@ -410,4 +469,4 @@ def test_later_sparse_loss_does_not_reduce_6330_projected_targets(): item = request(6330, grad=True) plan = r._plan_flat_forward([item]) assert item.target_tokens.numel() == 6330 - assert r._plan_head_workspace_bytes(plan) == 3 * 512 * 248320 * 2 + assert r._plan_head_workspace_bytes(plan) == 7 * 512 * 248320 * 2 diff --git a/tests/unit/test_trainer_rank_ignored_mixed_head.py b/tests/unit/test_trainer_rank_ignored_mixed_head.py index 9ac5d500b..d38316adf 100644 --- a/tests/unit/test_trainer_rank_ignored_mixed_head.py +++ b/tests/unit/test_trainer_rank_ignored_mixed_head.py @@ -17,10 +17,10 @@ def test_ignored_rows_reactivated_by_same_item_output_keep_backward_floor(extra) item = replace(request(128, grad=True, ignored=True), **extra) assert ( r._plan_head_workspace_bytes(r._plan_flat_forward([item])) - == 3 * 128 * 248320 * 2 + == 7 * 128 * 248320 * 2 ) - assert r._estimate_flat_forward([item], exact=True)[-1] == 3 * 128 * 248320 * 2 - assert r._estimate_flat_forward([item])[-1] == 3 * 128 * 248320 * 2 + assert r._estimate_flat_forward([item], exact=True)[-1] == 7 * 128 * 248320 * 2 + assert r._estimate_flat_forward([item])[-1] == 7 * 128 * 248320 * 2 assert r._head_target_chunk_rows([item], lower_bound=True) == 0 diff --git a/tests/unit/test_trainer_rank_mixed_head_memory.py b/tests/unit/test_trainer_rank_mixed_head_memory.py index 86e01c1c4..8ec7b7a0d 100644 --- a/tests/unit/test_trainer_rank_mixed_head_memory.py +++ b/tests/unit/test_trainer_rank_mixed_head_memory.py @@ -22,7 +22,7 @@ def test_adding_output_cannot_erase_existing_target_admission_floor(extra): after = r._plan_cost(plan).required print({"extra": extra, "before": before, "after": after}) assert after >= before - assert r._plan_head_workspace_bytes(plan) == 3 * 129 * 248320 * 2 + assert r._plan_head_workspace_bytes(plan) == 7 * 129 * 248320 * 2 r._available_memory_bytes = lambda: before - 1 assert not r._memory_check(plan).fits @@ -42,12 +42,20 @@ def test_sparse_target_prices_its_full_mixed_chunk_and_short_tail(extra): dense = 512 * 248320 * 2 assert ( r._group_head_workspace_bytes(512, req, grad_enabled=True, positions=full) - == 3 * dense + == 7 * dense ) assert ( r._group_head_workspace_bytes(512, req, grad_enabled=True, positions=tail) - == dense + == 7 * dense ) + # Optional eager demand is not an unconditional rejection lower bound. + for positions, expected in ((full, 3 * dense), (tail, dense)): + assert ( + r._group_head_workspace_bytes( + 512, req, grad_enabled=True, positions=positions, lower_bound=True + ) + == expected + ) @pytest.mark.parametrize("extra", [{"logits": True}, {"top_k": 2}]) @@ -82,7 +90,7 @@ def test_shared_multilabel_union_matches_actual_and_split_bounds(extra): ) assert ( r._plan_head_workspace_bytes(r._plan_flat_forward(req, memory_minimal=True)) - == 3 * 4 * 248320 * 2 + == 7 * 4 * 248320 * 2 ) @@ -91,9 +99,11 @@ def test_ignored_device_labels_and_no_target_keep_distinct_guards(extra): r = rank() ignored = replace(request(128, grad=True, ignored=True), **extra) dense = 128 * 248320 * 2 - assert r._plan_head_workspace_bytes(r._plan_flat_forward([ignored])) == 3 * dense + assert r._plan_head_workspace_bytes(r._plan_flat_forward([ignored])) == 7 * dense no_target = replace(ignored, target_tokens=None) - assert r._plan_head_workspace_bytes(r._plan_flat_forward([no_target])) == dense + assert r._plan_head_workspace_bytes(r._plan_flat_forward([no_target])) == ( + 7 * dense if "top_k" in extra else dense + ) device = replace( ignored, target_tokens=torch.empty(128, device="meta", dtype=torch.long) ) @@ -109,7 +119,7 @@ def test_ignored_device_labels_and_no_target_keep_distinct_guards(extra): @pytest.mark.parametrize("mutation", ["no_grad", "custom_scale", "mup", "head_hook"]) -def test_mixed_path_preserves_source_scaling_and_gradient_guards(mutation): +def test_mixed_path_prices_no_grad_but_preserves_source_scaling_guards(mutation): r = rank() item = replace(request(128, grad=True), logits=True) model = r.runtime.model[0] @@ -121,7 +131,8 @@ def test_mixed_path_preserves_source_scaling_and_gradient_guards(mutation): model.config.use_mup = True else: model.output_layer.register_forward_hook(lambda *args: None) - expected = 0 if mutation == "head_hook" else 128 * 248320 * 2 + multiplier = 0 if mutation == "head_hook" else 7 if mutation == "no_grad" else 1 + expected = multiplier * 128 * 248320 * 2 assert r._plan_head_workspace_bytes(r._plan_flat_forward([item])) == expected