//! Copying a selection out of a row-major buffer one contiguous run at a time. //! //! A selection's elements, in output order, fall into runs that are adjacent //! in the source: a whole block along the last dimension, blocks that touch //! (`stride == block`), and whole rows when the inner dimensions are selected //! in full. Copying run by run turns a 256 x 256 hyperslab of a 1024-wide //! dataset into 256 `memcpy`s of 1 KiB, where the old extractor recursed and //! bounds-checked once per element. #[cfg(not(feature = "std"))] use alloc::{vec, vec::Vec}; use crate::data_read::NativeElement; use crate::error::FormatError; use crate::selection::Selection; /// Row-major element strides of `dims` (the last dimension has stride 1). fn strides(dims: &[u64]) -> Vec { let mut s = vec![1u64; dims.len()]; for d in (0..dims.len().saturating_sub(1)).rev() { s[d] = s[d + 1].wrapping_mul(dims[d + 1]); } s } /// Merges adjacent runs before handing them on. struct Coalesce { start: u64, len: u64, emit: F, } impl Coalesce { #[inline] fn push(&mut self, start: u64, len: u64) { if len == 0 { return; } if self.len > 0 && self.start.wrapping_add(self.len) == start { self.len += len; return; } self.flush(); self.start = start; self.len = len; } fn flush(&mut self) { if self.len > 0 { (self.emit)(self.start, self.len); self.len = 0; } } } /// Call `emit(first_element, element_count)` for each run of a hyperslab's /// elements that is contiguous in a row-major dataset of shape `dims`, in /// the order the selection returns them. Adjacent runs are merged. /// /// Coordinates at or past a dimension's extent are skipped, as the /// element-wise extractor always did; callers that want them to be an error /// validate the selection first. The four vectors must have `dims.len()` /// entries. pub(crate) fn hyperslab_runs( dims: &[u64], start: &[u64], stride: &[u64], count: &[u64], block: &[u64], emit: impl FnMut(u64, u64), ) { let rank = dims.len(); let mut out = Coalesce { start: 0, len: 0, emit, }; if rank == 0 { out.push(0, 1); out.flush(); return; } if (0..rank).any(|d| count[d] == 0 || block[d] == 0) { return; } let strides = strides(dims); let last = rank - 1; // Odometer over the outer dimensions: (block index, offset in block). let mut ci = vec![0u64; last]; let mut bi = vec![0u64; last]; 'outer: loop { // Base offset of this row, or skip it if a coordinate is out of range. let mut base = 0u64; let mut in_range = true; for d in 0..last { let coord = start[d] .saturating_add(ci[d].saturating_mul(stride[d])) .saturating_add(bi[d]); if coord >= dims[d] { in_range = false; break; } base = base.wrapping_add(coord.wrapping_mul(strides[d])); } if in_range && (stride[last] == block[last] || count[last] == 1) { // Blocks that touch (the common unit-stride case: block 1, // stride 1) are one range; don't split it into per-element runs. let s = start[last]; let e = s .saturating_add(count[last].saturating_mul(block[last])) .min(dims[last]); if s < e { out.push(base.wrapping_add(s), e - s); } } else if in_range { for c in 0..count[last] { let s = start[last].saturating_add(c.saturating_mul(stride[last])); if s >= dims[last] { continue; } let e = s.saturating_add(block[last]).min(dims[last]); out.push(base.wrapping_add(s), e - s); } } // Advance the odometer, last outer dimension fastest. let mut d = last; loop { if d == 0 { break 'outer; } d -= 1; bi[d] += 1; if bi[d] < block[d] { break; } bi[d] = 0; ci[d] += 1; if ci[d] < count[d] { break; } ci[d] = 0; } } out.flush(); } /// The selected elements of `src` — a row-major dataset of shape `dims` and /// `elem_size`-byte elements — copied into a fresh `Vec`, one `memcpy` per /// contiguous run, with no zero-filling of the output first. /// /// For `T` other than `u8`, `elem_size` must equal `size_of::()`. The /// selection must be a validated hyperslab, point list or `None` (`All` is the /// caller's to handle); `src` must hold exactly the dataset. Anything that /// would read outside `src` is an error, never a partial result. pub(crate) fn gather( src: &[u8], dims: &[u64], elem_size: usize, selection: &Selection, ) -> Result, FormatError> { let t_size = core::mem::size_of::(); if elem_size == 0 || (t_size != 1 && t_size != elem_size) { return Err(FormatError::DataSizeMismatch { expected: t_size, actual: elem_size, }); } let n_elements = match selection { Selection::None => 0, Selection::Hyperslab { count, block, .. } => count .iter() .zip(block) .try_fold(1u64, |acc, (&c, &b)| acc.checked_mul(c.checked_mul(b)?)) .ok_or_else(|| FormatError::Overflow("hyperslab count x block overflows".into()))?, Selection::Points(points) => points.len() as u64, Selection::All => { return Err(FormatError::SelectionOutOfBounds( "gather does not take Selection::All".into(), )); } }; let out_bytes = crate::chunked_read::checked_byte_len(n_elements, elem_size)?; let out_len = out_bytes / t_size; let mut out: Vec = crate::bulk_alloc::vec_for_bulk(out_len); let dst = out.as_mut_ptr().cast::(); let mut written = 0usize; let mut failed = false; let mut copy_run = |first: u64, n: u64| { if failed { return; } let range = usize::try_from(first) .ok() .and_then(|f| f.checked_mul(elem_size)) .zip( usize::try_from(n) .ok() .and_then(|n| n.checked_mul(elem_size)), ) .and_then(|(at, len)| Some((at, len, at.checked_add(len)?))); match range { Some((at, len, end)) if end <= src.len() && written + len <= out_bytes => { // SAFETY: `src[at..end]` is in bounds (checked above), and // `dst + written .. + len` lies within `out`'s capacity of // `out_bytes` bytes (checked above); `out` is a fresh // allocation, so the regions do not overlap. unsafe { core::ptr::copy_nonoverlapping(src.as_ptr().add(at), dst.add(written), len) }; written += len; } _ => failed = true, } }; let mut bad_point = false; match selection { Selection::Hyperslab { start, stride, count, block, } => { let rank = dims.len(); if [start.len(), stride.len(), count.len(), block.len()] != [rank; 4] { return Err(FormatError::SelectionOutOfBounds( "hyperslab rank does not match dataset rank".into(), )); } hyperslab_runs(dims, start, stride, count, block, &mut copy_run); } Selection::Points(points) => { let strides = strides(dims); let mut runs = Coalesce { start: 0, len: 0, emit: &mut copy_run, }; for p in points { if p.len() != dims.len() || p.iter().zip(dims).any(|(c, n)| c >= n) { bad_point = true; break; } let at = p .iter() .zip(&strides) .fold(0u64, |acc, (c, s)| acc.wrapping_add(c.wrapping_mul(*s))); runs.push(at, 1); } runs.flush(); } Selection::None | Selection::All => {} } if failed || bad_point || written != out_bytes { return Err(FormatError::SelectionOutOfBounds( "selection addresses elements outside the dataset".into(), )); } // SAFETY: all `out_bytes` bytes, i.e. `out_len` values of `T`, were // written above, and every bit pattern is a valid `T` (`NativeElement`). unsafe { out.set_len(out_len) }; Ok(out) } #[cfg(test)] mod tests { use super::*; fn runs(dims: &[u64], sel: [&[u64]; 4]) -> Vec<(u64, u64)> { let mut v = Vec::new(); hyperslab_runs(dims, sel[0], sel[1], sel[2], sel[3], |s, n| v.push((s, n))); v } #[test] fn runs_merge_blocks_and_whole_rows() { // A box: one run per row. assert_eq!( runs(&[4, 10], [&[1, 2], &[1, 1], &[2, 3], &[1, 1]]), vec![(12, 3), (22, 3)] ); // Whole rows: one run. assert_eq!( runs(&[4, 10], [&[1, 0], &[1, 1], &[3, 10], &[1, 1]]), vec![(10, 30)] ); // stride == block: blocks merge. assert_eq!( runs(&[1, 10], [&[0, 1], &[1, 2], &[1, 4], &[1, 2]]), vec![(1, 8)] ); // Strided with blocks along both dimensions. assert_eq!( runs(&[6, 10], [&[0, 1], &[3, 4], &[2, 2], &[2, 2]]), vec![ (1, 2), (5, 2), (11, 2), (15, 2), (31, 2), (35, 2), (41, 2), (45, 2) ] ); // Empty. assert!(runs(&[4, 10], [&[0, 0], &[1, 1], &[0, 3], &[1, 1]]).is_empty()); // Scalar. assert_eq!(runs(&[], [&[], &[], &[], &[]]), vec![(0, 1)]); } #[test] fn gather_matches_element_order_and_rejects_out_of_range() { let dims = [3u64, 4]; let src: Vec = (0..12u16).flat_map(|v| v.to_le_bytes()).collect(); let sel = Selection::Hyperslab { start: vec![0, 1], stride: vec![2, 2], count: vec![2, 2], block: vec![1, 1], }; let got: Vec = gather(&src, &dims, 2, &sel).unwrap(); let want: Vec = [1u16, 3, 9, 11] .iter() .flat_map(|v| v.to_le_bytes()) .collect(); assert_eq!(got, want); let pts = Selection::Points(vec![vec![2, 3], vec![0, 0], vec![0, 1]]); let got: Vec = gather(&src, &dims, 2, &pts).unwrap(); let want: Vec = [11u16, 0, 1].iter().flat_map(|v| v.to_le_bytes()).collect(); assert_eq!(got, want); // Past the extent, or a source shorter than the dataset: an error. let bad = Selection::Points(vec![vec![3, 0]]); assert!(gather::(&src, &dims, 2, &bad).is_err()); let past = Selection::Hyperslab { start: vec![2, 0], stride: vec![1, 1], count: vec![2, 4], block: vec![1, 1], }; assert!(gather::(&src, &dims, 2, &past).is_err()); assert!(gather::(&src[..20], &dims, 2, &pts).is_err()); } }