Oregami
Repositories/oxedyne/fe2o3

oxedyne/fe2o3/fe2o3_infer/src/kern.rs

38.5 KiB, 67 runs

created by r1870400018:19731, which is this file's identity for as long as the history lasts, whatever it is later renamed to

download · who wrote it · its history

1//! Safe `f32` kernels, and the single runtime dispatch that selects between a
2//! fused-multiply-add code path and a baseline one.
3//!
4//! Every kernel here is ordinary safe Rust. The only `unsafe` token in the
5//! crate is in [`run`], where a `#[target_feature]` function is called after
6//! its features have been checked at runtime; that is a sanctioned exception
7//! and is documented at the call site.
8//!
9//! # Why two code paths
10//!
11//! Rust will not contract `a * b + c` into a fused multiply-add on its own --
12//! strict IEEE semantics forbid it -- so the only route to `vfmadd` from safe
13//! code is [`f32::mul_add`]. On a target that has no FMA instruction,
14//! `mul_add` becomes a libm call and costs about thirty times the arithmetic
15//! it replaces. Both bodies therefore exist, selected by a const generic, and
16//! the `mul_add` body is reachable only from a function compiled with the
17//! feature enabled.
18
19use crate::tensor::Tensor;
20
21use oxedyne_fe2o3_core::prelude::*;
22
23/// Height of the register tile in the blocked matrix kernel.
24///
25/// This is not a free parameter. Adjacent values differ by more than an order
26/// of magnitude in throughput, because the code generator decides all-or-nothing
27/// whether the `[[f32; NR]; MR]` accumulator lives in vector registers or spills
28/// to the stack. The regression guard in `tests/guard.rs` exists to catch a
29/// compiler upgrade that moves the boundary.
30pub const MR: usize = 6;
31
32/// Width of the register tile in the blocked matrix kernel. See [`MR`].
33pub const NR: usize = 16;
34
35/// Rows of `A` held in one cache block.
36pub const MC: usize = 256;
37
38/// Depth of one cache block.
39pub const KC: usize = 512;
40
41/// Columns of `B` held in one cache block.
42pub const NC: usize = 1024;
43
44/// The instruction set the kernels were dispatched onto.
45#[derive(Clone, Copy, Debug, Eq, PartialEq)]
46pub enum Cpu {
47 /// No vector feature was detected, so `mul_add` must not be reached.
48 Baseline,
49 /// An `x86-64` part carrying both AVX2 and FMA.
50 Avx2Fma,
51}
52
53impl Cpu {
54 /// Detects the best available instruction set on this machine.
55 pub fn detect() -> Self {
56 #[cfg(target_arch = "x86_64")]
57 {
58 if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
59 return Self::Avx2Fma;
60 }
61 }
62 Self::Baseline
63 }
64
65 /// Whether this path may use [`f32::mul_add`].
66 pub fn has_fma(&self) -> bool {
67 matches!(self, Self::Avx2Fma)
68 }
69}
70
71impl Default for Cpu {
72 fn default() -> Self {
73 Self::detect()
74 }
75}
76
77/// Reusable packing buffers, so a graph run allocates once rather than per layer.
78#[derive(Debug, Default)]
79pub struct Scratch {
80 /// Packed panels of `A`.
81 ap: Vec<f32>,
82 /// Packed panels of `B`.
83 bp: Vec<f32>,
84}
85
86impl Scratch {
87 /// Creates an empty scratch buffer.
88 pub fn new() -> Self {
89 Self { ap: Vec::new(), bp: Vec::new() }
90 }
91
92 /// Grows both buffers to at least the requested sizes.
93 #[inline(always)]
94 fn ensure(&mut self, na: usize, nb: usize) {
95 if self.ap.len() < na {
96 self.ap.resize(na, 0.0);
97 }
98 if self.bp.len() < nb {
99 self.bp.resize(nb, 0.0);
100 }
101 }
102}
103
104/// How a resize takes its value once the source position is known.
105#[derive(Clone, Copy, Debug, PartialEq)]
106pub enum Sample {
107 /// The value at the floor of the source position.
108 Nearest,
109 /// The weighted mean of the four samples around it.
110 Bilinear,
111}
112
113/// How a resize maps an output position back into the source.
114///
115/// The two conventions differ by half a sample and no more, and that is enough
116/// to move every box a network built on one of them predicts. Both models this
117/// crate carries name theirs in the graph, so neither is guessed at.
118#[derive(Clone, Copy, Debug, PartialEq)]
119pub enum Coord {
120 /// `src = dst / scale`. ONNX calls it `asymmetric`.
121 Asymmetric,
122 /// `src = (dst + ½) / scale − ½`, clamped at zero. ONNX calls it
123 /// `pytorch_half_pixel`, and it degenerates to zero when the output has a
124 /// single row or column.
125 HalfPixel,
126}
127
128impl Coord {
129 /// Maps one output index to a position in the source.
130 pub fn source(&self, dst: usize, scale: f32, out: usize) -> f32 {
131 match self {
132 Self::Asymmetric => dst as f32 / scale,
133 Self::HalfPixel => if out > 1 {
134 ((dst as f32 + 0.5) / scale - 0.5).max(0.0)
135 } else {
136 0.0
137 },
138 }
139 }
140}
141
142/// One unit of numerical work, named so that a single dispatch point can carry
143/// every kernel across the feature boundary.
144///
145/// Passing the work as a value rather than calling each kernel directly is what
146/// keeps the crate to one `unsafe` token: the whole set is monomorphised twice,
147/// once inside a `#[target_feature]` function and once outside it, and the
148/// caller picks between them by matching on [`Cpu`].
149pub enum Task<'a> {
150 /// `c[m, n] = bias + a[m, k] · b[k, n]`, all row-major.
151 Gemm {
152 /// Rows of `a` and of `c`.
153 m: usize,
154 /// Columns of `b` and of `c`.
155 n: usize,
156 /// Shared inner extent.
157 k: usize,
158 /// Left operand, `[m, k]`.
159 a: &'a [f32],
160 /// Right operand, `[k, n]`.
161 b: &'a [f32],
162 /// Destination, `[m, n]`, overwritten.
163 c: &'a mut [f32],
164 /// Optional per-column bias, prefilled into `c` before accumulation.
165 bias: Option<&'a [f32]>,
166 /// Packing buffers.
167 scratch: &'a mut Scratch,
168 },
169 /// `c[n] = a[k] · bt[n, k]ᵀ`, the matrix--vector case an ONNX `Gemm` with
170 /// `transB=1` presents, where each output is a contiguous dot product.
171 MatVecT {
172 /// Number of outputs.
173 n: usize,
174 /// Length of each dot product.
175 k: usize,
176 /// The vector, `[k]`.
177 a: &'a [f32],
178 /// Weights, `[n, k]`.
179 bt: &'a [f32],
180 /// Destination, `[n]`.
181 c: &'a mut [f32],
182 },
183 /// Gathers `[oh·ow, kh·kw·ch]` convolution patches out of an `NHWC` plane.
184 Im2Col {
185 /// Channels.
186 ch: usize,
187 /// Input height.
188 h: usize,
189 /// Input width.
190 w: usize,
191 /// Kernel height.
192 kh: usize,
193 /// Kernel width.
194 kw: usize,
195 /// Vertical stride.
196 sy: usize,
197 /// Horizontal stride.
198 sx: usize,
199 /// Padding above.
200 pt: usize,
201 /// Padding to the left.
202 pl: usize,
203 /// Output height.
204 oh: usize,
205 /// Output width.
206 ow: usize,
207 /// Source plane, `[h, w, ch]`.
208 x: &'a [f32],
209 /// Destination, `[oh·ow, kh·kw·ch]`.
210 out: &'a mut [f32],
211 },
212 /// Depthwise convolution in `NHWC`, one weight plane per channel.
213 Depthwise {
214 /// Channels.
215 ch: usize,
216 /// Input height.
217 h: usize,
218 /// Input width.
219 w: usize,
220 /// Kernel height.
221 kh: usize,
222 /// Kernel width.
223 kw: usize,
224 /// Vertical stride.
225 sy: usize,
226 /// Horizontal stride.
227 sx: usize,
228 /// Padding above.
229 pt: usize,
230 /// Padding to the left.
231 pl: usize,
232 /// Output height.
233 oh: usize,
234 /// Output width.
235 ow: usize,
236 /// Source plane, `[h, w, ch]`.
237 x: &'a [f32],
238 /// Weights, `[kh·kw, ch]`.
239 wt: &'a [f32],
240 /// Optional per-channel bias.
241 bias: Option<&'a [f32]>,
242 /// Destination, `[oh, ow, ch]`.
243 y: &'a mut [f32],
244 },
245 /// Per-channel affine map, `x = scale·x + bias`, over a channels-last buffer.
246 Scale {
247 /// Channels, the innermost extent.
248 ch: usize,
249 /// Buffer, rewritten in place.
250 x: &'a mut [f32],
251 /// Per-channel multiplier.
252 scale: &'a [f32],
253 /// Per-channel offset.
254 bias: &'a [f32],
255 },
256 /// Parametric rectified linear unit with a per-channel slope, written
257 /// branchlessly so that the loop still vectorises.
258 PRelu {
259 /// Channels, the innermost extent.
260 ch: usize,
261 /// Buffer, rewritten in place.
262 x: &'a mut [f32],
263 /// Per-channel negative slope.
264 slope: &'a [f32],
265 },
266 /// Rectified linear unit.
267 Relu {
268 /// Buffer, rewritten in place.
269 x: &'a mut [f32],
270 },
271 /// Leaky rectified linear unit, one slope for every channel.
272 Leaky {
273 /// Buffer, rewritten in place.
274 x: &'a mut [f32],
275 /// Negative slope.
276 slope: f32,
277 },
278 /// Logistic sigmoid.
279 Sigmoid {
280 /// Buffer, rewritten in place.
281 x: &'a mut [f32],
282 },
283 /// Maximum pool in `NHWC`, over any kernel, stride and padding.
284 ///
285 /// Padding contributes nothing rather than zero: a zero would win the
286 /// maximum wherever the real samples are negative, which after a leaky
287 /// rectifier they routinely are.
288 MaxPool {
289 /// Channels.
290 ch: usize,
291 /// Input height.
292 h: usize,
293 /// Input width.
294 w: usize,
295 /// Kernel height.
296 kh: usize,
297 /// Kernel width.
298 kw: usize,
299 /// Vertical stride.
300 sy: usize,
301 /// Horizontal stride.
302 sx: usize,
303 /// Padding above.
304 pt: usize,
305 /// Padding to the left.
306 pl: usize,
307 /// Output height.
308 oh: usize,
309 /// Output width.
310 ow: usize,
311 /// Source, `[h, w, ch]`.
312 x: &'a [f32],
313 /// Destination, `[oh, ow, ch]`.
314 y: &'a mut [f32],
315 },
316 /// Resampling of the two spatial axes in `NHWC`, up or down.
317 Resize {
318 /// Channels.
319 ch: usize,
320 /// Input height.
321 h: usize,
322 /// Input width.
323 w: usize,
324 /// Output height.
325 oh: usize,
326 /// Output width.
327 ow: usize,
328 /// How a value is taken once the source position is known.
329 sample: Sample,
330 /// How an output position maps back to a source position.
331 coord: Coord,
332 /// Source, `[h, w, ch]`.
333 x: &'a [f32],
334 /// Destination, `[oh, ow, ch]`.
335 y: &'a mut [f32],
336 },
337 /// Element-wise sum, accumulated into the first operand.
338 Add {
339 /// Accumulator.
340 x: &'a mut [f32],
341 /// Addend.
342 y: &'a [f32],
343 },
344}
345
346/// Runs one unit of work on the given instruction set.
347///
348/// This is the crate's only dispatch point and its only `unsafe` token.
349pub fn run(cpu: Cpu, task: Task<'_>) {
350 match cpu {
351 #[cfg(target_arch = "x86_64")]
352 Cpu::Avx2Fma => {
353 // The single sanctioned `unsafe` in this crate. Calling a
354 // `#[target_feature]` function from an unfeatured context requires
355 // it even though the body is entirely safe, and `Cpu::detect` has
356 // already established that this machine has AVX2 and FMA.
357 #[allow(unsafe_code)]
358 unsafe { dispatch_avx2_fma(task) }
359 },
360 #[cfg(not(target_arch = "x86_64"))]
361 Cpu::Avx2Fma => dispatch_baseline(task),
362 Cpu::Baseline => dispatch_baseline(task),
363 }
364}
365
366/// The kernel set compiled for AVX2 and FMA.
367///
368/// Nothing here is `unsafe`. A `#[target_feature]` function may call another
369/// carrying the same features without one, so this frame names the specialised
370/// wrappers below and the token stays at the single boundary in [`run`].
371///
372/// Each kernel keeps its own function rather than being inlined into this one.
373/// That is not a stylistic choice: folding the whole set into one body costs
374/// about three quarters of the matrix throughput, because the register
375/// allocator then has the entire dispatch to satisfy and gives up on holding
376/// the accumulator tile in vector registers.
377#[cfg(target_arch = "x86_64")]
378#[target_feature(enable = "avx2,fma")]
379fn dispatch_avx2_fma(task: Task<'_>) {
380 match task {
381 Task::Gemm { m, n, k, a, b, c, bias, scratch } =>
382 gemm_tf(m, n, k, a, b, c, bias, scratch),
383 Task::MatVecT { n, k, a, bt, c } =>
384 matvec_t_tf(n, k, a, bt, c),
385 Task::Im2Col { ch, h, w, kh, kw, sy, sx, pt, pl, oh, ow, x, out } =>
386 im2col_tf(ch, h, w, kh, kw, sy, sx, pt, pl, oh, ow, x, out),
387 Task::Depthwise { ch, h, w, kh, kw, sy, sx, pt, pl, oh, ow, x, wt, bias, y } =>
388 depthwise_tf(ch, h, w, kh, kw, sy, sx, pt, pl, oh, ow, x, wt, bias, y),
389 Task::Scale { ch, x, scale, bias } =>
390 scale_bias_tf(ch, x, scale, bias),
391 Task::PRelu { ch, x, slope } =>
392 prelu_tf(ch, x, slope),
393 Task::Relu { x } =>
394 relu_tf(x),
395 Task::Leaky { x, slope } =>
396 leaky_tf(x, slope),
397 Task::Sigmoid { x } =>
398 sigmoid_tf(x),
399 Task::MaxPool { ch, h, w, kh, kw, sy, sx, pt, pl, oh, ow, x, y } =>
400 maxpool_tf(ch, h, w, kh, kw, sy, sx, pt, pl, oh, ow, x, y),
401 Task::Resize { ch, h, w, oh, ow, sample, coord, x, y } =>
402 resize_tf(ch, h, w, oh, ow, sample, coord, x, y),
403 Task::Add { x, y } =>
404 add_tf(x, y),
405 }
406}
407
408/// The kernel set compiled for the baseline target, where `mul_add` is a
409/// library call and must not be reached.
410fn dispatch_baseline(task: Task<'_>) {
411 match task {
412 Task::Gemm { m, n, k, a, b, c, bias, scratch } =>
413 gemm::<false>(m, n, k, a, b, c, bias, scratch),
414 Task::MatVecT { n, k, a, bt, c } =>
415 matvec_t::<false>(n, k, a, bt, c),
416 Task::Im2Col { ch, h, w, kh, kw, sy, sx, pt, pl, oh, ow, x, out } =>
417 im2col(ch, h, w, kh, kw, sy, sx, pt, pl, oh, ow, x, out),
418 Task::Depthwise { ch, h, w, kh, kw, sy, sx, pt, pl, oh, ow, x, wt, bias, y } =>
419 depthwise::<false>(ch, h, w, kh, kw, sy, sx, pt, pl, oh, ow, x, wt, bias, y),
420 Task::Scale { ch, x, scale, bias } =>
421 scale_bias::<false>(ch, x, scale, bias),
422 Task::PRelu { ch, x, slope } =>
423 prelu(ch, x, slope),
424 Task::Relu { x } =>
425 relu(x),
426 Task::Leaky { x, slope } =>
427 leaky(x, slope),
428 Task::Sigmoid { x } =>
429 sigmoid(x),
430 Task::MaxPool { ch, h, w, kh, kw, sy, sx, pt, pl, oh, ow, x, y } =>
431 maxpool(ch, h, w, kh, kw, sy, sx, pt, pl, oh, ow, x, y),
432 Task::Resize { ch, h, w, oh, ow, sample, coord, x, y } =>
433 resize(ch, h, w, oh, ow, sample, coord, x, y),
434 Task::Add { x, y } =>
435 add(x, y),
436 }
437}
438
439/// The blocked matrix product, compiled for AVX2 and FMA.
440#[cfg(target_arch = "x86_64")]
441#[target_feature(enable = "avx2,fma")]
442fn gemm_tf(
443 m: usize,
444 n: usize,
445 k: usize,
446 a: &[f32],
447 b: &[f32],
448 c: &mut [f32],
449 bias: Option<&[f32]>,
450 scratch: &mut Scratch,
451) {
452 gemm::<true>(m, n, k, a, b, c, bias, scratch)
453}
454
455/// The matrix--vector product, compiled for AVX2 and FMA.
456#[cfg(target_arch = "x86_64")]
457#[target_feature(enable = "avx2,fma")]
458fn matvec_t_tf(n: usize, k: usize, a: &[f32], bt: &[f32], c: &mut [f32]) {
459 matvec_t::<true>(n, k, a, bt, c)
460}
461
462/// The patch gather, compiled for AVX2.
463#[cfg(target_arch = "x86_64")]
464#[target_feature(enable = "avx2,fma")]
465#[allow(clippy::too_many_arguments)]
466fn im2col_tf(
467 ch: usize,
468 h: usize,
469 w: usize,
470 kh: usize,
471 kw: usize,
472 sy: usize,
473 sx: usize,
474 pt: usize,
475 pl: usize,
476 oh: usize,
477 ow: usize,
478 x: &[f32],
479 out: &mut [f32],
480) {
481 im2col(ch, h, w, kh, kw, sy, sx, pt, pl, oh, ow, x, out)
482}
483
484/// The depthwise convolution, compiled for AVX2 and FMA.
485#[cfg(target_arch = "x86_64")]
486#[target_feature(enable = "avx2,fma")]
487#[allow(clippy::too_many_arguments)]
488fn depthwise_tf(
489 ch: usize,
490 h: usize,
491 w: usize,
492 kh: usize,
493 kw: usize,
494 sy: usize,
495 sx: usize,
496 pt: usize,
497 pl: usize,
498 oh: usize,
499 ow: usize,
500 x: &[f32],
501 wt: &[f32],
502 bias: Option<&[f32]>,
503 y: &mut [f32],
504) {
505 depthwise::<true>(ch, h, w, kh, kw, sy, sx, pt, pl, oh, ow, x, wt, bias, y)
506}
507
508/// The per-channel affine map, compiled for AVX2 and FMA.
509#[cfg(target_arch = "x86_64")]
510#[target_feature(enable = "avx2,fma")]
511fn scale_bias_tf(ch: usize, x: &mut [f32], sc: &[f32], bi: &[f32]) {
512 scale_bias::<true>(ch, x, sc, bi)
513}
514
515/// The parametric rectifier, compiled for AVX2.
516#[cfg(target_arch = "x86_64")]
517#[target_feature(enable = "avx2,fma")]
518fn prelu_tf(ch: usize, x: &mut [f32], slope: &[f32]) {
519 prelu(ch, x, slope)
520}
521
522/// The rectifier, compiled for AVX2.
523#[cfg(target_arch = "x86_64")]
524#[target_feature(enable = "avx2,fma")]
525fn relu_tf(x: &mut [f32]) {
526 relu(x)
527}
528
529/// The sigmoid, compiled for AVX2.
530#[cfg(target_arch = "x86_64")]
531#[target_feature(enable = "avx2,fma")]
532fn sigmoid_tf(x: &mut [f32]) {
533 sigmoid(x)
534}
535
536/// The maximum pool, compiled for AVX2.
537#[cfg(target_arch = "x86_64")]
538#[target_feature(enable = "avx2,fma")]
539#[allow(clippy::too_many_arguments)]
540fn maxpool_tf(
541 ch: usize,
542 h: usize,
543 w: usize,
544 kh: usize,
545 kw: usize,
546 sy: usize,
547 sx: usize,
548 pt: usize,
549 pl: usize,
550 oh: usize,
551 ow: usize,
552 x: &[f32],
553 y: &mut [f32],
554) {
555 maxpool(ch, h, w, kh, kw, sy, sx, pt, pl, oh, ow, x, y)
556}
557
558/// The resampling, compiled for AVX2.
559#[cfg(target_arch = "x86_64")]
560#[target_feature(enable = "avx2,fma")]
561#[allow(clippy::too_many_arguments)]
562fn resize_tf(
563 ch: usize,
564 h: usize,
565 w: usize,
566 oh: usize,
567 ow: usize,
568 sample: Sample,
569 coord: Coord,
570 x: &[f32],
571 y: &mut [f32],
572) {
573 resize(ch, h, w, oh, ow, sample, coord, x, y)
574}
575
576/// The leaky rectifier, compiled for AVX2.
577#[cfg(target_arch = "x86_64")]
578#[target_feature(enable = "avx2,fma")]
579fn leaky_tf(x: &mut [f32], slope: f32) {
580 leaky(x, slope)
581}
582
583/// The element-wise sum, compiled for AVX2.
584#[cfg(target_arch = "x86_64")]
585#[target_feature(enable = "avx2,fma")]
586fn add_tf(x: &mut [f32], y: &[f32]) {
587 add(x, y)
588}
589
590/// Fused multiply-add in whichever form the compiled path may use.
591///
592/// Under `FMA` this is `f32::mul_add`, which lowers to a single instruction and
593/// rounds once. Without it, the plain form, because `mul_add` would become a
594/// library call.
595#[inline(always)]
596fn fma<const FMA: bool>(a: f32, b: f32, c: f32) -> f32 {
597 if FMA {
598 a.mul_add(b, c)
599 } else {
600 a * b + c
601 }
602}
603
604/// Packs a `kc × nc` block of `b` into `NR`-wide panels, zero padded.
605#[inline(always)]
606fn pack_b(
607 b: &[f32],
608 ldb: usize,
609 p0: usize,
610 kc: usize,
611 j0: usize,
612 nc: usize,
613 out: &mut [f32],
614) {
615 let panels = (nc + NR - 1) / NR;
616 for p in 0..panels {
617 let jbase = j0 + p * NR;
618 let nv = core::cmp::min(NR, nc - p * NR);
619 let dst = &mut out[p * kc * NR..(p + 1) * kc * NR];
620 for kk in 0..kc {
621 let src = &b[(p0 + kk) * ldb + jbase..(p0 + kk) * ldb + jbase + nv];
622 let slot = &mut dst[kk * NR..kk * NR + NR];
623 for j in 0..nv {
624 slot[j] = src[j];
625 }
626 for j in nv..NR {
627 slot[j] = 0.0;
628 }
629 }
630 }
631}
632
633/// Packs an `mc × kc` block of `a` into `MR`-tall panels, zero padded.
634#[inline(always)]
635fn pack_a(
636 a: &[f32],
637 lda: usize,
638 i0: usize,
639 mc: usize,
640 p0: usize,
641 kc: usize,
642 out: &mut [f32],
643) {
644 let panels = (mc + MR - 1) / MR;
645 for p in 0..panels {
646 let ibase = i0 + p * MR;
647 let mv = core::cmp::min(MR, mc - p * MR);
648 let dst = &mut out[p * kc * MR..(p + 1) * kc * MR];
649 for i in 0..mv {
650 let src = &a[(ibase + i) * lda + p0..(ibase + i) * lda + p0 + kc];
651 for kk in 0..kc {
652 dst[kk * MR + i] = src[kk];
653 }
654 }
655 for i in mv..MR {
656 for kk in 0..kc {
657 dst[kk * MR + i] = 0.0;
658 }
659 }
660 }
661}
662
663/// The register-tile microkernel, `c[0..MR, 0..NR] += ap · bp`.
664///
665/// The accumulator is a fixed-size array so that the code generator can hold it
666/// in vector registers, and `chunks_exact` proves the inner extents without an
667/// index that could fail.
668#[inline(always)]
669fn micro<const FMA: bool>(
670 kc: usize,
671 ap: &[f32],
672 bp: &[f32],
673 c: &mut [f32],
674 ldc: usize,
675 i0: usize,
676 j0: usize,
677 mv: usize,
678 nv: usize,
679) {
680 let mut acc = [[0.0f32; NR]; MR];
681 let asub = &ap[..kc * MR];
682 let bsub = &bp[..kc * NR];
683 for (achunk, bchunk) in asub.chunks_exact(MR).zip(bsub.chunks_exact(NR)) {
684 for i in 0..MR {
685 let av = achunk[i];
686 for j in 0..NR {
687 acc[i][j] = fma::<FMA>(av, bchunk[j], acc[i][j]);
688 }
689 }
690 }
691 for i in 0..mv {
692 let base = (i0 + i) * ldc + j0;
693 let row = &mut c[base..base + nv];
694 let src = &acc[i];
695 for j in 0..nv {
696 row[j] += src[j];
697 }
698 }
699}
700
701/// Blocked, packed general matrix product.
702#[inline(always)]
703fn gemm<const FMA: bool>(
704 m: usize,
705 n: usize,
706 k: usize,
707 a: &[f32],
708 b: &[f32],
709 c: &mut [f32],
710 bias: Option<&[f32]>,
711 scratch: &mut Scratch,
712) {
713 match bias {
714 Some(bs) => {
715 for row in c.chunks_exact_mut(n) {
716 row.copy_from_slice(&bs[..n]);
717 }
718 },
719 None => {
720 for v in c.iter_mut() {
721 *v = 0.0;
722 }
723 },
724 }
725 let mcb = core::cmp::min(MC, m);
726 let kcb = core::cmp::min(KC, k);
727 let ncb = core::cmp::min(NC, n);
728 scratch.ensure(
729 ((mcb + MR - 1) / MR) * kcb * MR,
730 ((ncb + NR - 1) / NR) * kcb * NR,
731 );
732 let mut jc = 0;
733 while jc < n {
734 let nn = core::cmp::min(ncb, n - jc);
735 let mut pc = 0;
736 while pc < k {
737 let kk = core::cmp::min(kcb, k - pc);
738 pack_b(b, n, pc, kk, jc, nn, &mut scratch.bp);
739 let mut ic = 0;
740 while ic < m {
741 let mm = core::cmp::min(mcb, m - ic);
742 pack_a(a, k, ic, mm, pc, kk, &mut scratch.ap);
743 let jpan = (nn + NR - 1) / NR;
744 let ipan = (mm + MR - 1) / MR;
745 for jp in 0..jpan {
746 let nv = core::cmp::min(NR, nn - jp * NR);
747 let bpan = &scratch.bp[jp * kk * NR..(jp + 1) * kk * NR];
748 for ip in 0..ipan {
749 let mv = core::cmp::min(MR, mm - ip * MR);
750 let apan = &scratch.ap[ip * kk * MR..(ip + 1) * kk * MR];
751 micro::<FMA>(
752 kk,
753 apan,
754 bpan,
755 c,
756 n,
757 ic + ip * MR,
758 jc + jp * NR,
759 mv,
760 nv,
761 );
762 }
763 }
764 ic += mcb;
765 }
766 pc += kcb;
767 }
768 jc += ncb;
769 }
770}
771
772/// Matrix--vector product against a transposed weight matrix.
773///
774/// Eight partial accumulators break the latency chain of the multiply-add unit.
775/// This layer is bandwidth bound rather than compute bound, so the win over the
776/// general kernel comes from reading each weight exactly once.
777#[inline(always)]
778fn matvec_t<const FMA: bool>(n: usize, k: usize, a: &[f32], bt: &[f32], c: &mut [f32]) {
779 for j in 0..n {
780 let w = &bt[j * k..j * k + k];
781 let mut s = [0.0f32; 8];
782 let mut it_a = a.chunks_exact(8);
783 let mut it_w = w.chunks_exact(8);
784 for (ca, cw) in it_a.by_ref().zip(it_w.by_ref()) {
785 for l in 0..8 {
786 s[l] = fma::<FMA>(ca[l], cw[l], s[l]);
787 }
788 }
789 let mut tail = 0.0f32;
790 for (x, y) in it_a.remainder().iter().zip(it_w.remainder().iter()) {
791 tail = fma::<FMA>(*x, *y, tail);
792 }
793 c[j] = ((s[0] + s[1]) + (s[2] + s[3])) + ((s[4] + s[5]) + (s[6] + s[7])) + tail;
794 }
795}
796
797/// Gathers convolution patches out of a channels-last plane.
798///
799/// Only a kernel larger than one by one needs this. A one by one convolution in
800/// `NHWC` is already the `[m = h·w, k = ch]` matrix the product consumes, which
801/// is why twenty-six of the twenty-seven convolutions in a MobileFaceNet-shaped
802/// embedder skip it entirely.
803#[inline(always)]
804fn im2col(
805 ch: usize,
806 h: usize,
807 w: usize,
808 kh: usize,
809 kw: usize,
810 sy: usize,
811 sx: usize,
812 pt: usize,
813 pl: usize,
814 oh: usize,
815 ow: usize,
816 x: &[f32],
817 out: &mut [f32],
818) {
819 let kk = kh * kw * ch;
820 for oy in 0..oh {
821 let iy0 = (oy * sy) as isize - pt as isize;
822 for ox in 0..ow {
823 let ix0 = (ox * sx) as isize - pl as isize;
824 let row = &mut out[(oy * ow + ox) * kk..(oy * ow + ox) * kk + kk];
825 for ky in 0..kh {
826 let iy = iy0 + ky as isize;
827 for kx in 0..kw {
828 let ix = ix0 + kx as isize;
829 let dst = &mut row[(ky * kw + kx) * ch..(ky * kw + kx) * ch + ch];
830 if iy < 0 || iy as usize >= h || ix < 0 || ix as usize >= w {
831 for v in dst.iter_mut() {
832 *v = 0.0;
833 }
834 } else {
835 let base = (iy as usize * w + ix as usize) * ch;
836 dst.copy_from_slice(&x[base..base + ch]);
837 }
838 }
839 }
840 }
841 }
842}
843
844/// Depthwise convolution in `NHWC`.
845///
846/// The channel loop is innermost, which makes it unit stride in the activation,
847/// the weights and the output at once. The same arithmetic written channels-first
848/// runs about nine times slower, because nothing there vectorises.
849#[inline(always)]
850fn depthwise<const FMA: bool>(
851 ch: usize,
852 h: usize,
853 w: usize,
854 kh: usize,
855 kw: usize,
856 sy: usize,
857 sx: usize,
858 pt: usize,
859 pl: usize,
860 oh: usize,
861 ow: usize,
862 x: &[f32],
863 wt: &[f32],
864 bias: Option<&[f32]>,
865 y: &mut [f32],
866) {
867 for oy in 0..oh {
868 let iy0 = (oy * sy) as isize - pt as isize;
869 for ox in 0..ow {
870 let ix0 = (ox * sx) as isize - pl as isize;
871 let out = &mut y[(oy * ow + ox) * ch..(oy * ow + ox) * ch + ch];
872 match bias {
873 Some(bs) => out.copy_from_slice(&bs[..ch]),
874 None => {
875 for v in out.iter_mut() {
876 *v = 0.0;
877 }
878 },
879 }
880 for ky in 0..kh {
881 let iy = iy0 + ky as isize;
882 if iy < 0 || iy as usize >= h {
883 continue;
884 }
885 for kx in 0..kw {
886 let ix = ix0 + kx as isize;
887 if ix < 0 || ix as usize >= w {
888 continue;
889 }
890 let base = (iy as usize * w + ix as usize) * ch;
891 let src = &x[base..base + ch];
892 let kv = &wt[(ky * kw + kx) * ch..(ky * kw + kx) * ch + ch];
893 for c in 0..ch {
894 out[c] = fma::<FMA>(src[c], kv[c], out[c]);
895 }
896 }
897 }
898 }
899 }
900}
901
902/// Per-channel affine map over a channels-last buffer.
903#[inline(always)]
904fn scale_bias<const FMA: bool>(ch: usize, x: &mut [f32], sc: &[f32], bi: &[f32]) {
905 let sc = &sc[..ch];
906 let bi = &bi[..ch];
907 for row in x.chunks_exact_mut(ch) {
908 for j in 0..ch {
909 row[j] = fma::<FMA>(sc[j], row[j], bi[j]);
910 }
911 }
912}
913
914/// Parametric rectified linear unit, branchless.
915///
916/// Written as a comparison the loop keeps `v.max(0) + slope·v.min(0)`, because
917/// the obvious `if v >= 0` form does not vectorise and costs an order of
918/// magnitude over a whole network.
919#[inline(always)]
920fn prelu(ch: usize, x: &mut [f32], slope: &[f32]) {
921 let sl = &slope[..ch];
922 for row in x.chunks_exact_mut(ch) {
923 for j in 0..ch {
924 let v = row[j];
925 row[j] = v.max(0.0) + sl[j] * v.min(0.0);
926 }
927 }
928}
929
930/// Rectified linear unit.
931#[inline(always)]
932fn relu(x: &mut [f32]) {
933 for v in x.iter_mut() {
934 *v = v.max(0.0);
935 }
936}
937
938/// Leaky rectified linear unit, written branchlessly so the loop vectorises.
939#[inline(always)]
940fn leaky(x: &mut [f32], slope: f32) {
941 for v in x.iter_mut() {
942 // `max` and `min` split the value into its positive and negative parts,
943 // which costs two instructions and no branch.
944 *v = v.max(0.0) + slope * v.min(0.0);
945 }
946}
947
948/// Logistic sigmoid.
949#[inline(always)]
950fn sigmoid(x: &mut [f32]) {
951 for v in x.iter_mut() {
952 *v = 1.0 / (1.0 + (-*v).exp());
953 }
954}
955
956/// Maximum pool in `NHWC`, over any kernel, stride and padding.
957///
958/// A padded position contributes nothing at all rather than a zero. Zero is not
959/// the identity of a maximum: after a leaky rectifier a whole window can be
960/// negative, and a padding zero would then be the answer.
961#[inline(always)]
962#[allow(clippy::too_many_arguments)]
963fn maxpool(
964 ch: usize,
965 h: usize,
966 w: usize,
967 kh: usize,
968 kw: usize,
969 sy: usize,
970 sx: usize,
971 pt: usize,
972 pl: usize,
973 oh: usize,
974 ow: usize,
975 x: &[f32],
976 y: &mut [f32],
977) {
978 for oy in 0..oh {
979 for ox in 0..ow {
980 let o = (oy * ow + ox) * ch;
981 let out = &mut y[o..o + ch];
982 for v in out.iter_mut() {
983 *v = f32::NEG_INFINITY;
984 }
985 // Where this window starts in the source, before padding is removed.
986 let top = (oy * sy) as isize - pt as isize;
987 let left = (ox * sx) as isize - pl as isize;
988 for ky in 0..kh {
989 let iy = top + ky as isize;
990 if iy < 0 || iy as usize >= h {
991 continue;
992 }
993 for kx in 0..kw {
994 let ix = left + kx as isize;
995 if ix < 0 || ix as usize >= w {
996 continue;
997 }
998 let s = (iy as usize * w + ix as usize) * ch;
999 let src = &x[s..s + ch];
1000 for i in 0..ch {
1001 out[i] = out[i].max(src[i]);
1002 }
1003 }
1004 }
1005 }
1006 }
1007}
1008
1009/// Resampling of the two spatial axes in `NHWC`, up or down.
1010#[inline(always)]
1011#[allow(clippy::too_many_arguments)]
1012fn resize(
1013 ch: usize,
1014 h: usize,
1015 w: usize,
1016 oh: usize,
1017 ow: usize,
1018 sample: Sample,
1019 coord: Coord,
1020 x: &[f32],
1021 y: &mut [f32],
1022) {
1023 let scale_y = oh as f32 / h as f32;
1024 let scale_x = ow as f32 / w as f32;
1025 for oy in 0..oh {
1026 let sy = coord.source(oy, scale_y, oh);
1027 for ox in 0..ow {
1028 let sx = coord.source(ox, scale_x, ow);
1029 let o = (oy * ow + ox) * ch;
1030 match sample {
1031 Sample::Nearest => {
1032 let iy = (sy.floor() as usize).min(h - 1);
1033 let ix = (sx.floor() as usize).min(w - 1);
1034 let s = (iy * w + ix) * ch;
1035 y[o..o + ch].copy_from_slice(&x[s..s + ch]);
1036 },
1037 Sample::Bilinear => {
1038 let y0 = sy.floor().max(0.0) as usize;
1039 let x0 = sx.floor().max(0.0) as usize;
1040 let y1 = (y0 + 1).min(h - 1);
1041 let x1 = (x0 + 1).min(w - 1);
1042 let y0 = y0.min(h - 1);
1043 let x0 = x0.min(w - 1);
1044 let fy = sy - sy.floor();
1045 let fx = sx - sx.floor();
1046 // The four corners, weighted by how far the source position
1047 // sits between them.
1048 let (wa, wb) = ((1.0 - fy) * (1.0 - fx), (1.0 - fy) * fx);
1049 let (wc, wd) = (fy * (1.0 - fx), fy * fx);
1050 let a = (y0 * w + x0) * ch;
1051 let b = (y0 * w + x1) * ch;
1052 let c = (y1 * w + x0) * ch;
1053 let d = (y1 * w + x1) * ch;
1054 for i in 0..ch {
1055 y[o + i] = wa * x[a + i] + wb * x[b + i]
1056 + wc * x[c + i] + wd * x[d + i];
1057 }
1058 },
1059 }
1060 }
1061 }
1062}
1063
1064/// Element-wise sum, accumulated into the first operand.
1065#[inline(always)]
1066fn add(x: &mut [f32], y: &[f32]) {
1067 for (a, b) in x.iter_mut().zip(y.iter()) {
1068 *a += *b;
1069 }
1070}
1071
1072/// Cuts an activation along its channels, into runs of the given widths.
1073///
1074/// Channels are the innermost axis, so each part is a strided gather rather than
1075/// a slice; this is the price the channels-last layout charges for an operator
1076/// that names an axis, and it is paid a handful of times per model.
1077pub fn split_channels(t: &Tensor, widths: &[usize]) -> Outcome<Vec<Tensor>> {
1078 let (n, h, w, c) = res!(t.nhwc());
1079 let total = widths.iter().sum::<usize>();
1080 if total != c {
1081 return Err(err!(
1082 "A split of {:?} covers {} channels, but the activation has {}.",
1083 widths, total, c;
1084 Invalid, Input, Mismatch));
1085 }
1086 let rows = n * h * w;
1087 let mut parts = Vec::with_capacity(widths.len());
1088 let mut base = 0;
1089 for width in widths {
1090 let mut out = vec![0.0f32; rows * width];
1091 for r in 0..rows {
1092 let src = r * c + base;
1093 out[r * width..(r + 1) * width].copy_from_slice(&t.data[src..src + width]);
1094 }
1095 parts.push(res!(Tensor::new(vec![n, h, w, *width], out)));
1096 base += width;
1097 }
1098 Ok(parts)
1099}
1100
1101/// Joins activations along their channels, in the order given.
1102pub fn concat_channels(parts: &[&Tensor]) -> Outcome<Tensor> {
1103 let first = match parts.first() {
1104 Some(t) => *t,
1105 None => return Err(err!("A concatenation was given no operands."; Invalid, Input, Missing)),
1106 };
1107 let (n, h, w, _) = res!(first.nhwc());
1108 let mut widths = Vec::with_capacity(parts.len());
1109 for p in parts {
1110 let (pn, ph, pw, pc) = res!(p.nhwc());
1111 if (pn, ph, pw) != (n, h, w) {
1112 return Err(err!(
1113 "A concatenation of {:?} and {:?} disagrees away from the channels.",
1114 first.dims, p.dims;
1115 Invalid, Input, Mismatch));
1116 }
1117 widths.push(pc);
1118 }
1119 let c = widths.iter().sum::<usize>();
1120 let rows = n * h * w;
1121 let mut out = vec![0.0f32; rows * c];
1122 let mut base = 0;
1123 for (p, width) in parts.iter().zip(widths.iter()) {
1124 for r in 0..rows {
1125 let dst = r * c + base;
1126 out[dst..dst + width].copy_from_slice(&p.data[r * width..(r + 1) * width]);
1127 }
1128 base += width;
1129 }
1130 Tensor::new(vec![n, h, w, c], out)
1131}
1132
1133/// Interleaves the channels of a grouped activation.
1134///
1135/// This is ShuffleNet's channel shuffle. A model writes it as a reshape into
1136/// `[n, g, c/g, h, w]`, a transpose of the two new axes and a reshape back,
1137/// which in a channels-first layout moves every value; here the spatial axes are
1138/// untouched and it is a permutation of the innermost axis alone. Output channel
1139/// `i` takes input channel `(i mod g)·(c/g) + i div g`.
1140pub fn shuffle_channels(t: &Tensor, groups: usize) -> Outcome<Tensor> {
1141 let (n, h, w, c) = res!(t.nhwc());
1142 if groups == 0 || c % groups != 0 {
1143 return Err(err!(
1144 "A shuffle into {} groups does not divide {} channels.", groups, c;
1145 Invalid, Input, Mismatch));
1146 }
1147 let per = c / groups;
1148 // The gather, worked out once and reused for every position.
1149 let take = (0..c).map(|i| (i % groups) * per + i / groups).collect::<Vec<_>>();
1150 let rows = n * h * w;
1151 let mut out = vec![0.0f32; t.len()];
1152 for r in 0..rows {
1153 let (src, dst) = (r * c, r * c);
1154 for (i, from) in take.iter().enumerate() {
1155 out[dst + i] = t.data[src + from];
1156 }
1157 }
1158 Tensor::new(t.dims.clone(), out)
1159}
1160
1161/// Rewrites an `[n, h, w, c]` activation as the `[n, c·h·w]` row an ONNX
1162/// `Flatten` produces, which is channels-first order.
1163///
1164/// This is the one place the channels-last layout has to be undone, and it is
1165/// cheap: one gather per embedding, not per layer.
1166pub fn flatten_nchw(t: &Tensor) -> Outcome<Tensor> {
1167 let (n, h, w, c) = res!(t.nhwc());
1168 let plane = h * w;
1169 let mut out = vec![0.0f32; t.len()];
1170 for bi in 0..n {
1171 for ci in 0..c {
1172 let dst = (bi * c + ci) * plane;
1173 for p in 0..plane {
1174 out[dst + p] = t.data[(bi * plane + p) * c + ci];
1175 }
1176 }
1177 }
1178 Tensor::new(vec![n, c * plane], out)
1179}
1180
1181#[cfg(test)]
1182mod tests {
1183 use super::*;
1184
1185 /// Reference product in `f64`, so the comparison is against something the
1186 /// kernel does not share code with.
1187 fn reference(m: usize, n: usize, k: usize, a: &[f32], b: &[f32]) -> Vec<f32> {
1188 let mut c = vec![0.0f32; m * n];
1189 for i in 0..m {
1190 for j in 0..n {
1191 let mut s = 0.0f64;
1192 for p in 0..k {
1193 s += a[i * k + p] as f64 * b[p * n + j] as f64;
1194 }
1195 c[i * n + j] = s as f32;
1196 }
1197 }
1198 c
1199 }
1200
1201 /// A cheap reproducible generator, so a test needs no dependency.
1202 fn fill(n: usize, seed: u64) -> Vec<f32> {
1203 let mut s = seed;
1204 let mut v = Vec::with_capacity(n);
1205 for _ in 0..n {
1206 s = s.wrapping_mul(6364136223846793005).wrapping_add(1442695040888963407);
1207 v.push(((s >> 40) as f32 / 8_388_608.0) - 1.0);
1208 }
1209 v
1210 }
1211
1212 #[test]
1213 fn gemm_matches_a_wider_reference() -> Outcome<()> {
1214 for &(m, n, k) in &[(1, 1, 1), (6, 16, 5), (7, 17, 33), (196, 512, 128), (49, 71, 200)] {
1215 let a = fill(m * k, 1);
1216 let b = fill(k * n, 2);
1217 let want = reference(m, n, k, &a, &b);
1218 for cpu in [Cpu::Baseline, Cpu::detect()] {
1219 let mut c = vec![0.0f32; m * n];
1220 let mut s = Scratch::new();
1221 run(cpu, Task::Gemm {
1222 m, n, k,
1223 a: &a,
1224 b: &b,
1225 c: &mut c,
1226 bias: None,
1227 scratch: &mut s,
1228 });
1229 for i in 0..m * n {
1230 let d = (c[i] - want[i]).abs();
1231 if d > 1e-3 {
1232 return Err(err!(
1233 "On {}x{}x{} under {:?}, element {} read {} against {}.",
1234 m, n, k, cpu, i, c[i], want[i];
1235 Invalid, Mismatch));
1236 }
1237 }
1238 }
1239 }
1240 Ok(())
1241 }
1242
1243 #[test]
1244 fn both_paths_agree() -> Outcome<()> {
1245 let (m, n, k) = (37, 53, 71);
1246 let a = fill(m * k, 11);
1247 let b = fill(k * n, 12);
1248 let bias = fill(n, 13);
1249 let mut c0 = vec![0.0f32; m * n];
1250 let mut c1 = vec![0.0f32; m * n];
1251 let mut s = Scratch::new();
1252 run(Cpu::Baseline, Task::Gemm {
1253 m, n, k, a: &a, b: &b, c: &mut c0, bias: Some(&bias), scratch: &mut s });
1254 run(Cpu::detect(), Task::Gemm {
1255 m, n, k, a: &a, b: &b, c: &mut c1, bias: Some(&bias), scratch: &mut s });
1256 for i in 0..m * n {
1257 if (c0[i] - c1[i]).abs() > 1e-4 {
1258 return Err(err!(
1259 "The baseline and dispatched paths disagree at {}: {} against {}.",
1260 i, c0[i], c1[i];
1261 Invalid, Mismatch));
1262 }
1263 }
1264 Ok(())
1265 }
1266
1267 #[test]
1268 fn matvec_matches_the_general_kernel() -> Outcome<()> {
1269 let (n, k) = (128, 501);
1270 let a = fill(k, 21);
1271 let bt = fill(n * k, 22);
1272 let mut want = vec![0.0f32; n];
1273 for j in 0..n {
1274 let mut s = 0.0f64;
1275 for p in 0..k {
1276 s += a[p] as f64 * bt[j * k + p] as f64;
1277 }
1278 want[j] = s as f32;
1279 }
1280 for cpu in [Cpu::Baseline, Cpu::detect()] {
1281 let mut c = vec![0.0f32; n];
1282 run(cpu, Task::MatVecT { n, k, a: &a, bt: &bt, c: &mut c });
1283 for j in 0..n {
1284 if (c[j] - want[j]).abs() > 1e-3 {
1285 return Err(err!(
1286 "Under {:?}, output {} read {} against {}.", cpu, j, c[j], want[j];
1287 Invalid, Mismatch));
1288 }
1289 }
1290 }
1291 Ok(())
1292 }
1293
1294 #[test]
1295 fn depthwise_matches_a_direct_loop() -> Outcome<()> {
1296 let (ch, h, w) = (5, 7, 9);
1297 let x = fill(h * w * ch, 31);
1298 let wt = fill(9 * ch, 32);
1299 for stride in [1usize, 2] {
1300 let oh = (h + 2 - 3) / stride + 1;
1301 let ow = (w + 2 - 3) / stride + 1;
1302 let mut want = vec![0.0f32; oh * ow * ch];
1303 for oy in 0..oh {
1304 for ox in 0..ow {
1305 for c in 0..ch {
1306 let mut s = 0.0f64;
1307 for ky in 0..3isize {
1308 for kx in 0..3isize {
1309 let iy = (oy * stride) as isize - 1 + ky;
1310 let ix = (ox * stride) as isize - 1 + kx;
1311 if iy < 0 || iy as usize >= h || ix < 0 || ix as usize >= w {
1312 continue;
1313 }
1314 s += x[(iy as usize * w + ix as usize) * ch + c] as f64
1315 * wt[(ky as usize * 3 + kx as usize) * ch + c] as f64;
1316 }
1317 }
1318 want[(oy * ow + ox) * ch + c] = s as f32;
1319 }
1320 }
1321 }
1322 for cpu in [Cpu::Baseline, Cpu::detect()] {
1323 let mut y = vec![0.0f32; oh * ow * ch];
1324 run(cpu, Task::Depthwise {
1325 ch, h, w,
1326 kh: 3,
1327 kw: 3,
1328 sy: stride,
1329 sx: stride,
1330 pt: 1,
1331 pl: 1,
1332 oh, ow,
1333 x: &x,
1334 wt: &wt,
1335 bias: None,
1336 y: &mut y,
1337 });
1338 for i in 0..y.len() {
1339 if (y[i] - want[i]).abs() > 1e-5 {
1340 return Err(err!(
1341 "Depthwise stride {} under {:?} differs at {}: {} against {}.",
1342 stride, cpu, i, y[i], want[i];
1343 Invalid, Mismatch));
1344 }
1345 }
1346 }
1347 }
1348 Ok(())
1349 }
1350
1351 #[test]
1352 fn a_padded_pool_is_not_won_by_its_padding() -> Outcome<()> {
1353 // Every sample is negative, which is what a leaky rectifier hands on.
1354 // A padding zero would beat all of them at every edge.
1355 let (ch, h, w) = (1, 3, 3);
1356 let x = vec![-9.0, -8.0, -7.0, -6.0, -5.0, -4.0, -3.0, -2.0, -1.0];
1357 let (oh, ow) = (2, 2);
1358 let mut y = vec![0.0f32; oh * ow * ch];
1359 run(Cpu::detect(), Task::MaxPool {
1360 ch, h, w,
1361 kh: 3, kw: 3, sy: 2, sx: 2, pt: 1, pl: 1,
1362 oh, ow,
1363 x: &x, y: &mut y,
1364 });
1365 // Each window covers a corner quadrant of the plane, the rest being pad.
1366 let want = [-5.0, -4.0, -2.0, -1.0];
1367 for (i, w) in want.iter().enumerate() {
1368 req!(y[i], *w, "The pool took the padding rather than a sample.");
1369 }
1370 Ok(())
1371 }
1372
1373 #[test]
1374 fn a_bilinear_doubling_puts_the_quarters_where_they_belong() -> Outcome<()> {
1375 // Two samples, 0 and 4, doubled under the half-pixel rule. The output
1376 // centres fall at source -0.25, 0.25, 0.75 and 1.25, and the first is
1377 // clamped to zero, so the values are 0, 1, 3, 4.
1378 let (ch, h, w) = (1, 1, 2);
1379 let x = vec![0.0, 4.0];
1380 let mut y = vec![0.0f32; 4];
1381 run(Cpu::detect(), Task::Resize {
1382 ch, h, w, oh: 1, ow: 4,
1383 sample: Sample::Bilinear, coord: Coord::HalfPixel,
1384 x: &x, y: &mut y,
1385 });
1386 let want = [0.0, 1.0, 3.0, 4.0];
1387 for (i, v) in want.iter().enumerate() {
1388 let close = (y[i] - *v).abs() < 1e-5;
1389 req!(close, true, "Bilinear at {} gave {}, wanted {}.", i, y[i], v);
1390 }
1391
1392 // The asymmetric rule reads the same source at every output, which is
1393 // what separates the two conventions and is worth failing on.
1394 let mut z = vec![0.0f32; 4];
1395 run(Cpu::detect(), Task::Resize {
1396 ch, h, w, oh: 1, ow: 4,
1397 sample: Sample::Bilinear, coord: Coord::Asymmetric,
1398 x: &x, y: &mut z,
1399 });
1400 let differs = z.iter().zip(y.iter()).any(|(a, b)| (a - b).abs() > 1e-5);
1401 req!(differs, true, "The two coordinate rules gave the same answer.");
1402 Ok(())
1403 }
1404
1405 #[test]
1406 fn a_resize_to_the_same_size_changes_nothing() -> Outcome<()> {
1407 let (ch, h, w) = (2, 3, 5);
1408 let x = fill(h * w * ch, 17);
1409 for sample in [Sample::Nearest, Sample::Bilinear] {
1410 for coord in [Coord::Asymmetric, Coord::HalfPixel] {
1411 let mut y = vec![0.0f32; x.len()];
1412 run(Cpu::detect(), Task::Resize {
1413 ch, h, w, oh: h, ow: w, sample, coord,
1414 x: &x, y: &mut y,
1415 });
1416 for i in 0..x.len() {
1417 let same = (y[i] - x[i]).abs() < 1e-5;
1418 req!(same, true,
1419 "{:?}/{:?} moved value {} from {} to {}.",
1420 sample, coord, i, x[i], y[i]);
1421 }
1422 }
1423 }
1424 Ok(())
1425 }
1426
1427 #[test]
1428 fn pooling_and_doubling_invert_a_constant() -> Outcome<()> {
1429 let (ch, h, w) = (3, 4, 6);
1430 let x = fill(h * w * ch, 41);
1431 let mut pooled = vec![0.0f32; (h / 2) * (w / 2) * ch];
1432 run(Cpu::detect(), Task::MaxPool {
1433 ch, h, w,
1434 kh: 2, kw: 2, sy: 2, sx: 2, pt: 0, pl: 0,
1435 oh: h / 2, ow: w / 2,
1436 x: &x, y: &mut pooled,
1437 });
1438 let mut back = vec![0.0f32; h * w * ch];
1439 run(Cpu::detect(), Task::Resize {
1440 ch, h: h / 2, w: w / 2, oh: h, ow: w,
1441 sample: Sample::Nearest, coord: Coord::Asymmetric,
1442 x: &pooled, y: &mut back,
1443 });
1444 // Each pooled maximum is repeated over the two by two block it came from.
1445 for oy in 0..h / 2 {
1446 for ox in 0..w / 2 {
1447 for c in 0..ch {
1448 let want = pooled[(oy * (w / 2) + ox) * ch + c];
1449 for dy in 0..2 {
1450 for dx in 0..2 {
1451 let got = back[((oy * 2 + dy) * w + ox * 2 + dx) * ch + c];
1452 req!(got, want);
1453 }
1454 }
1455 }
1456 }
1457 }
1458 Ok(())
1459 }
1460}