use crate::ScalarBase;
use std::any::{Any, TypeId};
use std::cell::RefCell;
use std::collections::HashMap;
use std::mem::MaybeUninit;
use strided_kernel::MaybeSendSync;
use strided_view::{RawStridedMut, RawStridedRef, StridedArray, StridedView, StridedViewMut};
pub struct ContiguousOperand<T: Copy + 'static> {
ptr: *const T,
row_stride: isize,
col_stride: isize,
batch_strides: Vec<isize>,
conj: bool,
pub(crate) _buf: Option<StridedArray<T>>,
buf_is_pooled: bool,
}
pub struct ContiguousOperandMut<T: Copy + 'static> {
ptr: *mut T,
row_stride: isize,
col_stride: isize,
batch_strides: Vec<isize>,
needs_writeback: bool,
pub(crate) _buf: Option<StridedArray<T>>,
buf_is_pooled: bool,
}
#[allow(dead_code)]
pub(crate) struct UninitContiguousOperand<'a, 'b, T: Copy + 'static> {
destination: &'a mut RawStridedMut<'b, MaybeUninit<T>>,
ptr: *mut MaybeUninit<T>,
row_stride: isize,
col_stride: isize,
batch_strides: Vec<isize>,
temp: Option<StridedArray<MaybeUninit<T>>>,
writeback: Option<strided_kernel::CopyPlan>,
}
thread_local! {
static BUFFER_POOL: RefCell<HashMap<TypeId, Box<dyn Any>>> = RefCell::new(HashMap::new());
}
const MAX_POOL_PER_TYPE: usize = 16;
const MAX_POOLED_BYTES: usize = 64 * 1024 * 1024;
fn take_pooled_vec_uninit<T: Copy + 'static>(len: usize) -> Vec<T> {
BUFFER_POOL.with(|pool| {
let mut pool = pool.borrow_mut();
let entry = pool
.entry(TypeId::of::<T>())
.or_insert_with(|| Box::new(Vec::<Vec<T>>::new()));
let vecs = entry
.downcast_mut::<Vec<Vec<T>>>()
.expect("buffer pool type mismatch");
let mut best_idx = None;
let mut best_cap = usize::MAX;
for (idx, v) in vecs.iter().enumerate() {
let cap = v.capacity();
if cap >= len && cap < best_cap {
best_idx = Some(idx);
best_cap = cap;
}
}
let mut data = best_idx
.map(|idx| vecs.swap_remove(idx))
.unwrap_or_else(|| Vec::with_capacity(len));
if data.capacity() < len {
data.reserve(len - data.capacity());
}
unsafe { data.set_len(len) };
data
})
}
fn return_pooled_vec<T: Copy + 'static>(mut data: Vec<T>) {
let bytes = data.capacity().saturating_mul(std::mem::size_of::<T>());
if bytes == 0 || bytes > MAX_POOLED_BYTES {
return;
}
data.clear();
BUFFER_POOL.with(|pool| {
let mut pool = pool.borrow_mut();
let entry = pool
.entry(TypeId::of::<T>())
.or_insert_with(|| Box::new(Vec::<Vec<T>>::new()));
let vecs = entry
.downcast_mut::<Vec<Vec<T>>>()
.expect("buffer pool type mismatch");
if vecs.len() >= MAX_POOL_PER_TYPE {
if let Some((min_idx, min_cap)) = vecs
.iter()
.enumerate()
.map(|(i, v)| (i, v.capacity()))
.min_by_key(|(_, cap)| *cap)
{
if min_cap < data.capacity() {
vecs.swap_remove(min_idx);
vecs.push(data);
}
}
} else {
vecs.push(data);
}
});
}
fn alloc_col_major_uninit_with_pool<T: Copy + 'static>(
dims: &[usize],
) -> strided_view::Result<(StridedArray<T>, bool)> {
let total = dims
.iter()
.try_fold(1usize, |total, &dim| total.checked_mul(dim))
.ok_or(strided_view::StridedError::OffsetOverflow)?
.max(1);
let bytes = total.saturating_mul(std::mem::size_of::<T>());
if bytes == 0 || bytes > MAX_POOLED_BYTES {
return Ok((alloc_col_major_uninit(dims)?, false));
}
let data = take_pooled_vec_uninit::<T>(total);
let arr = unsafe { StridedArray::col_major_from_buffer_uninit(data, dims) };
Ok((arr, true))
}
fn alloc_maybe_pooled<T: Copy + 'static>(
dims: &[usize],
use_pool: bool,
) -> strided_view::Result<(StridedArray<T>, bool)> {
if use_pool {
alloc_col_major_uninit_with_pool(dims)
} else {
Ok((alloc_col_major_uninit(dims)?, false))
}
}
#[cfg(test)]
fn pooled_count_for_type<T: 'static>() -> usize {
BUFFER_POOL.with(|pool| {
let mut pool = pool.borrow_mut();
let Some(entry) = pool.get_mut(&TypeId::of::<T>()) else {
return 0;
};
entry
.downcast_mut::<Vec<Vec<T>>>()
.map_or(0, |vecs| vecs.len())
})
}
impl<T: Copy + 'static> ContiguousOperand<T> {
#[inline]
pub fn ptr(&self) -> *const T {
self.ptr
}
#[inline]
pub fn row_stride(&self) -> isize {
self.row_stride
}
#[inline]
pub fn col_stride(&self) -> isize {
self.col_stride
}
#[inline]
pub fn batch_strides(&self) -> &[isize] {
&self.batch_strides
}
#[inline]
pub fn conj(&self) -> bool {
self.conj
}
#[cfg(test)]
#[inline]
pub(crate) fn has_buf(&self) -> bool {
self._buf.is_some()
}
}
impl<T: Copy + 'static> ContiguousOperandMut<T> {
#[inline]
pub fn ptr(&self) -> *mut T {
self.ptr
}
#[inline]
pub fn row_stride(&self) -> isize {
self.row_stride
}
#[inline]
pub fn col_stride(&self) -> isize {
self.col_stride
}
#[inline]
pub fn batch_strides(&self) -> &[isize] {
&self.batch_strides
}
#[cfg(test)]
#[inline]
pub(crate) fn has_buf(&self) -> bool {
self._buf.is_some()
}
#[cfg(test)]
#[inline]
pub(crate) fn needs_writeback(&self) -> bool {
self.needs_writeback
}
}
impl<T: Copy + Send + Sync> ContiguousOperandMut<T> {
pub fn finalize_into(self, dest: &mut StridedViewMut<T>) -> crate::Result<()> {
if self.needs_writeback {
if let Some(ref buf) = self._buf {
strided_perm::copy_into(dest, &buf.view())?;
}
}
Ok(())
}
pub fn finalize_raw_into(self, dest: &mut RawStridedMut<'_, T>) -> crate::Result<()> {
if self.needs_writeback {
if let Some(ref buf) = self._buf {
let mut dest_view = dest.as_view_mut();
strided_perm::copy_into(&mut dest_view, &buf.view())?;
}
}
Ok(())
}
}
impl<T: Copy + 'static> Drop for ContiguousOperand<T> {
fn drop(&mut self) {
if self.buf_is_pooled {
if let Some(arr) = self._buf.take() {
return_pooled_vec(arr.into_data());
}
}
}
}
impl<T: Copy + 'static> Drop for ContiguousOperandMut<T> {
fn drop(&mut self) {
if self.buf_is_pooled {
if let Some(arr) = self._buf.take() {
return_pooled_vec(arr.into_data());
}
}
}
}
struct ContiguityCheck {
fused_g1: Option<(usize, isize)>,
fused_g2: Option<(usize, isize)>,
needs_copy: bool,
}
fn try_fuse_col_major_group(dims: &[usize], strides: &[isize]) -> Option<(usize, isize)> {
if dims.len() != strides.len() {
return None;
}
let total = dims
.iter()
.try_fold(1usize, |acc, &dim| acc.checked_mul(dim))?;
if dims.is_empty() {
return Some((1, 0));
}
let mut base_stride = None;
let mut expected_stride = None;
for (&dim, &stride) in dims.iter().zip(strides.iter()) {
if dim <= 1 {
continue;
}
if stride == 0 {
return None;
}
if let Some(expected) = expected_stride {
if stride != expected {
return None;
}
} else {
base_stride = Some(stride);
}
let dim = isize::try_from(dim).ok()?;
expected_stride = Some(stride.checked_mul(dim)?);
}
let stride = base_stride.unwrap_or_else(|| {
strides
.iter()
.copied()
.min_by_key(|stride| stride.unsigned_abs())
.unwrap_or(0)
});
Some((total, stride))
}
fn check_contiguity(
group1_dims: &[usize],
group1_strides: &[isize],
group2_dims: &[usize],
group2_strides: &[isize],
requires_unit_stride: bool,
) -> ContiguityCheck {
let fused_g1 = try_fuse_col_major_group(group1_dims, group1_strides);
let fused_g2 = try_fuse_col_major_group(group2_dims, group2_strides);
let mut needs_copy = fused_g1.is_none() || fused_g2.is_none();
if requires_unit_stride && !needs_copy {
let (_, rs) = fused_g1.unwrap();
let (_, cs) = fused_g2.unwrap();
if rs != 0 && rs != 1 && cs != 0 && cs != 1 {
needs_copy = true;
}
}
ContiguityCheck {
fused_g1,
fused_g2,
needs_copy,
}
}
fn col_major_layout(
buf: &StridedArray<impl Copy>,
n_group1: usize,
n_inner: usize,
) -> (isize, isize, Vec<isize>) {
let m: usize = buf.dims()[..n_group1].iter().product::<usize>().max(1);
let row_stride = if m == 0 { 0 } else { 1isize };
let col_stride = m as isize;
let batch_strides = buf.strides()[n_inner..].to_vec();
(row_stride, col_stride, batch_strides)
}
pub(crate) fn alloc_col_major_uninit<T: Copy>(
dims: &[usize],
) -> strided_view::Result<StridedArray<T>> {
let total = dims
.iter()
.try_fold(1usize, |total, &dim| total.checked_mul(dim))
.ok_or(strided_view::StridedError::OffsetOverflow)?
.max(1);
let mut data = Vec::with_capacity(total);
unsafe { data.set_len(total) };
let mut strides = vec![0isize; dims.len()];
if !dims.is_empty() {
strides[0] = 1;
for i in 1..dims.len() {
strides[i] = strides[i - 1] * dims[i - 1] as isize;
}
}
Ok(StridedArray::from_parts(data, dims, &strides, 0)?)
}
pub fn prepare_input_view<T: ScalarBase + 'static>(
view: &StridedView<T>,
n_group1: usize,
n_group2: usize,
conj: bool,
requires_unit_stride: bool,
use_pool: bool,
materialize_conj_fn: Option<fn(T) -> T>,
) -> crate::Result<ContiguousOperand<T>> {
let dims = view.dims();
let strides = view.strides();
let n_inner = n_group1 + n_group2;
if let Some(conj_fn) = materialize_conj_fn {
if conj {
let (mut buf, buf_is_pooled) = alloc_maybe_pooled(dims, use_pool)?;
strided_kernel::map_into(&mut buf.view_mut(), view, conj_fn)?;
let ptr = buf.view().ptr();
let (row_stride, col_stride, batch_strides) = col_major_layout(&buf, n_group1, n_inner);
return Ok(ContiguousOperand {
ptr,
row_stride,
col_stride,
batch_strides,
conj: false,
_buf: Some(buf),
buf_is_pooled,
});
}
}
let check = check_contiguity(
&dims[..n_group1],
&strides[..n_group1],
&dims[n_group1..n_inner],
&strides[n_group1..n_inner],
requires_unit_stride,
);
if check.needs_copy {
let (mut buf, buf_is_pooled) = alloc_maybe_pooled(dims, use_pool)?;
strided_kernel::copy_into_col_major(&mut buf.view_mut(), view)?;
let ptr = buf.view().ptr();
let (row_stride, col_stride, batch_strides) = col_major_layout(&buf, n_group1, n_inner);
Ok(ContiguousOperand {
ptr,
row_stride,
col_stride,
batch_strides,
conj,
_buf: Some(buf),
buf_is_pooled,
})
} else {
let (_, rs) = check.fused_g1.unwrap();
let (_, cs) = check.fused_g2.unwrap();
Ok(ContiguousOperand {
ptr: view.ptr(),
row_stride: rs,
col_stride: cs,
batch_strides: strides[n_inner..].to_vec(),
conj,
_buf: None,
buf_is_pooled: false,
})
}
}
pub fn prepare_input_raw<T: ScalarBase + 'static>(
view: &RawStridedRef<'_, T>,
n_group1: usize,
n_group2: usize,
conj: bool,
requires_unit_stride: bool,
use_pool: bool,
materialize_conj_fn: Option<fn(T) -> T>,
) -> crate::Result<ContiguousOperand<T>> {
let dims = view.dims();
let strides = view.strides();
let n_inner = n_group1 + n_group2;
if let Some(conj_fn) = materialize_conj_fn {
if conj {
let (mut buf, buf_is_pooled) = alloc_maybe_pooled(dims, use_pool)?;
strided_kernel::map_into(&mut buf.view_mut(), &view.as_view(), conj_fn)?;
let ptr = buf.view().ptr();
let (row_stride, col_stride, batch_strides) = col_major_layout(&buf, n_group1, n_inner);
return Ok(ContiguousOperand {
ptr,
row_stride,
col_stride,
batch_strides,
conj: false,
_buf: Some(buf),
buf_is_pooled,
});
}
}
let check = check_contiguity(
&dims[..n_group1],
&strides[..n_group1],
&dims[n_group1..n_inner],
&strides[n_group1..n_inner],
requires_unit_stride,
);
if check.needs_copy {
let (mut buf, buf_is_pooled) = alloc_maybe_pooled(dims, use_pool)?;
strided_kernel::copy_into_col_major(&mut buf.view_mut(), &view.as_view())?;
let ptr = buf.view().ptr();
let (row_stride, col_stride, batch_strides) = col_major_layout(&buf, n_group1, n_inner);
Ok(ContiguousOperand {
ptr,
row_stride,
col_stride,
batch_strides,
conj,
_buf: Some(buf),
buf_is_pooled,
})
} else {
let (_, rs) = check.fused_g1.unwrap();
let (_, cs) = check.fused_g2.unwrap();
Ok(ContiguousOperand {
ptr: view.ptr(),
row_stride: rs,
col_stride: cs,
batch_strides: strides[n_inner..].to_vec(),
conj,
_buf: None,
buf_is_pooled: false,
})
}
}
pub fn prepare_input_owned<T: ScalarBase + 'static>(
arr: StridedArray<T>,
n_group1: usize,
n_group2: usize,
conj: bool,
requires_unit_stride: bool,
use_pool: bool,
materialize_conj_fn: Option<fn(T) -> T>,
) -> crate::Result<ContiguousOperand<T>> {
let dims = arr.dims().to_vec();
let strides = arr.strides().to_vec();
let n_inner = n_group1 + n_group2;
if let Some(conj_fn) = materialize_conj_fn {
if conj {
let (mut buf, buf_is_pooled) = alloc_maybe_pooled(&dims, use_pool)?;
strided_kernel::map_into(&mut buf.view_mut(), &arr.view(), conj_fn)?;
let ptr = buf.view().ptr();
let (row_stride, col_stride, batch_strides) = col_major_layout(&buf, n_group1, n_inner);
return Ok(ContiguousOperand {
ptr,
row_stride,
col_stride,
batch_strides,
conj: false,
_buf: Some(buf),
buf_is_pooled,
});
}
}
let check = check_contiguity(
&dims[..n_group1],
&strides[..n_group1],
&dims[n_group1..n_inner],
&strides[n_group1..n_inner],
requires_unit_stride,
);
if check.needs_copy {
let (mut buf, buf_is_pooled) = alloc_maybe_pooled(&dims, use_pool)?;
strided_kernel::copy_into_col_major(&mut buf.view_mut(), &arr.view())?;
let ptr = buf.view().ptr();
let (row_stride, col_stride, batch_strides) = col_major_layout(&buf, n_group1, n_inner);
Ok(ContiguousOperand {
ptr,
row_stride,
col_stride,
batch_strides,
conj,
_buf: Some(buf),
buf_is_pooled,
})
} else {
let (_, rs) = check.fused_g1.unwrap();
let (_, cs) = check.fused_g2.unwrap();
let ptr = arr.view().ptr();
Ok(ContiguousOperand {
ptr,
row_stride: rs,
col_stride: cs,
batch_strides: strides[n_inner..].to_vec(),
conj,
_buf: Some(arr),
buf_is_pooled: false,
})
}
}
pub fn prepare_output_view<T: ScalarBase + 'static>(
view: &mut StridedViewMut<T>,
n_group1: usize,
n_group2: usize,
beta: T,
requires_unit_stride: bool,
use_pool: bool,
) -> crate::Result<ContiguousOperandMut<T>> {
let dims = view.dims().to_vec();
let strides = view.strides().to_vec();
let n_inner = n_group1 + n_group2;
let check = check_contiguity(
&dims[..n_group1],
&strides[..n_group1],
&dims[n_group1..n_inner],
&strides[n_group1..n_inner],
requires_unit_stride,
);
if check.needs_copy {
let (mut buf, buf_is_pooled) = alloc_maybe_pooled(&dims, use_pool)?;
if beta != T::zero() {
strided_kernel::copy_into_col_major(&mut buf.view_mut(), &view.as_view())?;
}
let ptr = buf.view_mut().as_mut_ptr();
let (row_stride, col_stride, batch_strides) = col_major_layout(&buf, n_group1, n_inner);
Ok(ContiguousOperandMut {
ptr,
row_stride,
col_stride,
batch_strides,
needs_writeback: true,
_buf: Some(buf),
buf_is_pooled,
})
} else {
let (_, rs) = check.fused_g1.unwrap();
let (_, cs) = check.fused_g2.unwrap();
Ok(ContiguousOperandMut {
ptr: view.as_mut_ptr(),
row_stride: rs,
col_stride: cs,
batch_strides: strides[n_inner..].to_vec(),
needs_writeback: false,
_buf: None,
buf_is_pooled: false,
})
}
}
pub fn prepare_output_raw<T: ScalarBase + 'static>(
view: &mut RawStridedMut<'_, T>,
n_group1: usize,
n_group2: usize,
beta: T,
requires_unit_stride: bool,
use_pool: bool,
) -> crate::Result<ContiguousOperandMut<T>> {
let dims = view.dims().to_vec();
let strides = view.strides().to_vec();
let n_inner = n_group1 + n_group2;
let check = check_contiguity(
&dims[..n_group1],
&strides[..n_group1],
&dims[n_group1..n_inner],
&strides[n_group1..n_inner],
requires_unit_stride,
);
if check.needs_copy {
let (mut buf, buf_is_pooled) = alloc_maybe_pooled(&dims, use_pool)?;
if beta != T::zero() {
strided_kernel::copy_into_col_major(&mut buf.view_mut(), &view.as_view())?;
}
let ptr = buf.view_mut().as_mut_ptr();
let (row_stride, col_stride, batch_strides) = col_major_layout(&buf, n_group1, n_inner);
Ok(ContiguousOperandMut {
ptr,
row_stride,
col_stride,
batch_strides,
needs_writeback: true,
_buf: Some(buf),
buf_is_pooled,
})
} else {
let (_, rs) = check.fused_g1.unwrap();
let (_, cs) = check.fused_g2.unwrap();
Ok(ContiguousOperandMut {
ptr: view.as_mut_ptr(),
row_stride: rs,
col_stride: cs,
batch_strides: strides[n_inner..].to_vec(),
needs_writeback: false,
_buf: None,
buf_is_pooled: false,
})
}
}
#[allow(dead_code)]
impl<'a, 'b, T: Copy + MaybeSendSync + 'static> UninitContiguousOperand<'a, 'b, T> {
#[inline]
pub(crate) fn ptr(&self) -> *mut MaybeUninit<T> {
self.ptr
}
#[inline]
pub(crate) fn row_stride(&self) -> isize {
self.row_stride
}
#[inline]
pub(crate) fn col_stride(&self) -> isize {
self.col_stride
}
#[inline]
pub(crate) fn batch_strides(&self) -> &[isize] {
&self.batch_strides
}
pub(crate) fn finalize(self) -> crate::Result<()> {
let Self {
destination,
temp,
writeback,
..
} = self;
let (Some(temp), Some(writeback)) = (temp, writeback) else {
return Ok(());
};
let dims = temp.dims().to_vec();
let strides = temp.strides().to_vec();
let data = temp.into_data();
let len = data.len();
let cap = data.capacity();
let ptr = data.as_ptr().cast_mut().cast::<T>();
std::mem::forget(data);
let initialized = unsafe {
StridedArray::from_parts(Vec::from_raw_parts(ptr, len, cap), &dims, &strides, 0)
}?;
let source = RawStridedRef::new(
initialized.data(),
initialized.dims(),
initialized.strides(),
initialized.view().offset(),
)?;
writeback.execute_uninit(destination, &source)?;
Ok(())
}
}
#[allow(dead_code)]
pub(crate) fn prepare_output_raw_uninit<'a, 'b, T: ScalarBase + MaybeSendSync + 'static>(
destination: &'a mut RawStridedMut<'b, MaybeUninit<T>>,
n_group1: usize,
n_group2: usize,
requires_unit_stride: bool,
) -> crate::Result<UninitContiguousOperand<'a, 'b, T>> {
let dims = destination.dims().to_vec();
let strides = destination.strides().to_vec();
let n_inner = n_group1 + n_group2;
let check = check_contiguity(
&dims[..n_group1],
&strides[..n_group1],
&dims[n_group1..n_inner],
&strides[n_group1..n_inner],
requires_unit_stride,
);
if !check.needs_copy {
let Some((_, row_stride)) = check.fused_g1 else {
return Err(strided_view::StridedError::PlanLayoutMismatch.into());
};
let Some((_, col_stride)) = check.fused_g2 else {
return Err(strided_view::StridedError::PlanLayoutMismatch.into());
};
let ptr = destination.as_mut_ptr();
return Ok(UninitContiguousOperand {
destination,
ptr: ptr.cast(),
row_stride,
col_stride,
batch_strides: strides[n_inner..].to_vec(),
temp: None,
writeback: None,
});
}
let temp = alloc_col_major_uninit::<MaybeUninit<T>>(&dims)?;
let ptr = temp.view().ptr().cast_mut();
let (row_stride, col_stride, batch_strides) = col_major_layout(&temp, n_group1, n_inner);
let writeback = strided_kernel::CopyPlan::compile(&dims, &strides, temp.strides())?;
Ok(UninitContiguousOperand {
destination,
ptr,
row_stride,
col_stride,
batch_strides,
temp: Some(temp),
writeback: Some(writeback),
})
}
#[cfg(test)]
mod tests_generic_backend {
use super::*;
use crate::backend::{Backend, NaiveBackend};
#[test]
fn test_input_for_backend_contiguous() {
let a = StridedArray::<f64>::col_major(&[2, 3]);
let view = a.view();
let op = prepare_input_view(
&view,
1,
1,
false,
<NaiveBackend as Backend<f64>>::REQUIRES_UNIT_STRIDE,
false,
None,
)
.unwrap();
assert!(op._buf.is_none());
assert_eq!(op.row_stride(), 1);
assert_eq!(op.col_stride(), 2);
assert!(!op.conj());
}
#[test]
fn test_input_for_backend_non_contiguous() {
let data = vec![0.0f64; 100];
let a = StridedArray::<f64>::from_parts(data, &[2, 3, 4], &[20, 4, 1], 0).unwrap();
let view = a.view();
let op = prepare_input_view(
&view,
2,
1,
false,
<NaiveBackend as Backend<f64>>::REQUIRES_UNIT_STRIDE,
false,
None,
)
.unwrap();
assert!(op._buf.is_some());
assert_eq!(op.row_stride(), 1);
assert_eq!(op.col_stride(), 6);
}
#[test]
fn test_output_for_backend_contiguous() {
let mut c = StridedArray::<f64>::col_major(&[2, 3]);
let mut view = c.view_mut();
let op = prepare_output_view(
&mut view,
1,
1,
0.0,
<NaiveBackend as Backend<f64>>::REQUIRES_UNIT_STRIDE,
false,
)
.unwrap();
assert!(!op.needs_writeback);
assert!(op._buf.is_none());
assert_eq!(op.row_stride(), 1);
assert_eq!(op.col_stride(), 2);
}
#[test]
fn test_output_for_backend_non_contiguous_beta_zero() {
let data = vec![0.0f64; 100];
let mut c = StridedArray::<f64>::from_parts(data, &[2, 3, 4], &[20, 4, 1], 0).unwrap();
let mut view = c.view_mut();
let op = prepare_output_view(
&mut view,
2,
1,
0.0,
<NaiveBackend as Backend<f64>>::REQUIRES_UNIT_STRIDE,
false,
)
.unwrap();
assert!(op.needs_writeback);
assert!(op._buf.is_some());
assert_eq!(op.row_stride(), 1);
assert_eq!(op.col_stride(), 6);
}
#[test]
fn test_output_for_backend_non_contiguous_beta_nonzero_and_finalize() {
let mut data = vec![0.0f64; 30];
data[0] = 10.0;
data[1] = 20.0;
data[10] = 40.0;
let mut c = StridedArray::<f64>::from_parts(data, &[2, 3, 1], &[10, 1, 1], 0).unwrap();
let mut view = c.view_mut();
let op = prepare_output_view(
&mut view,
2,
1,
1.0,
<NaiveBackend as Backend<f64>>::REQUIRES_UNIT_STRIDE,
false,
)
.unwrap();
assert!(op.needs_writeback);
let buf = op._buf.as_ref().unwrap();
assert_eq!(buf.get(&[0, 0, 0]), 10.0);
assert_eq!(buf.get(&[0, 1, 0]), 20.0);
assert_eq!(buf.get(&[1, 0, 0]), 40.0);
op.finalize_into(&mut view).unwrap();
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::backend::{ActiveBackend, Backend};
const UNIT_STRIDE: bool = <ActiveBackend as Backend<f64>>::REQUIRES_UNIT_STRIDE;
#[test]
fn test_borrowed_contiguous_no_copy() {
let a = StridedArray::<f64>::col_major(&[2, 3]);
let view = a.view();
let op = prepare_input_view(&view, 1, 1, false, UNIT_STRIDE, true, None).unwrap();
assert!(!op.has_buf());
assert_eq!(op.row_stride(), 1);
assert_eq!(op.col_stride(), 2);
assert!(!op.conj());
}
#[test]
fn test_borrowed_transposed_matrix_no_copy() {
let data = vec![0.0f64; 6];
let a_t = StridedArray::<f64>::from_parts(data, &[2, 3], &[3, 1], 0).unwrap();
let view = a_t.view();
let op = prepare_input_view(&view, 1, 1, false, UNIT_STRIDE, true, None).unwrap();
assert!(!op.has_buf());
assert_eq!(op.row_stride(), 3);
assert_eq!(op.col_stride(), 1);
}
#[test]
fn test_borrowed_batched_transposed_matrix_no_copy() {
let data = vec![0.0f64; 2 * 3 * 5];
let a_t = StridedArray::<f64>::from_parts(data, &[2, 3, 5], &[3, 1, 6], 0).unwrap();
let view = a_t.view();
let op = prepare_input_view(&view, 1, 1, false, UNIT_STRIDE, true, None).unwrap();
assert!(!op.has_buf());
assert_eq!(op.row_stride(), 3);
assert_eq!(op.col_stride(), 1);
assert_eq!(op.batch_strides(), &[6]);
}
#[test]
fn test_borrowed_non_contiguous_copies() {
let data = vec![0.0f64; 100];
let a = StridedArray::<f64>::from_parts(data, &[2, 3, 4], &[20, 4, 1], 0).unwrap();
let view = a.view();
let op = prepare_input_view(&view, 2, 1, false, UNIT_STRIDE, true, None).unwrap();
assert!(op.has_buf());
assert_eq!(op.row_stride(), 1);
assert_eq!(op.col_stride(), 6);
}
#[test]
fn test_owned_contiguous_no_copy() {
let a = StridedArray::<f64>::col_major(&[2, 3]);
let op = prepare_input_owned(a, 1, 1, false, UNIT_STRIDE, true, None).unwrap();
assert!(op.has_buf());
assert_eq!(op.row_stride(), 1);
assert_eq!(op.col_stride(), 2);
}
#[test]
fn test_owned_non_contiguous_copies() {
let data = vec![0.0f64; 100];
let a = StridedArray::<f64>::from_parts(data, &[2, 3, 4], &[20, 4, 1], 0).unwrap();
let op = prepare_input_owned(a, 2, 1, false, UNIT_STRIDE, true, None).unwrap();
assert!(op.has_buf());
assert_eq!(op.row_stride(), 1);
assert_eq!(op.col_stride(), 6);
}
#[test]
fn test_output_view_contiguous() {
let mut c = StridedArray::<f64>::col_major(&[2, 3]);
let mut view = c.view_mut();
let op = prepare_output_view(&mut view, 1, 1, 0.0, UNIT_STRIDE, true).unwrap();
assert!(!op.needs_writeback());
assert!(!op.has_buf());
assert_eq!(op.row_stride(), 1);
assert_eq!(op.col_stride(), 2);
}
#[test]
fn test_output_view_non_contiguous_beta_zero() {
let data = vec![0.0f64; 100];
let mut c = StridedArray::<f64>::from_parts(data, &[2, 3, 4], &[20, 4, 1], 0).unwrap();
let mut view = c.view_mut();
let op = prepare_output_view(&mut view, 2, 1, 0.0, UNIT_STRIDE, true).unwrap();
assert!(op.needs_writeback());
assert!(op.has_buf());
assert_eq!(op.row_stride(), 1);
assert_eq!(op.col_stride(), 6);
}
#[test]
fn test_output_view_non_contiguous_beta_nonzero_and_finalize() {
let mut data = vec![0.0f64; 30];
data[0] = 10.0;
data[1] = 20.0;
data[2] = 30.0;
data[10] = 40.0;
data[11] = 50.0;
data[12] = 60.0;
let mut c = StridedArray::<f64>::from_parts(data, &[2, 3, 1], &[10, 1, 1], 0).unwrap();
assert_eq!(c.get(&[0, 0, 0]), 10.0);
assert_eq!(c.get(&[1, 1, 0]), 50.0);
let mut view = c.view_mut();
let mut op = prepare_output_view(&mut view, 2, 1, 1.0, UNIT_STRIDE, true).unwrap();
assert!(op.needs_writeback());
assert!(op.has_buf());
let buf = op._buf.as_ref().unwrap();
assert_eq!(buf.get(&[0, 0, 0]), 10.0);
assert_eq!(buf.get(&[1, 1, 0]), 50.0);
{
let result_data = vec![100.0f64; 6];
let result =
StridedArray::<f64>::from_parts(result_data, &[2, 3, 1], &[3, 1, 1], 0).unwrap();
strided_kernel::copy_into(&mut op._buf.as_mut().unwrap().view_mut(), &result.view())
.unwrap();
op.ptr = op._buf.as_mut().unwrap().view_mut().as_mut_ptr();
}
op.finalize_into(&mut view).unwrap();
assert_eq!(c.get(&[0, 0, 0]), 100.0);
assert_eq!(c.get(&[0, 1, 0]), 100.0);
assert_eq!(c.get(&[0, 2, 0]), 100.0);
assert_eq!(c.get(&[1, 0, 0]), 100.0);
assert_eq!(c.get(&[1, 1, 0]), 100.0);
assert_eq!(c.get(&[1, 2, 0]), 100.0);
}
#[test]
fn test_prepare_input_view_temp_buffer_is_recycled() {
let before = pooled_count_for_type::<f64>();
let data = vec![0.0f64; 100];
let a = StridedArray::<f64>::from_parts(data, &[2, 3, 4], &[20, 4, 1], 0).unwrap();
let view = a.view();
{
let op = prepare_input_view(&view, 2, 1, false, UNIT_STRIDE, true, None).unwrap();
assert!(op.has_buf());
}
let after = pooled_count_for_type::<f64>();
assert!(after >= before.saturating_add(1));
}
#[test]
fn test_output_raw_uninit_direct_finalize() {
let mut storage = vec![MaybeUninit::<f64>::uninit(); 6];
let mut raw = RawStridedMut::new(&mut storage, &[2, 3], &[1, 2], 0).unwrap();
let op = prepare_output_raw_uninit(&mut raw, 1, 1, false).unwrap();
assert!(op.temp.is_none());
assert!(op.writeback.is_none());
for j in 0..3 {
for i in 0..2 {
let offset = i as isize * op.row_stride() + j as isize * op.col_stride();
unsafe {
op.ptr()
.offset(offset)
.write(MaybeUninit::new((10 + i + 2 * j) as f64));
}
}
}
op.finalize().unwrap();
drop(raw);
let values: Vec<f64> = storage
.into_iter()
.map(|x| unsafe { x.assume_init() })
.collect();
assert_eq!(values, vec![10.0, 11.0, 12.0, 13.0, 14.0, 15.0]);
}
#[test]
fn test_output_raw_uninit_temp_finalize_writes_back() {
let mut storage = vec![MaybeUninit::<f64>::uninit(); 30];
let mut raw = RawStridedMut::new(&mut storage, &[2, 3, 1], &[10, 1, 1], 0).unwrap();
let op = prepare_output_raw_uninit(&mut raw, 2, 1, true).unwrap();
assert!(op.temp.is_some());
assert!(op.writeback.is_some());
for i in 0..6 {
let offset = i as isize * op.row_stride();
unsafe {
op.ptr()
.offset(offset)
.write(MaybeUninit::new((20 + i) as f64));
}
}
op.finalize().unwrap();
drop(raw);
unsafe {
assert_eq!(storage[0].assume_init_ref(), &20.0);
assert_eq!(storage[1].assume_init_ref(), &22.0);
assert_eq!(storage[2].assume_init_ref(), &24.0);
assert_eq!(storage[10].assume_init_ref(), &21.0);
assert_eq!(storage[11].assume_init_ref(), &23.0);
assert_eq!(storage[12].assume_init_ref(), &25.0);
}
}
}