hurray_core/layout/addressing/
strided.rs1use crate::layout::StridedLayout;
6use crate::{Error, Result, Shape};
7
8use super::{validate_index, ElementAddress};
9
10impl ElementAddress for StridedLayout {
14 fn element_offset(&self, index: &[u64], shape: &Shape) -> Result<u64> {
15 validate_index(index, shape)?;
16 if self.strides.len() != shape.rank() {
17 return Err(Error::IndexRankMismatch {
18 index_rank: self.strides.len(),
19 shape_rank: shape.rank(),
20 });
21 }
22 let mut sum: i64 = 0;
23 for (&idx_val, &stride) in index.iter().zip(self.strides.iter()) {
24 let term = (idx_val as i64)
25 .checked_mul(stride)
26 .ok_or(Error::AddressOverflow)?;
27 sum = sum.checked_add(term).ok_or(Error::AddressOverflow)?;
28 }
29 Ok(sum as u64)
30 }
31}
32
33#[cfg(test)]
34mod tests {
35 use super::*;
36 use crate::{Error, Shape, DYNAMIC};
37
38 #[test]
43 fn strided_row_major_equivalent() {
44 let shape = Shape::new(vec![3, 4]).unwrap();
45 let layout = StridedLayout::new(vec![4, 1]);
46 assert_eq!(layout.element_offset(&[1, 2], &shape).unwrap(), 6);
47 }
48
49 #[test]
51 fn strided_col_major_equivalent() {
52 let shape = Shape::new(vec![3, 4]).unwrap();
53 let layout = StridedLayout::new(vec![1, 3]);
54 assert_eq!(layout.element_offset(&[1, 2], &shape).unwrap(), 7);
55 }
56
57 #[test]
61 fn strided_negative_stride_first_element_is_zero() {
62 let shape = Shape::new(vec![3, 4]).unwrap();
63 let layout = StridedLayout::new(vec![-4, 1]);
64 assert_eq!(layout.element_offset(&[0, 0], &shape).unwrap(), 0);
65 }
66
67 #[test]
71 fn strided_negative_stride_last_element() {
72 let shape = Shape::new(vec![3, 4]).unwrap();
73 let layout = StridedLayout::new(vec![-4, 1]);
74 let offset = layout.element_offset(&[2, 3], &shape).unwrap();
75 assert_eq!(offset, (-5i64) as u64);
77 }
78
79 #[test]
82 fn strided_zero_stride_broadcast() {
83 let shape = Shape::new(vec![5, 4]).unwrap();
84 let layout = StridedLayout::new(vec![0, 1]);
85 assert_eq!(layout.element_offset(&[3, 2], &shape).unwrap(), 2);
86 assert_eq!(layout.element_offset(&[0, 2], &shape).unwrap(), 2);
88 assert_eq!(layout.element_offset(&[4, 2], &shape).unwrap(), 2);
89 }
90
91 #[test]
94 fn strided_strides_rank_mismatch() {
95 let shape = Shape::new(vec![3, 4]).unwrap();
96 let layout = StridedLayout::new(vec![4, 1, 0]);
98 let err = layout.element_offset(&[1, 2], &shape).unwrap_err();
99 assert!(
100 matches!(err, Error::IndexRankMismatch { .. }),
101 "expected IndexRankMismatch, got {err:?}"
102 );
103 }
104
105 #[test]
107 fn strided_index_out_of_range() {
108 let shape = Shape::new(vec![3, 4]).unwrap();
109 let layout = StridedLayout::new(vec![4, 1]);
110 let err = layout.element_offset(&[3, 0], &shape).unwrap_err();
111 assert!(
112 matches!(
113 err,
114 Error::IndexOutOfRange {
115 dim: 0,
116 index: 3,
117 size: 3
118 }
119 ),
120 "expected IndexOutOfRange, got {err:?}"
121 );
122 }
123
124 #[test]
126 fn strided_dynamic_dim() {
127 let shape = Shape::new(vec![3, DYNAMIC]).unwrap();
128 let layout = StridedLayout::new(vec![1, 1]);
129 let err = layout.element_offset(&[0, 1], &shape).unwrap_err();
130 assert!(
131 matches!(err, Error::DynamicDimInIndexing { dim: 1 }),
132 "expected DynamicDimInIndexing, got {err:?}"
133 );
134 }
135}