Skip to main content

ruprim/block/
raking.rs

1use ruda_kernel::dsl as kernel_dsl;
2use ruda_kernel::dsl::prelude::*;
3use crate::collective::record::{RudaRecord, RudaRecordShared, RudaReference};
4
5/// Raking-grid geometry with explicit subgroup and shared-bank dimensions.
6#[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    /// First record in a raking thread's contiguous, aligned segment.
78    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}