Skip to main content

taconite_gaic/
lib.rs

1// SPDX-FileCopyrightText: Copyright (C) 2026 Brishen Hawkins
2// SPDX-License-Identifier: Apache-2.0
3
4//! GAIC (Grid Anchor based Image Cropping, VGG16 backbone) on an AMD XDNA
5//! NPU, replaying the bundle `iron/applications/gaic/export_gaic.py`
6//! writes, through [`taconite`].
7//!
8//! [`Gaic::features`] runs the backbone once per image — the 13 VGG16
9//! convs, each one flm.GEMM over an im2col *view* (see
10//! `iron/applications/gaic/gaic_npu.py`) — and reduces it to the 32-channel
11//! 1/16-scale map every crop is scored from; [`Gaic::score`] scores any
12//! number of candidate boxes against it (RoI + RoD align on the host, the
13//! 5184 -> 768 FC on the NPU, 768 -> 128 -> 1 on the host). [`anchors`]
14//! makes GAIC's candidate sets and [`preprocess`] the network input, both
15//! matching GAIC-Pytorch's demo exactly.
16//!
17//! Host glue per conv, in one threaded pass: the GEMM's output (bf16,
18//! pixel-major `[H][W + 2][OC]`, two junk columns a row) -> + bias, ReLU,
19//! [2x2 max-pool] -> the next conv's A source written straight into the
20//! shared input buffer. The activation buffers are host_only (uncached)
21//! BOs, so rows move through them with whole-row copies.
22//!
23//! One [`Gaic`] per process: its kernels stay resident as 9 of the NPU's 16
24//! hardware contexts. `Send`, not `Sync`.
25
26pub mod align;
27pub mod anchors;
28pub mod bundle;
29mod par;
30pub mod preprocess;
31
32use std::collections::{BTreeMap, HashMap, VecDeque};
33use std::fmt;
34use std::path::Path;
35use std::time::{Duration, Instant};
36
37use taconite::{bf16_to_f32, f32_to_bf16};
38// The NPU path: XRT (feature `xrt`, the default), or the driver's ioctls
39// with no XRT (feature `direct`, which wins when both are on).
40#[cfg(feature = "direct")]
41use taconite::direct::{Buffer, Kernel, Run, Session};
42#[cfg(all(feature = "xrt", not(feature = "direct")))]
43use taconite::{Buffer, Kernel, Run, Session};
44#[cfg(not(any(feature = "xrt", feature = "direct")))]
45compile_error!("no NPU path: enable feature `xrt` (the default) or `direct`");
46
47use par::par_rows;
48
49use bundle::{Bundle, ConvSpec};
50
51#[derive(Debug)]
52pub enum Error {
53    /// The bundle is missing, malformed, or does not match this runtime.
54    Bundle(String),
55    /// XRT: the device, a kernel load, a buffer or a run.
56    Npu(taconite::Error),
57    /// An input the model cannot take (size, box count).
58    Input(String),
59}
60
61impl fmt::Display for Error {
62    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
63        match self {
64            Error::Bundle(m) => write!(f, "GAIC bundle: {m}"),
65            Error::Npu(e) => write!(f, "{e}"),
66            Error::Input(m) => write!(f, "GAIC input: {m}"),
67        }
68    }
69}
70
71impl std::error::Error for Error {}
72
73impl From<taconite::Error> for Error {
74    fn from(e: taconite::Error) -> Self {
75        Error::Npu(e)
76    }
77}
78
79/// Launches kept in flight per layer: enough to hide the host's turnaround
80/// between chunks, few enough to stay clear of the driver's command queue.
81const IN_FLIGHT: usize = 8;
82
83/// Wall time per stage, accumulated since the last [`Gaic::reset_timing`].
84#[derive(Debug, Clone, Default)]
85pub struct Timing {
86    pub stages: BTreeMap<&'static str, Duration>,
87    pub dispatches: usize,
88}
89
90impl Timing {
91    fn add(&mut self, stage: &'static str, d: Duration) {
92        *self.stages.entry(stage).or_default() += d;
93    }
94    pub fn total(&self) -> Duration {
95        self.stages.values().sum()
96    }
97}
98
99impl fmt::Display for Timing {
100    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
101        for (k, v) in &self.stages {
102            write!(f, "{k} {:.1} ms, ", v.as_secs_f64() * 1e3)?;
103        }
104        write!(f, "{} dispatches", self.dispatches)
105    }
106}
107
108/// The reduced feature map of one image: `[reddim][h][w]` f32 at 1/16 of
109/// the network input (`input_w x input_h`), what every box is scored from.
110#[derive(Debug, Clone)]
111pub struct Features {
112    pub input_w: usize,
113    pub input_h: usize,
114    pub h: usize,
115    pub w: usize,
116    pub channels: usize,
117    pub map: Vec<f32>,
118}
119
120struct Layer {
121    spec: ConvSpec,
122    /// The GEMM's K (the window layer's is padded past P x 9C).
123    k: usize,
124    a_elems: usize,
125    c_elems: usize,
126    b: Buffer,
127    bias: Vec<f32>,
128}
129
130/// Per input size: the shared A-source / output buffers and every chunk's
131/// sub-buffer views of them (XRT caches a run per argument tuple, so these
132/// live as long as the size does).
133struct Plan {
134    w: usize,
135    h: usize,
136    y: Buffer,
137    c: Buffer,
138    layers: Vec<LayerPlan>,
139}
140
141struct LayerPlan {
142    /// The conv's input (and output) size, after any pool before it.
143    h: usize,
144    w: usize,
145    /// Output pixels computed: h x (w + 2), padded to whole dispatches.
146    pixels: usize,
147    chunks: Vec<(Buffer, Buffer)>,
148}
149
150pub struct Gaic {
151    bundle: Bundle,
152    session: Session,
153    kernels: HashMap<String, Kernel>,
154    layers: Vec<Layer>,
155    f3: usize,
156    f4: usize,
157    dimred_w: Vec<f32>,
158    dimred_b: Vec<f32>,
159    fc1_a: Buffer,
160    fc1_b: Buffer,
161    fc1_c: Buffer,
162    fc1_bias: Vec<f32>,
163    fc2_w: Vec<f32>,
164    fc2_b: Vec<f32>,
165    fc3_w: Vec<f32>,
166    fc3_b: f32,
167    plan: Option<Plan>,
168    threads: usize,
169    pub timing: Timing,
170}
171
172fn round_up(x: usize, m: usize) -> usize {
173    x.div_ceil(m) * m
174}
175
176/// `a . b` with eight partial sums, so the loop vectorises.
177#[inline]
178fn dot(a: &[f32], b: &[f32]) -> f32 {
179    let mut acc = [0f32; 8];
180    let (ca, cb) = (a.chunks_exact(8), b.chunks_exact(8));
181    let (ra, rb) = (ca.remainder(), cb.remainder());
182    for (x, y) in ca.zip(cb) {
183        for i in 0..8 {
184            acc[i] += x[i] * y[i];
185        }
186    }
187    let mut s = acc.iter().sum::<f32>();
188    for (x, y) in ra.iter().zip(rb) {
189        s += x * y;
190    }
191    s
192}
193
194impl Gaic {
195    /// Opens the NPU and loads a bundle: every kernel into a resident
196    /// hardware context, every weight into a device buffer.
197    pub fn load(bundle: &Path) -> Result<Self, Error> {
198        let bundle = Bundle::load(bundle)?;
199        let session = Session::open(0)?;
200        let mut kernels = HashMap::new();
201        for k in &bundle.kernels {
202            let ops = 2 * (k.m * k.k * k.n) as u64;
203            kernels.insert(k.key.clone(), session.load_kernel(&k.xclbin, &k.insts, Some(&k.name), ops)?);
204        }
205        // Packed weights; their sizes were checked against the kernels at load.
206        let upload = |name: &str| -> Result<Buffer, Error> {
207            let data = bundle.store.bytes(name)?;
208            let mut b = session.alloc(data.len())?;
209            b.write(data)?;
210            Ok(b)
211        };
212        let mut layers = Vec::new();
213        for c in &bundle.convs {
214            let k = bundle.kernel(&c.kernel)?;
215            layers.push(Layer {
216                spec: c.clone(),
217                k: k.k,
218                a_elems: k.a_elems,
219                c_elems: k.c_elems,
220                b: upload(&c.b)?,
221                bias: bundle.f32_vec(&c.bias)?,
222            });
223        }
224        let idx = |name: &str| bundle.convs.iter().position(|c| c.name == name).unwrap();
225        let (f3, f4) = (idx(&bundle.f3), idx(&bundle.f4));
226        let fk = bundle.kernel(&bundle.fc1.kernel)?;
227        let fc1_b = upload(&bundle.fc1.b)?;
228        let fc1_a = session.alloc_of::<u16>(fk.a_elems)?;
229        let fc1_c = session.alloc_of::<u16>(fk.c_elems)?;
230        let m = &bundle;
231        let dimred_w = m.f32_vec(&m.dimred_w)?;
232        let dimred_b = m.f32_vec(&m.dimred_b)?;
233        let fc1_bias = m.f32_vec(&m.fc1.bias)?;
234        let fc2_w = m.f32_vec(&m.fc2_w)?;
235        let fc2_b = m.f32_vec(&m.fc2_b)?;
236        let fc3_w = m.f32_vec(&m.fc3_w)?;
237        let fc3_b = m.f32_vec(&m.fc3_b)?[0];
238        let threads = par::default_threads();
239        Ok(Self {
240            bundle,
241            session,
242            kernels,
243            layers,
244            f3,
245            f4,
246            dimred_w,
247            dimred_b,
248            fc1_a,
249            fc1_b,
250            fc1_c,
251            fc1_bias,
252            fc2_w,
253            fc2_b,
254            fc3_w,
255            fc3_b,
256            plan: None,
257            threads,
258            timing: Timing::default(),
259        })
260    }
261
262    /// The loaded bundle: its layer table, and the tensor store (weights and
263    /// the self-check's `ref.*` references).
264    pub fn bundle(&self) -> &Bundle {
265        &self.bundle
266    }
267
268    pub fn reset_timing(&mut self) {
269        self.timing = Timing::default();
270    }
271
272    /// Host threads the glue uses (default: the machine's, up to 16).
273    pub fn set_threads(&mut self, n: usize) {
274        self.threads = n.max(1);
275    }
276
277    fn plan_for(&mut self, w: usize, h: usize) -> Result<(), Error> {
278        if self.plan.as_ref().is_some_and(|p| p.w == w && p.h == h) {
279            return Ok(());
280        }
281        self.plan = None; // free the previous size's buffers first
282        let (mut hh, mut ww) = (h, w);
283        let mut dims = Vec::new();
284        let (mut y_max, mut c_max) = (0, 0);
285        for l in &self.layers {
286            let s = &l.spec;
287            if s.pool_before {
288                hh /= 2;
289                ww /= 2;
290            }
291            let pixels = round_up(hh * (ww + 2), s.m_chunk * s.p);
292            let a_src = if s.window { pixels / s.p * l.k } else { (pixels + 2) * s.d };
293            y_max = y_max.max(a_src);
294            c_max = c_max.max(pixels * s.oc);
295            dims.push((hh, ww, pixels));
296        }
297        let y = self.session.alloc_of::<u16>(y_max)?;
298        let c = self.session.alloc_of::<u16>(c_max)?;
299        let mut layers = Vec::new();
300        for (l, &(h, w, pixels)) in self.layers.iter().zip(&dims) {
301            let s = &l.spec;
302            let chunk_px = s.m_chunk * s.p;
303            let step = if s.window { s.m_chunk * l.k } else { chunk_px * s.d };
304            let chunks = (0..pixels / chunk_px)
305                .map(|i| Ok((y.sub_of::<u16>(i * step, l.a_elems)?, c.sub_of::<u16>(i * chunk_px * s.oc, l.c_elems)?)))
306                .collect::<Result<Vec<_>, Error>>()?;
307            layers.push(LayerPlan { h, w, pixels, chunks });
308        }
309        self.plan = Some(Plan { w, h, y, c, layers });
310        Ok(())
311    }
312
313    /// The backbone and DimRed for one image: `chw` is the normalized
314    /// `[3][h][w]` input ([`preprocess::preprocess`]); `w` and `h` must be
315    /// multiples of 32 (GAIC's own resize makes them so).
316    pub fn features(&mut self, chw: &[f32], w: usize, h: usize) -> Result<Features, Error> {
317        if w % 32 != 0 || h % 32 != 0 || w < 64 || h < 64 {
318            return Err(Error::Input(format!("{w}x{h}: both sides must be multiples of 32, at least 64")));
319        }
320        if chw.len() != 3 * w * h {
321            return Err(Error::Input(format!("{} values for a 3x{h}x{w} input", chw.len())));
322        }
323        self.plan_for(w, h)?;
324        let threads = self.threads;
325
326        // The stem's input: the image as bf16 [h][w][3].
327        let t0 = Instant::now();
328        let mut act = vec![0u16; w * h * 3];
329        par_rows(&mut act, w * 3, threads, |r0, rows| {
330            for (i, px) in rows.chunks_exact_mut(3).enumerate() {
331                let p = r0 * w + i;
332                for c in 0..3 {
333                    px[c] = f32_to_bf16(chw[c * w * h + p]);
334                }
335            }
336        });
337        self.timing.add("glue", t0.elapsed());
338
339        let (mut f3, mut f4) = (Vec::new(), Vec::new());
340        let n = self.layers.len();
341        for i in 0..n {
342            let t0 = Instant::now();
343            {
344                let plan = self.plan.as_mut().unwrap();
345                let lp = &plan.layers[i];
346                build_a(&mut plan.y, &act, &self.layers[i], lp, threads);
347            }
348            self.timing.add("build_a", t0.elapsed());
349
350            // Syncs go through each chunk's sub-buffer: the shared buffers
351            // are sized for the largest layer, and syncing all of them for
352            // a small one was most of the host's time.
353            let t0 = Instant::now();
354            let plan = self.plan.as_ref().unwrap();
355            let lp = &plan.layers[i];
356            let layer = &self.layers[i];
357            let kernel = &self.kernels[&layer.spec.kernel];
358            let mut runs: VecDeque<(Run, &Buffer)> = VecDeque::new();
359            for (a, c) in &lp.chunks {
360                if runs.len() == IN_FLIGHT {
361                    let (r, c) = runs.pop_front().unwrap();
362                    r.wait()?;
363                    c.sync_from_device()?;
364                }
365                a.sync_to_device()?;
366                runs.push_back((kernel.start(&[a, &layer.b, c])?, c));
367            }
368            for (r, c) in runs {
369                r.wait()?;
370                c.sync_from_device()?;
371            }
372            self.timing.dispatches += lp.chunks.len();
373            self.timing.add("npu", t0.elapsed());
374
375            // Read in place: measured as fast as a bulk copy out first.
376            let t0 = Instant::now();
377            let s = &layer.spec;
378            let out = &plan.c.as_slice::<u16>()[..lp.pixels * s.oc];
379            let pool = self.layers.get(i + 1).is_some_and(|l| l.spec.pool_before);
380            if i == self.f3 || i == self.f4 {
381                let f = epilogue_f32(out, lp.h, lp.w, s.oc, &layer.bias, threads);
382                if i == self.f3 { f3 = f } else { f4 = f }
383            }
384            if i + 1 < n {
385                act = epilogue_bf16(out, lp.h, lp.w, s.oc, &layer.bias, pool, threads);
386            }
387            self.timing.add("epilogue", t0.elapsed());
388        }
389
390        let t0 = Instant::now();
391        let plan = self.plan.as_ref().unwrap();
392        let (l3, l4) = (&plan.layers[self.f3], &plan.layers[self.f4]);
393        let ch = self.layers[self.f4].spec.oc;
394        let (h5, w5) = (l4.h / 2, l4.w / 2);
395        let f5 = maxpool_f32(&f4, l4.h, l4.w, ch);
396        let r = self.bundle.reddim;
397        let cin = self.bundle.dimred_in;
398        let proj = |f: &[f32], off: usize| -> Vec<f32> {
399            let px = f.len() / ch;
400            let mut g = vec![0f32; px * r];
401            par_rows(&mut g, r, threads, |p0, rows| {
402                for (j, o) in rows.chunks_exact_mut(r).enumerate() {
403                    let x = &f[(p0 + j) * ch..(p0 + j + 1) * ch];
404                    for (k, v) in o.iter_mut().enumerate() {
405                        *v = dot(x, &self.dimred_w[k * cin + off..k * cin + off + ch]);
406                    }
407                }
408            });
409            g
410        };
411        let (h4, w4) = (l4.h, l4.w);
412        let g3 = interp_align_corners(&proj(&f3, 0), l3.h, l3.w, r, h4, w4);
413        let g4 = proj(&f4, ch);
414        let g5 = interp_align_corners(&proj(&f5, 2 * ch), h5, w5, r, h4, w4);
415        let mut map = vec![0f32; r * h4 * w4];
416        for p in 0..h4 * w4 {
417            for k in 0..r {
418                let i = p * r + k;
419                map[k * h4 * w4 + p] = (g3[i] + g4[i]) + (0.5 * g5[i] + self.dimred_b[k]);
420            }
421        }
422        self.timing.add("dimred", t0.elapsed());
423        Ok(Features { input_w: w, input_h: h, h: h4, w: w4, channels: r, map })
424    }
425
426    /// Scores `boxes` (`[x1, y1, x2, y2]` in the input's pixels) against an
427    /// image's [`Features`]; higher is a better crop.
428    pub fn score(&mut self, f: &Features, boxes: &[[f32; 4]]) -> Result<Vec<f32>, Error> {
429        let m = &self.bundle;
430        let (s, scale) = (m.align_size, m.spatial_scale);
431        let (k, k_pad, n1, rows) = (m.fc1.k, m.fc1.k_pad, m.fc1.n, m.fc1.m);
432        let threads = self.threads;
433        let mut scores = Vec::with_capacity(boxes.len());
434        for group in boxes.chunks(rows) {
435            let t0 = Instant::now();
436            let mut a = vec![0u16; rows * k_pad];
437            par_rows(&mut a[..group.len() * k_pad], k_pad, threads, |b0, out| {
438                let mut feat = vec![0f32; k];
439                for (j, row) in out.chunks_exact_mut(k_pad).enumerate() {
440                    align::box_features(&f.map, f.channels, f.h, f.w, group[b0 + j], s, scale, &mut feat);
441                    for (d, &v) in row.iter_mut().zip(&feat) {
442                        *d = f32_to_bf16(v);
443                    }
444                }
445            });
446            self.fc1_a.write(&a)?;
447            self.timing.add("align", t0.elapsed());
448
449            let t0 = Instant::now();
450            self.kernels[&m.fc1.kernel].run(&[&self.fc1_a, &self.fc1_b, &self.fc1_c])?;
451            self.fc1_c.sync_from_device()?;
452            self.timing.dispatches += 1;
453            self.timing.add("npu", t0.elapsed());
454
455            let t0 = Instant::now();
456            let c = self.fc1_c.as_slice::<u16>()[..group.len() * n1].to_vec();
457            for row in c.chunks_exact(n1) {
458                let h1: Vec<f32> =
459                    row.iter().zip(&self.fc1_bias).map(|(&v, &b)| (bf16_to_f32(v) + b).max(0.0)).collect();
460                let h2: Vec<f32> = (0..m.fc2_out)
461                    .map(|o| (dot(&h1, &self.fc2_w[o * m.fc2_in..(o + 1) * m.fc2_in]) + self.fc2_b[o]).max(0.0))
462                    .collect();
463                scores.push(dot(&h2, &self.fc3_w) + self.fc3_b);
464            }
465            self.timing.add("fc", t0.elapsed());
466        }
467        Ok(scores)
468    }
469}
470
471/// Writes layer `l`'s A source into `y`, from its input activation `act`
472/// (bf16 `[h][w][C]`). Pixel `q` of the zero-bordered input (row pitch
473/// `Wp = w + 2`) contributes `Y[q] = [X_pad[q] | X_pad[q + Wp] | X_pad[q + 2 Wp]]`
474/// (3C, zero past `q = h Wp`); the view layers store Y itself, D wide,
475/// the window layer P whole windows `[Y[p] | Y[p+1] | Y[p+2]]` a row.
476fn build_a(y: &mut Buffer, act: &[u16], l: &Layer, lp: &LayerPlan, threads: usize) {
477    let s = &l.spec;
478    let (h, w, c) = (lp.h, lp.w, s.c);
479    let wp = w + 2;
480    let n = h * wp;
481    // Y[q][dy*C .. (dy+1)*C] into `dst` (3C long).
482    let y_row = |q: usize, dst: &mut [u16]| {
483        if q >= n {
484            dst.fill(0);
485            return;
486        }
487        let (r, x) = (q / wp, q % wp);
488        for dy in 0..3 {
489            let d = &mut dst[dy * c..(dy + 1) * c];
490            let rr = r + dy;
491            if rr >= 1 && rr <= h && x >= 1 && x <= w {
492                let src = ((rr - 1) * w + (x - 1)) * c;
493                d.copy_from_slice(&act[src..src + c]);
494            } else {
495                d.fill(0);
496            }
497        }
498    };
499    const BLOCK: usize = 64;
500    let (row_len, rows) = if s.window { (l.k, lp.pixels / s.p) } else { (s.d, lp.pixels + 2) };
501    let dst = &mut y.as_mut_slice::<u16>()[..rows * row_len];
502    par_rows(dst, row_len, threads, |r0, piece| {
503        // Rows are assembled in cached memory and stored a block at a time:
504        // the destination is an uncached host_only mapping.
505        let mut buf = vec![0u16; BLOCK * row_len];
506        for (b, out) in piece.chunks_mut(BLOCK * row_len).enumerate() {
507            let tmp = &mut buf[..out.len()];
508            for (j, row) in tmp.chunks_exact_mut(row_len).enumerate() {
509                let m = r0 + b * BLOCK + j;
510                if s.window {
511                    for q in 0..s.p {
512                        for dx in 0..3 {
513                            let o = q * 9 * c + dx * 3 * c;
514                            y_row(m * s.p + q + dx, &mut row[o..o + 3 * c]);
515                        }
516                    }
517                    row[s.p * 9 * c..].fill(0);
518                } else {
519                    y_row(m, &mut row[..3 * c]);
520                    row[3 * c..].fill(0);
521                }
522            }
523            out.copy_from_slice(tmp);
524        }
525    });
526}
527
528/// GEMM output (bf16 `[pixels][OC]`, `pixels` = h x (w + 2) + padding) ->
529/// the next conv's input, bf16 `relu(out + bias)` `[h][w][OC]`, 2x2
530/// max-pooled when `pool` (pool and ReLU commute: both monotone).
531fn epilogue_bf16(out: &[u16], h: usize, w: usize, oc: usize, bias: &[f32], pool: bool, threads: usize) -> Vec<u16> {
532    let wp = w + 2;
533    // Input pixel (y, x)'s OC values.
534    let px = |y: usize, x: usize| &out[(y * wp + x) * oc..(y * wp + x + 1) * oc];
535    let (ho, wo) = if pool { (h / 2, w / 2) } else { (h, w) };
536    let mut act = vec![0u16; ho * wo * oc];
537    par_rows(&mut act, wo * oc, threads, |y0, rows| {
538        for (j, row) in rows.chunks_exact_mut(wo * oc).enumerate() {
539            let y = y0 + j;
540            for (x, d) in row.chunks_exact_mut(oc).enumerate() {
541                if pool {
542                    let (a, b) = (px(2 * y, 2 * x), px(2 * y, 2 * x + 1));
543                    let (c, e) = (px(2 * y + 1, 2 * x), px(2 * y + 1, 2 * x + 1));
544                    for o in 0..oc {
545                        let m = bf16_to_f32(a[o]).max(bf16_to_f32(b[o])).max(bf16_to_f32(c[o])).max(bf16_to_f32(e[o]));
546                        d[o] = f32_to_bf16((m + bias[o]).max(0.0));
547                    }
548                } else {
549                    for ((d, &s), &b) in d.iter_mut().zip(px(y, x)).zip(bias) {
550                        *d = f32_to_bf16((bf16_to_f32(s) + b).max(0.0));
551                    }
552                }
553            }
554        }
555    });
556    act
557}
558
559/// As [`epilogue_bf16`] without the pool, kept in f32: f3 / f4.
560fn epilogue_f32(out: &[u16], h: usize, w: usize, oc: usize, bias: &[f32], threads: usize) -> Vec<f32> {
561    let wp = w + 2;
562    let mut f = vec![0f32; h * w * oc];
563    par_rows(&mut f, w * oc, threads, |y0, rows| {
564        for (j, row) in rows.chunks_exact_mut(w * oc).enumerate() {
565            let src = &out[(y0 + j) * wp * oc..((y0 + j) * wp + w) * oc];
566            for (d, s) in row.chunks_exact_mut(oc).zip(src.chunks_exact(oc)) {
567                for ((d, &s), &b) in d.iter_mut().zip(s).zip(bias) {
568                    *d = (bf16_to_f32(s) + b).max(0.0);
569                }
570            }
571        }
572    });
573    f
574}
575
576fn maxpool_f32(f: &[f32], h: usize, w: usize, c: usize) -> Vec<f32> {
577    let (ho, wo) = (h / 2, w / 2);
578    let mut out = vec![0f32; ho * wo * c];
579    for y in 0..ho {
580        for x in 0..wo {
581            for o in 0..c {
582                let at = |yy: usize, xx: usize| f[(yy * w + xx) * c + o];
583                out[(y * wo + x) * c + o] =
584                    at(2 * y, 2 * x).max(at(2 * y, 2 * x + 1)).max(at(2 * y + 1, 2 * x)).max(at(2 * y + 1, 2 * x + 1));
585            }
586        }
587    }
588    out
589}
590
591/// `[h][w][c]` -> `[oh][ow][c]`, bilinear with align_corners=True, the way
592/// torch's CPU upsample_bilinear2d computes it (f32 source index, the far
593/// neighbour clamped at the edge).
594fn interp_align_corners(g: &[f32], h: usize, w: usize, c: usize, oh: usize, ow: usize) -> Vec<f32> {
595    let axis = |i: usize, o: usize| -> Vec<(usize, usize, f32, f32)> {
596        let scale = if o > 1 { (i as f32 - 1.0) / (o as f32 - 1.0) } else { 0.0 };
597        (0..o)
598            .map(|d| {
599                let src = scale * d as f32;
600                let i0 = src as usize;
601                let i1 = if i0 < i - 1 { i0 + 1 } else { i0 };
602                let l1 = src - i0 as f32;
603                (i0, i1, 1.0 - l1, l1)
604            })
605            .collect()
606    };
607    let (ys, xs) = (axis(h, oh), axis(w, ow));
608    let mut out = vec![0f32; oh * ow * c];
609    for (y, &(y0, y1, hy0, hy1)) in ys.iter().enumerate() {
610        for (x, &(x0, x1, wx0, wx1)) in xs.iter().enumerate() {
611            for k in 0..c {
612                let at = |yy: usize, xx: usize| g[(yy * w + xx) * c + k];
613                out[(y * ow + x) * c + k] =
614                    hy0 * (wx0 * at(y0, x0) + wx1 * at(y0, x1)) + hy1 * (wx0 * at(y1, x0) + wx1 * at(y1, x1));
615            }
616        }
617    }
618    out
619}