Skip to main content

taconite_sapiens2/
lib.rs

1// SPDX-FileCopyrightText: Copyright (C) 2026 Brishen Hawkins
2// SPDX-License-Identifier: Apache-2.0
3
4//! Sapiens2-Pose (`facebook/sapiens2-pose-0.4b`, `-1b`) on an AMD XDNA NPU:
5//! a person box in an image -> 308 keypoints (body, feet, hands, face).
6//!
7//! The bundle `iron/applications/sapiens2_pose/export_sapiens2.py` writes
8//! holds every compiled IRON kernel and the weights (the NPU ones
9//! pre-packed); this crate replays the forward the Python app runs
10//! (`sapiens2_common.py` / `sapiens2_npu.py`):
11//!
12//! | | NPU | host (here) |
13//! |---|---|---|
14//! | backbone (ViT, 3081 tokens; 0.4b: 24 layers, 1024 wide, 1b: 40 layers, 1536 wide) | every projection (`flm.GEMM`s: patch embedding, qkv, o, the SwiGLU gate+up, down) and the attention (the MHA operator) | the box crop (`preprocess.rs`), RMSNorms, q / k norms, 2D RoPE, residual adds |
15//! | head (2 transposed convs to 256 x 192, 3 1 x 1 convs, the predictor) | every convolution as an `flm.GEMM` (a transposed conv as one GEMM over its input's 2 x 2 windows) | the window layout, InstanceNorm + SiLU |
16//! | keypoints | | argmax + DARK refinement, back through the crop (`post.rs`) |
17//!
18//! [`Sapiens2::pose`] gives a box's keypoints in image coordinates and
19//! their heatmap scores, as HF's `post_process_pose_estimation`.
20
21use std::fmt;
22use std::path::Path;
23
24use taconite_bundle::{Manifest, Store};
25
26pub use taconite::Timing;
27
28pub mod model;
29pub mod npu;
30pub mod post;
31pub mod preprocess;
32
33use model::Model;
34use npu::Npu;
35pub use post::Keypoint;
36pub use preprocess::BBox;
37
38/// The bundle format: 2 when some GEMM leaves its bias to the host
39/// (`<i>.down.bias`, `d<j>.bias`: 1b); 1 (0.4b) is read too.
40pub const VERSION: u32 = 2;
41
42#[derive(Debug)]
43pub enum Error {
44    Bundle(String),
45    Npu(String),
46    Input(String),
47}
48
49impl fmt::Display for Error {
50    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
51        match self {
52            Error::Bundle(m) => write!(f, "bundle: {m}"),
53            Error::Npu(m) => write!(f, "NPU: {m}"),
54            Error::Input(m) => write!(f, "input: {m}"),
55        }
56    }
57}
58
59impl std::error::Error for Error {}
60
61impl From<taconite_bundle::Error> for Error {
62    fn from(e: taconite_bundle::Error) -> Self {
63        Error::Bundle(e.to_string())
64    }
65}
66
67impl From<taconite::Error> for Error {
68    fn from(e: taconite::Error) -> Self {
69        Error::Npu(e.to_string())
70    }
71}
72
73/// The model's constants (the manifest's params).
74#[derive(Debug, Clone)]
75pub struct Config {
76    pub d: usize,
77    pub i: usize,
78    pub layers: usize,
79    pub heads: usize,
80    pub hd: usize,
81    pub kv_heads: Vec<usize>,
82    pub regs: usize,
83    pub patch: usize,
84    /// crop height
85    pub h: usize,
86    /// crop width
87    pub w: usize,
88    pub eps: f32,
89    pub in_eps: f32,
90    /// the transposed convolutions' output channels
91    pub up: Vec<usize>,
92    /// the 1 x 1 convolutions' output channels
93    pub convs: Vec<usize>,
94    /// keypoints
95    pub k: usize,
96    pub box_pad: f32,
97    pub blur: usize,
98    pub mean: [f32; 3],
99    pub std: [f32; 3],
100    /// mirrored keypoint pairs (left, right)
101    pub flip_pairs: Vec<(usize, usize)>,
102}
103
104impl Config {
105    fn load(m: &Manifest) -> Result<Self, Error> {
106        let p = |k: &str| m.param_as::<usize>(k);
107        let f = |k: &str| m.param_as::<f32>(k);
108        let three = |k: &str| -> Result<[f32; 3], Error> {
109            let v: Vec<f32> = m.list(k)?;
110            v.try_into().map_err(|_| Error::Bundle(format!("{k}: 3 values")))
111        };
112        let flip_pairs = m
113            .param("flip_pairs")?
114            .split(',')
115            .map(|s| {
116                let (a, b) = s.split_once(':').ok_or_else(|| Error::Bundle(format!("flip pair {s}")))?;
117                Ok((a.parse().map_err(|_| Error::Bundle(s.into()))?, b.parse().map_err(|_| Error::Bundle(s.into()))?))
118            })
119            .collect::<Result<_, Error>>()?;
120        Ok(Config {
121            d: p("D")?,
122            i: p("I")?,
123            layers: p("layers")?,
124            heads: p("heads")?,
125            hd: p("hd")?,
126            kv_heads: m.list("kv_heads")?,
127            regs: p("regs")?,
128            patch: p("patch")?,
129            h: p("H")?,
130            w: p("W")?,
131            eps: f("eps")?,
132            in_eps: f("in_eps")?,
133            up: m.list("up")?,
134            convs: m.list("convs")?,
135            k: p("K")?,
136            box_pad: f("box_pad")?,
137            blur: p("blur")?,
138            mean: three("mean")?,
139            std: three("std")?,
140            flip_pairs,
141        })
142    }
143
144    pub fn gh(&self) -> usize {
145        self.h / self.patch
146    }
147
148    pub fn gw(&self) -> usize {
149        self.w / self.patch
150    }
151
152    /// CLS + register tokens
153    pub fn prefix(&self) -> usize {
154        1 + self.regs
155    }
156
157    pub fn qkv_width(&self, i: usize) -> usize {
158        (self.heads + 2 * self.kv_heads[i]) * self.hd
159    }
160
161    /// Heatmap height and width.
162    pub fn heatmap(&self) -> (usize, usize) {
163        let s = 1 << self.up.len();
164        (self.gh() * s, self.gw() * s)
165    }
166}
167
168pub struct Sapiens2 {
169    pub cfg: Config,
170    pub npu: Npu,
171    pub model: Model,
172    pub manifest: Manifest,
173    pub store: Store,
174    /// where the last call spent its time (`npu:<kernel>`, host stages)
175    pub timing: Timing,
176}
177
178impl Sapiens2 {
179    /// Loads a bundle: opens the NPU and uploads the packed weights
180    /// (kernels load on first use; [`Npu::preload`] loads them now).
181    pub fn load(dir: &Path) -> Result<Self, Error> {
182        let manifest = Manifest::load(dir, VERSION).or_else(|e| Manifest::load(dir, 1).map_err(|_| e))?;
183        let store = Store::load(dir)?;
184        let cfg = Config::load(&manifest)?;
185        let npu = Npu::open(&manifest)?;
186        let model = Model::load(&cfg, &store, &npu)?;
187        Ok(Sapiens2 { cfg, npu, model, manifest, store, timing: Timing::default() })
188    }
189
190    /// Hardware contexts the bundle's kernels use.
191    pub fn contexts(&self) -> usize {
192        self.npu.contexts
193    }
194
195    /// Frees every NPU hardware context the model holds, for another model
196    /// in the process (NPU2 has 16 across every process); the kernels load
197    /// again on the next call. The number of contexts freed.
198    pub fn release_contexts(&mut self) -> usize {
199        self.npu.release()
200    }
201
202    /// RGB8 `[h, w, 3]` and a box -> the model's input `[3, H, W]`.
203    pub fn preprocess(&self, rgb: &[u8], w: usize, h: usize, b: &BBox) -> Result<Vec<f32>, Error> {
204        if rgb.len() != w * h * 3 || w < 2 || h < 2 {
205            return Err(Error::Input(format!("{} bytes for a {w} x {h} RGB image", rgb.len())));
206        }
207        if !(b.w > 0.0 && b.h > 0.0) {
208            return Err(Error::Input(format!("an empty box {b:?}")));
209        }
210        let c = &self.cfg;
211        Ok(preprocess::crop(rgb, w, h, b, c.w, c.h, c.box_pad, c.mean, c.std))
212    }
213
214    /// pixel_values `[3, H, W]` -> the normalized patch features `[P, D]`.
215    pub fn features(&mut self, pixels: &[f32]) -> Result<Vec<f32>, Error> {
216        self.model.backbone(&self.cfg, &self.npu, pixels, &mut self.timing)
217    }
218
219    /// pixel_values `[3, H, W]` -> heatmaps `[K, h, w]`.
220    pub fn heatmaps(&mut self, pixels: &[f32]) -> Result<Vec<f32>, Error> {
221        let c = &self.cfg;
222        if pixels.len() != 3 * c.h * c.w {
223            return Err(Error::Input(format!("{} pixel values (want 3 x {} x {})", pixels.len(), c.h, c.w)));
224        }
225        self.timing.clear();
226        let f = self.model.backbone(&self.cfg, &self.npu, pixels, &mut self.timing)?;
227        self.model.head(&self.cfg, &self.npu, &f, &mut self.timing)
228    }
229
230    /// Heatmaps of a box's crop -> its keypoints in image coordinates.
231    pub fn keypoints(&mut self, heatmaps: &[f32], b: &BBox) -> Vec<Keypoint> {
232        let c = &self.cfg;
233        let (hh, hw) = c.heatmap();
234        let t0 = std::time::Instant::now();
235        let p = post::decode(heatmaps, c.k, hh, hw, c.blur);
236        let kp = post::to_image(&p, &b.window(c.w, c.h, c.box_pad), hh, hw);
237        self.timing.add("host.keypoints", t0.elapsed());
238        kp
239    }
240
241    /// RGB8 `[h, w, 3]` and a person box -> the 308 keypoints.
242    pub fn pose(&mut self, rgb: &[u8], w: usize, h: usize, b: &BBox) -> Result<Vec<Keypoint>, Error> {
243        let t0 = std::time::Instant::now();
244        let px = self.preprocess(rgb, w, h, b)?;
245        let dt = t0.elapsed();
246        let hm = self.heatmaps(&px)?;
247        self.timing.add("host.crop", dt);
248        Ok(self.keypoints(&hm, b))
249    }
250}
251
252/// Cosine similarity.
253pub fn cosine(a: &[f32], b: &[f32]) -> f32 {
254    let d = |x: &[f32], y: &[f32]| x.iter().zip(y).map(|(p, q)| (*p as f64) * (*q as f64)).sum::<f64>();
255    (d(a, b) / (d(a, a) * d(b, b)).sqrt()) as f32
256}