Skip to main content

taconite_sam3/
detr.rs

1// SPDX-FileCopyrightText: Copyright (C) 2026 Brishen Hawkins
2// SPDX-License-Identifier: Apache-2.0
3
4//! The DETR encoder (NPU, as `iron/applications/sam3/detr_npu.py`), the
5//! DETR decoder (host, as HF `Sam3DetrDecoder`, with all six layers' vision
6//! keys and values from one NPU GEMM) and the dot-product scoring.
7//!
8//! Encoder, per layer: `[q|k|v] = [LN1(x) + pos | LN1(x)] B` (heads of 32
9//! zero-padded to the MHA's 64, q pre-scaled by sqrt 2), MHA, o; the prompt
10//! cross-attention folded into two GEMMs (`h (W_q,h K_h^T)` for the scores,
11//! `p (V_h W_o,h)` for the output; B matrices packed here per prompt) with
12//! the softmax here; the ReLU MLP (fc2 reading fc1's output in place).
13
14use 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; // nn.LayerNorm's default, every LayerNorm here
24
25/// The decoder's outputs (and per-layer values for checking).
26pub struct Decoded {
27    /// `[Q, 256]` the last layer's normalised query states.
28    pub hidden: Vec<f32>,
29    /// `[Q, 4]` the final boxes, normalised xyxy.
30    pub boxes: Vec<f32>,
31    /// `[Q]` classification logits.
32    pub logits: Vec<f32>,
33    pub presence: f32,
34    /// `[layers][Q, 4]` cxcywh boxes after each layer.
35    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    /// `relu(l1) -> [relu(l2) ->] l_last`, HF `Sam3DecoderMLP`.
61    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    /// The cross-attention `prefix` (`d.<i>.ca` or `m.ca`) folded over this
74    /// prompt and packed into slots `<slot>.s` / `<slot>.c`: the scores
75    /// `h Bs + bs` (`N` = heads x prompt tokens) and the output `p Bc + b_o`.
76    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"))?; // [L, D]
82        let v = self.lin(&text.feats, d, &format!("{prefix}.v"))?;
83        let wq = st.f32(&format!("{prefix}.q.w"))?; // [D, D] out x in
84        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; // input feature
94                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    /// The folded prompt cross-attention on `x [T, D]` (already normalised
127    /// as `h`): `x_res += softmax(h Bs + bs) Bc + b_o`, per head over the
128    /// valid prompt tokens. Uses the d_s / d_c kernels and `slot`'s weights.
129    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    /// The DETR encoder: `fpn2 [T, 256]` (the 72 x 72 level, raster) and
162    /// the prompt -> `[T, 256]`.
163    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    /// Box sine embedding (HF `Sam3SinePositionEmbedding.encode_boxes`):
222    /// `[Q, 4]` cxcywh -> `[Q, 4 * 128]`, ordered (y, x, w, h).
223    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    /// Box relative position bias in its separable form, `(by, bx)`, each
240    /// `[H, 1 + Q, g]` (row 0, the presence token's, zero): the bias of
241    /// query `q` at key `(iy, ix)` is `by[h, q, iy] + bx[h, q, ix]` -- for
242    /// each query box and grid row/column, the log-scaled distances to the
243    /// box's edges through a small MLP per axis. The `[H, 1 + Q, T]` sum is
244    /// never built; `cpu::attention_dec` adds the two terms on the way.
245    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        // per axis: inputs [Q * g, 2] -> [Q * g, H]
254        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    /// The DETR decoder + scoring: `enc [T, 256]` and the prompt ->
283    /// [`Decoded`].
284    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        // every layer's vision keys (from enc + pos) and values (from enc)
289        {
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)?; // [T, layers x (K | V)] bf16
305        let hd = d / nh;
306        // this layer's keys transposed, [H, hd, T], and values, [H, T, hd]
307        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            // self-attention (q, k with the query positions)
337            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            // text cross-attention
350            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            // vision cross-attention, with the box bias
361            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            // the transpose in blocks of 8 key dimensions: one strided read
373            // of 8 bf16 per token, 8 sequential write streams
374            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            // MLP (post-norm)
398            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            // box refinement on the queries, presence from token 0
407            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    /// HF `Sam3DotProductScoring`: queries against the mean-pooled,
426    /// MLP-refined prompt.
427    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}