Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
27 commits
Select commit Hold shift + click to select a range
b253f67
Price the recomputed layer's attention or GDN activations in the chec…
bradhilton Sep 25, 2026
7cf9aa0
Count the recomputed mixer in generic admission without a MoE component
bradhilton Sep 25, 2026
a256cd3
Ground the GDN recompute width in its measured sites and exchange widths
bradhilton Sep 25, 2026
91c98a8
Price GDN l2norm outputs at value-head width
bradhilton Sep 25, 2026
cf56c61
Price one dispatched H-wide input per routed row under HybridEP
bradhilton Sep 25, 2026
d583cd2
Format the HybridEP converted-stage test
bradhilton Sep 25, 2026
210acec
Name the EP1 all-to-all's second routed copy as its expert-sorted rows
bradhilton Sep 25, 2026
c59a692
Price what is live at the recompute peak; make the EP allowance EP-de…
bradhilton Sep 25, 2026
9d03f33
Price HybridEP routed rows on the EP group's balanced share
bradhilton Sep 25, 2026
8bc0c98
Keep the per-boundary gradient unless MoE covers all recompute
bradhilton Sep 25, 2026
01f89b2
Require priced FC1 stages before trusting a slot's MoE coverage
bradhilton Sep 25, 2026
ae10b7e
Rewrap the EP allowance comment
bradhilton Sep 25, 2026
fbe6943
Merge main (0e0c31b32) into the recomputed-mixer floor
bradhilton Sep 26, 2026
3d1d18e
Keep boundary gradients above CP2; price TE beside the combine extent
bradhilton Sep 26, 2026
f0670b9
Price the adapter gradients a recompute backward holds at its peak
bradhilton Sep 27, 2026
ecf2ce2
Count head-tied and custom checkpoint gradients as live throughout
bradhilton Sep 27, 2026
067048e
Order slot names safely and leave base-model groups out of adapter slots
bradhilton Sep 27, 2026
789db66
Merge the adapter-gradient floor (#1002 at ecf2ce21b) into the recomp…
bradhilton Sep 27, 2026
c35b238
Test that base-model gradient groups price and defer nothing
bradhilton Sep 28, 2026
f1fe9e7
Merge main (762c89dfa) into the adapter-gradient floor
bradhilton Sep 28, 2026
ea0ae1e
Merge the adapter-gradient floor (#1002 at f1fe9e712, main 762c89dfa)
bradhilton Sep 28, 2026
e54f4bc
Walk gradient groups' backward one after another, in the worst order
bradhilton Sep 28, 2026
b2b66e4
Merge #1002's sequential gradient groups (e54f4bc25)
bradhilton Sep 28, 2026
0df629c
Price gradient groups' worst backward order in closed form
bradhilton Sep 28, 2026
3700b37
Merge #1002's closed-form group order (0df629cf5)
bradhilton Sep 28, 2026
2a64195
Test that a layer-count mismatch prices no adapter-gradient extra
bradhilton Sep 28, 2026
fd5ef3a
Merge the adapter-gradient floor's layer-mismatch test (#1002 2a6419528)
bradhilton Sep 28, 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
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_profile_warm.py \
tests/unit/test_trainer_rank_tp_floor.py \
tests/unit/test_trainer_rank_checkpoint_gradient_memory.py \
tests/unit/test_trainer_rank_adapter_gradient_memory.py \
tests/unit/test_trainer_rank_slot_memory.py \
tests/unit/test_trainer_rank_moe_memory.py \
tests/unit/test_trainer_rank_head_memory.py \
Expand Down Expand Up @@ -279,6 +280,7 @@ jobs:
--ignore=tests/unit/test_trainer_rank_profile_warm.py \
--ignore=tests/unit/test_trainer_rank_tp_floor.py \
--ignore=tests/unit/test_trainer_rank_checkpoint_gradient_memory.py \
--ignore=tests/unit/test_trainer_rank_adapter_gradient_memory.py \
--ignore=tests/unit/test_trainer_rank_slot_memory.py \
--ignore=tests/unit/test_trainer_rank_moe_memory.py \
--ignore=tests/unit/test_trainer_rank_head_memory.py \
Expand Down
39 changes: 39 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,45 @@ def context_parallel_rank_model_token_counts(
)


def context_parallel_model_token_total(
*,
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,
) -> int:
"""Return the CP group's model rows in its larger physical layout.

Dispatch runs at least one row on every rank; an empty rank's padding row
passes through the model too.
"""
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,
)
)
total = sum(
max(1, count) for count in bundle.token_layout_index.token_counts_by_rank
)
if not build_gdn_execution_spec:
return total
decision = _plan_gdn_global_execution(
planning_key=planning_key,
bundle=bundle,
topology=topology,
gdn_planner_config=gdn_planner_config,
)
return max(total, sum(max(1, count) for count in decision.gdn_token_counts_by_rank))


def _normalized_chunk_size(
*,
valid_tokens: int,
Expand Down
831 changes: 749 additions & 82 deletions src/art/trainer_rank/_impl.py

Large diffs are not rendered by default.

Loading
Loading