1use ruda_kernel::dsl as kernel_dsl;
2use ruda_kernel::dsl::prelude::*;
3use crate::collective::record::{RudaRead, RudaReadExpand};
4
5#[ruda]
6pub trait RudaDigitExtractor<K: RudaType>: RudaType {
7 fn digit(&self, key: K) -> u32;
8}
9
10#[ruda]
15pub fn rank_keys<K: RudaPrimitive, E: RudaDigitExtractor<K>>(
16 keys: &Array<K>, ranks: &mut Array<u32>, extractor: &E,
17 digit_scratch: &mut SharedMemory<u32>, index_scratch: &mut SharedMemory<u32>,
18 scan_scratch: &mut SharedMemory<u32>, digit_offsets: &mut SharedMemory<u32>,
19 valid_items: usize, #[comptime] threads: usize, #[comptime] items_per_thread: usize,
20 #[comptime] digit_bits: u32, #[comptime] descending: bool,
21) {
22 rank_access::<K, Array<K>, E>(keys, ranks, extractor, digit_scratch, index_scratch, scan_scratch,
23 digit_offsets, valid_items, threads, items_per_thread, digit_bits, descending);
24}
25
26#[ruda]
27pub fn rank_access<K: RudaType, I: RudaRead<K>, E: RudaDigitExtractor<K>>(
28 keys: &I, ranks: &mut Array<u32>, extractor: &E,
29 digit_scratch: &mut SharedMemory<u32>, index_scratch: &mut SharedMemory<u32>,
30 scan_scratch: &mut SharedMemory<u32>, digit_offsets: &mut SharedMemory<u32>,
31 valid_items: usize, #[comptime] threads: usize, #[comptime] items_per_thread: usize,
32 #[comptime] digit_bits: u32, #[comptime] descending: bool,
33) {
34 let start = UNIT_POS as usize * items_per_thread;
35 let mut digits = Array::<u32>::new(items_per_thread);
36 let mut indices = Array::<u32>::new(items_per_thread);
37 #[unroll]
38 for item in 0..items_per_thread {
39 if start + item < valid_items {
40 digits[item] = extractor.digit(keys.read(item));
41 indices[item] = (start + item) as u32;
42 }
43 }
44 crate::block::radix::sort_pairs::<u32, u32>(&mut digits, &mut indices, digit_scratch, index_scratch,
45 scan_scratch, valid_items, threads, items_per_thread, 0u32, digit_bits, descending);
46 let bins = 1usize << digit_bits;
47 #[unroll]
48 for item in 0..items_per_thread {
49 if start + item < valid_items {
50 let mut digit = digits[item];
51 if descending { digit = bins as u32 - 1 - digit; }
52 digit_scratch[start + item] = digit;
53 scan_scratch[indices[item] as usize] = (start + item) as u32;
54 }
55 }
56 sync_ruda();
57 #[unroll]
58 for item in 0..items_per_thread {
59 if start + item < valid_items { ranks[item] = scan_scratch[start + item]; }
60 }
61 let mut bucket = UNIT_POS as usize;
62 while bucket <= bins {
63 let mut low = 0usize;
64 let mut high = valid_items;
65 while low < high {
66 let middle = low + (high - low) / 2;
67 if (digit_scratch[middle] as usize) < bucket { low = middle + 1; } else { high = middle; }
68 }
69 digit_offsets[bucket] = low as u32;
70 bucket += threads;
71 }
72 sync_ruda();
73}