NVIDIA / NVIDIA/Megatron-LM

Add trainable QSA sparse attention with a TileLang backend

Open
#7,060 4 comments 1 reaction 1 assignee Claimed by @BestJuly View on GitHub
community-request enhancement waiting-on-customer
Dominant language
Python
Stars
17.9k
Forks
4.5k
Avg merge
4d 6h
Merged PRs (30d)
271

Description

**Is your feature request related to a problem? Please describe.**

Megatron Core currently has experimental sparse-attention support for DSA/MLA-style models, but it does not provide the QSA (Qwen Sparse Attention) training path used by Qwen3.8-Flash-Next.

Supporting QSA training requires several components that are not covered by the existing DSA implementation:

- A trainable QSA indexer.
- Exact block TopK selection.
- Sparse GQA attention forward and backward.
- Sparse KL supervision for training the indexer.
- Integration with activation recomputation and Megatron parallelism.
- Bounded-memory execution without materializing the full `[sequence, sequence / block_size]` score matrix.

@NVIDIA/mcore-oncall Could you please advise whether the proposed architecture and PR split below align with the Megatron Core direction?

**Describe the solution you'd like**

I have implemented a prototype QSA training backend under
`megatron/core/transformer/experimental_attention_variant/`.

The implementation contains:

- A PyTorch QSA reference implementation used as a correctness oracle.
- A trainable QSA indexer with official score scaling.
- Deterministic exact block TopK selection.
- TileLang Sparse-GQA forward and backward kernels.
- A streaming indexer runner with bounded memory usage.
- Sparse Stage-2 teacher/KL loss for indexer training.
- MCore `SelfAttention` integration.
- BSHD and packed THD/variable-length paths.
- Tensor/context parallel integration and activation-recomputation support.
- Hardware-aware TileLang pipeline-stage selection based on available shared memory.

The deterministic TopK implementation uses a lexicographic
`(score, block_id)` key. This is required because ReLU indexer scores commonly
produce ties at zero; nondeterministic tie handling can select different routes
during activation recomputation or across tensor-parallel ranks.

The current prototype has been validated with:

- PyTorch oracle forward/backward comparisons.
- QSA indexer score and route comparisons against mcore-bridge.
- Numerical and gradient tests for Sparse-GQA and sparse KL.
- Repeated deterministic TopK tests, including tied-score cases.
- Activation-recomputation route consistency.
- Multi-rank TP/CP tests.
- A real Qwen3.8-Flash-Next checkpoint training smoke test using the BSHD path
on 8 GPUs, with all QSA layers using the TileLang backend and no dense/PyTorch
attention fallback.

The intended scope is training. Inference and KV-cache support are not included
in the initial proposal.

To keep reviews manageable, I propose splitting the work into a PR series:

1. QSA mathematical reference, indexer contracts, and sparse KL loss.
2. TileLang streaming indexer and deterministic exact TopK.
3. TileLang Sparse-GQA forward/backward kernels and BSHD integration.
4. THD/variable-length and distributed TP/CP integration.

I will rebase each change on the latest `main`, reuse the shared TileLang
utilities already present in Megatron Core, include focused unit/functional
tests, and preserve attribution for adapted Apache-2.0 kernel implementations.

**Describe alternatives you've considered**

1. Reusing the existing DSA implementation directly.

DSA provides useful organization and distributed-training patterns, but its
MLA score representation and attention kernels do not match QSA's standard
GQA layout and block routing semantics.

2. Using the existing mcore-bridge dense-mask implementation.

This is useful as a routing reference, but it materializes a dense attention
mask and does not provide the bounded-memory TileLang Sparse-GQA training
path or trainable indexer KL required for long-context training.

3. Materializing all indexer block scores and applying `torch.topk`.

This is simple for correctness testing but has quadratic memory growth and
is unsuitable for the target long-context training workload.

**Additional context**

The implementation was developed for the QSA layers in
Qwen3.8-Flash-Next. It was organized using the existing experimental DSA
attention variant as the closest Megatron Core reference.

The exact TopK selection structure is adapted from TileLang's Apache-2.0
DeepSeek-V3.2 TopK example, with deterministic tie-breaking added for training
and activation recomputation.

Known initial-scope limitations:

- Inference and KV-cache support are not included.
- The final THD/TP/CP support matrix will be stated explicitly in the relevant
PR and guarded by fail-fast validation for unsupported combinations.

Before preparing the PR series, I would appreciate feedback on:

1. Whether QSA should remain an experimental attention variant.
2. Whether the proposed four-PR split is appropriate.
3. Whether BSHD should be merged first, followed by THD/CP support.
4. Which GPU architectures and distributed configurations are required for
acceptance.

Contributor guide

Open the contributing guide

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.