use std::fmt::{Debug, Display};
use crate::{
layouts::{Data, HostDataMut, HostDataRef},
source::Source,
};
use bytemuck::Pod;
use rand_distr::num_traits::Zero;
pub trait ZnxInfos {
fn n(&self) -> usize;
fn log_n(&self) -> usize {
(usize::BITS - (self.n() - 1).leading_zeros()) as _
}
fn size(&self) -> usize;
fn poly_count(&self) -> usize;
}
pub trait VecZnxInfos: ZnxInfos {
fn cols(&self) -> usize;
}
pub trait MatZnxInfos: ZnxInfos {
fn rows(&self) -> usize;
fn cols_in(&self) -> usize;
fn cols_out(&self) -> usize;
}
pub(crate) fn raw_scalars<S: Pod>(data: &[u8], span: usize) -> &[S] {
let ptr: *const u8 = data.as_ptr();
assert!(
(ptr as usize).is_multiple_of(align_of::<S>()),
"buffer not aligned to align_of::<Scalar>() = {}",
align_of::<S>()
);
assert!(
span.checked_mul(size_of::<S>())
.expect("element view byte size overflows usize")
<= data.len(),
"element view ({} scalars of {} bytes) exceeds the {}-byte buffer: this container has no element view for its word type",
span,
size_of::<S>(),
data.len()
);
unsafe { std::slice::from_raw_parts(ptr as *const S, span) }
}
pub(crate) fn raw_scalars_mut<S: Pod>(data: &mut [u8], span: usize) -> &mut [S] {
let len: usize = data.len();
let ptr: *mut u8 = data.as_mut_ptr();
assert!(
(ptr as usize).is_multiple_of(align_of::<S>()),
"buffer not aligned to align_of::<Scalar>() = {}",
align_of::<S>()
);
assert!(
span.checked_mul(size_of::<S>())
.expect("element view byte size overflows usize")
<= len,
"element view ({} scalars of {} bytes) exceeds the {}-byte buffer: this container has no element view for its word type",
span,
size_of::<S>(),
len
);
unsafe { std::slice::from_raw_parts_mut(ptr as *mut S, span) }
}
pub(crate) fn element_view_span<T: ZnxInfos + ?Sized>(infos: &T) -> usize {
infos
.n()
.checked_mul(infos.poly_count())
.expect("element view scalar count overflows usize")
}
pub trait DataView {
type D: Data;
fn data(&self) -> &Self::D;
}
pub trait DataViewMut: DataView {
fn data_mut(&mut self) -> &mut Self::D;
}
pub trait ZnxView: VecZnxInfos + DataView<D: HostDataRef> {
type Scalar: Copy + Zero + Display + Debug + Pod;
#[doc(hidden)]
fn validate_element_view(&self) {}
fn as_ptr(&self) -> *const Self::Scalar {
self.validate_element_view();
let ptr: *const u8 = self.data().as_ref().as_ptr();
assert!(
(ptr as usize).is_multiple_of(align_of::<Self::Scalar>()),
"buffer not aligned to align_of::<Scalar>() = {}",
align_of::<Self::Scalar>()
);
ptr as *const Self::Scalar
}
fn raw(&self) -> &[Self::Scalar] {
self.validate_element_view();
raw_scalars(self.data().as_ref(), element_view_span(self))
}
fn at_ptr(&self, i: usize, j: usize) -> *const Self::Scalar {
self.validate_element_view();
assert!(i < self.cols(), "cols: {} >= self.cols(): {}", i, self.cols());
assert!(j < self.size(), "size: {} >= self.size(): {}", j, self.size());
let offset: usize = j
.checked_mul(self.cols())
.and_then(|x| x.checked_add(i))
.and_then(|x| x.checked_mul(self.n()))
.expect("element view offset overflows usize");
assert!(
offset
.checked_add(self.n())
.and_then(|x| x.checked_mul(size_of::<Self::Scalar>()))
.expect("element view byte size overflows usize")
<= self.data().as_ref().len(),
"element view of block ({}, {}) exceeds the {}-byte buffer: this container has no element view for its word type",
i,
j,
self.data().as_ref().len()
);
unsafe { self.as_ptr().add(offset) }
}
fn at(&self, i: usize, j: usize) -> &[Self::Scalar] {
unsafe { std::slice::from_raw_parts(self.at_ptr(i, j), self.n()) }
}
}
pub trait ZnxViewMut: ZnxView + DataViewMut<D: HostDataMut> {
fn as_mut_ptr(&mut self) -> *mut Self::Scalar {
self.validate_element_view();
let ptr: *mut u8 = self.data_mut().as_mut().as_mut_ptr();
assert!(
(ptr as usize).is_multiple_of(align_of::<Self::Scalar>()),
"buffer not aligned to align_of::<Scalar>() = {}",
align_of::<Self::Scalar>()
);
ptr as *mut Self::Scalar
}
fn raw_mut(&mut self) -> &mut [Self::Scalar] {
self.validate_element_view();
let span: usize = element_view_span(self);
raw_scalars_mut(self.data_mut().as_mut(), span)
}
fn at_mut_ptr(&mut self, i: usize, j: usize) -> *mut Self::Scalar {
self.validate_element_view();
assert!(i < self.cols(), "cols: {} >= self.cols(): {}", i, self.cols());
assert!(j < self.size(), "size: {} >= self.size(): {}", j, self.size());
let offset: usize = j
.checked_mul(self.cols())
.and_then(|x| x.checked_add(i))
.and_then(|x| x.checked_mul(self.n()))
.expect("element view offset overflows usize");
assert!(
offset
.checked_add(self.n())
.and_then(|x| x.checked_mul(size_of::<Self::Scalar>()))
.expect("element view byte size overflows usize")
<= self.data().as_ref().len(),
"element view of block ({}, {}) exceeds the {}-byte buffer: this container has no element view for its word type",
i,
j,
self.data().as_ref().len()
);
unsafe { self.as_mut_ptr().add(offset) }
}
fn at_mut(&mut self, i: usize, j: usize) -> &mut [Self::Scalar] {
unsafe { std::slice::from_raw_parts_mut(self.at_mut_ptr(i, j), self.n()) }
}
}
impl<T> ZnxViewMut for T where T: ZnxView + DataViewMut<D: HostDataMut> {}
pub trait ZnxZero
where
Self: Sized,
{
fn zero(&mut self);
fn zero_at(&mut self, i: usize, j: usize);
}
pub trait FillUniform {
fn fill_uniform(&mut self, log_bound: usize, source: &mut Source);
}