Skip to content

[PyTorch] Add no-load-balance THD all-gather CP - #3221

Merged
sudhakarsingh27 merged 16 commits into
NVIDIA:mainfrom
sudhakarsingh27:sudhakars/relaxed-load-balancing-all-gather
Aug 27, 2026
Merged

sudhakarsingh27 merged 16 commits into
NVIDIA:mainfrom
sudhakarsingh27:sudhakars/relaxed-load-balancing-all-gather

Conversation

@sudhakarsingh27

@sudhakarsingh27 sudhakarsingh27 commented Jul 18, 2026 •

Copy link
Copy Markdown
Member

How to use

Select the token partition strategy explicitly and use the same value for both
context-parallel attention configuration and input partitioning:

from transformer_engine.pytorch import CPAttentionLoadBalancingStrategy

strategy = CPAttentionLoadBalancingStrategy.NO_LOAD_BALANCE

layer.set_context_parallel_group(
    cp_group,
    cp_global_ranks,
    cp_stream,
    cp_comm_type="all_gather",
    load_balancing_strategy=strategy,
)

input_ids, labels, position_ids = get_batch_on_this_cp_rank(
    cu_seqlens_padded,
    input_ids,
    labels,
    position_ids,
    cp_group=cp_group,
    load_balancing_strategy=strategy,
)

CPAttentionLoadBalancingStrategy.DUAL_CHUNK_SWAP remains the default.

What changed

  • Add the public CPAttentionLoadBalancingStrategy enum with
    DUAL_CHUNK_SWAP and experimental NO_LOAD_BALANCE strategies.
  • Propagate the selected strategy through TransformerLayer,
    MultiheadAttention, DotProductAttention, and the context-parallel attention
    backends.
  • Assign one contiguous physical-token chunk to each CP rank for
    NO_LOAD_BALANCE, while preserving logical document boundaries through THD
    metadata.
  • Capture the selected strategy for backward so forward and backward use the
    same token layout.
  • Reuse the existing native THD partition-index implementation for CUDA batch
    slicing while retaining the CPU dataloader fallback.
  • Add focused utility, propagation, FusedAttention, and gated unpadded
    FlashAttention 3 coverage.

Why

The default per-document DualChunkSwap partition divides every sequence into
2 * cp_size chunks and performs two attention steps per rank. Partitioning
the complete physical buffer into cp_size contiguous chunks allows one
attention step per rank and supports packed documents whose individual lengths
are not divisible by 2 * cp_size.

This strategy intentionally trades causal load balance for fewer attention
calls.

Scope and constraints

The experimental strategy currently requires THD self-attention,
cp_comm_type="all_gather", causal attention without a sliding window
(window_size=(-1, 0)), equal local Q/K/V physical lengths, and either
FusedAttention or FlashAttention 3. FlashAttention 3 requires
pad_between_seqs=False. FP8 and CUDA graph capture are not supported.

Input partitioning and attention configuration must use the same strategy.

Validation

  • 21 focused context-parallel utility tests passed.
  • CP2 BF16 THD all-gather FusedAttention forward/backward passed for both
    DUAL_CHUNK_SWAP and NO_LOAD_BALANCE.
  • TransformerLayer strategy propagation and the default
    DUAL_CHUNK_SWAP behavior passed focused smoke checks.
  • Autograd forward/backward arity validation and git diff --check passed.
  • Pylint passed with a 10.00/10 score on the changed production files.

@sudhakarsingh27
sudhakarsingh27 force-pushed the sudhakars/relaxed-load-balancing-all-gather branch from 50b86ff to 9645f28 Compare July 18, 2026 00:39
Allow causal THD attention to shard the complete packed token buffer with mirrored context-parallel chunks while retaining document boundaries through sequence metadata. This supports workloads whose individual documents are not divisible by twice the CP size.

Keep the existing per-document partition as the default and reject backend or attention combinations that the prototype has not validated, so existing paths remain unchanged.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
Reuse the existing THD CUDA partition and reorder kernels by representing the complete physical token buffer as one partitioning sequence. Preserve CPU and mixed-device dataloader behavior with a reference fallback.

