1use 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 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 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 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 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")?; let st = &self.store;
83 let (wi, bi) = (st.f32("m.inst.w")?, st.f32("m.inst.b")?); 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 const PB: usize = 64;
117 let blocks = px.div_ceil(PB);
118 let mut planes = vec![0f32; (q + 1) * blocks * PB]; 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 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
149fn 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 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}