vit-encoder-neuron-kernels
NKI kernels that run the whole transformer stack of RoPE vision-transformer encoders on AWS Trainium2 (trn2). Each block is 2-3 coarse kernels instead of dozens of small ops: attention plus the residual add and LayerNorm, and the MLP / SwiGLU FFN plus the next block's LayerNorm. The residual stream stays in fp32 between kernels.
Drop-in layers are included for two Hugging Face Transformers models:
| Layer | Replaces | Models |
|---|---|---|
NeuronDINOv3ViTModel |
DINOv3ViTModel.forward |
DINOv3 ViT-S/16, S+/16, B/16, L/16, H+/16 |
NeuronVJEPA2Encoder |
VJEPA2Encoder.forward |
V-JEPA2 ViT-L (video), up to 2048 tokens per clip |
The building blocks (BlockWeights / pack_block / EncoderRunner) are model-agnostic. Any
pre-norm ViT with the shape constraints below can use them, but only the two models above
have been validated.
Performance (trn2.3xlarge, bf16, torch.compile(backend="neuron"))
Measured 2026-10 on PyTorch Native Beta 6 (torch 2.13.0, torch-neuronx 2.13.3, NKI 0.7.0b1).
Real Hugging Face weights. "Stock" is the same model compiled with
torch.compile(backend="neuron") without these kernels: DINOv3 cast to bf16, V-JEPA2 as its fp32
checkpoint (stock V-JEPA2 is the encoder as shipped, and its accuracy is in the table below).
Per instance (one process per logical core, all cores busy):
| Model | Resolution | Stock (img/s) | These kernels (img/s) | Speedup |
|---|---|---|---|---|
| DINOv3 ViT-B/16 | 224 | 2,901.5 | 4,105.2 | 1.41x |
| DINOv3 ViT-L/16 | 224 | 1,061.2 | 1,767.5 | 1.67x |
| DINOv3 ViT-L/16 | 448 | 201.1 | 491.9 | 2.45x |
| DINOv3 ViT-H+/16 | 224 | 400.0 | 626.6 | 1.57x |
| DINOv3 ViT-H+/16 | 448 | 106.8 | 200.3 | 1.88x |
| V-JEPA2 ViT-L, 16 frames x 256 px | - | 36.8 clips/s | 157.3 clips/s | 4.27x |
Configuration: LNC=1 (8 logical cores) for every row except V-JEPA2 stock, whose best was LNC=2 (4 cores). DINOv3 kernels run at batch 8 (224 px) or batch 2 (448 px). DINOv3 stock runs at batch 4 (224 px) or batch 1 (448 px), its best measured. V-JEPA2 runs at batch 1 per core for both.
Single core, latency-oriented (LNC=2, batch 1):
| Model | Resolution | Stock | These kernels | Speedup |
|---|---|---|---|---|
| DINOv3 ViT-B/16 | 224 | 251.0 img/s | 510.7 img/s | 2.03x |
| DINOv3 ViT-L/16 | 224 | 97.1 img/s | 220.3 img/s (4.5 ms) | 2.27x |
| DINOv3 ViT-L/16 | 448 | 25.1 img/s | 83.0 img/s | 3.31x |
| DINOv3 ViT-H+/16 | 224 | 21.4 img/s | 88.2 img/s | 4.12x |
| DINOv3 ViT-H+/16 | 448 | 17.7 img/s | 32.4 img/s | 1.83x |
| V-JEPA2 ViT-L, 16 frames x 256 px | - | 101.5 ms/clip | 25.8 ms/clip | 3.93x |
The gain grows with resolution: attention cost grows with token count, and these kernels keep attention and its surroundings on chip. At 224 px with large batches the gap narrows. For example, stock ViT-L/16 at LNC=2 batch 4 is 252 img/s and the kernels give 321 img/s (1.27x).
Accuracy
Every result above passes a per-image check against the fp32 CPU model (rel_l2 <= 5e-2). On DINOv3 the kernels are closer to fp32 than stock bf16, because the residual stream and LayerNorm statistics stay in fp32:
| Model | Stock bf16 rel_l2 vs fp32 | These kernels |
|---|---|---|
| DINOv3 ViT-B/16 | 1.4-1.8e-2 | 0.9-1.7e-2 |
| DINOv3 ViT-L/16 | 1.3-1.6e-2 | 0.8-1.0e-2 |
| DINOv3 ViT-H+/16 | 1.0-1.7e-2 | 0.6-1.0e-2 |
| V-JEPA2 ViT-L | 1.4e-2 (fp32 model, compiler-chosen precision) | 3.9e-2 (bf16 matmuls) |
For V-JEPA2 the stock path compiles the fp32 checkpoint and keeps more of it in higher precision, so it is closer to fp32 than the kernels (which run every matmul in bf16).
Quick start
import types, torch
from kernels import get_kernel
from transformers import AutoModel
k = get_kernel("jburtoft/vit-encoder-neuron-kernels", version=1, trust_remote_code=True)
model = AutoModel.from_pretrained("facebook/dinov3-vitl16-pretrain-lvd1689m",
dtype=torch.float32).eval().to("neuron")
model.forward = types.MethodType(k.NeuronDINOv3ViTModel.forward, model)
model = torch.compile(model, backend="neuron")
with torch.no_grad():
out = model(pixel_values.to("neuron")) # [B, 3, H, W] fp32, any batch size
feats = out.last_hidden_state # [B, 1 + 4 registers + patches, 1024]
cls = out.pooler_output
V-JEPA2:
model = AutoModel.from_pretrained("facebook/vjepa2-vitl-fpc64-256", dtype=torch.float32).eval().to("neuron")
model.encoder.forward = types.MethodType(k.NeuronVJEPA2Encoder.forward, model.encoder)
enc = torch.compile(model.encoder, backend="neuron")
emb = enc(clip.to("neuron")).last_hidden_state # clip [B, 16, 3, 256, 256] -> [B, 2048, 1024]
Notes:
- Load the model in fp32. The layer folds LayerNorm affines and LayerScale into the weights once, on the first device forward, and casts the weights to bf16 itself.
torch.compile(backend="neuron")is part of the configuration. The layers setcan_torch_compile = True.- Set
NEURON_LOGICAL_NC_CONFIGbefore importing, and to the same value the process runs with:2(default) or1. The kernels are built for that core configuration. For LNC=1 also passNEURON_CC_FLAGS=--lnc=1. NKI_ENABLE_TRACE_CACHE=0is recommended when you edit kernels. The cross-process trace cache can otherwise reuse a stale NEFF.- On a vLLM-Neuron venv (CUDA-tagged torch),
import libtorch_neuronx_litebeforeget_kernelso thetorch-neuronvariant is selected. On PyTorch Native this is not needed.
Which LNC to use: LNC=1 (8 logical cores, one process per core) gives the most throughput per instance. LNC=2 (4 larger cores) gives the lowest latency per request.
Kernels
| Kernel | What it does |
|---|---|
attention_block |
LN'd hidden -> Q/K/V projection + RoPE + attention + output projection, one 8-head (D=64) or 4-head (D=128) group per program; any number of images per launch; padded keys skipped |
attention_block_l1 |
LNC=1: one program over all head groups, one output (no per-group partials) |
attention_block_l1_ln |
attention_block_l1 plus residual add and LayerNorm in the epilogue |
res_layernorm_fold |
residual add (sum of up to 4 branches) + affine-free LayerNorm |
mlp |
GELU MLP (fc1 -> GELU -> fc2) |
mlp_res_ln |
mlp plus residual add and the next block's LayerNorm in the epilogue |
swiglu_stream |
SwiGLU FFN that streams the weights in chunks of the intermediate dimension (so large FFNs fit on chip) |
swiglu_stream_ln |
swiglu_stream plus residual add and the next block's LayerNorm |
Host-side packing (pack_weights, pack_tables, lane_permutation, pack_block) supports both
RoPE pairings via rope=:
"half": pairs lanejwithj + D/2(rotate-half; DINOv3, EVA-02, HF Llama style)."interleaved": pairs lane2jwith2j + 1(V-JEPA2 style).
Supported shapes and limits
- Bidirectional (encoder) self-attention with RoPE. No causal masks, KV cache, cross-attention, or learned / absolute position embeddings. A model without RoPE could pass identity tables (cos=1, sin=0), but that is untested.
- Pre-norm blocks with standard LayerNorm (not RMSNorm), optional LayerScale.
- hidden size a multiple of 128; head_dim 64 or 128; any head count (padded to a multiple of 8 at D=64 or 4 at D=128).
- <= 2048 tokens per image after padding to a multiple of 128 x LNC (e.g. 720 x 720 px at patch 16). Longer inputs fall back to the stock forward.
- FFN: GELU MLP (
nl.gelu, the erf form) or SwiGLU (SiLU gate). - Inference only (no autograd). The layers refuse in training mode.
- ViT-7B-sized models (hidden 4096) are not covered by the drop-in layer. They need tensor parallelism. The streaming SwiGLU kernel is slower than the compiler's FFN at that width.
- Validated end to end on DINOv3 ViT-B/L/H+ and V-JEPA2 ViT-L only.
Requirements
- trn2 (verified on trn2.3xlarge, LNC=1 and LNC=2).
- PyTorch Native (TorchNeuron) with NKI 0.7 and
nkilib. Verified on Beta 6 (torch 2.13.0, torch-neuronx 2.13.3, NKI 0.7.0b1, neuronx-cc 2.0.404056). - Transformers with DINOv3 / V-JEPA2 (verified on 4.57.6);
kernelsforget_kernel.
Repository layout
build/torch-neuron/
├── __init__.py # public API
├── metadata.json
├── layers.py # NeuronDINOv3ViTModel, NeuronVJEPA2Encoder
├── encoder.py # BlockWeights, pack_block, rope_tables, EncoderRunner
└── nki_kernels/
├── attention_block.py # attention (per head group) + host-side packing
├── attention_block_l1.py # attention, all head groups (LNC=1)
├── attention_block_l1_ln.py # ... + residual + LayerNorm epilogue
├── attention_inner.py # modified copy of nkilib attention_cte (vendor attention core)
├── res_layernorm.py # residual + affine-free LayerNorm
├── mlp.py # GELU MLP
├── mlp_res_ln.py # GELU MLP + residual + next LayerNorm
└── swiglu_stream.py # weight-streaming SwiGLU (+ residual/LayerNorm variant)
Attribution and license
Apache-2.0 (see LICENSE, NOTICE). The attention, MLP and residual/LayerNorm kernels derive from
the V-JEPA2 encoder kernels in
jburtoft/vjepa2-neuron-kernels,
developed by Antonio Mena (@menaman123). This repo generalizes
them: both RoPE pairings, any head count, head_dim 128, batching, padded-key skipping, an LNC=1
all-group variant, and the fused residual/LayerNorm epilogues. swiglu_stream.py is new.
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.
DINOv3 and V-JEPA2 model weights are distributed under their own licenses.
- Downloads last month
- -