superintervals 0.3.2__tar.gz → 0.3.4__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.1
2
2
  Name: superintervals
3
- Version: 0.3.2
3
+ Version: 0.3.4
4
4
  Summary: Rapid interval intersections
5
5
  Author: Kez Cleal
6
6
  Author-email: Kez Cleal <clealk@cardiff.ac.uk>
@@ -8,7 +8,7 @@ build-backend = "setuptools.build_meta"
8
8
 
9
9
  [project]
10
10
  name = "superintervals"
11
- version = "0.3.2"
11
+ version = "0.3.4"
12
12
  description = "Rapid interval intersections"
13
13
  dependencies = ['Cython']
14
14
  authors = [{name = "Kez Cleal", email = "clealk@cardiff.ac.uk"}]
@@ -13,7 +13,6 @@ print('PAKCAGES', find_packages(where='src')) # Add this line for debugging
13
13
 
14
14
  setup(
15
15
  name='superintervals',
16
- version='0.3.0',
17
16
  description="Rapid interval intersections",
18
17
  author="Kez Cleal",
19
18
  author_email="clealk@cardiff.ac.uk",
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.1
2
2
  Name: superintervals
3
- Version: 0.3.2
3
+ Version: 0.3.4
4
4
  Summary: Rapid interval intersections
5
5
  Author: Kez Cleal
6
6
  Author-email: Kez Cleal <clealk@cardiff.ac.uk>
@@ -1,4 +1,4 @@
1
- // Version 0.3.1
1
+ // Version 0.3.4
2
2
  #pragma once
3
3
 
4
4
  #include <algorithm>
@@ -43,10 +43,19 @@ struct Interval {
43
43
  * @tparam T The data type associated with each interval
44
44
  */
45
45
  template<typename S, typename T>
46
+
46
47
  class IntervalMap {
47
48
  public:
48
49
  std::vector<S> starts;
49
- std::vector<S> ends;
50
+ // std::vector<S> ends;
51
+
52
+ #ifdef __AVX2__
53
+ alignas(32) std::vector<S> ends;
54
+ #elif defined(__ARM_NEON)
55
+ alignas(16) std::vector<S> ends;
56
+ #else
57
+ alignas(sizeof(S)) std::vector<S> ends;
58
+ #endif
50
59
  std::vector<size_t> branch;
51
60
  std::vector<T> data;
52
61
  bool start_sorted, end_sorted;
@@ -282,7 +291,6 @@ class IntervalMap {
282
291
  // First do exponential search if we have room
283
292
  size_t search_right = right;
284
293
  size_t bound = 1;
285
- // Exponential search to find a smaller range
286
294
  while (left > 0 && value < starts[left]) {
287
295
  search_right = left;
288
296
  left = (bound <= left) ? left - bound : 0;
@@ -414,69 +422,159 @@ class IntervalMap {
414
422
  }
415
423
  size_t found = 0;
416
424
 
417
- #ifdef SI_NOSIMD
425
+ #if defined(SI_NOSIMD)
426
+
418
427
  constexpr size_t block = 16;
419
- #elif defined(__AVX2__)
428
+ constexpr simd_kind active_simd = simd_kind::none;
429
+
430
+ #elif defined(__AVX2__)
431
+
420
432
  __m256i start_vec = _mm256_set1_epi32(start);
421
433
  constexpr size_t simd_width = 256 / (sizeof(S) * 8);
422
- constexpr size_t block = simd_width * 4;
423
- #elif defined(__ARM_NEON__) || defined(__aarch64__)
434
+ constexpr size_t block = simd_width * 4; // 2 cache lines
435
+ constexpr simd_kind active_simd = simd_kind::avx2;
436
+
437
+ #elif defined(__ARM_NEON__) || defined(__aarch64__)
438
+
424
439
  int32x4_t start_vec = vdupq_n_s32(start);
425
440
  constexpr size_t simd_width = 128 / (sizeof(S) * 8);
426
441
  uint32x4_t ones = vdupq_n_u32(1);
427
- constexpr size_t block = simd_width * 4;
428
- #endif
442
+ constexpr size_t block = simd_width * 8; // 2 cache lines
443
+ constexpr simd_kind active_simd = simd_kind::neon;
444
+
445
+ #endif
429
446
 
430
447
  while (i > 0) {
431
448
  if (start <= ends[i]) {
432
449
  ++found;
433
450
  --i;
434
- #ifdef SI_NOSIMD
435
- while (i > block) { // Rely on compiler auto vectorize
436
- size_t count = 0;
437
- for (size_t j = i; j > i - block; --j) {
438
- count += (start <= ends[j]) ? 1 : 0;
439
- }
440
- found += count;
441
- i -= block;
442
- if (count < block && start > ends[i + 1]) { // check for a branch
443
- break;
451
+ // Types with width !=4 will use the no-simd path here
452
+ if constexpr (active_simd == simd_kind::none || sizeof(S) != 4) {
453
+ while (i > block) {
454
+ size_t count = 0;
455
+ for (size_t j = i; j > i - block; --j) {
456
+ count += (start <= ends[j]) ? 1 : 0;
457
+ }
458
+ found += count;
459
+ i -= block;
460
+ if (count < block && start > ends[i + 1]) { // check for a branch
461
+ break;
462
+ }
444
463
  }
445
464
  }
446
-
447
- #elif defined(__AVX2__)
448
- while (i > block) {
449
- size_t count = 0;
450
- for (size_t j = i; j > i - block; j -= simd_width) {
451
- __m256i ends_vec = _mm256_loadu_si256((__m256i*)(&ends[j - simd_width + 1]));
452
- __m256i cmp_mask = _mm256_cmpgt_epi32(start_vec, ends_vec);
453
- int mask = _mm256_movemask_epi8(~cmp_mask);
454
- count += _mm_popcnt_u32(mask);
455
- }
456
- found += count / 4; // Each comparison result is 4 bits
457
- i -= block;
458
- if (count < block) {
459
- break;
465
+ #if defined(__AVX2__)
466
+ else if constexpr (active_simd == simd_kind::avx2) {
467
+ // while (i > block) {
468
+ // size_t count = 0;
469
+ // for (size_t j = i - block + 1; j < i; j += simd_width) {
470
+ // __m256i ends_vec = _mm256_loadu_si256((__m256i*)(&ends[j - simd_width + 1]));
471
+ // __m256i cmp_mask = _mm256_cmpgt_epi32(start_vec, ends_vec);
472
+ // int mask = _mm256_movemask_ps(_mm256_castsi256_ps(cmp_mask));
473
+ // count += 8 - _mm_popcnt_u32(mask);
474
+ // }
475
+ // found += count;
476
+ // i -= block;
477
+ // if (count < block) {
478
+ // break;
479
+ // }
480
+ // }
481
+
482
+ while (i > block) {
483
+ size_t j = i - block + 1;
484
+
485
+ // Load all 4 vectors
486
+ __m256i ends_vec0 = _mm256_load_si256((__m256i*)(&ends[j]));
487
+ __m256i ends_vec1 = _mm256_load_si256((__m256i*)(&ends[j + simd_width]));
488
+ __m256i ends_vec2 = _mm256_load_si256((__m256i*)(&ends[j + 2 * simd_width]));
489
+ __m256i ends_vec3 = _mm256_load_si256((__m256i*)(&ends[j + 3 * simd_width]));
490
+
491
+ // Compare all vectors
492
+ __m256i cmp_mask0 = _mm256_cmpgt_epi32(start_vec, ends_vec0);
493
+ __m256i cmp_mask1 = _mm256_cmpgt_epi32(start_vec, ends_vec1);
494
+ __m256i cmp_mask2 = _mm256_cmpgt_epi32(start_vec, ends_vec2);
495
+ __m256i cmp_mask3 = _mm256_cmpgt_epi32(start_vec, ends_vec3);
496
+
497
+ // Extract masks
498
+ int mask0 = _mm256_movemask_ps(_mm256_castsi256_ps(cmp_mask0));
499
+ int mask1 = _mm256_movemask_ps(_mm256_castsi256_ps(cmp_mask1));
500
+ int mask2 = _mm256_movemask_ps(_mm256_castsi256_ps(cmp_mask2));
501
+ int mask3 = _mm256_movemask_ps(_mm256_castsi256_ps(cmp_mask3));
502
+
503
+ // Count and accumulate
504
+ size_t count = (8 - _mm_popcnt_u32(mask0)) + (8 - _mm_popcnt_u32(mask1)) +
505
+ (8 - _mm_popcnt_u32(mask2)) + (8 - _mm_popcnt_u32(mask3));
506
+
507
+ found += count;
508
+ i -= block;
509
+ if (count < block) {
510
+ break;
511
+ }
460
512
  }
461
513
  }
462
- #elif defined(__ARM_NEON__) || defined(__aarch64__)
463
- while (i > block) {
464
- size_t count = 0;
465
- uint32x4_t mask, bool_mask;
466
- for (size_t j = i; j > i - block; j -= simd_width) { // Neon processes 4 int32 at a time
467
- int32x4_t ends_vec = vld1q_s32(&ends[j - simd_width + 1]);
468
- mask = vcleq_s32(start_vec, ends_vec); // True (0xFFFFFFFF) for elements where start_vec <= ends_vec
469
- bool_mask = vandq_u32(mask, ones);
470
- count += vaddvq_u32(bool_mask);
471
- }
472
- found += count;
473
- i -= block;
474
- // if (count < block && vgetq_lane_u32(mask, 0) == 0) { // check for overlap again, before checking for branch?
475
- if (count < block) { // check for overlap again, before checking for branch?
476
- break;
514
+ #elif defined(__ARM_NEON__) || defined(__aarch64__)
515
+ else { // NEON
516
+ // while (i > block) {
517
+ // size_t count = 0;
518
+ // uint32x4_t mask, bool_mask;
519
+ // for (size_t j = i - block + 1; j < i; j += simd_width) { // Neon 4 int32 at a time
520
+ // int32x4_t ends_vec = vld1q_s32(&ends[j]);
521
+ // mask = vcgtq_s32(start_vec, ends_vec); // start > ends[j]
522
+ // bool_mask = vaddq_u32(mask, ones);
523
+ // count += vaddvq_u32(bool_mask); // Sum all lanes
524
+ // }
525
+ // found += count;
526
+ // i -= block;
527
+ // if (count < block) { // check for overlap again, before checking for branch?
528
+ // break;
529
+ // }
530
+ // }
531
+ while (i > block) {
532
+ size_t j = i - block + 1;
533
+
534
+ // Load all 8 vectors
535
+ int32x4_t ends_vec0 = vld1q_s32(&ends[j]);
536
+ int32x4_t ends_vec1 = vld1q_s32(&ends[j + simd_width]);
537
+ int32x4_t ends_vec2 = vld1q_s32(&ends[j + 2 * simd_width]);
538
+ int32x4_t ends_vec3 = vld1q_s32(&ends[j + 3 * simd_width]);
539
+ int32x4_t ends_vec4 = vld1q_s32(&ends[j + 4 * simd_width]);
540
+ int32x4_t ends_vec5 = vld1q_s32(&ends[j + 5 * simd_width]);
541
+ int32x4_t ends_vec6 = vld1q_s32(&ends[j + 6 * simd_width]);
542
+ int32x4_t ends_vec7 = vld1q_s32(&ends[j + 7 * simd_width]);
543
+
544
+ // Compare all vectors
545
+ uint32x4_t mask0 = vcgtq_s32(start_vec, ends_vec0);
546
+ uint32x4_t mask1 = vcgtq_s32(start_vec, ends_vec1);
547
+ uint32x4_t mask2 = vcgtq_s32(start_vec, ends_vec2);
548
+ uint32x4_t mask3 = vcgtq_s32(start_vec, ends_vec3);
549
+ uint32x4_t mask4 = vcgtq_s32(start_vec, ends_vec4);
550
+ uint32x4_t mask5 = vcgtq_s32(start_vec, ends_vec5);
551
+ uint32x4_t mask6 = vcgtq_s32(start_vec, ends_vec6);
552
+ uint32x4_t mask7 = vcgtq_s32(start_vec, ends_vec7);
553
+
554
+ // Convert to boolean masks
555
+ uint32x4_t bool_mask0 = vaddq_u32(mask0, ones);
556
+ uint32x4_t bool_mask1 = vaddq_u32(mask1, ones);
557
+ uint32x4_t bool_mask2 = vaddq_u32(mask2, ones);
558
+ uint32x4_t bool_mask3 = vaddq_u32(mask3, ones);
559
+ uint32x4_t bool_mask4 = vaddq_u32(mask4, ones);
560
+ uint32x4_t bool_mask5 = vaddq_u32(mask5, ones);
561
+ uint32x4_t bool_mask6 = vaddq_u32(mask6, ones);
562
+ uint32x4_t bool_mask7 = vaddq_u32(mask7, ones);
563
+
564
+ // Sum all lanes and accumulate
565
+ size_t count = vaddvq_u32(bool_mask0) + vaddvq_u32(bool_mask1) +
566
+ vaddvq_u32(bool_mask2) + vaddvq_u32(bool_mask3) +
567
+ vaddvq_u32(bool_mask4) + vaddvq_u32(bool_mask5) +
568
+ vaddvq_u32(bool_mask6) + vaddvq_u32(bool_mask7);
569
+
570
+ found += count;
571
+ i -= block;
572
+ if (count < block) {
573
+ break;
574
+ }
477
575
  }
478
576
  }
479
- #endif
577
+ #endif
480
578
  } else {
481
579
  if (branch[i] == SIZE_MAX) {
482
580
  return found;
@@ -699,6 +797,8 @@ class IntervalMap {
699
797
 
700
798
  protected:
701
799
 
800
+ enum class simd_kind { none, avx2, neon };
801
+
702
802
  std::vector<Interval<S, T>> tmp;
703
803
 
704
804
  template<typename CompareFunc>
File without changes
File without changes
File without changes