Skip to main content

tract_linalg/x86_64_fma/
max.rs

1reduce_impl_wrap!(
2    f32,
3    x86_64_fma_max_f32_32n,
4    32,
5    8,
6    (),
7    f32::MIN,
8    #[inline(never)]
9    fn run(buf: &[f32], _: ()) -> f32 {
10        assert!(buf.len() % 32 == 0);
11        assert!(buf.len() > 0);
12        unsafe { x86_64_fma_max_f32_32n_run(buf) }
13    },
14    #[inline(never)]
15    fn reduce_two(a: f32, b: f32) -> f32 {
16        a.max(b)
17    }
18);
19
20#[target_feature(enable = "avx")]
21unsafe fn x86_64_fma_max_f32_32n_run(buf: &[f32]) -> f32 {
22    unsafe {
23        let len = buf.len();
24        let ptr = buf.as_ptr();
25        let mut acc = f32::MIN;
26        std::arch::asm!("
27            // reg-source vbroadcastss needs avx2; this kernel must stay avx-safe
28            vpermilps xmm0, xmm0, 0
29            vinsertf128 ymm0, ymm0, xmm0, 1
30            vmovaps ymm1, ymm0
31            vmovaps ymm2, ymm0
32            vmovaps ymm3, ymm0
33            2:
34                vmovaps ymm4, [{ptr}]
35                vmovaps ymm5, [{ptr} + 32]
36                vmovaps ymm6, [{ptr} + 64]
37                vmovaps ymm7, [{ptr} + 96]
38                vmaxps ymm0, ymm0, ymm4
39                vmaxps ymm1, ymm1, ymm5
40                vmaxps ymm2, ymm2, ymm6
41                vmaxps ymm3, ymm3, ymm7
42                add {ptr}, 128
43                sub {len}, 32
44                jnz 2b
45            vmaxps ymm0, ymm0, ymm1
46            vmaxps ymm2, ymm2, ymm3
47            vmaxps ymm0, ymm0, ymm2
48            vperm2f128 ymm1, ymm0, ymm0, 1      // copy second half (4xf32) of ymm0 to ymm1
49            vmaxps xmm0, xmm0, xmm1             // xmm0 contains 4 values to max
50            vpermilps xmm1, xmm0, 2 + (3 << 2)  // second 2x32 bit half moved to top
51            vmaxps xmm0, xmm0, xmm1             // xmm0 containes 2 values
52            vpermilps xmm1, xmm0, 1             // second f32 to top
53            vmaxps xmm0, xmm0, xmm1
54            ",
55        len = inout(reg) len => _,
56        ptr = inout(reg) ptr => _,
57        inout("ymm0") acc,
58        out("ymm1") _, out("ymm2") _, out("ymm3") _,
59        out("ymm4") _, out("ymm5") _, out("ymm6") _, out("ymm7") _
60        );
61        acc
62    }
63}
64
65#[cfg(test)]
66mod test_x86_64_fma_max_f32_32n {
67    use super::*;
68    crate::max_frame_tests!(is_x86_feature_detected!("avx"), f32, x86_64_fma_max_f32_32n);
69}
70
71// AVX-512 version: processes 64 f32 per loop iteration (4 zmm registers of 16
72// lanes each). Runtime-gated on avx512f (see x86_64_fma.rs::plug_avx512f); on
73// non-AVX512 CPUs this kernel is never registered and the FMA path above stays
74// in use. nr=64, 64-byte (16xf32) alignment.
75reduce_impl_wrap!(
76    f32,
77    x86_64_avx512_max_f32_64n,
78    64,
79    16,
80    (),
81    f32::MIN,
82    #[inline(never)]
83    fn run(buf: &[f32], _: ()) -> f32 {
84        assert!(buf.len() % 64 == 0);
85        assert!(buf.len() > 0);
86        unsafe { x86_64_avx512_max_f32_64n_run(buf) }
87    },
88    #[inline(never)]
89    fn reduce_two(a: f32, b: f32) -> f32 {
90        a.max(b)
91    }
92);
93
94#[target_feature(enable = "avx512f")]
95unsafe fn x86_64_avx512_max_f32_64n_run(buf: &[f32]) -> f32 {
96    unsafe {
97        let len = buf.len();
98        let ptr = buf.as_ptr();
99        let mut acc = f32::MIN;
100        std::arch::asm!("
101            vbroadcastss zmm0, xmm0
102            vmovaps zmm1, zmm0
103            vmovaps zmm2, zmm0
104            vmovaps zmm3, zmm0
105            2:
106                vmaxps zmm0, zmm0, [{ptr}]
107                vmaxps zmm1, zmm1, [{ptr} + 64]
108                vmaxps zmm2, zmm2, [{ptr} + 128]
109                vmaxps zmm3, zmm3, [{ptr} + 192]
110                add {ptr}, 256
111                sub {len}, 64
112                jnz 2b
113            vmaxps zmm0, zmm0, zmm1
114            vmaxps zmm2, zmm2, zmm3
115            vmaxps zmm0, zmm0, zmm2             // zmm0 holds 16 partial maxima
116            vextractf64x4 ymm1, zmm0, 1         // upper 256 bits (8xf32) of zmm0 -> ymm1 (avx512f)
117            vmaxps ymm0, ymm0, ymm1            // ymm0 holds 8 values
118            vextractf128 xmm1, ymm0, 1          // upper 4xf32 -> xmm1
119            vmaxps xmm0, xmm0, xmm1            // xmm0 holds 4 values
120            vpermilps xmm1, xmm0, 2 + (3 << 2)  // second 2x32 bit half moved to top
121            vmaxps xmm0, xmm0, xmm1            // xmm0 holds 2 values
122            vpermilps xmm1, xmm0, 1             // second f32 to top
123            vmaxps xmm0, xmm0, xmm1
124            ",
125        len = inout(reg) len => _,
126        ptr = inout(reg) ptr => _,
127        inout("zmm0") acc,
128        out("zmm1") _, out("zmm2") _, out("zmm3") _,
129        );
130        acc
131    }
132}
133
134#[cfg(test)]
135mod test_x86_64_avx512_max_f32_64n {
136    use super::*;
137    crate::max_frame_tests!(is_x86_feature_detected!("avx512f"), f32, x86_64_avx512_max_f32_64n);
138}