Skip to main content

multiply_accumulate4

Function multiply_accumulate4 

Source
pub unsafe fn multiply_accumulate4(
    rr: &mut __m256d,
    ri: &mut __m256d,
    ar: __m256d,
    ai: __m256d,
    br: __m256d,
    bi: __m256d,
)
Examples found in repository?
crates/competitive/src/math/fast_fourier_transform.rs (line 277)
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    }
More examples
Hide additional examples
crates/competitive/src/math/mint_fft_convolve.rs (line 100)
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}