Skip to main content

ruprim/block/
sort.rs

1use ruda_kernel::dsl as kernel_dsl;
2use ruda_kernel::dsl::prelude::*;
3use crate::collective::{RudaCompare, merge};
4
5/// Stable merge sort in blocked register order. Scratch arrays each contain
6/// `threads * items_per_thread` items and must not alias one another.
7#[ruda]
8pub fn merge_sort_keys<K: RudaPrimitive, C: RudaCompare<K>>(
9    keys: &mut Array<K>,
10    source: &mut SharedMemory<K>,
11    destination: &mut SharedMemory<K>,
12    compare: &C,
13    valid_items: usize,
14    #[comptime] threads: usize,
15    #[comptime] items_per_thread: usize,
16) {
17    let start = UNIT_POS as usize * items_per_thread;
18    #[unroll]
19    for item in 0..items_per_thread {
20        if start + item < valid_items {
21            source[start + item] = keys[item];
22        }
23    }
24    sync_ruda();
25    let mut run = 1usize;
26    while run < valid_items {
27        #[unroll]
28        for item in 0..items_per_thread {
29            let index = start + item;
30            if index < valid_items {
31                let rank = merge::rank::<K, C>(source, compare, index, valid_items, run, 0);
32                destination[rank] = source[index];
33            }
34        }
35        sync_ruda();
36        #[unroll]
37        for item in 0..items_per_thread {
38            let index = start + item;
39            if index < valid_items {
40                source[index] = destination[index];
41            }
42        }
43        sync_ruda();
44        run *= 2;
45    }
46    #[unroll]
47    for item in 0..items_per_thread {
48        if start + item < valid_items {
49            keys[item] = source[start + item];
50        }
51    }
52    sync_ruda();
53}
54
55/// Stable key/value merge sort; every value follows its original key.
56#[ruda]
57pub fn merge_sort_pairs<K: RudaPrimitive, V: RudaPrimitive, C: RudaCompare<K>>(
58    keys: &mut Array<K>,
59    values: &mut Array<V>,
60    source_keys: &mut SharedMemory<K>,
61    destination_keys: &mut SharedMemory<K>,
62    source_values: &mut SharedMemory<V>,
63    destination_values: &mut SharedMemory<V>,
64    compare: &C,
65    valid_items: usize,
66    #[comptime] threads: usize,
67    #[comptime] items_per_thread: usize,
68) {
69    let start = UNIT_POS as usize * items_per_thread;
70    #[unroll]
71    for item in 0..items_per_thread {
72        if start + item < valid_items {
73            source_keys[start + item] = keys[item];
74            source_values[start + item] = values[item];
75        }
76    }
77    sync_ruda();
78    let mut run = 1usize;
79    while run < valid_items {
80        #[unroll]
81        for item in 0..items_per_thread {
82            let index = start + item;
83            if index < valid_items {
84                let rank = merge::rank::<K, C>(source_keys, compare, index, valid_items, run, 0);
85                destination_keys[rank] = source_keys[index];
86                destination_values[rank] = source_values[index];
87            }
88        }
89        sync_ruda();
90        #[unroll]
91        for item in 0..items_per_thread {
92            let index = start + item;
93            if index < valid_items {
94                source_keys[index] = destination_keys[index];
95                source_values[index] = destination_values[index];
96            }
97        }
98        sync_ruda();
99        run *= 2;
100    }
101    #[unroll]
102    for item in 0..items_per_thread {
103        if start + item < valid_items {
104            keys[item] = source_keys[start + item];
105            values[item] = source_values[start + item];
106        }
107    }
108    sync_ruda();
109}