1use ruda_kernel::dsl as kernel_dsl;
2use ruda_kernel::dsl::prelude::*;
3use crate::collective::{RudaBinaryOp, RudaBinaryOpExpand, RudaCompare, RudaCompareExpand};
4use crate::collective::record::{RudaRecord, RudaRecordArray, RudaRecordShared, RudaRead, RudaWrite, RudaReadExpand, RudaWriteExpand};
5use crate::collective::decompose::{RudaDecomposer, RudaDecomposerExpand};
6use super::{RudaBlockPrefix, RudaBlockPrefixExpand};
7
8#[ruda]
12pub fn scan_with_prefix<T: RudaRecord, O: RudaBinaryOp<T>, P: RudaBlockPrefix<T>>(
13 input: &RudaRecordArray<T>, output: &mut RudaRecordArray<T>, scratch: &mut RudaRecordShared<T>,
14 op: &O, callback: &mut P, valid: usize, #[comptime] threads: usize, #[comptime] items: usize,
15 #[comptime] exclusive: bool,
16) -> T {
17 inclusive_scan::<T, O>(input, output, scratch, op, valid, threads, items);
18 let aggregate = scratch.read(valid - 1);
19 if UNIT_POS < PLANE_DIM {
20 let value = callback.prefix(aggregate);
21 if UNIT_POS == 0 { scratch.write(threads * items, value); }
22 }
23 sync_ruda();
24 let prefix = scratch.read(threads * items);
25 #[unroll]
26 for item in 0..items {
27 let index = UNIT_POS as usize * items + item;
28 if index < valid {
29 let mut value = prefix;
30 if exclusive {
31 if index > 0 { value = op.combine(prefix, scratch.read(index - 1)); }
32 } else { value = op.combine(prefix, scratch.read(index)); }
33 output.write(item, value);
34 }
35 }
36 sync_ruda();
37 aggregate
38}
39
40#[ruda]
43pub fn inclusive_scan<T: RudaRecord, O: RudaBinaryOp<T>>(
44 input: &RudaRecordArray<T>, output: &mut RudaRecordArray<T>,
45 scratch: &mut RudaRecordShared<T>, op: &O, valid: usize,
46 #[comptime] threads: usize, #[comptime] items: usize,
47) {
48 let start = UNIT_POS as usize * items;
49 #[unroll]
50 for item in 0..items {
51 if start + item < valid { scratch.write(start + item, input.read(item)); }
52 }
53 sync_ruda();
54 let mut distance = 1usize;
55 while distance < threads * items {
56 #[unroll]
57 for item in 0..items {
58 let index = start + item;
59 if index < valid {
60 let mut value = scratch.read(index);
61 if index >= distance { value = op.combine(scratch.read(index - distance), value); }
62 output.write(item, value);
63 }
64 }
65 sync_ruda();
66 #[unroll]
67 for item in 0..items {
68 if start + item < valid { scratch.write(start + item, output.read(item)); }
69 }
70 sync_ruda();
71 distance *= 2;
72 }
73 #[unroll]
74 for item in 0..items {
75 if start + item < valid { output.write(item, scratch.read(start + item)); }
76 }
77 sync_ruda();
78}
79
80#[ruda]
81pub fn exclusive_scan<T: RudaRecord, O: RudaBinaryOp<T>>(
82 input: &RudaRecordArray<T>, output: &mut RudaRecordArray<T>, scratch: &mut RudaRecordShared<T>,
83 initial: T, op: &O, valid: usize, #[comptime] threads: usize, #[comptime] items: usize,
84) -> T {
85 inclusive_scan::<T, O>(input, output, scratch, op, valid, threads, items);
86 let aggregate = scratch.read(valid - 1);
87 #[unroll]
88 for item in 0..items {
89 let index = UNIT_POS as usize * items + item;
90 if index < valid {
91 let mut value = initial;
92 if index > 0 { value = op.combine(initial, scratch.read(index - 1)); }
93 output.write(item, value);
94 }
95 }
96 sync_ruda();
97 aggregate
98}
99
100#[ruda]
101pub fn reduce<T: RudaRecord, O: RudaBinaryOp<T>>(
102 input: &RudaRecordArray<T>, scratch: &mut RudaRecordShared<T>, op: &O,
103 valid: usize, #[comptime] threads: usize, #[comptime] items: usize,
104) -> T {
105 let start = UNIT_POS as usize * items;
106 #[unroll]
107 for item in 0..items {
108 if start + item < valid { scratch.write(start + item, input.read(item)); }
109 }
110 sync_ruda();
111 let mut distance = 1usize;
112 while distance < threads * items {
113 #[unroll]
114 for item in 0..items {
115 let index = start + item;
116 if index % (distance * 2) == 0 && index + distance < valid {
117 scratch.write(index, op.combine(scratch.read(index), scratch.read(index + distance)));
118 }
119 }
120 sync_ruda();
121 distance *= 2;
122 }
123 let aggregate = scratch.read(0);
124 sync_ruda();
125 aggregate
126}
127
128#[ruda]
129pub fn merge_rank<K: RudaRecord, C: RudaCompare<K>>(
130 source: &RudaRecordShared<K>, compare: &C, index: usize, valid: usize, run: usize,
131 base: usize,
132) -> usize {
133 let group = index / (run * 2) * (run * 2);
134 let middle = min(group + run, valid);
135 let end = min(group + run * 2, valid);
136 let key = source.read(base + index);
137 let left = index < middle;
138 let mut low = if left { middle } else { group };
139 let mut high = if left { end } else { middle };
140 let begin = low;
141 while low < high {
142 let probe = low + (high - low) / 2;
143 let other = source.read(base + probe);
144 let before = if left { compare.before(other, key) } else { !compare.before(key, other) };
145 if before { low = probe + 1; } else { high = probe; }
146 }
147 group + (if left { index - group } else { index - middle }) + low - begin
148}
149
150#[ruda]
152pub fn merge_sort_pairs<K: RudaRecord, V: RudaRecord, C: RudaCompare<K>>(
153 keys: &mut RudaRecordArray<K>, values: &mut RudaRecordArray<V>,
154 source: &mut RudaRecordShared<K>, destination: &mut RudaRecordShared<K>,
155 source_values: &mut RudaRecordShared<V>, destination_values: &mut RudaRecordShared<V>,
156 compare: &C, valid: usize, #[comptime] items: usize,
157) {
158 let start = UNIT_POS as usize * items;
159 #[unroll]
160 for item in 0..items {
161 if start + item < valid {
162 source.write(start + item, keys.read(item));
163 source_values.write(start + item, values.read(item));
164 }
165 }
166 sync_ruda();
167 let mut run = 1usize;
168 while run < valid {
169 #[unroll]
170 for item in 0..items {
171 let index = start + item;
172 if index < valid {
173 let rank = merge_rank::<K, C>(source, compare, index, valid, run, 0);
174 destination.write(rank, source.read(index));
175 destination_values.write(rank, source_values.read(index));
176 }
177 }
178 sync_ruda();
179 #[unroll]
180 for item in 0..items {
181 let index = start + item;
182 if index < valid {
183 source.write(index, destination.read(index));
184 source_values.write(index, destination_values.read(index));
185 }
186 }
187 sync_ruda();
188 run *= 2;
189 }
190 #[unroll]
191 for item in 0..items {
192 if start + item < valid {
193 keys.write(item, source.read(start + item));
194 values.write(item, source_values.read(start + item));
195 }
196 }
197 sync_ruda();
198}
199
200#[ruda]
203pub fn radix_sort_pairs<K: RudaRecord, V: RudaRecord, D: RudaDecomposer<K>>(
204 keys: &mut RudaRecordArray<K>, values: &mut RudaRecordArray<V>,
205 scratch: &mut RudaRecordShared<K>, scratch_values: &mut RudaRecordShared<V>,
206 ranks: &mut SharedMemory<u32>, decomposer: &D, valid: usize,
207 #[comptime] threads: usize, #[comptime] items: usize,
208 #[comptime] begin_bit: usize, #[comptime] end_bit: usize, #[comptime] descending: bool,
209) {
210 let mut flags = Array::<u32>::new(items);
211 let mut prefixes = Array::<u32>::new(items);
212 let start = UNIT_POS as usize * items;
213 let sum = crate::collective::RudaSum {};
214 for bit in begin_bit..end_bit {
215 #[unroll]
216 for item in 0..items {
217 flags[item] = 0;
218 if start + item < valid {
219 let digit = decomposer.bit(keys.read(item), bit);
220 flags[item] = u32::cast_from(digit == descending);
221 }
222 }
223 crate::block::inclusive_scan::<u32, crate::collective::RudaSum>(
224 &flags, &mut prefixes, ranks, &sum, threads * items, threads, items);
225 let total = ranks[threads * items - 1] as usize;
226 #[unroll]
227 for item in 0..items {
228 let index = start + item;
229 if index < valid {
230 let prefix = prefixes[item] as usize;
231 let rank = if flags[item] != 0 { prefix - 1 } else { total + index - prefix };
232 scratch.write(rank, keys.read(item));
233 scratch_values.write(rank, values.read(item));
234 }
235 }
236 sync_ruda();
237 #[unroll]
238 for item in 0..items {
239 if start + item < valid {
240 keys.write(item, scratch.read(start + item));
241 values.write(item, scratch_values.read(start + item));
242 }
243 }
244 sync_ruda();
245 }
246}
247
248#[ruda]
249pub fn merge_sort_keys<K: RudaRecord, C: RudaCompare<K>>(
250 keys: &mut RudaRecordArray<K>, source: &mut RudaRecordShared<K>, destination: &mut RudaRecordShared<K>,
251 compare: &C, valid: usize, #[comptime] items: usize,
252) {
253 let start = UNIT_POS as usize * items;
254 #[unroll]
255 for item in 0..items {
256 if start + item < valid { source.write(start + item, keys.read(item)); }
257 }
258 sync_ruda();
259 let mut run = 1usize;
260 while run < valid {
261 #[unroll]
262 for item in 0..items {
263 if start + item < valid {
264 let rank = merge_rank::<K, C>(source, compare, start + item, valid, run, 0);
265 destination.write(rank, source.read(start + item));
266 }
267 }
268 sync_ruda();
269 #[unroll]
270 for item in 0..items {
271 if start + item < valid { source.write(start + item, destination.read(start + item)); }
272 }
273 sync_ruda();
274 run *= 2;
275 }
276 #[unroll]
277 for item in 0..items {
278 if start + item < valid { keys.write(item, source.read(start + item)); }
279 }
280 sync_ruda();
281}
282
283#[ruda]
284pub fn radix_sort_keys<K: RudaRecord, D: RudaDecomposer<K>>(
285 keys: &mut RudaRecordArray<K>, scratch: &mut RudaRecordShared<K>, ranks: &mut SharedMemory<u32>,
286 decomposer: &D, valid: usize, #[comptime] threads: usize, #[comptime] items: usize,
287 #[comptime] begin_bit: usize, #[comptime] end_bit: usize, #[comptime] descending: bool,
288) {
289 let mut flags = Array::<u32>::new(items);
290 let mut prefixes = Array::<u32>::new(items);
291 let start = UNIT_POS as usize * items;
292 let sum = crate::collective::RudaSum {};
293 for bit in begin_bit..end_bit {
294 #[unroll]
295 for item in 0..items {
296 flags[item] = 0;
297 if start + item < valid { flags[item] = u32::cast_from(decomposer.bit(keys.read(item), bit) == descending); }
298 }
299 crate::block::inclusive_scan::<u32, crate::collective::RudaSum>(
300 &flags, &mut prefixes, ranks, &sum, threads * items, threads, items);
301 let total = ranks[threads * items - 1] as usize;
302 #[unroll]
303 for item in 0..items {
304 let index = start + item;
305 if index < valid {
306 let prefix = prefixes[item] as usize;
307 let rank = if flags[item] != 0 { prefix - 1 } else { total + index - prefix };
308 scratch.write(rank, keys.read(item));
309 }
310 }
311 sync_ruda();
312 #[unroll]
313 for item in 0..items {
314 if start + item < valid { keys.write(item, scratch.read(start + item)); }
315 }
316 sync_ruda();
317 }
318}
319
320#[ruda]
321pub fn topk_ranks<K: RudaRecord, D: RudaDecomposer<K>>(
322 keys: &RudaRecordArray<K>, selected: &mut Array<u32>, ranks: &mut Array<u32>, scratch: &mut SharedMemory<u32>,
323 decomposer: &D, k: usize, valid: usize, #[comptime] threads: usize, #[comptime] items: usize,
324 #[comptime] begin_bit: usize, #[comptime] end_bit: usize, #[comptime] largest: bool,
325) {
326 let start = UNIT_POS as usize * items;
327 let capacity = threads * items;
328 let mut candidates = Array::<u32>::new(items);
329 let mut preferred = Array::<u32>::new(items);
330 let sum = crate::collective::RudaSum {};
331 let mut remaining = min(k, valid);
332 #[unroll]
333 for item in 0..items {
334 candidates[item] = u32::cast_from(start + item < valid);
335 selected[item] = 0;
336 }
337 for step in 0..end_bit - begin_bit {
338 let bit = end_bit - 1 - step;
339 #[unroll]
340 for item in 0..items {
341 preferred[item] = 0;
342 if candidates[item] != 0 { preferred[item] = u32::cast_from(decomposer.bit(keys.read(item), bit) == largest); }
343 }
344 crate::block::inclusive_scan::<u32, crate::collective::RudaSum>(&preferred, ranks, scratch, &sum, capacity, threads, items);
345 let count = scratch[capacity - 1] as usize;
346 let accept = count <= remaining;
347 #[unroll]
348 for item in 0..items {
349 if candidates[item] != 0 {
350 if accept {
351 if preferred[item] != 0 { selected[item] = 1; candidates[item] = 0; }
352 } else { candidates[item] = preferred[item]; }
353 }
354 }
355 if accept { remaining -= count; }
356 sync_ruda();
357 }
358 crate::block::inclusive_scan::<u32, crate::collective::RudaSum>(&candidates, ranks, scratch, &sum, capacity, threads, items);
359 #[unroll]
360 for item in 0..items {
361 if candidates[item] != 0 && ranks[item] as usize <= remaining { selected[item] = 1; }
362 }
363 crate::block::inclusive_scan::<u32, crate::collective::RudaSum>(selected, ranks, scratch, &sum, capacity, threads, items);
364}
365
366#[ruda]
367pub fn topk_pairs<K: RudaRecord, V: RudaRecord, D: RudaDecomposer<K>>(
368 keys: &mut RudaRecordArray<K>, values: &mut RudaRecordArray<V>, scratch: &mut RudaRecordShared<K>,
369 scratch_values: &mut RudaRecordShared<V>, scratch_ranks: &mut SharedMemory<u32>, decomposer: &D,
370 k: usize, valid: usize, #[comptime] threads: usize, #[comptime] items: usize,
371 #[comptime] begin_bit: usize, #[comptime] end_bit: usize, #[comptime] largest: bool,
372) {
373 let mut selected = Array::<u32>::new(items);
374 let mut ranks = Array::<u32>::new(items);
375 topk_ranks::<K, D>(keys, &mut selected, &mut ranks, scratch_ranks, decomposer, k, valid, threads, items, begin_bit, end_bit, largest);
376 #[unroll]
377 for item in 0..items {
378 if selected[item] != 0 {
379 let rank = ranks[item] as usize - 1;
380 scratch.write(rank, keys.read(item));
381 scratch_values.write(rank, values.read(item));
382 }
383 }
384 sync_ruda();
385 #[unroll]
386 for item in 0..items {
387 let index = UNIT_POS as usize * items + item;
388 if index < min(k, valid) { keys.write(item, scratch.read(index)); values.write(item, scratch_values.read(index)); }
389 }
390 sync_ruda();
391}
392
393#[ruda]
394pub fn topk_keys<K: RudaRecord, D: RudaDecomposer<K>>(
395 keys: &mut RudaRecordArray<K>, scratch: &mut RudaRecordShared<K>, scratch_ranks: &mut SharedMemory<u32>,
396 decomposer: &D, k: usize, valid: usize, #[comptime] threads: usize, #[comptime] items: usize,
397 #[comptime] begin_bit: usize, #[comptime] end_bit: usize, #[comptime] largest: bool,
398) {
399 let mut selected = Array::<u32>::new(items);
400 let mut ranks = Array::<u32>::new(items);
401 topk_ranks::<K, D>(keys, &mut selected, &mut ranks, scratch_ranks, decomposer, k, valid, threads, items, begin_bit, end_bit, largest);
402 #[unroll]
403 for item in 0..items {
404 if selected[item] != 0 { scratch.write(ranks[item] as usize - 1, keys.read(item)); }
405 }
406 sync_ruda();
407 #[unroll]
408 for item in 0..items {
409 let index = UNIT_POS as usize * items + item;
410 if index < min(k, valid) { keys.write(item, scratch.read(index)); }
411 }
412 sync_ruda();
413}