hurray_core/layout/addressing/mod.rs
1//! Layout address-computation traits and shared helpers.
2//!
3//! Implements the element-address formulas defined in
4//! `docs/spec/memory-layout.md § Element Address Computation` and per-layout
5//! spec files under `docs/spec/layouts/`.
6
7pub mod block_paged;
8pub mod col_major;
9pub mod coo;
10pub mod csc;
11pub mod csf;
12pub mod csr;
13pub mod hilbert;
14pub mod morton;
15pub mod row_major;
16pub mod strided;
17pub mod tiled;
18
19use crate::{ElementType, Error, Result, Shape, DYNAMIC};
20
21/// Implemented by every dense layout descriptor.
22///
23/// Converts a multi-dimensional logical index into a linear element offset.
24/// The offset is relative to `byte_offset` in the tensor descriptor, and may
25/// be negative for strided layouts (encoded as a signed two's-complement `u64`).
26///
27/// # Contract
28///
29/// Implementations MUST:
30/// - Reject `index.len() != shape.rank()` with [`Error::IndexRankMismatch`].
31/// - Reject any `DYNAMIC` dimension with [`Error::DynamicDimInIndexing`].
32/// - Reject any out-of-bounds index component with [`Error::IndexOutOfRange`].
33///
34/// # Examples
35///
36/// ```
37/// use hurray_core::layout::{LayoutDescriptor, addressing::ElementAddress};
38/// use hurray_core::Shape;
39///
40/// let shape = Shape::new(vec![3, 4]).unwrap();
41/// let offset = LayoutDescriptor::RowMajor
42/// .element_offset(&[1, 2], &shape)
43/// .unwrap();
44/// assert_eq!(offset, 6); // 1*4 + 2*1 = 6
45/// ```
46pub trait ElementAddress {
47 /// Returns the linear element offset (in logical elements) for the given index.
48 fn element_offset(&self, index: &[u64], shape: &Shape) -> Result<u64>;
49}
50
51// Sparse layouts do not implement the dense `ElementAddress` trait: their lookup needs the
52// tensor's index buffers, not just the shape. Each sparse layout instead exposes a
53// standalone `element_offset` function in its addressing submodule
54// (`coo`, `csr`, `csc`, `csf`), taking typed index-buffer slices and returning the storage
55// offset or `None` for a structural zero. (This is the shape ADR-014 OQ-014.1 deferred
56// pending a first consumer; CSF validated it and the others now follow it.)
57
58/// Converts a linear element offset to an absolute byte address within a buffer.
59///
60/// `element_offset` is treated as a **signed** two's-complement `u64`, so
61/// negative offsets from strided layouts work correctly.
62///
63/// # Arguments
64///
65/// - `element_offset` — output of [`ElementAddress::element_offset`].
66/// - `byte_offset` — the `byte_offset` field from the tensor descriptor.
67/// - `element_type` — tensor's element type (determines byte-width arithmetic).
68/// - `buffer_size` — total byte size of the buffer; used for bounds checking.
69///
70/// # Errors
71///
72/// - [`Error::AddressOverflow`] — intermediate arithmetic overflowed.
73/// - [`Error::ByteAddressOverflow`] — address falls outside `[0, buffer_size)`.
74///
75/// # Spec
76///
77/// Implements `docs/spec/memory-layout.md § Element Address Computation`.
78///
79/// # Examples
80///
81/// ```
82/// use hurray_core::ElementType;
83/// use hurray_core::layout::addressing::byte_address_from_element_offset;
84///
85/// // float32 at element offset 6 with byte_offset 0 in a 100-byte buffer.
86/// let addr = byte_address_from_element_offset(6, 0, ElementType::Float32, 100).unwrap();
87/// assert_eq!(addr, 24); // 6 * 4 = 24
88///
89/// // bool at element offset 9 (fits in 2nd byte): byte = 0 + floor(9/8) = 1.
90/// let addr = byte_address_from_element_offset(9, 0, ElementType::Bool, 100).unwrap();
91/// assert_eq!(addr, 1);
92/// ```
93pub fn byte_address_from_element_offset(
94 element_offset: u64,
95 byte_offset: u64,
96 element_type: ElementType,
97 buffer_size: u64,
98) -> Result<u64> {
99 let signed = element_offset as i64;
100 let bits = element_type.bit_width() as i64;
101
102 let byte_delta: i64 = if bits >= 8 {
103 // Whole-byte types: delta = signed_offset × (bits / 8).
104 signed.checked_mul(bits / 8).ok_or(Error::AddressOverflow)?
105 } else if bits == 6 {
106 // 6-bit types pack 4 elements per 3 bytes; group-level addressing.
107 // byte_delta = floor(signed / 4) × 3.
108 signed
109 .div_euclid(4)
110 .checked_mul(3)
111 .ok_or(Error::AddressOverflow)?
112 } else {
113 // Sub-byte power-of-two (bits ∈ {1, 2, 4}).
114 // Packing factor P = 8 / bits; byte_delta = floor(signed / P).
115 signed.div_euclid(8 / bits)
116 };
117
118 let byte_addr = (byte_offset as i64)
119 .checked_add(byte_delta)
120 .ok_or(Error::AddressOverflow)?;
121
122 if byte_addr < 0 || byte_addr as u64 >= buffer_size {
123 return Err(Error::ByteAddressOverflow { buffer_size });
124 }
125 Ok(byte_addr as u64)
126}
127
128/// Validates a multi-dimensional index against a shape.
129///
130/// Returns `Err` on rank mismatch, any DYNAMIC dimension, or any out-of-bounds
131/// index component. Called at the start of every `ElementAddress` impl.
132pub(crate) fn validate_index(index: &[u64], shape: &Shape) -> Result<()> {
133 if index.len() != shape.rank() {
134 return Err(Error::IndexRankMismatch {
135 index_rank: index.len(),
136 shape_rank: shape.rank(),
137 });
138 }
139 for (k, (&idx, &dim)) in index.iter().zip(shape.dims().iter()).enumerate() {
140 if dim == DYNAMIC {
141 return Err(Error::DynamicDimInIndexing { dim: k as u32 });
142 }
143 if idx >= dim {
144 return Err(Error::IndexOutOfRange {
145 dim: k as u32,
146 index: idx,
147 size: dim,
148 });
149 }
150 }
151 Ok(())
152}
153
154#[cfg(test)]
155mod tests {
156 use crate::layout::{CooLayout, LayoutDescriptor, MortonLayout, StridedLayout};
157 use crate::{ElementType, Error, Shape};
158
159 use super::byte_address_from_element_offset;
160
161 // ── LayoutDescriptor::element_offset dispatch ─────────────────────────────
162
163 // RowMajor dispatch: shape [3,4], element [1,2] → 1*4+2 = 6
164 #[test]
165 fn dispatch_row_major() {
166 let shape = Shape::new(vec![3, 4]).unwrap();
167 assert_eq!(
168 LayoutDescriptor::RowMajor
169 .element_offset(&[1, 2], &shape)
170 .unwrap(),
171 6
172 );
173 }
174
175 // ColMajor dispatch: shape [3,4], element [1,2] → 1*1+2*3 = 7
176 #[test]
177 fn dispatch_col_major() {
178 let shape = Shape::new(vec![3, 4]).unwrap();
179 assert_eq!(
180 LayoutDescriptor::ColMajor
181 .element_offset(&[1, 2], &shape)
182 .unwrap(),
183 7
184 );
185 }
186
187 // Strided dispatch: row-major strides [4,1], same result as RowMajor.
188 #[test]
189 fn dispatch_strided() {
190 let shape = Shape::new(vec![3, 4]).unwrap();
191 let layout = LayoutDescriptor::Strided(StridedLayout::new(vec![4, 1]));
192 assert_eq!(layout.element_offset(&[1, 2], &shape).unwrap(), 6);
193 }
194
195 // Morton dispatch: shape [4,4], morton_bits [2,2], element [2,3] → 14
196 #[test]
197 fn dispatch_morton() {
198 let shape = Shape::new(vec![4, 4]).unwrap();
199 let layout = LayoutDescriptor::Morton(MortonLayout::new(vec![2, 2]).unwrap());
200 assert_eq!(layout.element_offset(&[2, 3], &shape).unwrap(), 14);
201 }
202
203 // COO returns LayoutRequiresMultiBuffer with tag 0x06.
204 #[test]
205 fn dispatch_coo_returns_multi_buffer_error() {
206 let layout = LayoutDescriptor::Coo(CooLayout::new(0, false));
207 let shape = Shape::new(vec![4, 4]).unwrap();
208 let err = layout.element_offset(&[0, 0], &shape).unwrap_err();
209 assert!(
210 matches!(err, Error::LayoutRequiresMultiBuffer { layout_tag: 0x06 }),
211 "expected LayoutRequiresMultiBuffer{{0x06}}, got {err:?}"
212 );
213 }
214
215 // ── byte_address_from_element_offset ──────────────────────────────────────
216
217 // Spec docs/spec/memory-layout.md § Element Address Computation:
218 // Float32 at element offset 6, byte_offset 0, buffer 100 bytes → 6*4 = 24.
219 #[test]
220 fn byte_addr_float32_offset_6() {
221 let addr = byte_address_from_element_offset(6, 0, ElementType::Float32, 100).unwrap();
222 assert_eq!(addr, 24);
223 }
224
225 // Bool at element offset 9, byte_offset 0: floor(9/8) = 1.
226 #[test]
227 fn byte_addr_bool_offset_9() {
228 let addr = byte_address_from_element_offset(9, 0, ElementType::Bool, 100).unwrap();
229 assert_eq!(addr, 1);
230 }
231
232 // Int4 at element offset 7, byte_offset 0: floor(7/2) = 3.
233 #[test]
234 fn byte_addr_int4_offset_7() {
235 let addr = byte_address_from_element_offset(7, 0, ElementType::Int4, 100).unwrap();
236 assert_eq!(addr, 3);
237 }
238
239 // Uint4 at element offset 8, byte_offset 0: floor(8/2) = 4.
240 #[test]
241 fn byte_addr_uint4_offset_8() {
242 let addr = byte_address_from_element_offset(8, 0, ElementType::Uint4, 100).unwrap();
243 assert_eq!(addr, 4);
244 }
245
246 // Float64 at element offset 3, byte_offset 8: 8 + 3*8 = 32.
247 #[test]
248 fn byte_addr_float64_with_byte_offset() {
249 let addr = byte_address_from_element_offset(3, 8, ElementType::Float64, 200).unwrap();
250 assert_eq!(addr, 32);
251 }
252
253 // Negative strided offset:
254 // Float32, element_offset = (-5_i64 as u64), byte_offset = 40, buffer = 100.
255 // byte_delta = -5 * 4 = -20; byte_addr = 40 + (-20) = 20. Valid: 20 < 100.
256 #[test]
257 fn byte_addr_negative_strided_offset() {
258 let element_offset = (-5i64) as u64;
259 let addr = byte_address_from_element_offset(element_offset, 40, ElementType::Float32, 100)
260 .unwrap();
261 assert_eq!(addr, 20);
262 }
263
264 // ByteAddressOverflow when computed address >= buffer_size.
265 // Float32, offset 25, byte_offset 0, buffer 100: 25*4 = 100 >= 100 → overflow.
266 #[test]
267 fn byte_addr_overflow_at_exact_buffer_size() {
268 let err = byte_address_from_element_offset(25, 0, ElementType::Float32, 100).unwrap_err();
269 assert!(
270 matches!(err, Error::ByteAddressOverflow { buffer_size: 100 }),
271 "expected ByteAddressOverflow{{100}}, got {err:?}"
272 );
273 }
274
275 // ByteAddressOverflow when byte_offset insufficient to cover a negative delta.
276 // Float32, element_offset = (-5_i64 as u64), byte_offset = 4, buffer = 100.
277 // byte_delta = -20; byte_addr = 4 - 20 = -16 → negative → overflow.
278 #[test]
279 fn byte_addr_overflow_negative_offset_underflows_byte_offset() {
280 let element_offset = (-5i64) as u64;
281 let err = byte_address_from_element_offset(element_offset, 4, ElementType::Float32, 100)
282 .unwrap_err();
283 assert!(
284 matches!(err, Error::ByteAddressOverflow { .. }),
285 "expected ByteAddressOverflow, got {err:?}"
286 );
287 }
288
289 // First valid element (offset 0) always succeeds even with buffer_size=1.
290 #[test]
291 fn byte_addr_offset_zero_is_always_valid() {
292 let addr = byte_address_from_element_offset(0, 0, ElementType::Float32, 4).unwrap();
293 assert_eq!(addr, 0);
294 }
295
296 // Last valid float32 element in a 40-byte buffer (10 elements, offset 9):
297 // 9*4 = 36 < 40. ✓
298 #[test]
299 fn byte_addr_last_valid_float32_element() {
300 let addr = byte_address_from_element_offset(9, 0, ElementType::Float32, 40).unwrap();
301 assert_eq!(addr, 36);
302 }
303}