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
62routine_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}