1use taconite::bf16_to_f32;
15
16use crate::cpu::{self, Attn, Rows, W, ln_row, par_rows, sigmoid};
17use crate::npu::{pull, push};
18use crate::pack::pack_b;
19use crate::text::Text;
20use crate::vit::{add_bf16, layer_norm_bf16};
21use crate::{Error, Sam3, gemm, mha};
22
23const EPS: f32 = 1e-5; pub struct Decoded {
27 pub hidden: Vec<f32>,
29 pub boxes: Vec<f32>,
31 pub logits: Vec<f32>,
33 pub presence: f32,
34 pub layer_boxes: Vec<Vec<f32>>,
36 pub layer_presence: Vec<f32>,
37}
38
39fn inverse_sigmoid(x: f32) -> f32 {
40 let x = x.clamp(0.0, 1.0);
41 (x.max(1e-3) / (1.0 - x).max(1e-3)).ln()
42}
43
44fn xyxy(b: &[f32]) -> [f32; 4] {
45 [b[0] - 0.5 * b[2], b[1] - 0.5 * b[3], b[0] + 0.5 * b[2], b[1] + 0.5 * b[3]]
46}
47
48impl Sam3 {
49 pub(crate) fn lin(&self, x: &[f32], n_in: usize, name: &str) -> Result<Vec<f32>, Error> {
50 let st = &self.store;
51 let b = if st.has(&format!("{name}.b")) { Some(st.f32(&format!("{name}.b"))?) } else { None };
52 Ok(cpu::linear(x, n_in, W::F32(st.f32(&format!("{name}.w"))?), b))
53 }
54
55 fn ln(&self, x: &[f32], dim: usize, name: &str) -> Result<Vec<f32>, Error> {
56 let st = &self.store;
57 Ok(cpu::layer_norm(x, dim, st.f32(&format!("{name}.w"))?, st.f32(&format!("{name}.b"))?, EPS))
58 }
59
60 fn mlp(&self, x: &[f32], n_in: usize, name: &str, layers: usize) -> Result<Vec<f32>, Error> {
62 let mut h = self.lin(x, n_in, &format!("{name}.1"))?;
63 cpu::relu_(&mut h);
64 let d = h.len() / (x.len() / n_in);
65 let mut h = self.lin(&h, d, &format!("{name}.2"))?;
66 if layers == 3 {
67 cpu::relu_(&mut h);
68 h = self.lin(&h, d, &format!("{name}.3"))?;
69 }
70 Ok(h)
71 }
72
73 pub(crate) fn fold_cross(&mut self, prefix: &str, slot: &str, text: &Text) -> Result<(), Error> {
77 let c = &self.cfg;
78 let (d, nh, l) = (c.d_model, c.d_heads, c.text_len);
79 let hd = d / nh;
80 let st = &self.store;
81 let k = self.lin(&text.feats, d, &format!("{prefix}.k"))?; let v = self.lin(&text.feats, d, &format!("{prefix}.v"))?;
83 let wq = st.f32(&format!("{prefix}.q.w"))?; let bq = st.f32(&format!("{prefix}.q.b"))?;
85 let wo = st.f32(&format!("{prefix}.o.w"))?;
86 let bo = st.f32(&format!("{prefix}.o.b"))?;
87 let inv = 1.0 / (hd as f32).sqrt();
88 let mut bs_m = vec![0f32; d * nh * l];
89 let mut bs_b = vec![0f32; nh * l];
90 let mut bc_m = vec![0f32; nh * l * d];
91 par_rows(&mut bs_m, nh * l, |i0, piece| {
92 for (ii, row) in piece.chunks_mut(nh * l).enumerate() {
93 let i = i0 + ii; for h in 0..nh {
95 for j in 0..l {
96 let mut s = 0.0;
97 for e in 0..hd {
98 s += wq[(h * hd + e) * d + i] * k[j * d + h * hd + e];
99 }
100 row[h * l + j] = s * inv;
101 }
102 }
103 }
104 });
105 for h in 0..nh {
106 for j in 0..l {
107 bs_b[h * l + j] = (0..hd).map(|e| bq[h * hd + e] * k[j * d + h * hd + e]).sum::<f32>() * inv;
108 }
109 }
110 par_rows(&mut bc_m, d, |r0, piece| {
111 for (ri, row) in piece.chunks_mut(d).enumerate() {
112 let r = r0 + ri;
113 let (h, j) = (r / l, r % l);
114 for (o, out) in row.iter_mut().enumerate() {
115 *out = (0..hd).map(|e| v[j * d + h * hd + e] * wo[o * d + h * hd + e]).sum();
116 }
117 }
118 });
119 let ps = pack_b(self.npu.spec("d_s")?, &bs_m, Some(&bs_b));
120 let pc = pack_b(self.npu.spec("d_c")?, &bc_m, Some(bo));
121 self.slots.get_mut(&format!("{slot}.s")).unwrap().write(&ps)?;
122 self.slots.get_mut(&format!("{slot}.c")).unwrap().write(&pc)?;
123 Ok(())
124 }
125
126 pub(crate) fn cross_npu(&mut self, h: &[f32], x_res: &mut [f32], slot: &str, text: &Text) -> Result<(), Error> {
130 let c = &self.cfg;
131 let (d, nh, l) = (c.d_model, c.d_heads, c.text_len);
132 let t = h.len() / d;
133 let mut hb = vec![0u16; h.len()];
134 crate::narrow(h, &mut hb);
135 self.io.d_s.set_a(&hb)?;
136 gemm(&mut self.npu, &self.io.d_s, &self.slots[&format!("{slot}.s")], &mut self.timing)?;
137 let s = self.io.d_s.get_c(t)?;
138 let valid = &text.valid;
139 let mut pb = vec![0u16; t * nh * l];
140 par_rows(&mut pb, nh * l, |r0, piece| {
141 let mut row_f = vec![0f32; l];
142 for (ri, row) in piece.chunks_mut(nh * l).enumerate() {
143 let src = &s[(r0 + ri) * nh * l..][..nh * l];
144 for hh in 0..nh {
145 for j in 0..l {
146 row_f[j] = if valid[j] { bf16_to_f32(src[hh * l + j]) } else { f32::NEG_INFINITY };
147 }
148 cpu::softmax_(&mut row_f);
149 for j in 0..l {
150 row[hh * l + j] = taconite::f32_to_bf16(row_f[j]);
151 }
152 }
153 }
154 });
155 self.io.d_c.set_a(&pb)?;
156 gemm(&mut self.npu, &self.io.d_c, &self.slots[&format!("{slot}.c")], &mut self.timing)?;
157 add_bf16(x_res, &self.io.d_c.get_c(t)?, d, None);
158 Ok(())
159 }
160
161 pub fn detr_encoder(&mut self, fpn2: &[f32], text: &Text) -> Result<Vec<f32>, Error> {
164 let c = self.cfg.clone();
165 let (d, nh, t) = (c.d_model, c.d_heads, c.tokens());
166 let pd = self.npu.mhas["mha_d"].d * nh;
167 let mut x = fpn2.to_vec();
168 for i in 0..c.d_layers {
169 let p = |n: &str| format!("d.{i}.{n}");
170 self.fold_cross(&p("ca"), &format!("d.{i}"), text)?;
171 let st = &self.store;
172 let pos = st.f32("d.pos")?;
173 let (w1, b1) = (st.f32(&p("ln1.w"))?, st.f32(&p("ln1.b"))?);
174 let mut ab = vec![0u16; t * 2 * d];
175 par_rows(&mut ab, 2 * d, |r0, piece| {
176 let mut h = vec![0f32; d];
177 for (ri, row) in piece.chunks_mut(2 * d).enumerate() {
178 let r = r0 + ri;
179 ln_row(&x[r * d..(r + 1) * d], &mut h, w1, b1, EPS);
180 for j in 0..d {
181 row[j] = taconite::f32_to_bf16(h[j] + pos[r * d + j]);
182 row[d + j] = taconite::f32_to_bf16(h[j]);
183 }
184 }
185 });
186 self.io.d_qkv.set_a(&ab)?;
187 gemm(&mut self.npu, &self.io.d_qkv, &self.w[&p("qkv")], &mut self.timing)?;
188 let qkv = self.io.d_qkv.get_c(t)?;
189 let m = &mut self.io.mha_d;
190 let mut part = vec![0u16; t * pd];
191 for (j, buf) in [&mut m.q, &mut m.k, &mut m.v].into_iter().enumerate() {
192 par_rows(&mut part, pd, |r0, piece| {
193 for (ri, row) in piece.chunks_mut(pd).enumerate() {
194 row.copy_from_slice(&qkv[(r0 + ri) * 3 * pd + j * pd..][..pd]);
195 }
196 });
197 push(&part, buf)?;
198 }
199 mha(&mut self.npu, &self.io.mha_d, &mut self.timing)?;
200 let o = pull(&self.io.mha_d.o, t * pd)?;
201 self.io.d_o.set_a(&o)?;
202 gemm(&mut self.npu, &self.io.d_o, &self.w[&p("o")], &mut self.timing)?;
203 add_bf16(&mut x, &self.io.d_o.get_c(t)?, d, None);
204
205 let h = self.ln(&x, d, &p("ln2"))?;
206 self.cross_npu(&h, &mut x, &format!("d.{i}"), text)?;
207
208 let mut hb = vec![0u16; t * d];
209 {
210 let st = &self.store;
211 layer_norm_bf16(&x, d, st.f32(&p("ln3.w"))?, st.f32(&p("ln3.b"))?, EPS, &mut hb);
212 }
213 self.io.d_fc1.set_a(&hb)?;
214 gemm(&mut self.npu, &self.io.d_fc1, &self.w[&p("fc1")], &mut self.timing)?;
215 gemm(&mut self.npu, &self.io.d_fc2, &self.w[&p("fc2")], &mut self.timing)?;
216 add_bf16(&mut x, &self.io.d_fc2.get_c(t)?, d, Some(self.store.f32(&p("fc2.b"))?));
217 }
218 Ok(x)
219 }
220
221 fn encode_boxes(&self, boxes: &[f32]) -> Vec<f32> {
224 let f = self.cfg.d_model / 2;
225 let dim_t: Vec<f32> = (0..f).map(|i| 10000f32.powf(2.0 * (i / 2) as f32 / f as f32)).collect();
226 let scale = 2.0 * std::f32::consts::PI;
227 let mut out = Vec::with_capacity(boxes.len() / 4 * 4 * f);
228 for b in boxes.chunks(4) {
229 for coord in [b[1], b[0], b[2], b[3]] {
230 for (i, &d) in dim_t.iter().enumerate().take(f) {
231 let v = coord * scale / d;
232 out.push(if i % 2 == 0 { v.sin() } else { v.cos() });
233 }
234 }
235 }
236 out
237 }
238
239 fn rpb(&self, boxes: &[f32]) -> Result<(Vec<f32>, Vec<f32>), Error> {
246 let c = &self.cfg;
247 let (g, nh) = (c.grid, c.d_heads);
248 let q = boxes.len() / 4;
249 let enc = |v: f32| {
250 let v = v * 8.0;
251 v.signum() * (v.abs() + 1.0).log2() / 3.0
252 };
253 let axis = |lo: usize, name: &str| -> Result<Vec<f32>, Error> {
255 let mut inp = Vec::with_capacity(q * g * 2);
256 for b in boxes.chunks(4) {
257 let e = xyxy(b);
258 for i in 0..g {
259 let p = i as f32 / g as f32;
260 inp.push(enc(p - e[lo]));
261 inp.push(enc(p - e[lo + 2]));
262 }
263 }
264 self.mlp(&inp, 2, name, 2)
265 };
266 let ry = axis(1, "dec.rpb_y")?;
267 let rx = axis(0, "dec.rpb_x")?;
268 let lq = q + 1;
269 let (mut by, mut bx) = (vec![0f32; nh * lq * g], vec![0f32; nh * lq * g]);
270 for h in 0..nh {
271 for qq in 0..q {
272 let dst = (h * lq + qq + 1) * g;
273 for i in 0..g {
274 by[dst + i] = ry[(qq * g + i) * nh + h];
275 bx[dst + i] = rx[(qq * g + i) * nh + h];
276 }
277 }
278 }
279 Ok((by, bx))
280 }
281
282 pub fn detr_decoder(&mut self, enc: &[f32], text: &Text) -> Result<Decoded, Error> {
285 let c = self.cfg.clone();
286 let (d, nh, t) = (c.d_model, c.d_heads, c.tokens());
287 let kvw = self.npu.spec("dec_kv")?.n;
288 {
290 let pos = self.store.f32("d.pos")?;
291 let mut ab = vec![0u16; t * 2 * d];
292 par_rows(&mut ab, 2 * d, |r0, piece| {
293 for (ri, row) in piece.chunks_mut(2 * d).enumerate() {
294 let r = r0 + ri;
295 for j in 0..d {
296 row[j] = taconite::f32_to_bf16(enc[r * d + j] + pos[r * d + j]);
297 row[d + j] = taconite::f32_to_bf16(enc[r * d + j]);
298 }
299 }
300 });
301 self.io.dec_kv.set_a(&ab)?;
302 }
303 gemm(&mut self.npu, &self.io.dec_kv, &self.w["dec.kv"], &mut self.timing)?;
304 let kv = self.io.dec_kv.get_c(t)?; let hd = d / nh;
306 let (mut kt, mut vh) = (vec![0f32; t * d], vec![0f32; t * d]);
308
309 let st = &self.store;
310 let mut refb: Vec<f32> = st.f32("dec.reference_points")?.iter().map(|&v| sigmoid(v)).collect();
311 let mut hs = st.f32("dec.presence_token")?.to_vec();
312 hs.extend_from_slice(st.f32("dec.query_embed")?);
313 let mut out = Decoded {
314 hidden: vec![],
315 boxes: vec![],
316 logits: vec![],
317 presence: 0.0,
318 layer_boxes: vec![],
319 layer_presence: vec![],
320 };
321 let mut normed = vec![];
322 for l in 0..c.dec_layers {
323 let p = |s: &str| format!("dec.{l}.{s}");
324 let t0 = std::time::Instant::now();
325 let sine = self.encode_boxes(&refb);
326 let qpos_q = self.mlp(&sine, 2 * d, "dec.ref_point_head", 2)?;
327 let mut qpos = vec![0f32; d];
328 qpos.extend_from_slice(&qpos_q);
329 let with_pos = |hs: &[f32]| -> Vec<f32> { hs.iter().zip(&qpos).map(|(a, b)| a + b).collect() };
330 self.timing.add("dec_qpos", t0.elapsed());
331 let t0 = std::time::Instant::now();
332 let (by, bx) = self.rpb(&refb)?;
333 self.timing.add("dec_rpb", t0.elapsed());
334 let t0 = std::time::Instant::now();
335
336 let qk = with_pos(&hs);
338 let q = self.lin(&qk, d, &p("sa.q"))?;
339 let k = self.lin(&qk, d, &p("sa.k"))?;
340 let v = self.lin(&hs, d, &p("sa.v"))?;
341 let a =
342 cpu::attention(&q, d, Rows::f32(&k, d), Rows::f32(&v, d), &Attn { heads: nh, ..Default::default() });
343 let mut o = self.lin(&a, d, &p("sa.o"))?;
344 cpu::add_(&mut o, &hs);
345 hs = self.ln(&o, d, &p("sa_ln"))?;
346 self.timing.add("dec_sa", t0.elapsed());
347 let t0 = std::time::Instant::now();
348
349 let q = self.lin(&with_pos(&hs), d, &p("tca.q"))?;
351 let k = self.lin(&text.feats, d, &p("tca.k"))?;
352 let v = self.lin(&text.feats, d, &p("tca.v"))?;
353 let at = Attn { heads: nh, valid: Some(&text.valid), ..Default::default() };
354 let a = cpu::attention(&q, d, Rows::f32(&k, d), Rows::f32(&v, d), &at);
355 let mut o = self.lin(&a, d, &p("tca.o"))?;
356 cpu::add_(&mut o, &hs);
357 hs = self.ln(&o, d, &p("tca_ln"))?;
358 self.timing.add("dec_tca", t0.elapsed());
359
360 let q = self.lin(&with_pos(&hs), d, &p("vca.q"))?;
362 let t0 = std::time::Instant::now();
363 let (koff, voff) = (l * 2 * d, l * 2 * d + d);
364 par_rows(&mut vh, hd, |r0, piece| {
365 for (ri, row) in piece.chunks_mut(hd).enumerate() {
366 let (h, j) = ((r0 + ri) / t, (r0 + ri) % t);
367 for (o, &x) in row.iter_mut().zip(&kv[j * kvw + voff + h * hd..][..hd]) {
368 *o = bf16_to_f32(x);
369 }
370 }
371 });
372 const KB: usize = 8;
375 par_rows(&mut kt, KB * t, |b0, piece| {
376 for (bi, blk) in piece.chunks_mut(KB * t).enumerate() {
377 let (h, d0) = ((b0 + bi) / (hd / KB), ((b0 + bi) % (hd / KB)) * KB);
378 for j in 0..t {
379 let src = &kv[j * kvw + koff + h * hd + d0..][..KB];
380 for (dd, &x) in src.iter().enumerate() {
381 blk[dd * t + j] = bf16_to_f32(x);
382 }
383 }
384 }
385 });
386 self.timing.add("dec_kvconv", t0.elapsed());
387 let t0 = std::time::Instant::now();
388 let a = cpu::attention_dec(&q, d, nh, &kt, &vh, &by, &bx, c.grid);
389 self.timing.add("dec_vattn", t0.elapsed());
390 let t0 = std::time::Instant::now();
391 let mut o = self.lin(&a, d, &p("vca.o"))?;
392 cpu::add_(&mut o, &hs);
393 hs = self.ln(&o, d, &p("vca_ln"))?;
394 self.timing.add("dec_vo", t0.elapsed());
395 let t0 = std::time::Instant::now();
396
397 let mut f = self.lin(&hs, d, &p("fc1"))?;
399 cpu::relu_(&mut f);
400 let mut f = self.lin(&f, c.d_ffn, &p("fc2"))?;
401 cpu::add_(&mut f, &hs);
402 hs = self.ln(&f, d, &p("mlp_ln"))?;
403 self.timing.add("dec_mlp", t0.elapsed());
404 let t0 = std::time::Instant::now();
405
406 normed = self.ln(&hs[d..], d, "dec.out_ln")?;
408 let delta = self.mlp(&normed, d, "dec.box_head", 3)?;
409 for (b, dl) in refb.iter_mut().zip(&delta) {
410 *b = sigmoid(dl + inverse_sigmoid(*b));
411 }
412 let pres = self.ln(&hs[..d], d, "dec.presence_ln")?;
413 let pl = self.mlp(&pres, d, "dec.presence_head", 3)?[0].clamp(-10.0, 10.0);
414 out.layer_boxes.push(refb.clone());
415 out.layer_presence.push(pl);
416 self.timing.add("dec_heads", t0.elapsed());
417 }
418 out.presence = *out.layer_presence.last().unwrap();
419 out.boxes = refb.chunks(4).flat_map(xyxy).collect();
420 out.logits = self.score(&normed, text)?;
421 out.hidden = normed;
422 Ok(out)
423 }
424
425 fn score(&self, hidden: &[f32], text: &Text) -> Result<Vec<f32>, Error> {
428 let d = self.cfg.d_model;
429 let l = self.cfg.text_len;
430 let mut t = self.mlp(&text.feats, d, "score.text_mlp", 2)?;
431 cpu::add_(&mut t, &text.feats);
432 let t = self.ln(&t, d, "score.text_ln")?;
433 let nv = text.valid.iter().filter(|&&v| v).count().max(1) as f32;
434 let mut pooled = vec![0f32; d];
435 for j in 0..l {
436 if text.valid[j] {
437 cpu::add_(&mut pooled, &t[j * d..(j + 1) * d]);
438 }
439 }
440 pooled.iter_mut().for_each(|v| *v /= nv);
441 let pt = self.lin(&pooled, d, "score.text_proj")?;
442 let pq = self.lin(hidden, d, "score.query_proj")?;
443 let scale = 1.0 / (d as f32).sqrt();
444 Ok(pq.chunks(d).map(|q| (cpu::dot(q, &pt) * scale).clamp(-12.0, 12.0)).collect())
445 }
446}