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
71reduce_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}