rope-hf-neuron-kernels
Surfaces the nki-lib rope_hf kernel (rotary position embedding, HuggingFace
[batch, heads, seq, head_dim] layout) 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
rope_hffrom the installednki-libpackage (import namenkilib) and re-exports thinwrap_nkiwrappers. The kernel source is maintained in nki-lib; this repo is a discovery/loading shim.
When this kernel helps (read this first)
Benchmarked embedded in a real Qwen2.5-1.5B attention block, whole model
compiled with torch.compile(backend="neuron"), against neuronx-cc-fused native
PyTorch RoPE (the fair baseline). The verdict depends on sequence length:
| Sequence length | nki-lib rope_hf vs compiled native RoPE |
|---|---|
| Short (S ≤ 256) | WINS (~1.34x) |
| Mid (S = 512–2048) | loses (0.89–0.98x) — use native RoPE here |
| Long (S ≥ 4096) | WINS (1.18–1.62x) |
Use this kernel for short or long sequence workloads. In the mid-range,
the neuronx-cc-fused native RoPE is as fast or faster — the kernel's fixed
HOP-boundary cost isn't amortized there. At long S, native rotate_half RoPE's
intermediate materialization blows up (e.g. at S=4096 native jumps to 21 ms vs the
kernel's 13 ms) and the fused BF16 kernel pulls clearly ahead.
Requirements
nki— declared inmetadata.jsonpython-depends(kernels neuron-backend allow-list; NVIDIA's CUDA equivalent isnvidia-cutlass-dsl).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 Beta 5+;torch.neuronis registered natively there — no shim needed). - Trainium (trn2 recommended).
Usage
from kernels import get_kernel
rope = get_kernel("jburtoft/rope-hf-neuron-kernels", version=1, trust_remote_code=True)
# functional (paired Q/K), HuggingFace layout [B, H, S, D]:
q_rot, k_rot = rope.rope_hf_pair(q, k, cos, sin, lnc=2)
# or as an nn.Module inside a torch.compile(backend="neuron") model:
rope_mod = rope.RopeHF(lnc=2)
q_rot, k_rot = rope_mod(q, k, cos, sin)
cos/sinmay be[S, D](shared) or[B, S, D](per-batch).q_heads != k_heads(GQA) is supported.- Constraint:
Smust be divisible by128 * lnc(256 atlnc=2).
Validated results (trn2.3xlarge, LNC=2, PyTorch Native Beta 5)
Qwen2.5-1.5B attention block (RMSNorm + QKV proj + rope + SDPA + o_proj), 4 layers stacked, whole model compiled, BF16, prefill. nki-lib rope_hf (wrap_nki) vs compiled native RoPE, identical weights:
| S | native (ms) | nki-lib (ms) | speedup | parity (cos_sim) |
|---|---|---|---|---|
| 256 | 0.840 | 0.627 | 1.34x | 0.999963 |
| 512 | 1.067 | 1.092 | 0.98x | 0.999940 |
| 1024 | 1.879 | 2.018 | 0.93x | 0.999902 |
| 2048 | 3.029 | 3.396 | 0.89x | 0.999860 |
| 4096 | 21.045 | 13.027 | 1.62x | 1.0001 |
| 8192 | 38.101 | 27.088 | 1.41x | 1.0005 |
| 16384 | 96.274 | 75.954 | 1.27x | 1.0026 |
| 32768 | 286.150 | 242.753 | 1.18x | 1.0081 |
Parity is high throughout (BF16 accumulation drift grows slightly at very long S,
still well within tolerance). See examples/.
License
Apache-2.0 (same as nki-lib). The kernel is © Amazon.com, Inc. and distributed as part of nki-lib.
- Downloads last month
- -