pub unsafe fn store4(value: &mut Complex4, re: __m256d, im: __m256d)Examples found in repository?
crates/competitive/src/math/mint_fft_convolve.rs (line 113)
79unsafe fn dot_soa(a0: &mut [Complex4], a1: &mut [Complex4], b0: &mut [Complex4], b1: &[Complex4]) {
80 let n = a0.len() * 4;
81 RotateCache::ensure(n / 2);
82 RotateCache::with(|cache| {
83 for i in 0..a0.len() {
84 let (mut cr, mut ci) = load4(&b0[i]);
85 let (mut dr, mut di) = load4(&b1[i]);
86 let mut c0r = _mm256_setzero_pd();
87 let mut c0i = _mm256_setzero_pd();
88 let mut c1r = _mm256_setzero_pd();
89 let mut c1i = _mm256_setzero_pd();
90 let mut c2r = _mm256_setzero_pd();
91 let mut c2i = _mm256_setzero_pd();
92 let w = eval_twiddle(cache, 1, a0.len(), i);
93 let wr = _mm256_setr_pd(w.re, 1.0, 1.0, 1.0);
94 let wi = _mm256_setr_pd(w.im, 0.0, 0.0, 0.0);
95 for lane in 0..4 {
96 let ar = _mm256_set1_pd(a0[i].re[lane]);
97 let ai = _mm256_set1_pd(a0[i].im[lane]);
98 let br = _mm256_set1_pd(a1[i].re[lane]);
99 let bi = _mm256_set1_pd(a1[i].im[lane]);
100 multiply_accumulate4(&mut c0r, &mut c0i, ar, ai, cr, ci);
101 multiply_accumulate4(&mut c1r, &mut c1i, ar, ai, dr, di);
102 multiply_accumulate4(&mut c1r, &mut c1i, br, bi, cr, ci);
103 multiply_accumulate4(&mut c2r, &mut c2i, br, bi, dr, di);
104 if lane != 3 {
105 cr = _mm256_permute4x64_pd::<0x93>(cr);
106 ci = _mm256_permute4x64_pd::<0x93>(ci);
107 dr = _mm256_permute4x64_pd::<0x93>(dr);
108 di = _mm256_permute4x64_pd::<0x93>(di);
109 (cr, ci) = mul4(cr, ci, wr, wi);
110 (dr, di) = mul4(dr, di, wr, wi);
111 }
112 }
113 store4(&mut a0[i], c0r, c0i);
114 store4(&mut a1[i], c1r, c1i);
115 store4(&mut b0[i], c2r, c2i);
116 }
117 });
118}
119
120#[target_feature(enable = "avx2,fma")]
121unsafe fn split_u64_coefficients(values: &[u64], n: usize) -> [Vec<Complex4>; 5] {
122 let mut result: [Vec<Complex4>; 5] = std::array::from_fn(|_| {
123 let mut part = Vec::with_capacity(n / 4);
124 advise_huge_pages(&mut part);
125 part
126 });
127 for (i, chunk) in values.chunks(4).enumerate() {
128 let mut parts = [Complex4::default(); 5];
129 for (lane, mut value) in chunk.iter().copied().enumerate() {
130 for part in &mut parts {
131 let digit = ((value << 51) as i64) >> 51;
132 value = (value >> 13).wrapping_add(u64::from(digit < 0));
133 part.re[lane] = digit as f64;
134 }
135 }
136 for (result, part) in result.iter_mut().zip(parts) {
137 if i < n / 4 {
138 result.push(part);
139 } else {
140 result[i - n / 4].im = part.re;
141 }
142 }
143 }
144 for part in &mut result {
145 part.resize(n / 4, Complex4::default());
146 }
147 result
148}
149
150#[target_feature(enable = "avx2,fma")]
151unsafe fn dot_u64_soa(a: &mut [Vec<Complex4>; 5], b: &[Vec<Complex4>; 5]) {
152 let n = a[0].len() * 4;
153 RotateCache::ensure(n / 2);
154 RotateCache::with(|cache| {
155 for block in 0..a[0].len() {
156 let mut br = [_mm256_setzero_pd(); 5];
157 let mut bi = br;
158 let mut rr = br;
159 let mut ri = br;
160 for part in 0..5 {
161 (br[part], bi[part]) = load4(&b[part][block]);
162 }
163 let w = eval_twiddle(cache, 1, a[0].len(), block);
164 let wr = _mm256_setr_pd(w.re, 1.0, 1.0, 1.0);
165 let wi = _mm256_setr_pd(w.im, 0.0, 0.0, 0.0);
166 for lane in 0..4 {
167 let ar: [__m256d; 5] =
168 std::array::from_fn(|part| _mm256_set1_pd(a[part][block].re[lane]));
169 let ai: [__m256d; 5] =
170 std::array::from_fn(|part| _mm256_set1_pd(a[part][block].im[lane]));
171 for part in 0..5 {
172 for left in 0..=part {
173 multiply_accumulate4(
174 &mut rr[part],
175 &mut ri[part],
176 ar[left],
177 ai[left],
178 br[part - left],
179 bi[part - left],
180 );
181 }
182 }
183 if lane != 3 {
184 for part in 0..5 {
185 br[part] = _mm256_permute4x64_pd::<0x93>(br[part]);
186 bi[part] = _mm256_permute4x64_pd::<0x93>(bi[part]);
187 (br[part], bi[part]) = mul4(br[part], bi[part], wr, wi);
188 }
189 }
190 }
191 for part in 0..5 {
192 store4(&mut a[part][block], rr[part], ri[part]);
193 }
194 }
195 });
196}More examples
crates/competitive/src/math/fast_fourier_transform.rs (line 147)
109 pub unsafe fn fft_soa(a: &mut [Complex4]) {
110 let n = a.len() * 4;
111 RotateCache::ensure(n / 2);
112 RotateCache::with(|cache| {
113 let parity = n.trailing_zeros() & 1;
114 for leaf in (0..n).step_by(16) {
115 let mut level = (n + leaf).trailing_zeros();
116 level -= u32::from(level & 1 != parity);
117 while level >= 4 {
118 let len = 1usize << level;
119 let q = leaf >> level;
120 let width = len / 16;
121 let start = q * width * 4;
122 let (a, rest) = a[start..start + width * 4].split_at_mut(width);
123 let (b, rest) = rest.split_at_mut(width);
124 let (c, d) = rest.split_at_mut(width);
125 let w1 = eval_twiddle(cache, 4, n >> level, q);
126 let w2 = w1 * w1;
127 let w3 = w1 * w2;
128 let (w1r, w1i) = (_mm256_set1_pd(w1.re), _mm256_set1_pd(w1.im));
129 let (w2r, w2i) = (_mm256_set1_pd(w2.re), _mm256_set1_pd(w2.im));
130 let (w3r, w3i) = (_mm256_set1_pd(w3.re), _mm256_set1_pd(w3.im));
131 for i in 0..width {
132 let (ar, ai) = load4(&a[i]);
133 let (br, bi) = load4(&b[i]);
134 let (cr, ci) = load4(&c[i]);
135 let (dr, di) = load4(&d[i]);
136 let (br, bi) = mul4(br, bi, w1r, w1i);
137 let (cr, ci) = mul4(cr, ci, w2r, w2i);
138 let (dr, di) = mul4(dr, di, w3r, w3i);
139 let acr = _mm256_add_pd(ar, cr);
140 let aci = _mm256_add_pd(ai, ci);
141 let bdr = _mm256_add_pd(br, dr);
142 let bdi = _mm256_add_pd(bi, di);
143 let acd_r = _mm256_sub_pd(ar, cr);
144 let acd_i = _mm256_sub_pd(ai, ci);
145 let bdd_r = _mm256_sub_pd(br, dr);
146 let bdd_i = _mm256_sub_pd(bi, di);
147 store4(&mut a[i], _mm256_add_pd(acr, bdr), _mm256_add_pd(aci, bdi));
148 store4(&mut b[i], _mm256_sub_pd(acr, bdr), _mm256_sub_pd(aci, bdi));
149 store4(
150 &mut c[i],
151 _mm256_sub_pd(acd_r, bdd_i),
152 _mm256_add_pd(acd_i, bdd_r),
153 );
154 store4(
155 &mut d[i],
156 _mm256_add_pd(acd_r, bdd_i),
157 _mm256_sub_pd(acd_i, bdd_r),
158 );
159 }
160 level -= 2;
161 }
162 }
163 if parity != 0 {
164 let blocks = n / 8;
165 for k in 0..blocks {
166 let w = eval_twiddle(cache, 2, blocks, k);
167 let wr = _mm256_set1_pd(w.re);
168 let wi = _mm256_set1_pd(w.im);
169 let (ar, ai) = load4(&a[k * 2]);
170 let (br, bi) = load4(&a[k * 2 + 1]);
171 let (br, bi) = mul4(br, bi, wr, wi);
172 store4(&mut a[k * 2], _mm256_add_pd(ar, br), _mm256_add_pd(ai, bi));
173 store4(
174 &mut a[k * 2 + 1],
175 _mm256_sub_pd(ar, br),
176 _mm256_sub_pd(ai, bi),
177 );
178 }
179 }
180 });
181 }
182
183 #[target_feature(enable = "avx2,fma")]
184 pub unsafe fn ifft_soa(a: &mut [Complex4]) {
185 let n = a.len() * 4;
186 RotateCache::ensure(n / 2);
187 RotateCache::with(|cache| {
188 let parity = n.trailing_zeros() & 1;
189 if parity != 0 {
190 let blocks = n / 8;
191 for k in 0..blocks {
192 let w = eval_twiddle(cache, 2, blocks, k).conjugate();
193 let wr = _mm256_set1_pd(w.re);
194 let wi = _mm256_set1_pd(w.im);
195 let (ar, ai) = load4(&a[k * 2]);
196 let (br, bi) = load4(&a[k * 2 + 1]);
197 store4(&mut a[k * 2], _mm256_add_pd(ar, br), _mm256_add_pd(ai, bi));
198 let (br, bi) = mul4(_mm256_sub_pd(ar, br), _mm256_sub_pd(ai, bi), wr, wi);
199 store4(&mut a[k * 2 + 1], br, bi);
200 }
201 }
202 for leaf in (12..n).step_by(16) {
203 let max_level = (leaf + 3).trailing_ones();
204 let mut level = 4 + parity;
205 while level <= max_level {
206 let len = 1usize << level;
207 let q = leaf >> level;
208 let width = len / 16;
209 let start = q * width * 4;
210 let (a, rest) = a[start..start + width * 4].split_at_mut(width);
211 let (b, rest) = rest.split_at_mut(width);
212 let (c, d) = rest.split_at_mut(width);
213 let w1 = eval_twiddle(cache, 4, n >> level, q).conjugate();
214 let w2 = w1 * w1;
215 let w3 = w1 * w2;
216 let (w1r, w1i) = (_mm256_set1_pd(w1.re), _mm256_set1_pd(w1.im));
217 let (w2r, w2i) = (_mm256_set1_pd(w2.re), _mm256_set1_pd(w2.im));
218 let (w3r, w3i) = (_mm256_set1_pd(w3.re), _mm256_set1_pd(w3.im));
219 for i in 0..width {
220 let (ar, ai) = load4(&a[i]);
221 let (br, bi) = load4(&b[i]);
222 let (cr, ci) = load4(&c[i]);
223 let (dr, di) = load4(&d[i]);
224 let abr = _mm256_add_pd(ar, br);
225 let abi = _mm256_add_pd(ai, bi);
226 let cdr = _mm256_add_pd(cr, dr);
227 let cdi = _mm256_add_pd(ci, di);
228 let abd_r = _mm256_sub_pd(ar, br);
229 let abd_i = _mm256_sub_pd(ai, bi);
230 let cdd_r = _mm256_sub_pd(cr, dr);
231 let cdd_i = _mm256_sub_pd(ci, di);
232 store4(&mut a[i], _mm256_add_pd(abr, cdr), _mm256_add_pd(abi, cdi));
233 let (br, bi) = mul4(
234 _mm256_add_pd(abd_r, cdd_i),
235 _mm256_sub_pd(abd_i, cdd_r),
236 w1r,
237 w1i,
238 );
239 store4(&mut b[i], br, bi);
240 let (cr, ci) =
241 mul4(_mm256_sub_pd(abr, cdr), _mm256_sub_pd(abi, cdi), w2r, w2i);
242 store4(&mut c[i], cr, ci);
243 let (dr, di) = mul4(
244 _mm256_sub_pd(abd_r, cdd_i),
245 _mm256_add_pd(abd_i, cdd_r),
246 w3r,
247 w3i,
248 );
249 store4(&mut d[i], dr, di);
250 }
251 level += 2;
252 }
253 }
254 let scale = _mm256_set1_pd(4.0 / n as f64);
255 for value in a {
256 let (re, im) = load4(value);
257 store4(value, _mm256_mul_pd(re, scale), _mm256_mul_pd(im, scale));
258 }
259 });
260 }
261
262 #[target_feature(enable = "avx2,fma")]
263 unsafe fn dot_one_soa(a: &mut [Complex4], b: &[Complex4]) {
264 let n = a.len() * 4;
265 RotateCache::ensure(n / 2);
266 RotateCache::with(|cache| {
267 for i in 0..a.len() {
268 let (mut br, mut bi) = load4(&b[i]);
269 let mut rr = _mm256_setzero_pd();
270 let mut ri = _mm256_setzero_pd();
271 let w = eval_twiddle(cache, 1, a.len(), i);
272 let wr = _mm256_setr_pd(w.re, 1.0, 1.0, 1.0);
273 let wi = _mm256_setr_pd(w.im, 0.0, 0.0, 0.0);
274 for lane in 0..4 {
275 let ar = _mm256_set1_pd(a[i].re[lane]);
276 let ai = _mm256_set1_pd(a[i].im[lane]);
277 multiply_accumulate4(&mut rr, &mut ri, ar, ai, br, bi);
278 if lane != 3 {
279 br = _mm256_permute4x64_pd::<0x93>(br);
280 bi = _mm256_permute4x64_pd::<0x93>(bi);
281 (br, bi) = mul4(br, bi, wr, wi);
282 }
283 }
284 store4(&mut a[i], rr, ri);
285 }
286 });
287 }