Skip to main content

gather_avx512

Function gather_avx512 

Source
unsafe fn gather_avx512<const SCALE: i32>(
    base: *const i64,
    index: __m512i,
) -> __m512i
Examples found in repository?
crates/competitive/src/data_structure/wavelet_matrix.rs (line 373)
340    pub unsafe fn quad_avx512(
341        layers: &[super::WaveletMatrixQuadVector],
342        states: &mut [[usize; 4]],
343    ) {
344        for chunk in states.as_chunks_mut::<8>().0 {
345            let mut starts = [0u64; 8];
346            let mut ends = [0u64; 8];
347            let mut keys = [0u64; 8];
348            let mut result = [0u64; 8];
349            for (i, s) in chunk.iter().enumerate() {
350                starts[i] = s[0] as u64;
351                ends[i] = s[1] as u64;
352                keys[i] = s[2] as u64;
353            }
354            let mut start = _mm512_loadu_si512(starts.as_ptr().cast());
355            let mut end = _mm512_loadu_si512(ends.as_ptr().cast());
356            let mut key = _mm512_loadu_si512(keys.as_ptr().cast());
357            let mut code = _mm512_set1_epi64(0);
358            let one = _mm512_set1_epi64(1);
359            macro_rules! blend {
360                ($m:expr,$x:expr,$y:expr) => {
361                    _mm512_or_si512(_mm512_andnot_si512($m, $x), _mm512_and_si512($m, $y))
362                };
363            }
364            for layer in layers {
365                let base = layer.blocks.as_ptr().cast::<i64>();
366                let ranks = |pos| {
367                    let offset = _mm512_and_si512(pos, _mm512_set1_epi64(63));
368                    let mask = _mm512_sub_epi64(_mm512_sllv_epi64(one, offset), one);
369                    let i = _mm512_sllv_epi64(
370                        _mm512_srlv_epi64(pos, _mm512_set1_epi64(6)),
371                        _mm512_set1_epi64(2),
372                    );
373                    let lo = _mm512_and_si512(gather_avx512::<8>(base, i), mask);
374                    let hi =
375                        _mm512_and_si512(gather_avx512::<8>(base, _mm512_add_epi64(i, one)), mask);
376                    let a = _mm512_popcnt_epi64(lo);
377                    let b = _mm512_popcnt_epi64(hi);
378                    let c = _mm512_popcnt_epi64(_mm512_and_si512(lo, hi));
379                    let p01 = gather_avx512::<8>(base, _mm512_add_epi64(i, _mm512_set1_epi64(2)));
380                    let p23 = gather_avx512::<8>(base, _mm512_add_epi64(i, _mm512_set1_epi64(3)));
381                    let r0 = _mm512_add_epi64(
382                        _mm512_and_si512(p01, _mm512_set1_epi64(u32::MAX as i64)),
383                        _mm512_add_epi64(_mm512_sub_epi64(_mm512_sub_epi64(offset, a), b), c),
384                    );
385                    let r1 = _mm512_add_epi64(
386                        _mm512_srlv_epi64(p01, _mm512_set1_epi64(32)),
387                        _mm512_sub_epi64(a, c),
388                    );
389                    let r2 = _mm512_add_epi64(
390                        _mm512_and_si512(p23, _mm512_set1_epi64(u32::MAX as i64)),
391                        _mm512_sub_epi64(b, c),
392                    );
393                    let r3 = _mm512_sub_epi64(_mm512_sub_epi64(_mm512_sub_epi64(pos, r0), r1), r2);
394                    [r0, r1, r2, r3]
395                };
396                let l = ranks(start);
397                let r = ranks(end);
398                let low_count =
399                    _mm512_sub_epi64(_mm512_add_epi64(r[0], r[1]), _mm512_add_epi64(l[0], l[1]));
400                let high = greater_avx512(key, _mm512_sub_epi64(low_count, one));
401                key = _mm512_sub_epi64(key, _mm512_and_si512(high, low_count));
402                let l0 = blend!(high, l[0], l[2]);
403                let r0 = blend!(high, r[0], r[2]);
404                let l1 = blend!(high, l[1], l[3]);
405                let r1 = blend!(high, r[1], r[3]);
406                let count0 = _mm512_sub_epi64(r0, l0);
407                let low = greater_avx512(key, _mm512_sub_epi64(count0, one));
408                key = _mm512_sub_epi64(key, _mm512_and_si512(low, count0));
409                let base0 = blend!(
410                    high,
411                    _mm512_set1_epi64(layer.starts[0] as i64),
412                    _mm512_set1_epi64(layer.starts[2] as i64)
413                );
414                let base1 = blend!(
415                    high,
416                    _mm512_set1_epi64(layer.starts[1] as i64),
417                    _mm512_set1_epi64(layer.starts[3] as i64)
418                );
419                let offset = blend!(low, base0, base1);
420                start = _mm512_add_epi64(offset, blend!(low, l0, l1));
421                end = _mm512_add_epi64(offset, blend!(low, r0, r1));
422                code = _mm512_or_si512(
423                    _mm512_sllv_epi64(code, _mm512_set1_epi64(2)),
424                    _mm512_or_si512(
425                        _mm512_and_si512(high, _mm512_set1_epi64(2)),
426                        _mm512_and_si512(low, one),
427                    ),
428                );
429            }
430            _mm512_storeu_si512(result.as_mut_ptr().cast(), code);
431            for (i, s) in chunk.iter_mut().enumerate() {
432                s[3] = result[i] as usize;
433            }
434        }
435    }