ruprim/block/
run_length.rs1use ruda_kernel::dsl as kernel_dsl;
2use ruda_kernel::dsl::prelude::*;
3use crate::collective::RudaSum;
4use crate::collective::record::{RudaRead, RudaReadExpand, RudaWrite, RudaWriteExpand};
5
6#[ruda]
7pub fn prepare_access<T: RudaType, L: Numeric, I: RudaRead<T>, N: RudaRead<L>, S: RudaWrite<T>>(
8 run_values: &I, run_lengths: &N, values: &mut S, ends: &mut SharedMemory<u64>,
9 num_runs: usize, #[comptime] threads: usize, #[comptime] runs_per_thread: usize,
10) -> u64 {
11 let start = UNIT_POS as usize * runs_per_thread;
12 let mut lengths = Array::<u64>::new(runs_per_thread);
13 let mut prefixes = Array::<u64>::new(runs_per_thread);
14 #[unroll]
15 for item in 0..runs_per_thread {
16 lengths[item] = 0;
17 if start + item < num_runs {
18 values.write(start + item, run_values.read(item));
19 lengths[item] = u64::cast_from(run_lengths.read(item));
20 }
21 }
22 let sum = RudaSum {};
23 crate::block::inclusive_scan::<u64, RudaSum>(&lengths, &mut prefixes, ends, &sum,
24 threads * runs_per_thread, threads, runs_per_thread);
25 let total = ends[threads * runs_per_thread - 1];
26 sync_ruda();
27 total
28}
29
30#[ruda]
31pub fn decode_access<T: RudaType, C: Numeric, S: RudaRead<T>, W: RudaWrite<T>, O: RudaWrite<C>>(
32 values: &S, ends: &SharedMemory<u64>, output: &mut W, relative_offsets: &mut O,
33 window_offset: u64, num_runs: usize, #[comptime] items_per_thread: usize,
34) {
35 let start = UNIT_POS as u64 * items_per_thread as u64;
36 let mut total = 0u64;
37 if num_runs > 0 { total = ends[num_runs - 1]; }
38 #[unroll]
39 for item in 0..items_per_thread {
40 let relative = start + item as u64;
41 if window_offset < total && relative < total - window_offset {
42 let position = window_offset + relative;
43 let mut low = 0usize;
44 let mut high = num_runs;
45 while low < high {
46 let middle = low + (high - low) / 2;
47 if ends[middle] <= position { low = middle + 1; } else { high = middle; }
48 }
49 let mut begin = 0u64;
50 if low > 0 { begin = ends[low - 1]; }
51 output.write(item, values.read(low));
52 relative_offsets.write(item, C::cast_from(position - begin));
53 }
54 }
55}
56
57#[ruda]
60pub fn prepare<T: RudaPrimitive>(
61 run_values: &Array<T>, run_lengths: &Array<u64>,
62 values: &mut SharedMemory<T>, ends: &mut SharedMemory<u64>,
63 num_runs: usize, #[comptime] threads: usize, #[comptime] runs_per_thread: usize,
64) -> u64 {
65 let start = UNIT_POS as usize * runs_per_thread;
66 let mut lengths = Array::<u64>::new(runs_per_thread);
67 let mut prefixes = Array::<u64>::new(runs_per_thread);
68 #[unroll]
69 for item in 0..runs_per_thread {
70 let mut length = 0u64;
71 if start + item < num_runs {
72 values[start + item] = run_values[item];
73 length = run_lengths[item];
74 }
75 lengths[item] = length;
76 }
77 let sum = RudaSum {};
78 crate::block::inclusive_scan::<u64, RudaSum>(&lengths, &mut prefixes, ends, &sum,
79 threads * runs_per_thread, threads, runs_per_thread);
80 let total = ends[threads * runs_per_thread - 1];
81 sync_ruda();
82 total
83}
84
85#[ruda]
88pub fn decode_window<T: RudaPrimitive>(
89 values: &SharedMemory<T>, ends: &SharedMemory<u64>,
90 output: &mut Array<T>, relative_offsets: &mut Array<u64>,
91 window_offset: u64, num_runs: usize,
92 #[comptime] items_per_thread: usize,
93) {
94 let start = UNIT_POS as u64 * items_per_thread as u64;
95 let mut total = 0u64;
96 if num_runs > 0 { total = ends[num_runs - 1]; }
97 #[unroll]
98 for item in 0..items_per_thread {
99 let relative = start + item as u64;
100 if window_offset < total && relative < total - window_offset {
101 let position = window_offset + relative;
102 let mut low = 0usize;
103 let mut high = num_runs;
104 while low < high {
105 let middle = low + (high - low) / 2;
106 if ends[middle] <= position { low = middle + 1; } else { high = middle; }
107 }
108 let mut begin = 0u64;
109 if low > 0 { begin = ends[low - 1]; }
110 output[item] = values[low];
111 relative_offsets[item] = position - begin;
112 }
113 }
114}