1use 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
38pub 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#[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 pub h: usize,
86 pub w: usize,
88 pub eps: f32,
89 pub in_eps: f32,
90 pub up: Vec<usize>,
92 pub convs: Vec<usize>,
94 pub k: usize,
96 pub box_pad: f32,
97 pub blur: usize,
98 pub mean: [f32; 3],
99 pub std: [f32; 3],
100 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 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 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 pub timing: Timing,
176}
177
178impl Sapiens2 {
179 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 pub fn contexts(&self) -> usize {
192 self.npu.contexts
193 }
194
195 pub fn release_contexts(&mut self) -> usize {
199 self.npu.release()
200 }
201
202 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 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 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 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 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
252pub 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}