Skip to main content

ruprim/block/
io.rs

1use ruda_kernel::dsl as kernel_dsl;
2use ruda_kernel::dsl::prelude::*;
3use ruda_kernel::dsl::prelude::barrier::{Barrier, BarrierToken};
4
5/// Enqueue a collective asynchronous global-to-shared copy. All block threads
6/// use the same initialized block barrier and slices. Source length must not
7/// exceed the destination; destination is not readable until wait completes.
8#[ruda]
9pub fn copy_async<T: RudaPrimitive>(barrier: &Barrier, input: &Slice<T>, output: &mut SliceMut<T>) {
10    barrier.memcpy_async_cooperative(input, output);
11}
12
13/// Commit one phase after one or more copy_async calls. Every participating
14/// thread commits once; the returned token identifies that phase for wait.
15#[ruda]
16pub fn commit(barrier: &Barrier) -> BarrierToken { barrier.arrive() }
17
18#[ruda]
19pub fn wait(barrier: &Barrier, token: BarrierToken) { barrier.wait(token); }
20
21/// Load a tile in blocked or striped order, assigning `padding` to the tail.
22#[ruda]
23pub fn load<T: RudaPrimitive>(
24    input: &Array<T>,
25    output: &mut Array<T>,
26    offset: usize,
27    valid_items: usize,
28    padding: T,
29    #[comptime] threads: usize,
30    #[comptime] items_per_thread: usize,
31    #[comptime] striped: bool,
32) {
33    #[unroll]
34    for item in 0..items_per_thread {
35        let index = if striped { item * threads + UNIT_POS as usize } else { UNIT_POS as usize * items_per_thread + item };
36        let mut value = padding;
37        if index < valid_items {
38            value = input[offset + index];
39        }
40        output[item] = value;
41    }
42}
43
44/// Store only valid elements of a blocked or striped register tile.
45#[ruda]
46pub fn store<T: RudaPrimitive>(
47    input: &Array<T>,
48    output: &mut Array<T>,
49    offset: usize,
50    valid_items: usize,
51    #[comptime] threads: usize,
52    #[comptime] items_per_thread: usize,
53    #[comptime] striped: bool,
54) {
55    #[unroll]
56    for item in 0..items_per_thread {
57        let index = if striped { item * threads + UNIT_POS as usize } else { UNIT_POS as usize * items_per_thread + item };
58        if index < valid_items {
59            output[offset + index] = input[item];
60        }
61    }
62}
63
64/// Cooperatively load a tile into shared memory; every thread participates.
65#[ruda]
66pub fn load_to_shared<T: RudaPrimitive>(
67    input: &Array<T>,
68    output: &mut SharedMemory<T>,
69    offset: usize,
70    valid_items: usize,
71    padding: T,
72    #[comptime] threads: usize,
73    #[comptime] tile_items: usize,
74) {
75    let mut index = UNIT_POS as usize;
76    while index < tile_items {
77        let mut value = padding;
78        if index < valid_items {
79            value = input[offset + index];
80        }
81        output[index] = value;
82        index += threads;
83    }
84    sync_ruda();
85}