scatter-add-neuron-kernels

Surfaces the nki-lib scatter_add kernel on the Hugging Face Kernel Hub, loadable through the standard kernels / get_kernel interface.

This repo does not contain a copy of the kernel. It imports scatter_add from the installed nki-lib package (import name nkilib) and re-exports a wrap_nki wrapper. The kernel source is maintained in nki-lib; this repo is a discovery/loading shim.

Computes input[index[i], :] += src[i, :] (PyTorch scatter_add on dim 0). torch.Tensor.scatter_add_ is a severe performance cliff on Neuron — hundreds of ms for a few-MB update — and at some shapes the neuronx-cc compiler cannot lower it at all. This kernel replaces it. Common use: MoE token dispatch/combine (scatter expert outputs back to token positions), e.g. Mixtral / Qwen2.5-MoE.

Requirements

  • nki — declared in metadata.json python-depends (kernels neuron-backend allow-list).
  • nki_library (nkilib) — the actual kernel source, imported at load time. Not on the kernels allow-list, so not declared in metadata; an external runtime requirement. Pre-installed in the AWS PyTorch Native Beta container and on the AWS Neuron pip index (not public PyPI).
  • A torch.neuron-registered PyTorch build (PyTorch Native / TorchNeuron).
  • Trainium (trn2 recommended).

Usage

# PyTorch Native Beta 5: torch.neuron is registered natively — no shim needed.
from kernels import get_kernel
sa = get_kernel("jburtoft/scatter-add-neuron-kernels", version=1, trust_remote_code=True)

# input [N, D], index [K] int32, src [K, D]
out = sa.scatter_add(input, index, src, lnc=2)   # in-place accumulate + return

# or as an nn.Module inside a torch.compile(backend="neuron") model:
out = sa.ScatterAdd(lnc=2)(input, index, src)

Constraints: input/src 2-D, index 1-D, dim=0. Indices within any 128-row tile should be unique (standard scatter-add correctness rule).

Validated results (trn2.3xlarge, LNC=2, PyTorch Native Beta 5)

MoE-dispatch shapes, D=2048, BF16. Both baselines run on-device; the compiled baseline is torch.scatter_add inside torch.compile(backend="neuron"):

Shape (N×D, K) Parity nki-lib torch (compiled) torch (eager) vs compiled vs eager
N=2048, K=2048 0.999907 0.271 ms 1.023 ms 587.9 ms 3.8x 2171x
N=4096, K=8192 ~1.0 0.721 ms compiler cannot lower 7501 ms only working path 10401x

The kernel beats both eager and neuronx-cc-compiled torch.scatter_add. At the larger shape the compiler fails to lower torch.scatter_add (int64→index materialization), so the nki-lib kernel is the only working device path. See examples/.

Embedded in a real Mixtral-style MoE layer

The HF MoE combine step is final_hidden_states.index_add_(0, token_idx, current) per expert — exactly this kernel. Measured in a faithful Mixtral-style MoE layer (hidden=4096, 8 experts, top_k=2, static expert capacity), whole layer compiled with backend="neuron", nki scatter vs native index_add_, identical weights:

Tokens (N) capacity C native index_add_ nki scatter speedup parity
512 256 3.36 ms 1.43 ms 2.35x 0.999272
2048 768 8.27 ms 3.35 ms 2.47x 1.000046
4096 1536 17.93 ms 5.97 ms 3.01x 1.001044

The win grows with token count (2.35x → 3.01x). Note: the dynamic HF router (topk/where/data-dependent gather) does not compile on the Neuron backend — real Neuron MoE uses static expert capacity, which is what this measurement reflects.

License

Apache-2.0 (same as nki-lib). The kernel is © Amazon.com, Inc. and distributed as part of nki-lib.

Downloads last month
-
kernel
neuron
trainium
nki
scatter-add
Mixture of Experts
apache-2.0