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 sets can_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 where torch is a CUDA-tagged build, import libtorch_neuronx_lite (or torch_xla) before get_kernel(...) so torch.neuron is registered and the torch-neuron variant is selected. On a PyTorch Native container (torch.neuron already 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 under torch.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 rsqrt is 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 — ModularAllocator is imported from it).
  • Transformers with KernelConfig support; kernels library.
  • NKI_ENABLE_TRACE_CACHE=0 in 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
-
kernel
neuron
nki
trainium
trn2
vjepa2
video
vision
attention
encoder
apache-2.0