1use 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 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 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 {
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 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 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 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 {
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
171pub(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 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 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
298pub(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
313pub(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}