taconite-clip — CLIP ViT-H/14 on the NPU, from Rust
A Rust runtime for
laion/CLIP-ViT-H-14-laion2B-s32B-b79K:
image and text embeddings and zero-shot classification, with both
transformers on the AMD XDNA NPU (NPU2) through IRON kernels replayed with
taconite. It is the forward of
IRON's Python app
(iron/applications/clip_vit_h14), kernel for kernel.
# cats.jpg: cat 1.000, dog 0.000, car 0.000
# car.png: car 0.999, cat 0.000, dog 0.000
let mut clip = load?;
let px = clip.preprocess; // [3, 224, 224]
let img = clip.encode_images?; // [n, 1024]
let txt = clip.encode_texts?; // [m, 1024]
let logits = logits;
What runs where
| NPU | host (this crate) | |
|---|---|---|
| image | all 32 layers: qkv / o / fc1 (+GELU) / fc2 GEMMs, attention (16 heads of 80), residual adds + LayerNorms, each kernel reading the last one's output in place (tower.rs) |
decode (the image crate), resize + crop + normalise (preprocess.rs), the patch embedding (f32), CLS + position embedding, pre-/post-LayerNorm, projection |
| text | all 24 layers, the same kernels, causal attention | BPE tokenizer (sam3's), token + position embedding, final LayerNorm at the end token, projection |
8 hardware contexts (NPU2 has 16), all loaded at start-up; 4 images and 8
prompts a pass. The library is std-only; the CLI adds the image crate.
- Preprocessing is HF's
CLIPImageProcessor, bit for bit: torchvision's antialiased bicubic resize of the uint8 image (its fixed-point two-pass resampler, int16 weights), center crop, the fused(x - 255 mean) / (255 std)in f32. On a PNG every one of the 150528 values matches; JPEGs differ only through decoding (theimagecrate vs PIL, a level here and there). - The patch embedding runs on the host in f32, as the model's conv. As an NPU GEMM its bf16 output, rounded before the position embedding and pre-LayerNorm, cost ~0.0025 image-embedding cosine for ~25 ms.
- Images sit 264 rows apart (257 tokens rounded up to 8), so an image's embedding does not depend on the others in its pass; see the Python app's README.
Validation
clip check <bundle> on NPU2 (two reference sets of 4 images from the
float32 model on the Radeon 890M via ROCm):
tokenizer:
[ok] 8/8 texts tokenize as HF's CLIPTokenizer
reference set 0: 4 images x 8 prompts (cat, dog, car, person, laptop, kitchen, bicycle, pizza)
[ok] preprocess 0_car.png: 0 of 150528 values differ from HF's (max 0.0000)
[ok] image embeddings: cosine to float32 min 0.99944 (mean 0.99966); to the Python NPU app min 0.99984
[ok] text embeddings: cosine to float32 min 0.99912 (mean 0.99954); to the Python NPU app min 1.00000
[ok] zero-shot: top-1 agrees 4/4; logits max |diff| 0.350, probabilities 0.0002
[ok] 0_cats.jpg: cat 1.000 (reference cat 1.000)
...
reference set 1: 4 images x 12 prompts (tabby cat, siamese cat, kitten, ...)
[ok] zero-shot: top-1 agrees 4/4; logits max |diff| 0.344, probabilities 0.0057
[ok] 1_cats.jpg: tabby cat 0.884 (reference tabby cat 0.871)
...
ALL CHECKS PASSED
The text tower matches the Python app exactly; the image side differs from it only by the host's float arithmetic (patch embedding, LayerNorms).
Performance
Warm, 4 images (decoded from files) x 8 labels, clip classify, Ryzen AI 9
HX 370: ~1.3-1.5 s a forward (the spread is the machine's power state),
of which the NPU is ~1.2-1.4 s (vision fc1 0.27, qkv 0.12, o 0.10, fc2
0.12, MHA 0.08, AddLN 0.07; text 0.45) and the host ~80 ms (decode and
preprocess ~50, the patch embedding ~30). Loading the bundle takes 0.7 s
(the Python app compiles and packs for ~30 s). The float32 model on the
890M iGPU takes 2.75 s for the same forward.
Building and running
# the bundle (in the NPU container, see the Python app)
# the binary (links XRT; see below for the XRT-free build)
To run without XRT, build with --no-default-features --features cli,direct.
The kernels then go through the amdxdna driver's ioctls
(taconite::direct); the build needs no XRT headers or library, and the
binary links none. On NPU2, check
prints the same numbers as the XRT build, at the same speed (~1.29 s for 4
images × 8 labels).
The bundle (1.3 GiB) holds every kernel, the packed weights, the tokenizer,
the reference sets and their images; its format is taconite-bundle's plus
the records export_clip.py's docstring lists.
Not done
- Batches are fixed at export (4 images, 8 prompts a pass); a smaller request still runs the full pass.
- The tokenizer and host math come from the
taconite-sam3crate; a shared runtime-helpers crate would be the cleaner home.
License
Apache-2.0, except src/preprocess.rs, whose resampler ports torch's
antialiased uint8 resize (BSD-3-Clause, LICENSE-PYTORCH), itself after
Pillow's (MIT-CMU, LICENSE-PILLOW).