Skip to content
Merged
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
11 changes: 4 additions & 7 deletions src/art/trainer_rank/_gdn_memory.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
42 changes: 42 additions & 0 deletions tests/unit/test_trainer_rank_pending_memory.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading