diff --git a/src/art/trainer_rank/_gdn_memory.py b/src/art/trainer_rank/_gdn_memory.py index 38027dd76..f55e28473 100644 --- a/src/art/trainer_rank/_gdn_memory.py +++ b/src/art/trainer_rank/_gdn_memory.py @@ -332,14 +332,11 @@ def plan_floor(rank: Any, plan: Any) -> tuple[int, int]: if sum(s.length for s in group.packed.segments) != rows: raise ValueError("GDN packed rows disagree with segment geometry") retained += rows * layers * rank._hidden_size * 2 + moe = rank._moe_workspace_bytes( + rows, checkpoint_grad=True, slot_ref=group.slot_ref + ) workspace = max( workspace, - *( - rank._moe_workspace_bytes( - rows, checkpoint_grad=True, slot_ref=group.slot_ref - ) - + s.pending(rows, buckets) - for s in shapes - ), + *(moe + s.pending(rows, buckets) for s in shapes), ) return retained, workspace diff --git a/tests/unit/test_trainer_rank_pending_memory.py b/tests/unit/test_trainer_rank_pending_memory.py index 490352f76..359de8daa 100644 --- a/tests/unit/test_trainer_rank_pending_memory.py +++ b/tests/unit/test_trainer_rank_pending_memory.py @@ -146,6 +146,48 @@ def test_actual_constructor_cache_and_full_plan(pending_rank): assert lower.required <= rank._memory_check(plan).estimated_required_bytes +@pytest.mark.parametrize("shape_count", [1, 30]) +@pytest.mark.parametrize("slot_ref", [None, "selected-adapter"]) +def test_pending_floor_prices_moe_once_per_group(monkeypatch, shape_count, slot_ref): + shapes = tuple(g.Shape(2, 4, 8, 8, 4, i, 0) for i in range(shape_count)) + monkeypatch.setattr(g, "model_shapes", lambda rank, ref: (40, shapes)) + calls = [] + price = 100 + + def moe(rows, *, checkpoint_grad=False, slot_ref=None): + calls.append((rows, checkpoint_grad, slot_ref)) + return price * rows + + rank = SimpleNamespace(_hidden_size=16, _moe_workspace_bytes=moe) + groups = [] + for rows, grad, ref in [(65, True, slot_ref), (17, False, "reference")]: + segment = SimpleNamespace( + start=0, end=rows, length=rows, packed_start=0, group_id=0, parent_id=0 + ) + groups.append( + SimpleNamespace( + grad_enabled=grad, + slot_ref=ref, + packed=SimpleNamespace(tokens=torch.empty(rows), segments=[segment]), + ) + ) + plan = SimpleNamespace(groups=groups) + pending = max( + s.pending(65, g.cp1_buckets(groups[0].packed.segments)) for s in shapes + ) + # Preserve the original per-layer maximum and the separate no-grad group. + for price in (100, 200): + calls.clear() + assert g.plan_floor(rank, plan) == (65 * 40 * 16 * 2, 65 * price + pending) + assert calls == [(65, True, slot_ref), (17, False, "reference")] + # No gradient groups or unsupported model metadata must not price MoE. + calls.clear() + assert g.plan_floor(rank, SimpleNamespace(groups=groups[1:])) == (0, 0) + monkeypatch.setattr(g, "model_shapes", lambda rank, ref: None) + assert g.plan_floor(rank, plan) == (0, 0) + assert calls == [] + + @pytest.mark.parametrize("fits_after", (False, True)) def test_exact_pending_demand_survives_recovery(monkeypatch, pending_rank, fits_after): from test_trainer_rank_cache_recovery import _check_component_demand_recovery