Skip to main content

tract_linalg/x86_64/
max.rs

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