Skip to main content

ruprim/block/
record.rs

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/// Stateful prefix callback; scratch has threads * items + 1 slots. As with
9/// the scalar API, the first native subgroup calls the callback and thread
10/// zero's returned prefix is used. Valid input is positive.
11#[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/// Ordered blocked scan of records. All block threads participate; valid is
41/// positive and at most threads * items. Scratch has that many record slots.
42#[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/// Stable merge sorting; every key and value may itself be a nested record.
151#[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/// Stable LSD radix sorting over an arbitrary-width decomposition. Scratch
201/// arrays have threads * items entries; bit interval is half-open and valid.
202#[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}