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_kernel from the installed nki-lib package (import name nkilib) and re-exports wrap_nki wrappers. 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 in metadata.json python-depends (accepted by the 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).
  • 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] with H a multiple of 128. Reshape a [B, S, hidden] RowParallel output to [B*S, hidden]; fall back to dist.all_reduce when B*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
-
kernel
neuron
trainium
nki
collectives
tensor-parallel
apache-2.0