--- library_name: kernels license: apache-2.0 tags: - kernel - neuron - nki - trainium - inferentia - attention - sparse-attention - block-sparse-attention - msa - minimax-m3 - gqa - prefill - flash-attention --- # 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 ```python 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 ```python # 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`](https://huggingface.co/kernels/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](./LICENSE).