Rename the opt-in policy to packed_super_sequence to distinguish physical partitioning from THD packing, and add a narrowly gated matched-input benchmark path so the two policies can be compared without workload or timing asymmetry.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
Allow relaxed THD all-gather to assign one contiguous chunk per CP rank while preserving the mirrored policy for compatibility and comparison. Rank-major ownership needs no KV reorder and reduces each rank to one attention step; keep it opt-in while the performance tradeoffs are evaluated.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
Keep per-document partitioning as the default and packed-contiguous as the only packed opt-in so the experimental API has one global ownership contract. Delete the unused 2*CP global metadata and reorder paths, simplify packed metadata to one chunk and one attention step per rank, and retain a negative test for the retired selector.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
Keep the experimental packed-contiguous selection internal to context parallelism so public TransformerLayer and attention APIs retain their existing signatures. Per-document partitioning remains the default unless NVTE_EXPERIMENTAL_CP_AG_THD_PACKED_CONTIGUOUS=1 is set.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
Keep the initial upstream review focused on the production context-parallel implementation. The test development remains available in the preceding commits while the branch tip restores the existing test suite unchanged.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
@sudhakarsingh27
sudhakarsingh27 force-pushed the sudhakars/relaxed-load-balancing-all-gather branch from 9645f28 to 57f3cf8 Compare August 18, 2026 21:44
@sudhakarsingh27
sudhakarsingh27 marked this pull request as ready for review August 18, 2026 23:02
@sudhakarsingh27 sudhakarsingh27 self-assigned this Aug 18, 2026
@greptile-apps

greptile-apps Bot commented Aug 18, 2026 •

Copy link
Copy Markdown
Contributor

Greptile Summary

The PR adds an explicit context-parallel load-balancing strategy and introduces contiguous THD partitioning for no-load-balance all-gather attention.

  • Propagates the strategy from TransformerLayer through the attention modules and backend adapters.
  • Captures the forward strategy for consistent backward token restoration and step selection.
  • Adds THD metadata, partitioning helpers, validation, and distributed backend coverage.

Confidence Score: 5/5

The PR appears safe to merge.

No blocking failure remains.

Important Files Changed

Filename Overview
transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py Adds contiguous THD partitioning and consistently uses the forward-captured strategy for backward restoration, step selection, and gradient reduce-scatter layout.
transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py Stores and forwards the selected context-parallel strategy through backend dispatch, including checkpointed attention calls.
transformer_engine/pytorch/attention/multi_head_attention.py Propagates context-parallel strategy configuration to core dot-product attention.
transformer_engine/pytorch/transformer.py Exposes strategy configuration at the TransformerLayer boundary and forwards it to attention children.
transformer_engine/pytorch/attention/dot_product_attention/backends.py Passes the strategy into FlashAttention and FusedAttention context-parallel execution.
transformer_engine/pytorch/attention/dot_product_attention/utils.py Updates attention utility behavior to account for the selected context-parallel partition strategy.
transformer_engine/pytorch/constants.py Defines the public context-parallel load-balancing strategy enum.
transformer_engine/pytorch/init.py Re-exports the load-balancing strategy through the public PyTorch package.
tests/pytorch/attention/test_cp_utils.py Adds focused coverage for contiguous partition metadata, captured restoration mode, and CPU/CUDA rank slicing.
tests/pytorch/attention/test_attention_with_cp.py Adds distributed FusedAttention and FlashAttention 3 coverage for no-load-balance THD attention.
tests/pytorch/attention/run_attention_with_cp.py Extends the distributed attention harness to generate and partition inputs according to the selected strategy.

Sequence Diagram

sequenceDiagram
    participant U as Caller
    participant TL as TransformerLayer
    participant MHA as MultiheadAttention
    participant DPA as DotProductAttention
    participant CP as CP All-Gather Attention
    U->>TL: set_context_parallel_group(strategy)
    TL->>MHA: propagate CP configuration
    MHA->>DPA: propagate strategy
    DPA->>CP: forward(Q, K, V, strategy)
    CP->>CP: partition and build THD metadata
    CP->>CP: save strategy on autograd context
    CP-->>DPA: attention output
    Note over CP: Backward uses captured strategy
    CP->>CP: restore K/V and select step layout
    CP->>CP: unrestore dK/dV for reduce-scatter
