Skip to main content

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}