use super::SimdCow;
use crate::align::Alignment;
use crate::arch::SimdArch;
use crate::kernel::SimdKernel;
use crate::ops::{Dot, Sub};
use crate::scalar::{FloatElement, Scalar};
use crate::vec::AlignedVec;
use crate::view::SimdError;
extern crate alloc;
impl<'a, T: 'a, Arch, Align> SimdCow<'a, T, Arch, Align>
where
T: Scalar,
Arch: SimdArch + SimdKernel<T>,
Align: Alignment,
{
#[inline]
pub fn add_scalar_cow(&self, rhs: T) -> SimdCow<'static, T, Arch, Align> {
broadcast_op::<T, Arch, Align>(
self,
rhs,
|a, b| a + b,
|va, vsplat| unsafe { Arch::add(va, vsplat) },
)
}
#[inline]
pub fn sub_scalar_cow(&self, rhs: T) -> SimdCow<'static, T, Arch, Align> {
broadcast_op::<T, Arch, Align>(
self,
rhs,
|a, b| a - b,
|va, vsplat| unsafe { Arch::sub(va, vsplat) },
)
}
#[inline]
pub fn mul_scalar_cow(&self, rhs: T) -> SimdCow<'static, T, Arch, Align> {
broadcast_op::<T, Arch, Align>(
self,
rhs,
|a, b| a * b,
|va, vsplat| unsafe { Arch::mul(va, vsplat) },
)
}
#[inline]
pub fn div_scalar_cow(&self, rhs: T) -> SimdCow<'static, T, Arch, Align> {
broadcast_op::<T, Arch, Align>(
self,
rhs,
|a, b| a / b,
|va, vsplat| unsafe { Arch::div(va, vsplat) },
)
}
#[inline]
pub fn div_cow(
&self,
other: &SimdCow<'_, T, Arch, Align>,
) -> Result<SimdCow<'static, T, Arch, Align>, SimdError> {
self.zip_cow(other, crate::ops::Div)
}
#[inline]
pub fn sub_cow_op(
&self,
other: &SimdCow<'_, T, Arch, Align>,
) -> Result<SimdCow<'static, T, Arch, Align>, SimdError> {
self.zip_cow(other, Sub)
}
}
impl<'a, T: 'a, Arch, Align> SimdCow<'a, T, Arch, Align>
where
T: Scalar + FloatElement,
Arch: SimdArch + SimdKernel<T>,
Align: Alignment,
{
#[inline]
pub fn norm_sq(&self) -> T {
let v = self.view();
v.zip_reduce(&v, Dot).unwrap_or(T::ZERO)
}
#[inline]
pub fn norm(&self) -> T {
self.norm_sq().sqrt()
}
#[inline]
pub fn normalize(&self) -> SimdCow<'static, T, Arch, Align> {
let n = self.norm();
if n == T::ZERO {
return SimdCow::zeros(self.len());
}
let inv = T::ONE / n;
self.mul_scalar_cow(inv)
}
#[inline]
pub fn histogram_cow(&self, n_bins: usize, lo: T, hi: T) -> alloc::vec::Vec<usize>
where
T: PartialOrd,
{
assert!(n_bins > 0, "n_bins must be > 0");
assert!(lo < hi, "lo must be < hi");
let lo_w = lo.to_f64();
let bin_width = (hi.to_f64() - lo_w) / n_bins as f64;
let mut counts = alloc::vec![0usize; n_bins];
for &x in self.as_ref().iter() {
if x < lo || x >= hi {
continue;
}
let bin = (((x.to_f64() - lo_w) / bin_width) as usize).min(n_bins - 1);
counts[bin] += 1;
}
counts
}
}
#[inline(always)]
fn broadcast_op<T, Arch, Align>(
cow: &SimdCow<'_, T, Arch, Align>,
rhs: T,
scalar_op: impl Fn(T, T) -> T + Copy,
vector_op: impl Fn(Arch::Vector, Arch::Vector) -> Arch::Vector + Copy,
) -> SimdCow<'static, T, Arch, Align>
where
T: Scalar,
Arch: SimdArch + SimdKernel<T>,
Align: Alignment,
{
let data = cow.as_ref();
let len = data.len();
let mut out: AlignedVec<T, Align> = AlignedVec::with_capacity(len);
let lane_count = Arch::LANE_COUNT;
let simd_len = (len / lane_count) * lane_count;
let ptr_in = data.as_ptr();
let ptr_out = out.as_mut_ptr();
unsafe {
let vsplat = Arch::splat(rhs);
let load = |p: *const T| -> Arch::Vector {
if crate::align::is_aligned_for_arch::<Arch, Align>() {
Arch::load_aligned(p)
} else {
Arch::load_unaligned(p)
}
};
let store = |p: *mut T, v: Arch::Vector| {
if crate::align::is_aligned_for_arch::<Arch, Align>() {
Arch::store_aligned(p, v);
} else {
Arch::store_unaligned(p, v);
}
};
let mut i = 0usize;
while i < simd_len {
let va = load(ptr_in.add(i));
let vr = vector_op(va, vsplat);
store(ptr_out.add(i), vr);
i += lane_count;
}
for i in simd_len..len {
core::ptr::write(ptr_out.add(i), scalar_op(*ptr_in.add(i), rhs));
}
out.set_len(len);
}
SimdCow::Owned(out)
}