Skip to main content

taconite_sam3/
vit.rs

1// SPDX-FileCopyrightText: Copyright (C) 2026 Brishen Hawkins
2// SPDX-License-Identifier: Apache-2.0
3
4//! The ViT backbone (32 layers over 72 x 72 tokens), as
5//! `iron/applications/sam3/sam3_npu.py`'s `NpuViT`: tokens in window order
6//! for the whole backbone, q/k head dims in half-split order (both baked
7//! into the exported weights and tables), every Linear and the attention on
8//! the NPU, LayerNorms / RoPE / residuals here -- or, for bundles with
9//! `vit_device`, those on the NPU too (`vit_device`). The patch embedding (a
10//! 14 x 14 stride-14 conv) is a GEMM over the patches too (K = 588 padded
11//! to 768).
12//!
13//! The host works in cached memory throughout and moves whole buffers in
14//! and out of the device BOs (`push` / `pull`): their host mappings are
15//! uncached.
16
17use taconite::{bf16_to_f32, f32_to_bf16};
18
19use std::collections::HashMap;
20
21use crate::bundle::Store;
22use crate::cpu::{ln_row, par_rows};
23use crate::npu::{Buffer, Npu, pull, push};
24use crate::{Config, Error, Ios, Sam3, Timing, gemm, gemm_dev, mha, op};
25
26impl Sam3 {
27    /// `pixels [3, S, S]` -> the backbone's last hidden state `[T, 1024]`,
28    /// raster order.
29    pub fn vit(&mut self, pixels: &[f32]) -> Result<Vec<f32>, Error> {
30        if self.cfg.vit_device {
31            return vit_device(&mut self.npu, &mut self.io, &self.w, &self.store, &self.cfg, &mut self.timing, pixels);
32        }
33        let c = self.cfg.clone();
34        let (t, dim, g, ps, s) = (c.tokens(), c.vit_dim, c.grid, c.patch, c.image_size);
35        if pixels.len() != 3 * s * s {
36            return Err(Error::Input(format!("pixels must be [3, {s}, {s}]")));
37        }
38        let perm = self.store.i32("v.perm")?.to_vec();
39        let eps = c.vit_eps;
40
41        // patches, window order: row j = patch perm[j], columns (ch, ky, kx)
42        let ke = self.npu.spec("v_embed")?.k;
43        let mut patches = vec![0u16; t * ke];
44        par_rows(&mut patches, ke, |r0, piece| {
45            for (ri, row) in piece.chunks_mut(ke).enumerate() {
46                let p = perm[r0 + ri] as usize;
47                let (py, px) = (p / g, p % g);
48                for ch in 0..3 {
49                    for ky in 0..ps {
50                        let src = ch * s * s + (py * ps + ky) * s + px * ps;
51                        for kx in 0..ps {
52                            row[(ch * ps + ky) * ps + kx] = f32_to_bf16(pixels[src + kx]);
53                        }
54                    }
55                }
56            }
57        });
58        self.io.v_embed.set_a(&patches)?;
59        gemm(&mut self.npu, &self.io.v_embed, &self.w["v.embed"], &mut self.timing)?;
60        let emb = self.io.v_embed.get_c(t)?;
61        let st = &self.store;
62        let pos = st.f32("v.pos")?;
63        let (lw, lb) = (st.f32("v.ln_pre.w")?, st.f32("v.ln_pre.b")?);
64        let mut x = vec![0f32; t * dim];
65        par_rows(&mut x, dim, |r0, piece| {
66            let mut tmp = vec![0f32; dim];
67            for (ri, row) in piece.chunks_mut(dim).enumerate() {
68                let r = r0 + ri;
69                for j in 0..dim {
70                    tmp[j] = bf16_to_f32(emb[r * dim + j]) + pos[r * dim + j];
71                }
72                ln_row(&tmp, row, lw, lb, eps);
73            }
74        });
75
76        let (heads, hd) = (c.vit_heads, dim / c.vit_heads);
77        let half = hd / 2;
78        let ws2 = c.window * c.window;
79        let rope_win = (st.f32("v.rope.win.cos")?.to_vec(), st.f32("v.rope.win.sin")?.to_vec());
80        let rope_glob = (st.f32("v.rope.glob.cos")?.to_vec(), st.f32("v.rope.glob.sin")?.to_vec());
81        let mut a = vec![0u16; t * dim];
82        let mut mq = vec![0u16; t * dim];
83        for i in 0..c.vit_layers {
84            let p = |n: &str| format!("v.{i}.{n}");
85            // LN1 -> qkv
86            {
87                let st = &self.store;
88                layer_norm_bf16(&x, dim, st.f32(&p("ln1.w"))?, st.f32(&p("ln1.b"))?, eps, &mut a);
89            }
90            self.io.v_qkv.set_a(&a)?;
91            gemm(&mut self.npu, &self.io.v_qkv, &self.w[&p("qkv")], &mut self.timing)?;
92            let t0 = std::time::Instant::now();
93            let qkv = self.io.v_qkv.get_c(t)?;
94
95            // RoPE into the MHA's layout: [n_win * heads, 576, 64] for the
96            // windowed layers, [T, heads * 64] for the global ones
97            let global = c.vit_global.contains(&i);
98            let (cos, sin) = if global { (&rope_glob.0, &rope_glob.1) } else { (&rope_win.0, &rope_win.1) };
99            // MHA row r (64 wide) <-> (token, head)
100            let at = |r: usize| -> (usize, usize) {
101                if global {
102                    (r / heads, r % heads)
103                } else {
104                    let (wh, si) = (r / ws2, r % ws2);
105                    ((wh / heads) * ws2 + si, wh % heads)
106                }
107            };
108            let m = if global { &mut self.io.mha_glob } else { &mut self.io.mha_win };
109            for (part, buf) in [(0, &mut m.q), (1, &mut m.k), (2, &mut m.v)] {
110                par_rows(&mut mq, hd, |r0, piece| {
111                    for (ri, row) in piece.chunks_mut(hd).enumerate() {
112                        let (tok, h) = at(r0 + ri);
113                        let src = &qkv[tok * 3 * dim + part * dim + h * hd..][..hd];
114                        if part == 2 {
115                            row.copy_from_slice(src);
116                            continue;
117                        }
118                        for d in 0..half {
119                            let (x1, x2) = (bf16_to_f32(src[d]), bf16_to_f32(src[half + d]));
120                            let (co, si) = (cos[tok * half + d], sin[tok * half + d]);
121                            row[d] = f32_to_bf16(x1 * co - x2 * si);
122                            row[half + d] = f32_to_bf16(x2 * co + x1 * si);
123                        }
124                    }
125                });
126                push(&mq, buf)?;
127            }
128            let t1 = std::time::Instant::now();
129            mha(&mut self.npu, m, &mut self.timing)?;
130            let t2 = std::time::Instant::now();
131            // attention output -> o's A, token-major
132            let o = pull(&m.o, t * dim)?;
133            par_rows(&mut a, dim, |t0, piece| {
134                for (ti, row) in piece.chunks_mut(dim).enumerate() {
135                    let tok = t0 + ti;
136                    for h in 0..heads {
137                        let r = if global {
138                            tok * heads + h
139                        } else {
140                            let (w, si) = (tok / ws2, tok % ws2);
141                            (w * heads + h) * ws2 + si
142                        };
143                        row[h * hd..(h + 1) * hd].copy_from_slice(&o[r * hd..(r + 1) * hd]);
144                    }
145                }
146            });
147            self.io.v_o.set_a(&a)?;
148            self.timing.add("vit_rope_io", t1 - t0 + t2.elapsed());
149            gemm(&mut self.npu, &self.io.v_o, &self.w[&p("o")], &mut self.timing)?;
150            add_bf16(&mut x, &self.io.v_o.get_c(t)?, dim, None);
151
152            // LN2 -> fc1 (+gelu) -> fc2, the intermediate staying on the device
153            {
154                let st = &self.store;
155                layer_norm_bf16(&x, dim, st.f32(&p("ln2.w"))?, st.f32(&p("ln2.b"))?, eps, &mut a);
156            }
157            self.io.v_fc1.set_a(&a)?;
158            gemm(&mut self.npu, &self.io.v_fc1, &self.w[&p("fc1")], &mut self.timing)?;
159            gemm(&mut self.npu, &self.io.v_fc2, &self.w[&p("fc2")], &mut self.timing)?;
160            add_bf16(&mut x, &self.io.v_fc2.get_c(t)?, dim, Some(self.store.f32(&p("fc2.b"))?));
161        }
162        let mut out = vec![0f32; t * dim];
163        for (j, row) in x.chunks(dim).enumerate() {
164            let r = perm[j] as usize;
165            out[r * dim..(r + 1) * dim].copy_from_slice(row);
166        }
167        Ok(out)
168    }
169}
170
171/// The backbone with every layer on the device (bundles with
172/// `vit_device`, `iron/applications/sam3/sam3_npu.py`'s
173/// `_forward_device`): per layer
174    ///
175///   qkv GEMM -> RoPE(q | k) -> MHA -> o GEMM -> AddLN(+o) -> fc1 ->
176///   fc2 -> AddLN(+fc2)
177///
178/// each kernel reading the previous one's output buffer in place; the
179/// host embeds the patches and normalises once before layer 0 and
180/// reads the residual stream (f32 on the device with `vit_res_f32`,
181/// else bf16) after layer 31. Over the runtime's parts rather than
182/// `&mut Sam3` so that `Sam3::segment` can run the text encoder on the
183/// store meanwhile.
184pub(crate) fn vit_device(
185    npu: &mut Npu,
186    io: &mut Ios,
187    w: &HashMap<String, Buffer>,
188    store: &Store,
189    cfg: &Config,
190    timing: &mut Timing,
191    pixels: &[f32],
192) -> Result<Vec<f32>, Error> {
193    {
194        let c = cfg;
195        let (t, dim, g, ps, s) = (c.tokens(), c.vit_dim, c.grid, c.patch, c.image_size);
196        if pixels.len() != 3 * s * s {
197            return Err(Error::Input(format!("pixels must be [3, {s}, {s}]")));
198        }
199        let perm = store.i32("v.perm")?.to_vec();
200        let eps = c.vit_eps;
201        let t0 = std::time::Instant::now();
202
203        // patches (window order) -> embed GEMM -> + pos -> LN_pre: x0, then
204        // the first LayerNorm (no affine: folded into qkv) of bf16(x0)
205        let ke = npu.spec("v_embed")?.k;
206        let mut patches = vec![0u16; t * ke];
207        par_rows(&mut patches, ke, |r0, piece| {
208            for (ri, row) in piece.chunks_mut(ke).enumerate() {
209                let p = perm[r0 + ri] as usize;
210                let (py, px) = (p / g, p % g);
211                for ch in 0..3 {
212                    for ky in 0..ps {
213                        let src = ch * s * s + (py * ps + ky) * s + px * ps;
214                        for kx in 0..ps {
215                            row[(ch * ps + ky) * ps + kx] = f32_to_bf16(pixels[src + kx]);
216                        }
217                    }
218                }
219            }
220        });
221        io.v_embed.set_a(&patches)?;
222        gemm(npu, &io.v_embed, &w["v.embed"], timing)?;
223        let emb = io.v_embed.get_c(t)?;
224        let st = store;
225        let pos = st.f32("v.pos")?;
226        let (lw, lb) = (st.f32("v.ln_pre.w")?, st.f32("v.ln_pre.b")?);
227        let ones = vec![1f32; dim];
228        let zeros = vec![0f32; dim];
229        // the residual stream's first value: f32, or rounded to bf16 when the
230        // device keeps the stream in bf16 (then h is the LayerNorm of that)
231        let res_f32 = c.vit_res_f32;
232        let mut x0 = vec![0f32; t * dim];
233        let mut h = vec![0u16; t * dim];
234        par_rows(&mut x0, dim, |r0, piece| {
235            let mut tmp = vec![0f32; dim];
236            for (ri, row) in piece.chunks_mut(dim).enumerate() {
237                let r = r0 + ri;
238                for j in 0..dim {
239                    tmp[j] = bf16_to_f32(emb[r * dim + j]) + pos[r * dim + j];
240                }
241                ln_row(&tmp, row, lw, lb, eps);
242                if !res_f32 {
243                    for v in row.iter_mut() {
244                        *v = bf16_to_f32(f32_to_bf16(*v));
245                    }
246                }
247            }
248        });
249        par_rows(&mut h, dim, |r0, piece| {
250            let mut out = vec![0f32; dim];
251            for (ri, row) in piece.chunks_mut(dim).enumerate() {
252                let r = r0 + ri;
253                ln_row(&x0[r * dim..(r + 1) * dim], &mut out, &ones, &zeros, eps);
254                for (o, v) in row.iter_mut().zip(&out) {
255                    *o = f32_to_bf16(*v);
256                }
257            }
258        });
259        if res_f32 {
260            push(&x0, &mut io.xres[0])?;
261        } else {
262            let xb: Vec<u16> = x0.iter().map(|&v| f32_to_bf16(v)).collect();
263            push(&xb, &mut io.xres[0])?;
264        }
265        io.v_qkv.set_a(&h)?;
266        timing.add("vit_prologue", t0.elapsed());
267
268        let rope_out = io.rope_out.as_ref().ok_or_else(|| Error::Bundle("no RoPE buffer".into()))?;
269        for i in 0..c.vit_layers {
270            let p = |n: &str| format!("v.{i}.{n}");
271            let global = c.vit_global.contains(&i);
272            let (tab, mha_key) = if global { ("v.rope_tab.glob", "mha_glob") } else { ("v.rope_tab.win", "mha_win") };
273            gemm_dev(npu, &io.v_qkv, &w[&p("qkv")], timing)?;
274            op(npu, "rope", &[&io.v_qkv.c, &w[tab], rope_out], timing)?;
275            op(npu, mha_key, &[rope_out, rope_out, &io.v_qkv.c, &io.v_o.a], timing)?;
276            gemm_dev(npu, &io.v_o, &w[&p("o")], timing)?;
277            op(npu, "addln", &[&io.xres[0], &io.v_o.c, &io.xres[1], &io.v_fc1.a], timing)?;
278            gemm_dev(npu, &io.v_fc1, &w[&p("fc1")], timing)?;
279            gemm_dev(npu, &io.v_fc2, &w[&p("fc2")], timing)?;
280            op(npu, "addln", &[&io.xres[1], &io.v_fc2.c, &io.xres[0], &io.v_qkv.a], timing)?;
281        }
282        let t1 = std::time::Instant::now();
283        let xf: Vec<f32> = if c.vit_res_f32 {
284            pull(&io.xres[0], t * dim)?
285        } else {
286            pull::<u16>(&io.xres[0], t * dim)?.into_iter().map(bf16_to_f32).collect()
287        };
288        let mut out = vec![0f32; t * dim];
289        for (j, row) in xf.chunks(dim).enumerate() {
290            let r = perm[j] as usize;
291            out[r * dim..(r + 1) * dim].copy_from_slice(row);
292        }
293        timing.add("vit_readback", t1.elapsed());
294        Ok(out)
295    }
296}
297
298/// LayerNorm of `x`'s rows into `dst` (bf16), rows `dim` wide.
299pub(crate) fn layer_norm_bf16(x: &[f32], dim: usize, w: &[f32], b: &[f32], eps: f32, dst: &mut [u16]) {
300    let n = x.len();
301    par_rows(&mut dst[..n], dim, |r0, piece| {
302        let mut tmp = vec![0f32; dim];
303        for (ri, row) in piece.chunks_mut(dim).enumerate() {
304            let r = r0 + ri;
305            ln_row(&x[r * dim..(r + 1) * dim], &mut tmp, w, b, eps);
306            for (o, v) in row.iter_mut().zip(&tmp) {
307                *o = f32_to_bf16(*v);
308            }
309        }
310    });
311}
312
313/// `x += c (+ bias)`, `c` bf16 rows `dim` wide (at least as many as `x`'s).
314pub(crate) fn add_bf16(x: &mut [f32], c: &[u16], dim: usize, bias: Option<&[f32]>) {
315    par_rows(x, dim, |r0, piece| {
316        for (ri, row) in piece.chunks_mut(dim).enumerate() {
317            let src = &c[(r0 + ri) * dim..(r0 + ri + 1) * dim];
318            for j in 0..dim {
319                row[j] += bf16_to_f32(src[j]) + bias.map_or(0.0, |b| b[j]);
320            }
321        }
322    });
323}