Skip to main content

hurray_core/layout/addressing/
strided.rs

1//! Strided layout element offset computation.
2//!
3//! Spec: docs/spec/layouts/strided.md § Element Address
4
5use crate::layout::StridedLayout;
6use crate::{Error, Result, Shape};
7
8use super::{validate_index, ElementAddress};
9
10// Sum is computed as i64 to support negative strides; cast to u64 via
11// two's-complement reinterpretation so byte_address_from_element_offset
12// can reconstruct the signed value. ADR-014 Amendment.
13impl 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    // Spec docs/spec/layouts/strided.md § Element Address:
39    // offset = Σ index[k] × strides[k]  (strides are signed i64, in logical elements)
40    //
41    // Row-major strides [4,1] for shape [3,4]: element [1,2] → 1*4 + 2*1 = 6.
42    #[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    // Col-major strides [1,3] for shape [3,4]: element [1,2] → 1*1 + 2*3 = 7.
50    #[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    // Negative stride (reversed row dimension):
58    // shape [3,4], strides [-4, 1].
59    // element [0,0] → 0*(-4) + 0*1 = 0.
60    #[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    // Negative stride (reversed rows):
68    // shape [3,4], strides [-4, 1].
69    // element [2,3] → 2*(-4) + 3*1 = -5 → stored as i64 → u64 two's-complement.
70    #[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        // -5 as u64 via two's-complement.
76        assert_eq!(offset, (-5i64) as u64);
77    }
78
79    // Zero stride (broadcast dimension):
80    // strides [0,1] for shape [5,4]: element [3,2] → 3*0 + 2*1 = 2.
81    #[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        // All elements in the first dimension map to the same offset.
87        assert_eq!(layout.element_offset(&[0, 2], &shape).unwrap(), 2);
88        assert_eq!(layout.element_offset(&[4, 2], &shape).unwrap(), 2);
89    }
90
91    // Error: strides length does not match shape rank.
92    // validate_index passes (rank matches), but then the strides-rank check fires.
93    #[test]
94    fn strided_strides_rank_mismatch() {
95        let shape = Shape::new(vec![3, 4]).unwrap();
96        // 3 strides for rank-2 shape → mismatch.
97        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    // Error: index out of bounds.
106    #[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    // Error: DYNAMIC dimension in shape.
125    #[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}