1use ruda_kernel::dsl as kernel_dsl;
2use ruda_kernel::dsl::prelude::*;
3use crate::collective::record::{RudaRecord, RudaRecordShared, RudaReference};
4
5#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
7pub struct RudaRakingLayout {
8 pub shared_elements: usize,
9 pub max_raking_threads: usize,
10 pub segment_length: usize,
11 pub raking_threads: usize,
12 pub has_conflicts: bool,
13 pub conflict_degree: usize,
14 pub segment_padding: bool,
15 pub grid_elements: usize,
16 pub unguarded: bool,
17}
18
19impl RudaRakingLayout {
20 pub fn new(threads: usize, subgroup_width: usize, shared_banks: usize) -> Result<Self, &'static str> {
21 if threads == 0 || subgroup_width == 0 || shared_banks == 0 { return Err("raking dimensions must be positive"); }
22 let max_raking_threads = threads.min(subgroup_width);
23 let segment_length = threads.div_ceil(max_raking_threads);
24 let raking_threads = threads.div_ceil(segment_length);
25 let has_conflicts = shared_banks % segment_length == 0;
26 let conflict_degree = if has_conflicts {
27 max_raking_threads.checked_mul(segment_length).ok_or("raking size overflow")? / shared_banks
28 } else { 1 };
29 let segment_padding = segment_length % 2 == 0 && segment_length > 2;
30 let stride = segment_length.checked_add(usize::from(segment_padding)).ok_or("raking size overflow")?;
31 let grid_elements = raking_threads.checked_mul(stride).ok_or("raking size overflow")?;
32 Ok(Self {
33 shared_elements: threads, max_raking_threads, segment_length, raking_threads,
34 has_conflicts, conflict_degree, segment_padding, grid_elements,
35 unguarded: threads % raking_threads == 0,
36 })
37 }
38}
39
40#[ruda]
41pub fn placement_index(thread: usize, #[comptime] segment_length: usize, #[comptime] padding: bool) -> usize {
42 let mut index = thread;
43 if padding { index += thread / segment_length; }
44 index
45}
46
47#[ruda]
48pub fn raking_index(thread: usize, #[comptime] segment_length: usize, #[comptime] padding: bool) -> usize {
49 let stride = comptime![segment_length + usize::from(padding)];
50 thread * stride
51}
52
53#[derive(RudaType)]
54pub struct RudaRakingGrid<T: RudaRecord> {
55 storage: RudaRecordShared<T>,
56 #[ruda(comptime)]
57 segment_length: usize,
58 #[ruda(comptime)]
59 padding: bool,
60}
61
62#[ruda]
63impl<T: RudaRecord> RudaRakingGrid<T> {
64 pub fn new(#[comptime] layout: RudaRakingLayout) -> Self {
65 RudaRakingGrid::<T> {
66 storage: RudaRecordShared::<T>::new(comptime![layout.grid_elements]),
67 segment_length: comptime![layout.segment_length],
68 padding: comptime![layout.segment_padding],
69 }
70 }
71
72 pub fn placement(&self, thread: usize) -> RudaReference<T> {
73 let index = placement_index(thread, self.segment_length, self.padding);
74 RudaReference::<T>::new(self.storage.address(index))
75 }
76
77 pub fn segment(&self, thread: usize) -> RudaReference<T> {
79 let index = raking_index(thread, self.segment_length, self.padding);
80 RudaReference::<T>::new(self.storage.address(index))
81 }
82
83 pub fn item(&self, thread: usize, item: usize) -> RudaReference<T> {
84 let index = raking_index(thread, self.segment_length, self.padding) + item;
85 RudaReference::<T>::new(self.storage.address(index))
86 }
87}