Loading

Reviews (10): Last reviewed commit: "Limit deterministic THD guard to F16 bac..." | Re-trigger Greptile

Comment thread transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py Outdated
Name the single contiguous-chunk policy after its deliberate lack of causal load balancing so the performance tradeoff is explicit. Remove the CPU reference partitioner and require the existing CUDA path for per-document metadata.

Capture the selected layout through backward and add focused helper plus CP2 forward/backward coverage so mutable environment state cannot make the two passes use different token orders.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
@sudhakarsingh27 sudhakarsingh27 changed the title [PyTorch] Add experimental packed-contiguous THD all-gather CP [PyTorch] Add no-load-balance THD all-gather CP Aug 19, 2026
Comment thread transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py Outdated
Validate the experimental policy where context-parallel communication is selected so unsupported combinations fail before entering custom autograd. Passing the captured mode into the internal all-gather call also prevents a second environment lookup from selecting a different layout.

Restore the default CPU dataloader slicing behavior, keep one stream dependency per format path, and exercise padded feature execution so the lean experimental path does not regress existing callers.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
Comment thread transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py Outdated
Comment thread transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py Outdated
Comment thread transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py Outdated
Comment thread transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py Outdated
Refresh the PR on current upstream before addressing review feedback.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
Make token partitioning explicit so input slicing and attention cannot diverge through mutable process state.

Reuse native THD indices for CUDA while preserving CPU dataloader behavior, and cover the supported FusedAttention and unpadded FlashAttention 3 paths.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
Use the shorter public name because the strategy governs both input partitioning and attention execution. Consolidate redundant low-level partition tests into the existing end-to-end slicing coverage.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
Document the experimental contract at the public selection points. Preserve the legacy child-setter invocation for the default strategy so the additive API does not disrupt existing extension modules.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
@sudhakarsingh27

Copy link
Copy Markdown
Member Author

/te-ci pytorch L1

cyanguwa
cyanguwa previously approved these changes Aug 21, 2026
The focused no-load-balance test uses cp_2_0, which reaches the same known cuDNN deterministic THD backward workspace limit as the generic CP matrix. Apply the equivalent sm90 skip so the standalone coverage does not bypass that guard.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
@sudhakarsingh27

Copy link
Copy Markdown
Member Author

/te-ci pytorch

Apply the known sm90 cuDNN backward workspace restriction in the common selector so every caller avoids choosing FusedAttention for unsupported large THD training configurations. Keep inference, smaller shapes, nondeterministic execution, and other architectures unaffected.

Remove test-local skips and make the focused CP case use normal backend availability preflight so coverage follows the production selection path.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
Refresh the PR with current upstream attention fixes. Resolve the overlapping context-parallel changes by preserving upstream THD padding and stream-safety behavior while retaining strategy-specific single-chunk all-gather handling.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
Comment thread tests/pytorch/attention/test_cp_utils.py
cyanguwa
cyanguwa previously approved these changes Aug 27, 2026

@cyanguwa cyanguwa left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Two minor comments - please run the local CI at least and I'm happy to approve again. Thanks for the PR!

Comment thread tests/pytorch/attention/test_attention_with_cp.py
Comment thread transformer_engine/pytorch/attention/dot_product_attention/utils.py
The empirical large-workspace limitation applies to deterministic F16/BF16 THD backward. Keep the FP8 fused-attention backend available for the same geometry.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
@sudhakarsingh27
sudhakarsingh27 merged commit 902468d into NVIDIA:main Aug 27, 2026
11 of 16 checks passed
fheinecke pushed a commit that referenced this pull request Aug 31, 2026
* Add packed THD partitioning for all-gather CP

Allow causal THD attention to shard the complete packed token buffer with mirrored context-parallel chunks while retaining document boundaries through sequence metadata. This supports workloads whose individual documents are not divisible by twice the CP size.

