- MiniMax-M3 MSA Block-Sparse GQA Prefill Kernels for AWS Neuron
MiniMax-M3 MSA Block-Sparse GQA Prefill Kernels for AWS Neuron
A Trainium port of MiniMax's Multi-Scale Sparse Attention (MSA) prefill attention, implemented in NKI for trn2. Computes attention over a caller-selected top-K subset of KV blocks per query block -- the block-sparse pattern used by MiniMax-M3 and other Native-Sparse-Attention / DSA-family models.
Version: 1.0.0
What ships
Three kernels:
block_sparse_gqa_prefill-- SPMD (LNC=2) prefill with normalized output. The main kernel.block_sparse_gqa_prefill_state-- SPMD prefill emitting(o, l, m)flash state for context-parallel merge across KV shards.block_sparse_attention_single_q-- single-Q-block MHA reference kernel with broader shape coverage (useful for the NKI CPU simulator and for porting to other block-sparse layouts).
Plus host helpers (ops.py): a block-KV-index builder, a causal block-mask builder,
a numpy FP32 dense reference, and a cp_merge for combining state across CP ranks.
What it does
- Block-sparse GQA prefill attention: for each query block, attends to only the
TOPKKV blocks named inkv_indices, using a flash-attention online softmax. - Physically skips unselected KV blocks via an indirect-DMA gather (
oob_mode.skip), so cost scales withTOPK, not with total sequence length. - Handles
-1sentinel slots (early query blocks with fewer thanTOPKcausal predecessors, or CP ranks that do not own a block) with exact-infmasking. - Applies an intra-block causal mask on the diagonal block.
- Runs single-core, Q-parallel across the 4 logical cores of a trn2.3xlarge (data
parallelism over disjoint query slices), or context-parallel over KV shards using
the state kernel +
cp_merge.
What it does not do
- Decode. Prefill only. Decode-side attention uses a different kernel
(
attention_tkgin nki-lib) with its own KV-cache machinery. - Runtime
softmax_scale. The kernels use1/sqrt(D). The wrappers raiseNotImplementedErrorif a different scale is passed. - Batch > 1 in a single launch. The wrapper loops over the batch dimension.
- Non-reference shapes without recompile. The kernels specialize on the shapes
they see; a different
H_q/H_kv/D/BS/TOPKtriggers a recompile. Numerics are verified on the reference MiniMax-M3 shape and on Shape A (below). - Sequence lengths beyond 1M on a single trn2.3xlarge. See the ceiling table.
- Embedding in a large traced graph (e.g. an NxDI /
torch_xlamodel). These kernels are validated for direct invocation (the usage shown below). Embedding the prefill kernel inside a large traced model graph currently hits aneuronx-cclimitation: theTensorizer/MaskPropagationpass does not converge on the resulting graph (observed hanging for hours at 300% CPU on a 60-layer context-encoding graph, while the same graph with dense attention compiles in ~4 minutes). This is independent of loop structure -- it reproduces with both the shipped multi-launch kernel and annl.fori_loopsingle-launch variant that compiles in ~30s standalone -- so the cause appears intrinsic to how the block-sparse indirect-DMA gather interacts with that compiler pass. Direct invocation is unaffected.
Performance (MiniMax-M3 shape: H_q=64, H_kv=4, D=128, block_size=128, topk=16, BF16)
Measured on trn2.3xlarge, LNC=2, SDK 2.31 DLAMI 20260708 + PyTorch Native Beta 4 (torch-neuronx 2.11.3.0, neuronx-cc 2.26.6360, NKI 0.5.0). Warm median of 5 runs.
Single-core (one logical NeuronCore)
| seq_len | Blocks kept | NKI kernel |
|---|---|---|
| 8,192 | 25.0% | 95.7 ms |
| 16,384 | 12.5% | 189.1 ms |
| 32,768 | 6.25% | 378.8 ms |
| 65,536 | 3.12% | 755.1 ms |
Linear scaling at ~1.5 ms per query block, constant across seq_len (cost depends on
TOPK, not total length).
Q-parallel across 4 logical cores (recommended for long seq_len)
Run 4 processes, one per logical core (NEURON_RT_VISIBLE_CORES), each computing a
disjoint 1/4 slice of the query sequence against the full KV cache. Output is
bit-identical to single-core. See examples/05_qparallel.py.
| seq_len | Single-core | Q-parallel DP=4 |
|---|---|---|
| 65,536 | 755.1 ms | 195.2 ms |
| 1,048,576 (1M) | OOM | 4342 ms (241K tok/s) |
Sequence-length ceilings (single trn2.3xlarge, BF16)
| Config | Max seq_len | Fails at |
|---|---|---|
| Single-core (1 process) | 262,144 (256K) | 512K (HBM OOM) |
| Q-parallel DP=4 (4 processes) | 1,048,576 (1M) | 2M (host + device memory in the example driver) |
| CP=2 (state kernel, on-device gather) | 65,536 (64K) | 128K |
| CP=4 (state kernel, on-device gather) | 8,192 (8K) | 64K |
For the customer target range (256K-1M): 256K completes in ~1.1 s, 1M in ~4.3 s via
Q-parallel DP=4 on a single trn2.3xlarge. For seq_len beyond 1M, use a
trn2.48xlarge (more logical cores) or a ring-attention decomposition (see
ROADMAP.md).
Correctness
BF16 kernel vs FP32 numpy reference, compared over the full output tensor (sentinel-heavy early query blocks included):
- Shape A (
S=1024, H_q=8, H_kv=2, topk=4): cos_sim 0.9999753 - Shape B, MiniMax-M3 (
S=8192, H_q=64, H_kv=4, topk=16): cos_sim 0.9999754 - Q-parallel DP=4 vs single-core: bit-identical (max_abs_diff = 0.0)
- CP=2 merged vs single-node: cos_sim 0.9999971
- CP=4 multiprocess: cos_sim 1.0012 (BF16 noise floor)
Reproduce via examples/02_parity.py, 04_cp_merge.py, 05_qparallel.py --verify,
06_cp_multiprocess.py, 07_cp_ondevice.py.
Direct usage
import torch
from kernels import get_kernel
m3msa = get_kernel(
"jburtoft/minimax-m3-msa-neuron-kernels",
revision="v1.0.0",
trust_remote_code=True,
)
# MiniMax-M3 shape, BF16, single-node prefill.
B, S, H_q, H_kv, D = 1, 8192, 64, 4, 128
q = torch.randn(B, S, H_q, D, dtype=torch.bfloat16, device="neuron")
k = torch.randn(B, S, H_kv, D, dtype=torch.bfloat16, device="neuron")
v = torch.randn(B, S, H_kv, D, dtype=torch.bfloat16, device="neuron")
# kv_indices: [B, num_q_blocks, TOPK] int32, -1 sentinels for unused slots.
# Produce this from your Lightning-Indexer (or any block-level attention router).
kv_indices = m3msa.build_block_kv_indices(keep_mask, topk=16)
out = m3msa.block_sparse_gqa_prefill(q, k, v, kv_indices)
# out: [B, S, H_q, D] bfloat16
Context-parallel usage
# Each rank owns a KV shard. Indices for blocks not owned by this rank are set to
# the sentinel -1 in the local kv_indices.
o_r, l_r, m_r = m3msa.block_sparse_gqa_prefill_state(q, k_shard, v_shard, kv_indices_local)
# all_gather (o, l, m) across CP ranks, then merge:
states = [(o_r, l_r, m_r) for each rank] # e.g. via torch.distributed.all_gather
out = m3msa.cp_merge(states, output_dtype=torch.bfloat16)
See examples/04_cp_merge.py for a single-device CP=2 simulation with parity to a
single-node run.
Reference config
The kernels are compiled against the MiniMax-M3 reference shape (constants.py):
| Constant | Value | Note |
|---|---|---|
BS |
128 | block size (tokens per K/V block, also Q tile size) |
D |
128 | head dim |
H_Q |
64 | query heads per layer |
H_KV |
4 | KV heads per layer (GQA group = 16) |
TOPK |
16 | top-K KV blocks per Q block |
LNC_DEGREE |
2 | SPMD degree (LNC=2 on trn2.3xlarge) |
Requirements:
S_k % BS == 0,S_q % (BS * LNC_DEGREE) == 0(divisible by 256 by default)H_q % H_kv == 0(GQA)- self-attention prefill (
S_q == S_k), except Q-parallel where each rank passes its Q slice and the full K/V.
How it works
- Indirect DMA gather: for each Q block, gather the
TOPKselected KV blocks via.ap(vector_offset=..., indirect_dim=0)withoob_mode.skipfor-1sentinels. This physically skips unselected KV blocks. - Sentinel masking: a per-slot bias derived from
kv_indices(-infwhereidx < 0, else 0) is added to the scores, so sentinel slots get zero softmax weight -- exact-infmasking, no separate mask tensor needed. - Online softmax: running max + running sum + running output, FP32 in SBUF.
- Intra-block causal mask on the diagonal block.
- SPMD sharding:
kernel[LNC_DEGREE](args)runs one Q-tile per logical-core instance. - The main kernel returns normalized output; the state kernel returns
(o_unnormalized, l, m)for the CP merge:m_final = max(m_r); l_final = sum(l_r * exp(m_r - m_final)) o_final = sum(o_r * exp(m_r - m_final)) / l_final
Repository layout
build/torch-neuron/
├── __init__.py <- public API re-exports
├── metadata.json <- HF kernels metadata
├── constants.py <- BS, D, H_Q, H_KV, TOPK, LNC_DEGREE, sentinel
├── ops.py <- host helpers + numpy reference
└── nki_kernels/
├── block_sparse_gqa_prefill.py <- main SPMD kernel
├── block_sparse_gqa_prefill_wrapper.py
├── block_sparse_gqa_prefill_state.py <- state-output kernel for CP
├── block_sparse_gqa_prefill_state_wrapper.py <- + cp_merge()
└── block_sparse_attention_single_q.py <- single-Q reference kernel
examples/
├── 01_smoke_test.py 02_parity.py 03_benchmark.py
├── 04_cp_merge.py 05_qparallel.py 06_cp_multiprocess.py
├── 07_cp_ondevice.py 08_dp_cp_hybrid.py + launchers
Environment (verified working)
- Instance: trn2.3xlarge, LNC=2
- DLAMI:
Deep Learning AMI Neuron (Ubuntu 24.04) 20260708(SDK 2.31) - PyTorch Native Beta 4: torch-neuronx
2.11.3.0, neuronx-cc2.26.6360, NKI0.5.0 - Also correctness-verified under NKI 0.6.0 (SDK 2.32 pre-release); Beta 4 recommended.
Attribution
A Trainium port of the Multi-Scale Sparse Attention (MSA) pattern from the
MiniMax-M1 / MiniMax-M3 technical reports (MiniMax AI) and the reference GPU
implementation at MiniMaxAI/msa.
The block-sparse top-K-KV-block-selection pattern is related to DeepSeek DSA,
FlashAttention-2 (online-softmax state + CP merge), and NATTEN.
License
Apache 2.0. See LICENSE.
- Downloads last month
- -