Skip to main content

store4

Function store4 

Source
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
Hide additional 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    }