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 TOPK KV blocks named in kv_indices, using a flash-attention online softmax.
  • Physically skips unselected KV blocks via an indirect-DMA gather (oob_mode.skip), so cost scales with TOPK, not with total sequence length.
  • Handles -1 sentinel slots (early query blocks with fewer than TOPK causal predecessors, or CP ranks that do not own a block) with exact -inf masking.
  • 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_tkg in nki-lib) with its own KV-cache machinery.
  • Runtime softmax_scale. The kernels use 1/sqrt(D). The wrappers raise NotImplementedError if 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/TOPK triggers 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_xla model). These kernels are validated for direct invocation (the usage shown below). Embedding the prefill kernel inside a large traced model graph currently hits a neuronx-cc limitation: the Tensorizer/MaskPropagation pass 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 an nl.fori_loop single-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

  1. Indirect DMA gather: for each Q block, gather the TOPK selected KV blocks via .ap(vector_offset=..., indirect_dim=0) with oob_mode.skip for -1 sentinels. This physically skips unselected KV blocks.
  2. Sentinel masking: a per-slot bias derived from kv_indices (-inf where idx < 0, else 0) is added to the scores, so sentinel slots get zero softmax weight -- exact -inf masking, no separate mask tensor needed.
  3. Online softmax: running max + running sum + running output, FP32 in SBUF.
  4. Intra-block causal mask on the diagonal block.
  5. SPMD sharding: kernel[LNC_DEGREE](args) runs one Q-tile per logical-core instance.
  6. 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-cc 2.26.6360, NKI 0.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
-
kernel
neuron
nki
trainium
inferentia
attention
sparse-attention
block-sparse-attention
msa
minimax-m3
gqa
prefill
flash-attention
apache-2.0