Skip to main content

ruprim/block/
rank.rs

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/// Stable ranks of extracted digits, without rearranging the input keys.
11/// Extracted digits fit `digit_bits` (1..=31). Tile scratch arrays each contain
12/// the complete tile; `digit_offsets` contains `(1 << digit_bits) + 1` entries
13/// in output bucket order. Adjacent offsets also give bucket populations.
14#[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}