use std::{fmt, marker::PhantomData, num::NonZeroUsize, ptr::NonNull};
use thiserror::Error;
use crate::{
internal,
views::rowmajor::{self, Matrix},
Reborrow,
};
#[derive(Debug)]
pub struct Layout<T> {
nrows: usize,
ncols: usize,
cstride: usize,
_type: PhantomData<fn() -> T>,
}
impl<T> Layout<T> {
pub fn new(nrows: usize, ncols: usize, cstride: usize) -> Result<Self, LayoutError> {
LayoutError::check::<T>(nrows, ncols, cstride)?;
Ok(Self {
nrows,
ncols,
cstride,
_type: PhantomData,
})
}
pub fn nrows(&self) -> usize {
self.nrows
}
pub fn ncols(&self) -> usize {
self.ncols
}
pub fn cstride(&self) -> usize {
self.cstride
}
pub fn linear_length(&self) -> usize {
self.nrows.saturating_sub(1) * self.cstride + self.nrows.min(1) * self.ncols
}
}
impl<T> Clone for Layout<T> {
fn clone(&self) -> Self {
*self
}
}
impl<T> Copy for Layout<T> {}
impl<T> From<rowmajor::Layout<T>> for Layout<T> {
fn from(layout: rowmajor::Layout<T>) -> Self {
Self {
nrows: layout.nrows(),
ncols: layout.ncols(),
cstride: layout.ncols(),
_type: PhantomData,
}
}
}
fn linear_length(nrows: usize, ncols: usize, cstride: usize) -> Option<usize> {
nrows
.saturating_sub(1)
.checked_mul(cstride)
.and_then(|main| main.checked_add(nrows.min(1) * ncols))
}
#[derive(Debug)]
pub struct LayoutError(LayoutErrorInner);
impl LayoutError {
fn check<T>(nrows: usize, ncols: usize, cstride: usize) -> Result<usize, Self> {
if cstride < ncols {
Err(Self(LayoutErrorInner::InvalidStride { ncols, cstride }))
} else {
let linear_length = match linear_length(nrows, ncols, cstride) {
Some(len) => len,
None => {
return Err(Self(LayoutErrorInner::Overflow {
nrows,
cstride,
elsize: None,
}));
}
};
let elsize = std::mem::size_of::<T>();
let bytes = linear_length.saturating_mul(elsize);
if bytes > (isize::MAX as usize) {
Err(Self(LayoutErrorInner::Overflow {
nrows,
cstride,
elsize: NonZeroUsize::new(elsize),
}))
} else {
Ok(linear_length)
}
}
}
}
impl fmt::Display for LayoutError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
self.0.fmt(f)
}
}
impl std::error::Error for LayoutError {}
#[derive(Debug)]
enum LayoutErrorInner {
InvalidStride {
ncols: usize,
cstride: usize,
},
Overflow {
nrows: usize,
cstride: usize,
elsize: Option<NonZeroUsize>,
},
}
impl fmt::Display for LayoutErrorInner {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::InvalidStride { ncols, cstride } => write!(
f,
"column stride {} must be greater than or equal to number of columns {}",
cstride, ncols
),
Self::Overflow {
nrows,
cstride,
elsize,
} => match elsize {
Some(elsize) => write!(
f,
"a {}x{} strided matrix with element size {} exceeds isize::MAX bytes",
nrows, cstride, elsize
),
None => write!(
f,
"a {}x{} strided matrix has a length exceeding usize::MAX",
nrows, cstride
),
},
}
}
}
#[derive(Debug)]
pub struct Strided<'a, T> {
ptr: NonNull<T>,
layout: Layout<T>,
_lifetime: PhantomData<&'a [T]>,
}
impl<'a, T> Strided<'a, T> {
pub fn try_from_data(
data: &'a [T],
nrows: usize,
ncols: usize,
cstride: usize,
) -> Result<Self, TryFromError> {
let layout = Layout::new(nrows, ncols, cstride).map_err(TryFromError::LayoutError)?;
let expected = layout.linear_length();
if data.len() < expected {
Err(TryFromError::InvalidLength {
got: data.len(),
expected,
})
} else {
Ok(unsafe { Self::from_data_unchecked(data, layout) })
}
}
unsafe fn from_data_unchecked(data: &'a [T], layout: Layout<T>) -> Self {
debug_assert!(data.len() >= layout.linear_length());
Self {
ptr: internal::slice_to_nonnull(data),
layout,
_lifetime: PhantomData,
}
}
fn as_nonnull(&self) -> NonNull<T> {
self.ptr
}
pub fn layout(&self) -> Layout<T> {
self.layout
}
pub fn as_ptr(&self) -> *const T {
self.as_nonnull().as_ptr().cast_const()
}
pub fn ncols(&self) -> usize {
self.layout().ncols()
}
pub fn nrows(&self) -> usize {
self.layout().nrows()
}
pub fn cstride(&self) -> usize {
self.layout().cstride()
}
pub fn as_slice(&self) -> &[T] {
let layout = self.layout();
unsafe { std::slice::from_raw_parts(self.as_ptr(), layout.linear_length()) }
}
pub unsafe fn element_unchecked(&self, row: usize, col: usize) -> &T {
let layout = self.layout();
debug_assert!(row < layout.nrows());
debug_assert!(col < layout.ncols());
unsafe { &*self.as_ptr().add(layout.cstride() * row + col) }
}
pub fn get_element(&self, row: usize, col: usize) -> Option<&T> {
if row < self.nrows() && col < self.ncols() {
Some(unsafe { self.element_unchecked(row, col) })
} else {
None
}
}
pub fn element(&self, row: usize, col: usize) -> &T {
assert!(
row < self.nrows(),
"row {} is out of bounds for a matrix with {} rows",
row,
self.nrows()
);
assert!(
col < self.ncols(),
"col {} is out of bounds for a matrix with {} cols",
col,
self.ncols()
);
unsafe { self.element_unchecked(row, col) }
}
pub unsafe fn row_unchecked(&self, row: usize) -> &[T] {
let layout = self.layout();
debug_assert!(row < layout.nrows());
unsafe {
std::slice::from_raw_parts(self.as_ptr().add(layout.cstride() * row), layout.ncols())
}
}
pub fn get_row(&self, row: usize) -> Option<&[T]> {
if row < self.nrows() {
Some(unsafe { self.row_unchecked(row) })
} else {
None
}
}
pub fn row(&self, row: usize) -> &[T] {
assert!(
row < self.nrows(),
"row {} is out of bounds for a matrix with {} rows",
row,
self.nrows()
);
unsafe { self.row_unchecked(row) }
}
pub fn rows(&self) -> Rows<'_, T> {
Rows::new(*self)
}
}
#[derive(Debug, Error)]
pub enum TryFromError {
#[error(transparent)]
LayoutError(LayoutError),
#[error(
"argument of length {} is shorter than the expected length {}",
got,
expected
)]
InvalidLength { got: usize, expected: usize },
}
impl<'a, T> From<rowmajor::Ref<'a, T>> for Strided<'a, T> {
fn from(matrix: rowmajor::Ref<'a, T>) -> Self {
let layout = Layout::from(matrix.layout());
unsafe { Self::from_data_unchecked(matrix.into_slice(), layout) }
}
}
unsafe impl<T> Send for Strided<'_, T> where T: Sync {}
unsafe impl<T> Sync for Strided<'_, T> where T: Sync {}
impl<T> Clone for Strided<'_, T> {
fn clone(&self) -> Self {
*self
}
}
impl<T> Copy for Strided<'_, T> {}
impl<'a, T> Reborrow<'a> for Strided<'_, T> {
type Target = Strided<'a, T>;
fn reborrow(&'a self) -> Self::Target {
*self
}
}
#[derive(Debug)]
pub struct Rows<'a, T> {
ptr: NonNull<T>,
remaining: usize,
ncols: usize,
cstride: usize,
_lifetime: PhantomData<&'a T>,
}
impl<'a, T> Rows<'a, T> {
fn new(strided: Strided<'a, T>) -> Self {
let layout = strided.layout();
Self {
ptr: strided.as_nonnull(),
remaining: layout.nrows(),
ncols: layout.ncols(),
cstride: layout.cstride(),
_lifetime: PhantomData,
}
}
}
unsafe impl<T> Send for Rows<'_, T> where T: Sync {}
unsafe impl<T> Sync for Rows<'_, T> where T: Sync {}
impl<'a, T> Iterator for Rows<'a, T> {
type Item = &'a [T];
fn next(&mut self) -> Option<&'a [T]> {
self.remaining.checked_sub(1).map(|remaining| {
let item =
unsafe { std::slice::from_raw_parts(self.ptr.as_ptr().cast_const(), self.ncols) };
self.remaining = remaining;
if remaining != 0 {
self.ptr = unsafe { self.ptr.add(self.cstride) };
}
item
})
}
fn size_hint(&self) -> (usize, Option<usize>) {
(self.remaining, Some(self.remaining))
}
}
impl<T> ExactSizeIterator for Rows<'_, T> {}
impl<T> std::iter::FusedIterator for Rows<'_, T> {}
#[cfg(test)]
mod tests {
use super::*;
use crate::views::rowmajor::MatrixMut;
#[test]
fn test_linear_length() {
assert_eq!(linear_length(0, 1, 1).unwrap(), 0);
assert_eq!(linear_length(0, 2, 2).unwrap(), 0);
assert_eq!(linear_length(0, 2, 3).unwrap(), 0);
assert_eq!(linear_length(0, 2, 4).unwrap(), 0);
for row in 1..10 {
for col in 1..10 {
assert_eq!(linear_length(row, col, col).unwrap(), row * col);
}
}
assert_eq!(linear_length(1, 5, 10).unwrap(), 5);
assert_eq!(linear_length(1, 7, 99).unwrap(), 7);
for row in 2..10 {
for col in 0..10 {
for cstride in col..12 {
assert_eq!(
linear_length(row, col, cstride).unwrap(),
(row - 1) * cstride + col
);
}
}
}
assert!(linear_length(usize::MAX, 2, 2).is_none());
assert!(linear_length(2, usize::MAX, 2).is_none());
assert!(linear_length(2, 2, usize::MAX).is_none());
}
#[test]
fn test_layout_new() {
let layout = Layout::<usize>::new(3, 4, 4).unwrap();
assert_eq!(layout.nrows(), 3);
assert_eq!(layout.ncols(), 4);
assert_eq!(layout.cstride(), 4);
assert_eq!(layout.linear_length(), 12);
let layout = Layout::<usize>::new(3, 4, 6).unwrap();
assert_eq!(layout.linear_length(), 2 * 6 + 4);
assert!(Layout::<usize>::new(0, 0, 0).is_ok());
let err = Layout::<usize>::new(3, 4, 3).unwrap_err();
assert_eq!(
err.to_string(),
"column stride 3 must be greater than or equal to number of columns 4"
);
let err = Layout::<usize>::new(usize::MAX, usize::MAX, usize::MAX).unwrap_err();
assert_eq!(
err.to_string(),
format!(
"a {}x{} strided matrix has a length exceeding usize::MAX",
usize::MAX,
usize::MAX
)
);
let err = Layout::<usize>::new(isize::MAX as usize, 1, 1).unwrap_err();
assert_eq!(
err.to_string(),
format!(
"a {}x{} strided matrix with element size {} exceeds isize::MAX bytes",
isize::MAX,
1,
std::mem::size_of::<usize>(),
)
);
let length = isize::MAX as usize;
let layout = Layout::<u8>::new(length, 1, 1).unwrap();
assert_eq!(layout.linear_length(), length);
assert!(Layout::<u8>::new(length + 1, 1, 1).is_err());
assert!(Layout::<u16>::new(length, 1, 1).is_err());
}
#[test]
fn test_try_from_data_errors() {
let m = rowmajor::Owned::<usize>::from_element(10, 10, 0);
let nrows = m.nrows();
let ncols = m.ncols();
let err = Strided::try_from_data(m.as_slice(), 2, 2, 1).unwrap_err();
assert_eq!(
err.to_string(),
"column stride 1 must be greater than or equal to number of columns 2"
);
let err = Strided::try_from_data(m.as_slice(), nrows, ncols, ncols + 1).unwrap_err();
assert_eq!(
err.to_string(),
"argument of length 100 is shorter than the expected length 109",
);
}
#[test]
fn test_element_and_row_out_of_bounds() {
let m = create_test_matrix(3, 4);
let v = Strided::try_from_data(m.as_slice(), m.nrows(), m.ncols(), m.ncols()).unwrap();
assert!(v.get_element(2, 3).is_some());
assert!(v.get_row(2).is_some());
assert!(v.get_element(3, 0).is_none(), "row out-of-bounds");
assert!(v.get_element(0, 4).is_none(), "col out-of-bounds");
assert!(v.get_element(3, 4).is_none(), "both out-of-bounds");
assert!(v.get_row(3).is_none());
}
#[test]
#[should_panic(expected = "row 3 is out of bounds for a matrix with 3 rows")]
fn test_element_panics_on_row() {
let m = create_test_matrix(3, 4);
let v = Strided::try_from_data(m.as_slice(), m.nrows(), m.ncols(), m.ncols()).unwrap();
v.element(3, 0);
}
#[test]
#[should_panic(expected = "col 4 is out of bounds for a matrix with 4 cols")]
fn test_element_panics_on_col() {
let m = create_test_matrix(3, 4);
let v = Strided::try_from_data(m.as_slice(), m.nrows(), m.ncols(), m.ncols()).unwrap();
v.element(0, 4);
}
#[test]
#[should_panic(expected = "row 3 is out of bounds for a matrix with 3 rows")]
fn test_row_panics() {
let m = create_test_matrix(3, 4);
let v = Strided::try_from_data(m.as_slice(), m.nrows(), m.ncols(), m.ncols()).unwrap();
v.row(3);
}
#[test]
fn test_clone_copy_reborrow() {
let m = create_test_matrix(3, 4);
let v = Strided::try_from_data(m.as_slice(), m.nrows(), m.ncols(), m.ncols()).unwrap();
let copied = v;
let cloned = Clone::clone(&v);
assert_eq!(v.as_ptr(), copied.as_ptr());
assert_eq!(v.as_ptr(), cloned.as_ptr());
let reborrowed = v.reborrow();
assert_eq!(reborrowed.as_ptr(), v.as_ptr());
assert_eq!(reborrowed.nrows(), v.nrows());
assert_eq!(reborrowed.ncols(), v.ncols());
}
#[test]
fn test_send_sync() {
fn assert_send_sync<T: Send + Sync>() {}
assert_send_sync::<Strided<'_, u8>>();
assert_send_sync::<Rows<'_, u8>>();
}
#[test]
fn test_rows_iterator_properties() {
let m = create_test_matrix(4, 3);
let v = Strided::try_from_data(&m.as_slice()[1..], m.nrows(), m.ncols() - 1, m.ncols())
.unwrap();
let mut rows = v.rows();
assert_eq!(rows.len(), 4);
assert_eq!(rows.size_hint(), (4, Some(4)));
for expected_row in 0..4 {
let row = rows.next().unwrap();
assert_eq!(row, &m.row(expected_row)[1..]);
}
assert_eq!(rows.next(), None);
assert_eq!(rows.next(), None);
assert_eq!(rows.len(), 0);
}
#[test]
fn test_rows_iterator_zero_rows() {
let m = create_test_matrix(5, 5);
let v = Strided::try_from_data(m.as_slice(), 0, 4, 5).unwrap();
let mut rows = v.rows();
assert_eq!(rows.len(), 0);
assert_eq!(rows.next(), None);
}
#[test]
fn test_rows_iterator_zero_cols() {
let m = create_test_matrix(5, 5);
let v = Strided::try_from_data(m.as_slice(), 5, 0, 5).unwrap();
let rows = v.rows();
assert_eq!(rows.len(), 5);
assert_eq!(rows.size_hint(), (5, Some(5)));
let mut count = 0;
for r in rows {
assert!(r.is_empty());
count += 1;
}
assert_eq!(count, 5);
}
#[test]
fn test_rows_iterator_zero_cstride() {
let m = create_test_matrix(5, 5);
let v = Strided::try_from_data(m.as_slice(), 5, 0, 0).unwrap();
let rows = v.rows();
assert_eq!(rows.len(), 5);
assert_eq!(rows.size_hint(), (5, Some(5)));
let mut count = 0;
for r in rows {
assert!(r.is_empty());
count += 1;
}
assert_eq!(count, 5);
}
fn test_indexing(dut: Strided<'_, usize>, expected: rowmajor::Ref<'_, usize>) {
assert_eq!(dut.nrows(), expected.nrows());
assert_eq!(dut.ncols(), expected.ncols());
if dut.cstride() == dut.ncols() {
assert_eq!(dut.as_slice(), expected.as_slice());
} else {
assert_ne!(dut.as_slice(), expected.as_slice());
}
for row in 0..dut.nrows() {
for col in 0..dut.ncols() {
let e = *expected.element(row, col);
assert_eq!(
*dut.element(row, col),
e,
"failed on (row, col) = ({}, {})",
row,
col
);
assert_eq!(
*dut.get_element(row, col).unwrap(),
e,
"failed on (row, col) = ({}, {})",
row,
col
);
}
}
for row in 0..dut.nrows() {
assert_eq!(dut.row(row), expected.row(row), "failed on row {}", row);
assert_eq!(
dut.get_row(row).unwrap(),
expected.row(row),
"failed on row {}",
row
);
}
assert!(dut.rows().eq(expected.rows()));
}
fn create_test_matrix(nrows: usize, ncols: usize) -> rowmajor::Owned<usize> {
let mut i = 0;
rowmajor::Owned::from_fn(nrows, ncols, |_| {
let v = i;
i += 1;
v
})
}
#[test]
fn test_basic_indexing() {
let m = create_test_matrix(5, 3);
let ptr = m.as_ptr();
let v = Strided::try_from_data(m.as_slice(), m.nrows(), m.ncols(), m.ncols()).unwrap();
assert_eq!(v.as_ptr(), ptr, "base pointer was not preserved");
assert_eq!(v.nrows(), m.nrows());
assert_eq!(v.ncols(), m.ncols());
assert_eq!(v.cstride(), m.ncols());
test_indexing(v, m.as_view());
let v = Strided::try_from_data(
&(m.as_slice()[..(4 * m.ncols() + 2)]),
m.nrows(),
2,
m.ncols(),
)
.unwrap();
assert_eq!(v.as_ptr(), ptr, "base pointer was not preserved");
let mut expected = rowmajor::Owned::from_element(5, 2, 0);
for row in 0..expected.nrows() {
for col in 0..expected.ncols() {
*expected.element_mut(row, col) = *m.element(row, col);
}
}
test_indexing(v, expected.as_view());
let v = Strided::try_from_data(&(m.as_slice()[1..]), m.nrows(), 2, m.ncols()).unwrap();
let mut expected = rowmajor::Owned::from_element(5, 2, 0);
for row in 0..expected.nrows() {
for col in 0..expected.ncols() {
*expected.element_mut(row, col) = *m.element(row, col + 1);
}
}
test_indexing(v, expected.as_view());
}
#[test]
fn matrix_conversion() {
let m = create_test_matrix(3, 4);
let ptr = m.as_ptr();
let v: Strided<_> = m.as_view().into();
assert_eq!(v.as_ptr(), ptr);
assert_eq!(v.cstride(), m.ncols());
assert_eq!(v.layout().linear_length(), m.layout().num_elements());
test_indexing(v, m.as_view());
}
#[test]
fn test_zero_sized() {
let m = create_test_matrix(5, 5);
let v = Strided::try_from_data(m.as_slice(), 0, 4, 5).unwrap();
assert_eq!(v.nrows(), 0);
assert_eq!(v.ncols(), 4);
assert_eq!(v.cstride(), 5);
let v = Strided::try_from_data(m.as_slice(), 5, 0, 5).unwrap();
assert_eq!(v.nrows(), 5);
assert_eq!(v.ncols(), 0);
assert_eq!(v.cstride(), 5);
for row in 0..v.nrows() {
let empty: &[usize] = &[];
assert_eq!(v.get_row(row).unwrap(), empty);
}
}
#[test]
fn test_try_shrink_from() {
let m = rowmajor::Owned::<usize>::from_element(10, 10, 0);
let nrows = m.nrows();
let ncols = m.ncols();
let s = Strided::try_from_data(m.as_slice(), nrows, ncols, ncols).unwrap();
assert_eq!(s.as_slice(), m.as_slice());
let s = Strided::try_from_data(m.as_slice(), nrows, 5, ncols).unwrap();
assert_eq!(s.as_ptr(), m.as_ptr());
let s = Strided::try_from_data(m.as_slice(), nrows, ncols, ncols + 1);
assert!(s.is_err());
}
#[test]
fn test_invalid_stride_is_an_error_not_a_panic() {
let m = rowmajor::Owned::<usize>::from_element(4, 4, 0);
let err = Strided::try_from_data(m.as_slice(), 2, 2, 1).unwrap_err();
assert!(matches!(err, TryFromError::LayoutError(_)));
}
}