Add trainable QSA sparse attention with a TileLang backend
- 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
Assessment
This issue has not been assessed yet.