hurray_core/layout/addressing/
tiled.rs1use 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
25fn tiled_offset(layout: &TiledLayout, index: &[u64], dims: &[u64]) -> Result<u64> {
30 let rank = index.len();
31
32 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 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 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 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 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 tiled_offset(inner, &intra_idx, &layout.tile_shape)? as i64
78 }
79 _ => unreachable!("inner_layout validated in TiledLayout::new"),
80 };
81
82 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
92fn 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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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}