V-JEPA2 ViT-L Encoder NKI Kernels for AWS Trainium2
Version 1 (2026-09-25). Four NKI kernels replacing the per-layer torch ops of the
V-JEPA2 ViT-L video encoder on trn2:
4.35x over the stock torch path (104.2 → 23.9 ms/clip), MFU 10.1% → 44.1% (neuron-explorer
mfu_estimated_percent, device window). Output is bitwise identical to the reference
implementation this was verified against, and rel_l2 4.2e-2 vs the fp32 torch encoder (the
model's own bf16-matmul sensitivity).
What ships
| Component | Purpose |
|---|---|
NeuronVJEPA2Encoder (layer) |
KernelConfig-mappable drop-in for VJEPA2Encoder: the 4-kernel layer loop, lazy one-time weight fold, stock fallback off-device/off-geometry |
res_layernorm_fold |
K1/K3: residual add + affine-free LayerNorm (fp32 stream in, bf16 normalized out); the LN affines are folded into the neighboring matmuls |
attention_block |
K2: fused QKV projection (LN-gain folded, lane-permuted) + 3D RoPE + flash-attention inner + fused out-projection → [2, N, H] per-core partials (LNC=2) |
mlp |
K4: fc1 + GELU + fc2, bf16 matmuls, fp32 boundary, prologue-prestaged row blocks |
pack_*, build_rope_tables |
host-side weight folding (used internally by the layer; exported for direct kernel callers) |
Quick start
import torch
from transformers import AutoModel, KernelConfig
kernel_config = KernelConfig({
"VJEPA2Encoder": "jburtoft/vjepa2-neuron-kernels:NeuronVJEPA2Encoder",
})
model = AutoModel.from_pretrained(
"facebook/vjepa2-vitl-fpc64-256",
dtype=torch.float32,
kernel_config=kernel_config,
device_map="neuron",
).eval()
enc = torch.compile(model.encoder, backend="neuron") # the published numbers are of this graph
video = your_clip # [1, 16, 3, 256, 256] fp32
with torch.no_grad():
embeddings = enc(video).last_hidden_state # [1, 2048, 1024]
torch.compile(backend="neuron")is part of the configuration. Eager dispatch adds ~5.5 ms/forward. The layer setscan_torch_compile, so the kernels library will not silently swap it out under compile.
Resolving via
get_kernel/kernels: on a vLLM (0.24) venv wheretorchis a CUDA-tagged build,import libtorch_neuronx_lite(ortorch_xla) beforeget_kernel(...)sotorch.neuronis registered and thetorch-neuronvariant is selected. On a PyTorch Native container (torch.neuronalready registered) this import is unnecessary.
Performance (trn2.3xlarge, verified 2026-09-25)
| stock torch encoder | these kernels | |
|---|---|---|
| throughput, compiled graph | 104.18 ms/clip | 23.94 ms/clip (4.35x) |
device MFU (mfu_estimated_percent) |
10.1% | 44.08% |
| end-to-end serving, preprocess on device + double-buffered uint8 H2D | — | 24.59 ms/clip = 42.8% |
| end-to-end single stream, nothing overlapped | — | 29.86 ms = 35.7% |
MFU is neuron-explorer's own field: model FLOPs (1.6557e12, equal to the closed-form QK+PV+projections+MLP+patch count to 0.000%) over 157.3 TFLOP/s (one logical NeuronCore = 2 physical cores, bf16) times the window. One logical core of four: a single stream uses ~11% of the chip; four concurrent streams reach ~44% chip-wide.
Where the remaining 2.3x to physics goes (measured, not estimated): the softmax exp runs on the
Scalar engine at 1 element/lane/cycle (+2.6 ms over physics — unreachable without different
hardware), instruction issue overhead (3.8 ms), and the queueing of three engines co-bound
within 0.5% (3.7 ms). The exact-math kernel-level candidate classes are enumerated to
exhaustion; details in the source repository's cycle records.
Correctness
- Full-forward output bitwise identical (
torch.equal) to the verified reference implementation, including through this repo's KernelConfig-style forward graft undertorch.compile. - rel_l2 vs the fp32 torch encoder: 4.2028e-2 — identical to the stock-vs-fp32 difference (both paths run bf16 matmuls; neither is "more approximate" than the other).
- K1/K3's GpSimd
rsqrtis measurably closer to fp64 than the torch path's (6.9e-8 vs 3.9e-5 relative). - The vendor attention inner is run-to-run nondeterministic on some layers at a few bf16 positions (fp64 agreement unchanged) — an upstream kernel property, reported to the Neuron team. Bitwise comparisons must use a fixed captured sample.
Requirements
- Hardware: trn2 (verified on trn2.3xlarge, LNC=2).
- SDK: neuronx-cc 2.27.2878.0; torch-neuronx 2.11.3 (PyTorch Native); nki 0.6.0;
nkilib (bundled with the SDK —
ModularAllocatoris imported from it). - Transformers with
KernelConfigsupport;kernelslibrary. NKI_ENABLE_TRACE_CACHE=0in the environment (the runtime defaults the cross-process trace cache on; a stale NEFF silently captures a different instruction stream).
Repository layout
build/torch-neuron/
├── __init__.py # public API re-exports
├── metadata.json # kernels-library metadata (backend: neuron)
├── layers.py # NeuronVJEPA2Encoder (KernelConfig target)
└── nki_kernels/
├── res_layernorm.py # K1/K3
├── attention_block.py # K2 outer (QKV + RoPE + out-proj) and weight packing
├── attention_inner.py # K2 inner: MODIFIED copy of the Neuron SDK's nkilib attention_cte
├── mlp.py # K4
└── rope_tables.py # host-side 3D RoPE tables
Known limitations
- 16 frames / 2048 tokens / batch 1 / 256x256 only. The model's native 64-frame clip (8192 tokens) is NOT supported: three of the vendor attention kernel's SBUF buffers scale with keys-per-section (347 KiB/lane needed vs the 192 KiB budget, device-confirmed), so 8192 keys require the flash-sectioning path, whose cross-section combine this integration does not yet reproduce correctly. Off-geometry inputs fall back to the stock implementation.
- Inference only — no autograd rule; the layer refuses in training mode.
- The first device forward performs the one-time weight fold eagerly (~seconds); it is
deliberately kept outside
torch.compile(compiled folding rounds differently and breaks bitwise parity with the verified reference).
Attribution and license
Apache-2.0 (see LICENSE, NOTICE). nki_kernels/attention_inner.py is a modified copy of
the AWS Neuron SDK's nkilib core/attention/attention_cte.py (Copyright Amazon.com, Inc.
or its affiliates, Apache-2.0); the file header lists the modifications (four engine-placement
pins worth ~1.19x end-to-end, and removal of code paths unreachable for this configuration).
The other three kernels, the layer glue and the packing are original work.
- Downloads last month
- -