Skip to main content

ruprim/warp/
io.rs

1use ruda_kernel::dsl as kernel_dsl;
2use ruda_kernel::dsl::prelude::*;
3
4/// Load a logical warp tile into per-lane registers. Invalid items receive
5/// `padding`. The tile starts at `offset`; `valid_items` is its valid length.
6#[ruda]
7pub fn load<T: RudaPrimitive>(
8    input: &Array<T>,
9    output: &mut Array<T>,
10    offset: usize,
11    valid_items: usize,
12    padding: T,
13    #[comptime] items_per_lane: usize,
14    #[comptime] width: u32,
15    #[comptime] striped: bool,
16) {
17    let lane = (UNIT_POS_PLANE % width) as usize;
18    #[unroll]
19    for item in 0..items_per_lane {
20        let index = if striped { item * width as usize + lane } else { lane * items_per_lane + item };
21        let mut value = padding;
22        if index < valid_items {
23            value = input[offset + index];
24        }
25        output[item] = value;
26    }
27}
28
29/// Store a register tile in blocked or striped order, preserving invalid tails.
30#[ruda]
31pub fn store<T: RudaPrimitive>(
32    input: &Array<T>,
33    output: &mut Array<T>,
34    offset: usize,
35    valid_items: usize,
36    #[comptime] items_per_lane: usize,
37    #[comptime] width: u32,
38    #[comptime] striped: bool,
39) {
40    let lane = (UNIT_POS_PLANE % width) as usize;
41    #[unroll]
42    for item in 0..items_per_lane {
43        let index = if striped { item * width as usize + lane } else { lane * items_per_lane + item };
44        if index < valid_items {
45            output[offset + index] = input[item];
46        }
47    }
48}