all-reduce-hbm-neuron-kernels
Surfaces the nki-lib all_reduce_hbm_kernel (HBM-based all-reduce-SUM for
tensor-parallel layers on Trainium) 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
all_reduce_hbm_kernelfrom the installednki-libpackage (import namenkilib) and re-exportswrap_nkiwrappers. The kernel source is maintained in nki-lib; this repo is a discovery/loading shim.
Drop-in replacement for torch.distributed.all_reduce on the all-reduce-SUM that
ends every RowParallelLinear forward (attention o_proj, MLP down_proj).
Runs BF16 directly (no FP32 cast) and beats the framework collective at the
small/medium payloads of a single-node TP layer.
Requirements
nki— declared inmetadata.jsonpython-depends(accepted by thekernelsneuron-backend allow-list).nki_library(nkilib) — the actual kernel source; imported at load time. Not on thekernelsallow-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). - A Neuron distributed context:
dist.init_process_group(backend="neuron"). - Trainium (trn2 recommended), TP degree ≥ 2.
Usage
# PyTorch Native Beta 5: torch.neuron is registered natively — no shim needed.
import torch, torch.distributed as dist
from kernels import get_kernel
dist.init_process_group(backend="neuron")
world = dist.get_world_size()
ranks = list(range(world))
ar = get_kernel("jburtoft/all-reduce-hbm-neuron-kernels", version=1, trust_remote_code=True)
# inside RowParallelLinear.forward, after the local matmul (y: [B*S, hidden]):
y = ar.all_reduce(y, ranks, lnc=2) # out-of-place: use the return value
# trainable TP layer (identity backward):
y = ar.NkiAllReduce.apply(y, ranks, 2)
- Input must be 2-D
[H, W]withHa multiple of 128. Reshape a[B, S, hidden]RowParallel output to[B*S, hidden]; fall back todist.all_reducewhenB*S % 128 != 0. - Out-of-place — always use the returned tensor.
Validated example — Mistral-7B, TP=4
mistralai/Mistral-7B-v0.3
hidden=4096, RowParallelLinear all-reduce payloads, BF16, vs dist.all_reduce,
trn2.3xlarge TP=4 LNC=2, PyTorch Native Beta 5:
Payload ([B*S, 4096]) |
Size | dist.all_reduce |
nki-lib HBM | Speedup | Parity (cos_sim) |
|---|---|---|---|---|---|
| S=512 | 4.2 MB | 264–270 µs | 149–162 µs | 1.67–1.78x | 0.999868 |
| S=2048 | 16.8 MB | 466–478 µs | 392–394 µs | 1.19–1.21x | 0.999972 |
The win is largest at smaller payloads (framework launch overhead dominates) and
narrows as the op becomes bandwidth-bound (converges to the framework ~128 MB).
Prefer the framework collective for very large or cross-node bandwidth-bound
all-reduces (e.g. DP gradient all-reduce over EFA). See examples/.
Sibling HBM collectives in the same nki-lib module: all_gather_hbm_kernel,
reduce_scatter_hbm_kernel, all_to_all_hbm_kernel.
License
Apache-2.0 (same as nki-lib). The kernel is © Amazon.com, Inc. and distributed as part of nki-lib.
- Downloads last month
- -