1use ruda_kernel::dsl as kernel_dsl;
2use ruda_kernel::dsl::prelude::*;
3use crate::collective::{RudaCompare, merge};
4
5#[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#[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}