Skip to main content

taconite_sam3/
mask.rs

1// SPDX-FileCopyrightText: Copyright (C) 2026 Brishen Hawkins
2// SPDX-License-Identifier: Apache-2.0
3
4//! The mask decoder, as `iron/applications/sam3/mask_npu.py`:
5//!
6//! - the DETR encoder output attends to the prompt -- the same folded
7//!   cross-attention as the DETR encoder's, on the same two kernels;
8//! - the pixel decoder: nearest 2x upsample + skip, 3x3 conv (NPU), GroupNorm
9//!   + ReLU, twice (72 -> 144 -> 288);
10//! - the mask head folded into one GEMM over the 288 x 288 pixel embedding:
11//!   `pred_masks[q] = (W_inst^T m_q) . p + m_q . b_inst` for the 200 query
12//!   mask embeddings `m_q`, and the semantic head in column 200.
13
14use taconite::{bf16_to_f32, f32_to_bf16};
15
16use crate::cpu::{self, par_rows};
17use crate::pack::pack_b;
18use crate::text::Text;
19use crate::{Error, Sam3, gemm};
20
21impl Sam3 {
22    /// `(pred_masks [Q, S, S], semantic [S, S])` from the decoder's last
23    /// normalised query states, the FPN levels (channels-last) and the DETR
24    /// encoder output.
25    pub fn mask_decoder(
26        &mut self,
27        hidden: &[f32],
28        fpn: &[Vec<f32>; 3],
29        enc: &[f32],
30        text: &Text,
31    ) -> Result<(Vec<f32>, Vec<f32>), Error> {
32        let c = self.cfg.clone();
33        let (d, g) = (c.d_model, c.grid);
34        let q = hidden.len() / d;
35
36        // prompt cross-attention on the encoder output
37        let t0 = std::time::Instant::now();
38        self.fold_cross("m.ca", "m", text)?;
39        let h = cpu::layer_norm(enc, d, self.store.f32("m.ca_ln.w")?, self.store.f32("m.ca_ln.b")?, 1e-5);
40        let mut p = enc.to_vec();
41        self.cross_npu(&h, &mut p, "m", text)?;
42        self.timing.add("mask_cross", t0.elapsed());
43
44        // pixel decoder
45        let mut side = g;
46        for (i, skip) in [&fpn[1], &fpn[0]].into_iter().enumerate() {
47            let s2 = 2 * side;
48            let t0 = std::time::Instant::now();
49            let mut up = vec![0u16; s2 * s2 * d];
50            par_rows(&mut up, d, |r0, piece| {
51                for (ri, row) in piece.chunks_mut(d).enumerate() {
52                    let r = r0 + ri;
53                    let src = ((r / s2) / 2 * side + (r % s2) / 2) * d;
54                    for j in 0..d {
55                        row[j] = f32_to_bf16(p[src + j] + skip[r * d + j]);
56                    }
57                }
58            });
59            self.timing.add("mask_up", t0.elapsed());
60            let bias = self.store.f32(&format!("m.conv{i}.b"))?.to_vec();
61            let mut y = self.conv3x3(&up, s2, &format!("m.conv{i}"), &bias)?;
62            let t0 = std::time::Instant::now();
63            group_norm_relu(
64                &mut y,
65                d,
66                8,
67                self.store.f32(&format!("m.gn{i}.w"))?,
68                self.store.f32(&format!("m.gn{i}.b"))?,
69            );
70            self.timing.add("mask_gn", t0.elapsed());
71            p = y;
72            side = s2;
73        }
74        let t0 = std::time::Instant::now();
75
76        // folded mask + semantic head
77        let mut m = self.lin(hidden, d, "m.embed.0")?;
78        cpu::relu_(&mut m);
79        let mut m = self.lin(&m, d, "m.embed.1")?;
80        cpu::relu_(&mut m);
81        let m = self.lin(&m, d, "m.embed.2")?; // [Q, D]
82        let st = &self.store;
83        let (wi, bi) = (st.f32("m.inst.w")?, st.f32("m.inst.b")?); // [D out, D in]
84        let (ws, bs) = (st.f32("m.sem.w")?, st.f32("m.sem.b")?);
85        let spec = self.npu.spec("m_head")?.clone();
86        let n = spec.n;
87        if q + 1 > n {
88            return Err(Error::Input(format!("{q} queries + semantic > {n}")));
89        }
90        let mut bm = vec![0f32; d * n];
91        let mut bb = vec![0f32; n];
92        for qi in 0..q {
93            let mq = &m[qi * d..(qi + 1) * d];
94            for i in 0..d {
95                bm[i * n + qi] = (0..d).map(|o| wi[o * d + i] * mq[o]).sum();
96            }
97            bb[qi] = cpu::dot(mq, bi);
98        }
99        for i in 0..d {
100            bm[i * n + q] = ws[i];
101        }
102        bb[q] = bs[0];
103        let packed = pack_b(&spec, &bm, Some(&bb));
104        self.slots.get_mut("m.head").unwrap().write(&packed)?;
105        let px = side * side;
106        let mut pb = vec![0u16; p.len()];
107        crate::narrow(&p, &mut pb);
108        self.io.m_head.set_a(&pb)?;
109        self.timing.add("mask_head_prep", t0.elapsed());
110        gemm(&mut self.npu, &self.io.m_head, &self.slots["m.head"], &mut self.timing)?;
111        let t0 = std::time::Instant::now();
112        let out = self.io.m_head.get_c(px)?;
113        // [px, n] -> [q, px], in blocks of pixels: contiguous reads, and
114        // each plane written a run at a time (read a plane at a time it
115        // is a strided gather over the whole output)
116        const PB: usize = 64;
117        let blocks = px.div_ceil(PB);
118        let mut planes = vec![0f32; (q + 1) * blocks * PB]; // [q + 1, blocks * PB]
119        let stride = blocks * PB;
120        {
121            let ptr = planes.as_mut_ptr() as usize;
122            let mut tasks = vec![0u8; blocks];
123            par_rows(&mut tasks, 1, |b0, piece| {
124                for bi in 0..piece.len() {
125                    let p0 = (b0 + bi) * PB;
126                    let np = (px - p0).min(PB);
127                    for qi in 0..=q {
128                        // disjoint runs of one plane per task
129                        let dst = unsafe { std::slice::from_raw_parts_mut((ptr as *mut f32).add(qi * stride + p0), np) };
130                        for (pi, v) in dst.iter_mut().enumerate() {
131                            *v = bf16_to_f32(out[(p0 + pi) * n + qi]);
132                        }
133                    }
134                }
135            });
136        }
137        let mut masks = vec![0f32; q * px];
138        par_rows(&mut masks, px, |q0, piece| {
139            for (qi, plane) in piece.chunks_mut(px).enumerate() {
140                plane.copy_from_slice(&planes[(q0 + qi) * stride..][..px]);
141            }
142        });
143        let semantic = planes[q * stride..][..px].to_vec();
144        self.timing.add("mask_out", t0.elapsed());
145        Ok((masks, semantic))
146    }
147}
148
149/// `relu(GroupNorm(x))` in place, channels-last `x [H*W, C]`.
150fn group_norm_relu(x: &mut [f32], c: usize, groups: usize, w: &[f32], b: &[f32]) {
151    let cg = c / groups;
152    let px = x.len() / c;
153    // per-group sums, partial per task then combined (f64: 21M values)
154    let tasks = px.min(4 * cpu::threads()).max(1);
155    let per = px.div_ceil(tasks);
156    let mut partial = vec![[0f64; 2 * 8]; tasks];
157    assert!(groups <= 8, "GroupNorm with more than 8 groups");
158    par_rows(&mut partial, 1, |t0, piece| {
159        for (ti, st) in piece.iter_mut().enumerate() {
160            let r0 = (t0 + ti) * per;
161            let r1 = (r0 + per).min(px);
162            for row in x[r0 * c..r1 * c].chunks(c) {
163                for gi in 0..groups {
164                    let (mut s, mut ss) = (0f32, 0f32);
165                    for &v in &row[gi * cg..(gi + 1) * cg] {
166                        s += v;
167                        ss += v * v;
168                    }
169                    st[2 * gi] += s as f64;
170                    st[2 * gi + 1] += ss as f64;
171                }
172            }
173        }
174    });
175    let mut stats = vec![(0f64, 0f64); groups];
176    for st in &partial {
177        for (gi, s) in stats.iter_mut().enumerate() {
178            s.0 += st[2 * gi];
179            s.1 += st[2 * gi + 1];
180        }
181    }
182    let nn = (px * cg) as f64;
183    let norm: Vec<(f32, f32)> = stats
184        .iter()
185        .map(|&(s, ss)| {
186            let mean = s / nn;
187            let var = (ss / nn - mean * mean).max(0.0);
188            (mean as f32, (1.0 / (var + 1e-5).sqrt()) as f32)
189        })
190        .collect();
191    par_rows(x, c, |_, piece| {
192        for row in piece.chunks_mut(c) {
193            for (j, v) in row.iter_mut().enumerate() {
194                let (mean, inv) = norm[j / cg];
195                *v = ((*v - mean) * inv * w[j] + b[j]).max(0.0);
196            }
197        }
198    });
199}