unsafe fn greater_avx512(a: __m512i, b: __m512i) -> __m512iExamples found in repository?
crates/competitive/src/data_structure/wavelet_matrix.rs (line 400)
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 }