1use ruda_kernel::dsl as kernel_dsl;
2use ruda_kernel::dsl::prelude::*;
3use ruda_kernel::dsl::prelude::barrier::{Barrier, BarrierToken};
4
5#[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#[ruda]
16pub fn commit(barrier: &Barrier) -> BarrierToken { barrier.arrive() }
17
18#[ruda]
19pub fn wait(barrier: &Barrier, token: BarrierToken) { barrier.wait(token); }
20
21#[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#[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#[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}