Skip to main content

hurray_core/layout/addressing/
tiled.rs

1//! Tiled / blocked layout element offset computation.
2//!
3//! Spec: docs/spec/layouts/tiled.md § Element Address
4
5use crate::layout::TiledLayout;
6use crate::{Error, Result, Shape};
7
8use super::col_major::col_major_offset;
9use super::row_major::row_major_offset;
10use super::{validate_index, ElementAddress};
11
12impl ElementAddress for TiledLayout {
13    fn element_offset(&self, index: &[u64], shape: &Shape) -> Result<u64> {
14        validate_index(index, shape)?;
15        if self.tile_shape.len() != shape.rank() {
16            return Err(Error::IndexRankMismatch {
17                index_rank: self.tile_shape.len(),
18                shape_rank: shape.rank(),
19            });
20        }
21        tiled_offset(self, index, shape.dims())
22    }
23}
24
25/// Recursive tiled element offset computation.
26///
27/// `index` is the logical index into the tiled space; `dims` is the shape of
28/// that space. Called recursively for `inner_layout == 0x04`.
29fn tiled_offset(layout: &TiledLayout, index: &[u64], dims: &[u64]) -> Result<u64> {
30    let rank = index.len();
31
32    // 1. Tile index and intra-tile index per dimension.
33    let mut tile_idx = vec![0u64; rank];
34    let mut intra_idx = vec![0u64; rank];
35    for k in 0..rank {
36        tile_idx[k] = index[k] / layout.tile_shape[k];
37        intra_idx[k] = index[k] % layout.tile_shape[k];
38    }
39
40    // 2. Tile-grid shape: ceil(dims[k] / tile_shape[k]).
41    let tile_grid_dims: Vec<u64> = dims
42        .iter()
43        .zip(layout.tile_shape.iter())
44        .map(|(&s, &t)| s.div_ceil(t))
45        .collect();
46
47    // 3. Linear tile number (in tile units).
48    let tile_number: i64 = match layout.outer_layout {
49        0x01 => row_major_offset(&tile_idx, &tile_grid_dims)? as i64,
50        0x02 => col_major_offset(&tile_idx, &tile_grid_dims)? as i64,
51        0x03 => strided_offset(
52            &tile_idx,
53            layout.outer_strides.as_ref().unwrap().strides.as_slice(),
54        )?,
55        _ => unreachable!("outer_layout validated in TiledLayout::new"),
56    };
57
58    // 4. Total elements per tile.
59    let tile_size: i64 = layout.tile_shape.iter().try_fold(1i64, |acc, &d| {
60        acc.checked_mul(d as i64).ok_or(Error::AddressOverflow)
61    })?;
62
63    // 5. Intra-tile element offset.
64    let intra_offset: i64 = match layout.inner_layout {
65        0x01 => row_major_offset(&intra_idx, &layout.tile_shape)? as i64,
66        0x02 => col_major_offset(&intra_idx, &layout.tile_shape)? as i64,
67        0x03 => strided_offset(
68            &intra_idx,
69            layout.inner_strides.as_ref().unwrap().strides.as_slice(),
70        )?,
71        0x04 => {
72            let inner = layout
73                .inner_tiled
74                .as_ref()
75                .expect("inner_tiled is Some when inner_layout == 0x04");
76            // Recursive call: index into the outer tile's index space.
77            tiled_offset(inner, &intra_idx, &layout.tile_shape)? as i64
78        }
79        _ => unreachable!("inner_layout validated in TiledLayout::new"),
80    };
81
82    // 6. Final offset = tile_number × tile_size + intra_offset.
83    let total = tile_number
84        .checked_mul(tile_size)
85        .ok_or(Error::AddressOverflow)?
86        .checked_add(intra_offset)
87        .ok_or(Error::AddressOverflow)?;
88
89    Ok(total as u64)
90}
91
92/// Computes a strided linear offset from an index and a stride slice (in `i64`).
93///
94/// Used for both outer (tile-grid) and inner (within-tile) strided layouts.
95fn strided_offset(idx: &[u64], strides: &[i64]) -> Result<i64> {
96    let mut sum: i64 = 0;
97    for k in 0..idx.len() {
98        let term = (idx[k] as i64)
99            .checked_mul(strides[k])
100            .ok_or(Error::AddressOverflow)?;
101        sum = sum.checked_add(term).ok_or(Error::AddressOverflow)?;
102    }
103    Ok(sum)
104}
105
106#[cfg(test)]
107mod tests {
108    use super::*;
109    use crate::layout::{InnerStrides, OuterStrides, TiledLayout};
110    use crate::Shape;
111
112    // Spec docs/spec/layouts/tiled.md § Element Address:
113    //
114    // shape [6,8], tile_shape [2,4], outer=row-major (0x01), inner=row-major (0x01).
115    //
116    // Element [3,5]:
117    //   tile_idx    = [3/2, 5/4] = [1, 1]
118    //   intra_idx   = [3%2, 5%4] = [1, 1]
119    //   tile_grid   = [6/2, 8/4] = [3, 2]
120    //   tile_number = row_major([1,1], [3,2]) = 1*2+1 = 3
121    //   tile_size   = 2*4 = 8
122    //   intra_offset= row_major([1,1], [2,4]) = 1*4+1 = 5
123    //   final       = 3*8 + 5 = 29
124    #[test]
125    fn tiled_spec_example() {
126        let layout = TiledLayout::new(vec![2, 4], 0x01, 0x01, None, None, None).unwrap();
127        let shape = Shape::new(vec![6, 8]).unwrap();
128        assert_eq!(layout.element_offset(&[3, 5], &shape).unwrap(), 29);
129    }
130
131    // First element [0,0] is always at offset 0.
132    #[test]
133    fn tiled_first_element() {
134        let layout = TiledLayout::new(vec![2, 4], 0x01, 0x01, None, None, None).unwrap();
135        let shape = Shape::new(vec![6, 8]).unwrap();
136        assert_eq!(layout.element_offset(&[0, 0], &shape).unwrap(), 0);
137    }
138
139    // Last element [5,7] in shape [6,8], tile [2,4], row-major / row-major:
140    //   tile_idx    = [5/2, 7/4] = [2, 1]
141    //   intra_idx   = [5%2, 7%4] = [1, 3]
142    //   tile_grid   = [3, 2]
143    //   tile_number = row_major([2,1],[3,2]) = 2*2+1 = 5
144    //   tile_size   = 8
145    //   intra_offset= row_major([1,3],[2,4]) = 1*4+3 = 7
146    //   final       = 5*8 + 7 = 47
147    #[test]
148    fn tiled_last_element() {
149        let layout = TiledLayout::new(vec![2, 4], 0x01, 0x01, None, None, None).unwrap();
150        let shape = Shape::new(vec![6, 8]).unwrap();
151        assert_eq!(layout.element_offset(&[5, 7], &shape).unwrap(), 47);
152    }
153
154    // Column-major inner layout:
155    // shape [6,8], tile [2,4], outer=row-major (0x01), inner=col-major (0x02).
156    //
157    // Element [3,5]:
158    //   tile_idx = [1,1], intra_idx = [1,1]
159    //   tile_number = 3 (same as spec example, outer unchanged)
160    //   tile_size   = 8
161    //   intra_offset= col_major([1,1],[2,4]) = 1*1 + 1*2 = 3
162    //   final       = 3*8 + 3 = 27
163    #[test]
164    fn tiled_col_major_inner() {
165        let layout = TiledLayout::new(vec![2, 4], 0x01, 0x02, None, None, None).unwrap();
166        let shape = Shape::new(vec![6, 8]).unwrap();
167        assert_eq!(layout.element_offset(&[3, 5], &shape).unwrap(), 27);
168    }
169
170    // Recursive tiling:
171    // shape [8,8], outer tile [4,4], outer=row-major, inner=tiled (0x04).
172    // Inner tile [2,2], inner_outer=row-major, inner_inner=row-major.
173    //
174    // Element [3,5]:
175    //   outer tile_idx  = [3/4, 5/4] = [0, 1]
176    //   outer intra_idx = [3%4, 5%4] = [3, 1]
177    //   outer tile_grid = [2, 2]
178    //   outer tile_num  = row_major([0,1],[2,2]) = 0*2+1 = 1
179    //   outer tile_size = 4*4 = 16
180    //   inner (tile [2,2] in [4,4]):
181    //     inner tile_idx  = [3/2, 1/2] = [1, 0]
182    //     inner intra_idx = [3%2, 1%2] = [1, 1]
183    //     inner tile_grid = [2, 2]
184    //     inner tile_num  = row_major([1,0],[2,2]) = 1*2+0 = 2
185    //     inner tile_size = 2*2 = 4
186    //     inner intra     = row_major([1,1],[2,2]) = 1*2+1 = 3
187    //     inner_offset    = 2*4 + 3 = 11
188    //   final = 1*16 + 11 = 27
189    #[test]
190    fn tiled_recursive() {
191        let inner = TiledLayout::new(vec![2, 2], 0x01, 0x01, None, None, None).unwrap();
192        let outer =
193            TiledLayout::new(vec![4, 4], 0x01, 0x04, None, None, Some(Box::new(inner))).unwrap();
194        let shape = Shape::new(vec![8, 8]).unwrap();
195        assert_eq!(outer.element_offset(&[3, 5], &shape).unwrap(), 27);
196    }
197
198    // Strided outer layout:
199    // shape [4,4], tile [2,2], outer=strided (0x03) with strides [1,2],
200    // inner=row-major.
201    //
202    // Element [2,3]:
203    //   tile_idx  = [1, 1], intra_idx = [0, 1]
204    //   tile_number = strided([1,1], strides=[1,2]) = 1*1 + 1*2 = 3
205    //   tile_size   = 4
206    //   intra_offset= row_major([0,1],[2,2]) = 0*2+1 = 1
207    //   final       = 3*4 + 1 = 13
208    #[test]
209    fn tiled_strided_outer() {
210        let outer_strides = OuterStrides::new(vec![1, 2]);
211        let layout =
212            TiledLayout::new(vec![2, 2], 0x03, 0x01, Some(outer_strides), None, None).unwrap();
213        let shape = Shape::new(vec![4, 4]).unwrap();
214        assert_eq!(layout.element_offset(&[2, 3], &shape).unwrap(), 13);
215    }
216
217    // Strided inner layout:
218    // shape [4,4], tile [2,2], outer=row-major,
219    // inner=strided (0x03) with strides [1,2] (col-major for 2×2 tile).
220    //
221    // Element [2,3]:
222    //   tile_idx  = [1,1], intra_idx = [0,1]
223    //   tile_number = row_major([1,1],[2,2]) = 1*2+1 = 3
224    //   tile_size   = 4
225    //   intra_offset= strided([0,1], strides=[1,2]) = 0*1 + 1*2 = 2
226    //   final       = 3*4 + 2 = 14
227    #[test]
228    fn tiled_strided_inner() {
229        let inner_strides = InnerStrides::new(vec![1, 2]);
230        let layout =
231            TiledLayout::new(vec![2, 2], 0x01, 0x03, None, Some(inner_strides), None).unwrap();
232        let shape = Shape::new(vec![4, 4]).unwrap();
233        assert_eq!(layout.element_offset(&[2, 3], &shape).unwrap(), 14);
234    }
235
236    // Error: tile_shape rank differs from shape rank.
237    #[test]
238    fn tiled_tile_shape_rank_mismatch() {
239        let layout = TiledLayout::new(vec![2], 0x01, 0x01, None, None, None).unwrap();
240        let shape = Shape::new(vec![6, 8]).unwrap();
241        let err = layout.element_offset(&[3, 5], &shape).unwrap_err();
242        assert!(
243            matches!(err, Error::IndexRankMismatch { .. }),
244            "expected IndexRankMismatch, got {err:?}"
245        );
246    }
247
248    // Error: index out of bounds.
249    #[test]
250    fn tiled_index_out_of_range() {
251        let layout = TiledLayout::new(vec![2, 4], 0x01, 0x01, None, None, None).unwrap();
252        let shape = Shape::new(vec![6, 8]).unwrap();
253        let err = layout.element_offset(&[6, 0], &shape).unwrap_err();
254        assert!(
255            matches!(err, Error::IndexOutOfRange { dim: 0, .. }),
256            "expected IndexOutOfRange, got {err:?}"
257        );
258    }
259}