Skip to main content

ruprim/block/
run_length.rs

1use 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/// Prepare inclusive run ends once, then reuse them for any decoded window.
58/// Returns the total decoded length. The run lengths' sum must fit in U64.
59#[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/// Decode a blocked window, also returning each item's offset within its run.
86/// Entries beyond the total decoded length are left unchanged.
87#[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}