Skip to main content

ruprim/collective/
record.rs

1use core::marker::PhantomData;
2use ruda_kernel::dsl as kernel_dsl;
3use ruda_kernel::dsl::prelude::*;
4
5type ArrayExpand<T> = NativeExpand<Array<T>>;
6
7/// A device value with an explicit byte layout. Composite implementations
8/// lower field operations to scalar IR, retaining the record's native layout.
9#[ruda]
10pub trait RudaRecord: Copy + 'static + RudaType<ExpandType: Assign> {
11    const SIZE: usize;
12    const ALIGN: usize;
13    fn load(address: u64) -> Self;
14    fn store(address: u64, value: Self);
15    fn shuffle(value: Self, lane: u32) -> Self;
16}
17
18macro_rules! scalar_record {
19    ($($ty:ty),* $(,)?) => {$(
20        #[ruda]
21        impl RudaRecord for $ty {
22            const SIZE: usize = core::mem::size_of::<Self>();
23            const ALIGN: usize = core::mem::align_of::<Self>();
24            fn load(address: u64) -> Self { native_load::<Self>(address) }
25            fn store(address: u64, value: Self) { native_store(address, value); }
26            fn shuffle(value: Self, lane: u32) -> Self { plane_shuffle(value, lane) }
27        }
28    )*};
29}
30scalar_record!(bool, u8, u16, u32, u64, i8, i16, i32, i64, f32, f64, half::f16, half::bf16);
31
32/// Define a C-layout record, including nested records, for primitive operators.
33#[macro_export]
34macro_rules! ruda_record {
35    ($(#[$meta:meta])* $vis:vis struct $name:ident { $($field_vis:vis $field:ident : $ty:ty),* $(,)? }) => {
36        $(#[$meta])*
37        #[repr(C)]
38        #[derive(Clone, Copy, ruda_kernel::dsl::RudaType, ruda_kernel::dsl::prelude::RudaTypeMut, ruda_kernel::dsl::RudaLaunch)]
39        $vis struct $name { $($field_vis $field: $ty),* }
40        #[ruda_kernel::dsl::ruda]
41        impl $crate::collective::record::RudaRecord for $name {
42            const SIZE: usize = core::mem::size_of::<Self>();
43            const ALIGN: usize = core::mem::align_of::<Self>();
44            fn load(address: u64) -> Self {
45                $name { $($field: <$ty as $crate::collective::record::RudaRecord>::load(
46                    address + ruda_kernel::dsl::comptime![core::mem::offset_of!(Self, $field) as u64])),* }
47            }
48            fn store(address: u64, value: Self) {
49                $(<$ty as $crate::collective::record::RudaRecord>::store(
50                    address + ruda_kernel::dsl::comptime![core::mem::offset_of!(Self, $field) as u64], value.$field);)*
51            }
52            fn shuffle(value: Self, lane: u32) -> Self {
53                $name { $($field: <$ty as $crate::collective::record::RudaRecord>::shuffle(value.$field, lane)),* }
54            }
55        }
56    };
57}
58
59#[ruda]
60pub trait RudaRead<T: RudaType>: RudaType { fn read(&self, index: usize) -> T; }
61
62#[ruda]
63pub trait RudaWrite<T: RudaType>: RudaRead<T> { fn write(&mut self, index: usize, value: T); }
64
65#[ruda]
66impl<T: RudaPrimitive> RudaRead<T> for Array<T> {
67    fn read(&self, index: usize) -> T { self[index] }
68}
69#[ruda]
70impl<T: RudaPrimitive> RudaWrite<T> for Array<T> {
71    fn write(&mut self, index: usize, value: T) { self[index] = value; }
72}
73#[ruda]
74impl<T: RudaPrimitive> RudaRead<T> for SharedMemory<T> {
75    fn read(&self, index: usize) -> T { self[index] }
76}
77#[ruda]
78impl<T: RudaPrimitive> RudaWrite<T> for SharedMemory<T> {
79    fn write(&mut self, index: usize, value: T) { self[index] = value; }
80}
81
82#[derive(RudaType)]
83pub struct RudaRecordArray<T: RudaRecord> {
84    bytes: Array<u8>,
85    #[ruda(comptime)]
86    marker: PhantomData<T>,
87}
88
89#[ruda]
90impl<T: RudaRecord> RudaRecordArray<T> {
91    pub fn new(#[comptime] length: usize) -> Self {
92        let bytes = Array::<u8>::new(comptime![(length * T::SIZE + T::ALIGN).max(1)]);
93        RudaRecordArray::<T> { bytes, marker: PhantomData }
94    }
95    pub fn address(&self, index: usize) -> u64 {
96        let base = native_address(&self.bytes.to_slice(), 0);
97        let align = comptime![T::ALIGN as u64];
98        (base + align - 1) / align * align + index as u64 * comptime![T::SIZE as u64]
99    }
100}
101#[ruda]
102impl<T: RudaRecord> RudaRead<T> for RudaRecordArray<T> {
103    fn read(&self, index: usize) -> T { T::load(self.address(index)) }
104}
105#[ruda]
106impl<T: RudaRecord> RudaWrite<T> for RudaRecordArray<T> {
107    fn write(&mut self, index: usize, value: T) { T::store(self.address(index), value); }
108}
109
110#[derive(RudaType)]
111pub struct RudaRecordShared<T: RudaRecord> {
112    bytes: SharedMemory<u8>,
113    #[ruda(comptime)]
114    marker: PhantomData<T>,
115}
116
117#[ruda]
118impl<T: RudaRecord> RudaRecordShared<T> {
119    pub fn new(#[comptime] length: usize) -> Self {
120        let bytes = SharedMemory::<u8>::new_aligned(comptime![(length * T::SIZE).max(1)], comptime![T::ALIGN]);
121        RudaRecordShared::<T> { bytes, marker: PhantomData }
122    }
123    pub fn address(&self, index: usize) -> u64 {
124        let stride = comptime![T::SIZE];
125        let index = usize::cast_from(index);
126        native_address(&self.bytes.to_slice(), index * stride)
127    }
128}
129#[ruda]
130impl<T: RudaRecord> RudaRead<T> for RudaRecordShared<T> {
131    fn read(&self, index: usize) -> T { T::load(self.address(index)) }
132}
133#[ruda]
134impl<T: RudaRecord> RudaWrite<T> for RudaRecordShared<T> {
135    fn write(&mut self, index: usize, value: T) { T::store(self.address(index), value); }
136}
137
138/// Stable reference into an existing allocation, not a reference to a local copy.
139#[derive(Clone, RudaType, RudaTypeMut)]
140pub struct RudaReference<T: RudaRecord> {
141    pub address: u64,
142    #[ruda(comptime)]
143    marker: PhantomData<T>,
144}
145
146#[ruda]
147impl<T: RudaRecord> RudaReference<T> {
148    pub fn new(address: u64) -> Self { RudaReference::<T> { address, marker: PhantomData } }
149    pub fn read(&self) -> T { T::load(self.address) }
150    pub fn write(&self, value: T) { T::store(self.address, value); }
151}
152
153/// A device-native strided record range. Addresses and strides are in bytes.
154/// The caller retains the allocation and guarantees alignment and bounds for
155/// every access, including any references retained by an operation.
156#[derive(Clone, Copy, RudaType, RudaLaunch)]
157pub struct RudaNativeRecords {
158    pub address: u64,
159    pub stride: u64,
160}
161
162#[ruda]
163impl<T: RudaRecord> RudaRead<T> for RudaNativeRecords {
164    fn read(&self, index: usize) -> T { T::load(self.address + index as u64 * self.stride) }
165}
166#[ruda]
167impl<T: RudaRecord> RudaWrite<T> for RudaNativeRecords {
168    fn write(&mut self, index: usize, value: T) { T::store(self.address + index as u64 * self.stride, value); }
169}
170
171#[ruda]
172pub trait RudaAddress<T: RudaRecord>: RudaRead<T> {
173    fn reference(&self, index: usize) -> RudaReference<T>;
174}
175
176#[ruda]
177impl<T: RudaRecord> RudaAddress<T> for RudaNativeRecords {
178    fn reference(&self, index: usize) -> RudaReference<T> {
179        RudaReference::<T>::new(self.address + index as u64 * self.stride)
180    }
181}
182
183/// Recursive product: nesting has no fixed arity limit.
184#[derive(Clone, RudaType)]
185pub struct RudaPair<A: RudaType, B: RudaType> {
186    pub first: A,
187    pub second: B,
188}
189
190#[derive(Clone, RudaType, RudaLaunch)]
191pub struct RudaZip<I: RudaType, J: RudaType> {
192    pub first: I,
193    pub second: J,
194}
195
196#[ruda]
197impl<A: RudaType, B: RudaType, I: RudaRead<A>, J: RudaRead<B>> RudaRead<RudaPair<A, B>> for RudaZip<I, J> {
198    fn read(&self, index: usize) -> RudaPair<A, B> {
199        RudaPair::<A, B> { first: self.first.read(index), second: self.second.read(index) }
200    }
201}
202
203#[derive(Clone, RudaType, RudaLaunch)]
204pub struct RudaReferences<T: RudaRecord, I: RudaAddress<T>> {
205    pub input: I,
206    #[ruda(comptime)]
207    pub marker: PhantomData<T>,
208}
209
210#[ruda]
211impl<T: RudaRecord, I: RudaAddress<T>> RudaRead<RudaReference<T>> for RudaReferences<T, I> {
212    fn read(&self, index: usize) -> RudaReference<T> { self.input.reference(index) }
213}