Keep the existing per-document partition as the default and reject backend or attention combinations that the prototype has not validated, so existing paths remain unchanged.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>

* Refine THD super-sequence partitioning

Reuse the existing THD CUDA partition and reorder kernels by representing the complete physical token buffer as one partitioning sequence. Preserve CPU and mixed-device dataloader behavior with a reference fallback.

Rename the opt-in policy to packed_super_sequence to distinguish physical partitioning from THD packing, and add a narrowly gated matched-input benchmark path so the two policies can be compared without workload or timing asymmetry.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>

* Add contiguous THD all-gather partition

Allow relaxed THD all-gather to assign one contiguous chunk per CP rank while preserving the mirrored policy for compatibility and comparison. Rank-major ownership needs no KV reorder and reduces each rank to one attention step; keep it opt-in while the performance tradeoffs are evaluated.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>

* Remove mirrored packed THD partition

Keep per-document partitioning as the default and packed-contiguous as the only packed opt-in so the experimental API has one global ownership contract. Delete the unused 2*CP global metadata and reorder paths, simplify packed metadata to one chunk and one attention step per rank, and retain a negative test for the retired selector.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>

* Gate packed THD partition with env flag

Keep the experimental packed-contiguous selection internal to context parallelism so public TransformerLayer and attention APIs retain their existing signatures. Per-document partitioning remains the default unless NVTE_EXPERIMENTAL_CP_AG_THD_PACKED_CONTIGUOUS=1 is set.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>

* Remove experimental packed THD tests from PR

Keep the initial upstream review focused on the production context-parallel implementation. The test development remains available in the preceding commits while the branch tip restores the existing test suite unchanged.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>

* Rename experimental THD policy and add coverage

Name the single contiguous-chunk policy after its deliberate lack of causal load balancing so the performance tradeoff is explicit. Remove the CPU reference partitioner and require the existing CUDA path for per-document metadata.

Capture the selected layout through backward and add focused helper plus CP2 forward/backward coverage so mutable environment state cannot make the two passes use different token orders.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>

* Move THD no-load-balance checks to CP dispatch

Validate the experimental policy where context-parallel communication is selected so unsupported combinations fail before entering custom autograd. Passing the captured mode into the internal all-gather call also prevents a second environment lookup from selecting a different layout.

Restore the default CPU dataloader slicing behavior, keep one stream dependency per format path, and exercise padded feature execution so the lean experimental path does not regress existing callers.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>

* Expose CP attention load-balancing strategy

Make token partitioning explicit so input slicing and attention cannot diverge through mutable process state.

Reuse native THD indices for CUDA while preserving CPU dataloader behavior, and cover the supported FusedAttention and unpadded FlashAttention 3 paths.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>

* Rename CP load-balancing strategy

Use the shorter public name because the strategy governs both input partitioning and attention execution. Consolidate redundant low-level partition tests into the existing end-to-end slicing coverage.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>

* Mark no-load-balance CP strategy experimental

Document the experimental contract at the public selection points. Preserve the legacy child-setter invocation for the default strategy so the additive API does not disrupt existing extension modules.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>

* Skip known sm90 deterministic THD OOM

The focused no-load-balance test uses cp_2_0, which reaches the same known cuDNN deterministic THD backward workspace limit as the generic CP matrix. Apply the equivalent sm90 skip so the standalone coverage does not bypass that guard.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>

* Move deterministic THD guard to backend selection

Apply the known sm90 cuDNN backward workspace restriction in the common selector so every caller avoids choosing FusedAttention for unsupported large THD training configurations. Keep inference, smaller shapes, nondeterministic execution, and other architectures unaffected.

Remove test-local skips and make the focused CP case use normal backend availability preflight so coverage follows the production selection path.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>

* Limit deterministic THD guard to F16 backend

The empirical large-workspace limitation applies to deterministic F16/BF16 THD backward. Keep the FP8 fused-attention backend available for the same geometry.

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>

---------

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
(cherry picked from commit 902468d)
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants