subetha_pointers/
kstep_pointer.rs1use std::marker::PhantomData;
31
32#[derive(Debug)]
34pub struct KStepPointer<T> {
35 base: *const T,
36 k_step: u8,
37 _phantom: PhantomData<*const T>,
38}
39
40unsafe impl<T: Send> Send for KStepPointer<T> {}
41unsafe impl<T: Sync> Sync for KStepPointer<T> {}
42
43impl<T> KStepPointer<T> {
44 pub const SIGNATURE: subetha_core::AxisMask = subetha_core::AxisMask::from_axes(
48 &[subetha_core::Axis::Stride],
49 );
50
51 pub const unsafe fn new(base: *const T, k_step: u8) -> Self {
60 Self { base, k_step, _phantom: PhantomData }
61 }
62
63 pub const unsafe fn tight(base: *const T) -> Self {
70 unsafe { Self::new(base, 0) }
71 }
72
73 pub const unsafe fn cache_line(base: *const T) -> Self {
81 let t_size = std::mem::size_of::<T>();
83 let k = if t_size >= 64 { 0 }
84 else if t_size >= 32 { 1 }
85 else if t_size >= 16 { 2 }
86 else if t_size >= 8 { 3 }
87 else if t_size >= 4 { 4 }
88 else if t_size >= 2 { 5 }
89 else { 6 };
90 unsafe { Self::new(base, k) }
91 }
92
93 #[inline]
94 pub const fn base(&self) -> *const T { self.base }
95
96 #[inline]
97 pub const fn k_step(&self) -> u8 { self.k_step }
98
99 #[inline]
101 pub const fn stride(&self) -> usize {
102 std::mem::size_of::<T>() << self.k_step
103 }
104
105 #[inline]
111 pub unsafe fn at(&self, i: usize) -> *const T {
112 let offset_bytes = i * self.stride();
113 unsafe { (self.base as *const u8).add(offset_bytes) as *const T }
114 }
115
116 #[inline]
123 pub unsafe fn get(&self, i: usize) -> &T {
124 unsafe { &*self.at(i) }
125 }
126
127 pub unsafe fn iter(&self, count: usize) -> StridedIter<'_, T> {
133 StridedIter { ptr: *self, i: 0, count, _life: PhantomData }
134 }
135}
136
137impl<T> Clone for KStepPointer<T> {
138 fn clone(&self) -> Self { *self }
139}
140impl<T> Copy for KStepPointer<T> {}
141
142pub struct StridedIter<'a, T> {
144 ptr: KStepPointer<T>,
145 i: usize,
146 count: usize,
147 _life: PhantomData<&'a T>,
148}
149
150impl<'a, T> Iterator for StridedIter<'a, T> {
151 type Item = &'a T;
152 fn next(&mut self) -> Option<&'a T> {
153 if self.i >= self.count { return None; }
154 let r = unsafe { &*self.ptr.at(self.i) };
155 self.i += 1;
156 Some(r)
157 }
158 fn size_hint(&self) -> (usize, Option<usize>) {
159 let n = self.count - self.i;
160 (n, Some(n))
161 }
162}
163
164impl<'a, T> ExactSizeIterator for StridedIter<'a, T> {}
165
166#[cfg(test)]
167mod tests {
168 use super::*;
169
170 #[test]
171 fn tight_stride_walks_contiguous_array() {
172 let data: Vec<u64> = (0..10u64).collect();
173 let p = unsafe { KStepPointer::tight(data.as_ptr()) };
174 assert_eq!(p.k_step(), 0);
175 assert_eq!(p.stride(), 8);
176 for i in 0..10 {
177 assert_eq!(unsafe { *p.get(i) }, i as u64);
178 }
179 }
180
181 #[test]
182 fn k_step_1_skips_every_other() {
183 let data: Vec<u64> = (0..10u64).collect();
184 let p = unsafe { KStepPointer::new(data.as_ptr(), 1) };
185 assert_eq!(p.stride(), 16);
186 for i in 0..5 {
188 assert_eq!(unsafe { *p.get(i) }, (i * 2) as u64);
189 }
190 }
191
192 #[test]
193 fn k_step_2_skips_by_four() {
194 let data: Vec<u64> = (0..16u64).collect();
195 let p = unsafe { KStepPointer::new(data.as_ptr(), 2) };
196 assert_eq!(p.stride(), 32);
197 for i in 0..4 {
198 assert_eq!(unsafe { *p.get(i) }, (i * 4) as u64);
199 }
200 }
201
202 #[test]
203 fn cache_line_stride_picks_correct_k() {
204 let data: Vec<u64> = (0..64u64).collect();
206 let p = unsafe { KStepPointer::cache_line(data.as_ptr()) };
207 assert_eq!(p.k_step(), 3, "u64 cache_line should be k=3");
208 assert_eq!(p.stride(), 64);
209 for i in 0..8 {
211 assert_eq!(unsafe { *p.get(i) }, (i * 8) as u64);
212 }
213 }
214
215 #[test]
216 fn cache_line_stride_for_u8() {
217 let data: Vec<u8> = (0u8..64).collect();
219 let p = unsafe { KStepPointer::<u8>::cache_line(data.as_ptr()) };
220 assert_eq!(p.k_step(), 6);
221 assert_eq!(p.stride(), 64);
222 }
223
224 #[test]
225 fn strided_iter_yields_correct_elements() {
226 let data: Vec<u64> = (0..20u64).collect();
227 let p = unsafe { KStepPointer::new(data.as_ptr(), 2) };
228 let collected: Vec<u64> = unsafe { p.iter(5) }.copied().collect();
229 assert_eq!(collected, vec![0, 4, 8, 12, 16]);
230 }
231
232 #[test]
233 fn matrix_row_stride_workflow() {
234 let matrix: Vec<u64> = (0..16u64).collect();
236 let col0 = unsafe { KStepPointer::new(matrix.as_ptr(), 2) };
237 let col0_vals: Vec<u64> = unsafe { col0.iter(4) }.copied().collect();
238 assert_eq!(col0_vals, vec![0, 4, 8, 12]);
239 }
240
241 #[test]
242 fn layout_is_16_bytes() {
243 assert_eq!(std::mem::size_of::<KStepPointer<u64>>(), 16);
245 }
246}