Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
37 commits
Select commit Hold shift + click to select a range
cdf4d29
Mirror the bytes a recomputed CP attention keeps for backward
bradhilton Sep 25, 2026
5899a2d
Price CP2 recompute per rank on each layer type's layout
bradhilton Sep 25, 2026
0e7255c
Bound CP2 layout pricing from below and count flex's own LSE
bradhilton Sep 25, 2026
9b4e9d7
Guard the layout gate and state why the lower bound rounds down
bradhilton Sep 25, 2026
33185b4
Run the CP layout memory tests in the Megatron CI lane
bradhilton Sep 25, 2026
c1500ca
Charge single-head view copies and partial-tape indices; time layout …
bradhilton Sep 26, 2026
4368998
Keep partial-tape indices when the first stage drops its tape
bradhilton Sep 26, 2026
a4b6a26
Merge the rebased recomputed-mixer floor into per-rank CP2 layouts
bradhilton Sep 26, 2026
7583308
Merge the recomputed-mixer floor's review fixes into per-rank CP2 lay…
bradhilton Sep 26, 2026
b4a619e
Test the combine-extent floor through per-rank layouts
bradhilton Sep 26, 2026
0f59d7c
Pin that the layout path's combine floor adds no second TE growth
bradhilton Sep 27, 2026
8d7c3cf
Merge commit 'a0feb0c2577bf74fb14028f85fd0003a082175d6' into HEAD
bradhilton Sep 28, 2026
be1db0c
Price adapter gradients against each rank's layout boundaries
bradhilton Sep 28, 2026
42a2e2b
Pin TE growth and a stage between the combine floor and its TE
bradhilton Sep 28, 2026
0606917
Pair each CP rank's adapter-gradient extra with its own floor
bradhilton Sep 28, 2026
67f8ae7
Merge #963's sequential gradient groups; walk each rank's groups in turn
bradhilton Sep 28, 2026
22f67b0
Test split lower bounds and admission parity with pending gradients
bradhilton Sep 28, 2026
a487f98
Merge #963's closed-form group order
bradhilton Sep 28, 2026
92dd4e0
Resolve each slot's pending gradients once across ranks
bradhilton Sep 28, 2026
bd5a83a
Test the split lower bound with two gradient slots
bradhilton Sep 28, 2026
632dea9
Pin the two-slot lower bound's own extra
bradhilton Sep 28, 2026
7863832
Merge #963's layer-mismatch test
bradhilton Sep 28, 2026
977dea3
Merge #963's staged head and decoder pricing
bradhilton Sep 28, 2026
13dc898
Merge #963's head-stage row state
bradhilton Sep 28, 2026
a5e9d75
Merge #963's traced head staging
bradhilton Sep 28, 2026
c72ddf8
Merge #963's fused-statistics proof for head staging
bradhilton Sep 28, 2026
f1bbd63
Merge #963's TE-free profile readings
bradhilton Sep 28, 2026
34c2a57
Merge #963's strict staged head statistics and V1t revert
bradhilton Sep 28, 2026
d881fba
Merge #963's per-kernel fused-statistics proof
bradhilton Sep 28, 2026
6dc567b
Merge #963's strict binding per thread, group and staged admission
bradhilton Sep 28, 2026
653136d
Price CP2 recompute per rank on each layer type's layout, on the extr…
bradhilton Sep 29, 2026
c847095
Merge #963's replay hardening (e39b4603) into the CP layout port
bradhilton Sep 29, 2026
0684d70
Record #978's pre-extraction history (6dc567b03) as ported
bradhilton Sep 29, 2026
4c5a4d1
Merge #963's replay test coverage (bab7be10a) into the CP layout port
bradhilton Sep 29, 2026
e204ee8
Refuse CP layout facts live pricing cannot produce
bradhilton Sep 29, 2026
645c906
Merge #963's Megatron-lane entry (1763c67dc) into the CP layout port
bradhilton Sep 29, 2026
bf79e44
Test that an ungrouped estimate cannot record CP layouts
bradhilton Sep 29, 2026
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
6 changes: 5 additions & 1 deletion .github/workflows/prek.yml
Original file line number Diff line number Diff line change
Expand Up @@ -246,6 +246,8 @@ jobs:
tests/unit/test_trainer_rank_pending_memory.py \
tests/unit/test_trainer_rank_shared_memory.py \
tests/unit/test_trainer_rank_converted_memory.py \
tests/unit/test_trainer_rank_layout_memory.py \
tests/unit/test_context_parallel_retained_bytes.py \
tests/unit/test_trainer_rank_split.py \
tests/unit/test_megatron_compile_garbage.py \
tests/unit/test_trainer_rank_cache_recovery.py::test_dense_cp_exact_demand_fits_after_recovery \
Expand Down Expand Up @@ -297,4 +299,6 @@ jobs:
--ignore=tests/unit/test_trainer_rank_pending_memory.py \
--ignore=tests/unit/test_trainer_rank_shared_memory.py \
--ignore=tests/unit/test_megatron_compile_garbage.py \
--ignore=tests/unit/test_trainer_rank_converted_memory.py
--ignore=tests/unit/test_trainer_rank_converted_memory.py \
--ignore=tests/unit/test_trainer_rank_layout_memory.py \
--ignore=tests/unit/test_context_parallel_retained_bytes.py
107 changes: 107 additions & 0 deletions src/art/megatron/context_parallel/executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,7 @@
DkvReducePlan,
ExactMaskMetadata,
FlexMaskSpec,
RankRuntimePlan,
StageExecutionSpec,
StagePlan,
TokenRange,
Expand Down Expand Up @@ -1550,6 +1551,112 @@ def _merge_stage_output_grads_from_tape(
return stage_out_grads, stage_lse_grads


def minimum_retained_bytes_per_row(
*,
q_heads: int,
kv_heads: int,
head_dim: int,
value_head_dim: int,
element_size: int,
) -> int:
"""The least ``retained_stage_record_bytes`` keeps per own row.

Every rank with rows runs a local stage over all of them (each row attends
to itself); an aligned one keeps at least the contiguous copies of
multi-head views, flex's output and its two LSEs.
"""
copies = (q_heads > 1) * q_heads * head_dim + (kv_heads > 1) * kv_heads * (
head_dim + value_head_dim
)
return (copies + q_heads * value_head_dim) * element_size + 2 * q_heads * 4


def retained_stage_record_bytes(
rank_plan: RankRuntimePlan,
*,
q_heads: int,
kv_heads: int,
head_dim: int,
value_head_dim: int,
element_size: int,
block_size: SparseBlockSize,
) -> int:
"""Bytes a recomputed attention keeps for backward beyond its own-row tensors.

A size-only mirror of ``_forward_stage_records`` with ``record_for_backward``
and of ``_run_stage_attention``; keep the three in step. Every stage that
runs keeps the Q/K/V flex consumed: a padded copy when the execution length
differs, else a contiguous copy of the permuted ``q_flat``/``k_flat`` view;
partial-range gathers and remote fetch buffers are kept as the stage's
inputs. It keeps flex's output and LSE at the execution length, and
logical-length copies of them when padded, plus flex's own LSE beside the
normalized one the FLASH backend returns (counted on every backend). Every
producing stage after the first keeps a merge-tape clone of the
accumulators. Accumulators themselves are transient.
"""
own = int(rank_plan.local_valid_lengths[0]) if rank_plan.local_valid_lengths else 0
accum_size = 4 if element_size < 4 else element_size
q_row = q_heads * head_dim * element_size
k_row = kv_heads * head_dim * element_size
v_row = kv_heads * value_head_dim * element_size
out_row = q_heads * value_head_dim * element_size
lse_row = q_heads * 4
tape_row = q_heads * (value_head_dim + 1) * accum_size
total = 0
local_produced = False
tapes: list[int] = []
for stage in _ordered_stage_plans(rank_plan.stage_plans):
if not (stage.q_len > 0 and stage.k_len > 0 and stage.slices):
continue
q_len = _logical_stage_q_len(stage)
k_len = _logical_stage_k_len(stage)
q_pad, k_pad, _family = select_sparse_execution_family(
is_local_stage=bool(stage.is_local_stage),
q_len=int(stage.q_len),
k_len=int(stage.k_len),
block_size=block_size,
)
q_full = _ranges_cover_full_length(stage.owner_local_q_ranges, length=own)
# Queries: the full range is a permuted view of q_flat; a partial range
# is gathered into a contiguous tensor kept as the stage's input.
if not q_full:
total += q_row * q_len
if q_pad != q_len:
total += q_row * q_pad
elif q_full:
# A view of the projection output: flex copies it unless it is
# already contiguous, which one head alone does not guarantee
# (a fused QKV split keeps the projection's token stride).
total += q_row * q_len
# Keys and values: local ranges as for queries; remote ones land in
# contiguous head-major fetch buffers kept as the stage's inputs.
k_full = bool(stage.is_local_stage) and _ranges_cover_full_length(
stage.owner_local_k_ranges, length=own
)
if not k_full:
total += (k_row + v_row) * k_len
if k_pad != k_len:
total += (k_row + v_row) * k_pad
elif k_full:
total += (k_row + v_row) * k_len
total += (out_row + 2 * lse_row) * q_pad
if q_pad != q_len:
total += (out_row + lse_row) * q_len
tape = tape_row * (own if q_full else q_len)
if not q_full:
# Its int64 row index is kept even by the first producing stage.
total += 8 * q_len
if stage.is_local_stage:
local_produced = True
else:
tapes.append(tape)
# The first producing stage records no tape: the local stage when it ran,
# else whichever remote stage is ready first, so drop the smallest tape.
if tapes and not local_produced:
tapes.remove(min(tapes))
return total + sum(tapes)


def _forward_stage_records(
*,
q_flat: torch.Tensor,
Expand Down
49 changes: 49 additions & 0 deletions src/art/megatron/context_parallel/runtime.py
Original file line number Diff line number Diff line change
Expand Up @@ -433,6 +433,55 @@ def context_parallel_rank_model_token_counts(
)


def context_parallel_rank_layouts(
*,
group_ids: torch.Tensor,
parent_ids: torch.Tensor,
topology: ParallelTopology,
config: ContextParallelConfig,
original_seq_len: int,
build_gdn_execution_spec: bool,
gdn_planner_config: Any | None = None,
) -> tuple[tuple[int, ...], tuple[int, ...] | None, tuple[RankRuntimePlan, ...]]:
"""Each CP rank's attention rows, GDN rows and attention stage plan.

Uses the cached planning bundle and per-rank runtime plans that execution
builds, so a memory estimate sees the layouts the ranks will run.
"""
planning_key, bundle, _group_ids_cpu, _parent_ids_cpu = (
_get_or_build_planning_bundle(
group_ids=group_ids,
parent_ids=parent_ids,
topology=topology,
config=config,
original_seq_len=original_seq_len,
build_gdn_execution_spec=build_gdn_execution_spec,
)
)
attention = tuple(bundle.token_layout_index.token_counts_by_rank)
gdn = None
if build_gdn_execution_spec:
gdn = tuple(
_plan_gdn_global_execution(
planning_key=planning_key,
bundle=bundle,
topology=topology,
gdn_planner_config=gdn_planner_config,
).gdn_token_counts_by_rank
)
plans = tuple(
_get_or_build_bundle_rank_plan(
planning_key=planning_key,
bundle=bundle,
original_seq_len=original_seq_len,
target_rank=rank,
block_size=config.block_size,
)
for rank in range(len(attention))
)
return attention, gdn, plans


def context_parallel_model_token_total(
*,
group_ids: torch.Tensor,
Expand Down
46 changes: 41 additions & 5 deletions src/art/trainer_rank/_impl.py
Original file line number Diff line number Diff line change
Expand Up @@ -1069,6 +1069,17 @@ def signature(self) -> _MemorySignature:
_AnyForwardPlan = _FlatForwardPlan | _SplitForwardPlan


@dataclass(frozen=True)
class _GroupLayout:
"""One packed group's CP layouts on every rank, for layout-aware pricing."""

attention_rows: tuple[int, ...]
gdn_rows: tuple[int, ...] | None
# What each rank's recomputed CP attention keeps for backward beyond its
# own-row activations (``retained_stage_record_bytes``).
attention_retained: tuple[int, ...]


@dataclass(frozen=True)
class _SubforwardCost:
"""Memory terms of one candidate subforward while all graphs stay live.
Expand Down Expand Up @@ -3041,6 +3052,7 @@ def _subforward_cost(
gdn_segments: int = 0,
group_rows: tuple[tuple[int, bool], ...] = (),
group_routed_rows: tuple[int, ...] | None = None,
group_layouts: tuple[_GroupLayout, ...] | None = None,
slot_refs: tuple["LoRASlotRef | None", ...] | None = None,
head_workspace_bytes: int = 0,
head_backward_traced: bool = False,
Expand All @@ -3049,7 +3061,11 @@ def _subforward_cost(
hybridep_growth_bytes: int = 0,
) -> _SubforwardCost:
checkpoint_memory = self._checkpoint_memory_floor(
group_rows, slot_refs, gdn_segments, routed_rows=group_routed_rows
group_rows,
slot_refs,
gdn_segments,
routed_rows=group_routed_rows,
layouts=group_layouts,
)
required = self._estimate_required_memory_bytes_from_values(
packed_tokens=packed_tokens,
Expand All @@ -3059,6 +3075,7 @@ def _subforward_cost(
gdn_segments=gdn_segments,
group_rows=group_rows,
group_routed_rows=group_routed_rows,
group_layouts=group_layouts,
slot_refs=slot_refs,
head_workspace_bytes=head_workspace_bytes,
checkpoint_floor=checkpoint_floor,
Expand All @@ -3080,13 +3097,21 @@ def _subforward_cost(
)
# Input gradients live at the recomputed layer's peak; kept out of
# forward retention, including the cold fallback above.
# Per-rank layouts price each rank's own boundaries; the input-gradient
# allowance keeps the busiest rank's (``_checkpoint_input_gradient_bytes``).
gradient = self._checkpoint_input_gradient_bytes(
group_rows, slot_refs, retained=checkpoint_retained
group_rows,
slot_refs,
retained=checkpoint_retained if group_layouts is None else None,
)
gradient_slots = self._gradient_slots(group_rows, slot_refs)
adapter_gradient = (
self._checkpoint_adapter_gradient_bytes(
self._checkpoint_gradient_groups(group_rows, slot_refs)
self._checkpoint_adapter_gradient_extra(
(checkpoint_retained, checkpoint_workspace),
group_rows,
slot_refs,
group_routed_rows,
group_layouts,
)
if gradient
else 0
Expand All @@ -3099,7 +3124,7 @@ def _subforward_cost(
forward_required = required
if gradient:
head_stage = self._checkpoint_head_stage_bytes(
head_workspace_bytes, gradient, group_rows, slot_refs
head_workspace_bytes, gradient, group_rows, slot_refs, group_layouts
)
peak = checkpoint_workspace + adapter_gradient
if head_stage is not None:
Expand Down Expand Up @@ -4929,6 +4954,17 @@ def _cp_group_model_tokens(
)

_group_head_workspace_bytes = _memory._group_head_workspace_bytes
_layer_gdn_inputs = _memory._layer_gdn_inputs
_layout_checkpoint_floor = _memory._layout_checkpoint_floor
_layout_checkpoint_rank_floors = _memory._layout_checkpoint_rank_floors
_layout_pricing_supported = _memory._layout_pricing_supported
_minimum_layouts = _memory._minimum_layouts
_layout_layer_boundaries = _memory._layout_layer_boundaries
_layout_adapter_gradient_bytes = _memory._layout_adapter_gradient_bytes
_checkpoint_adapter_gradient_extra = _memory._checkpoint_adapter_gradient_extra
_recomputed_mixer_widths = _memory._recomputed_mixer_widths
_plan_group_layouts = _micro_batch_planner._plan_group_layouts
_compute_group_layouts = _micro_batch_planner._compute_group_layouts
_triton_min_rows = _memory._triton_min_rows
_head_backward_traced = _memory._head_backward_traced
_te_workspace_growth_bytes = _memory._te_workspace_growth_bytes
Expand Down
Loading
Loading