use crate::align::Alignment;
use crate::arch::SimdArch;
use crate::execution::ExecutionMode;
use crate::kernel::SimdKernel;
use crate::scalar::Scalar;
use crate::view::SimdView;
pub struct ZipChunks<'a, 'b, T, Arch: SimdArch, Align: Alignment, Mode: ExecutionMode> {
base_a: *const T,
base_b: *const T,
pos: usize,
simd_end: usize,
total_a: usize,
total_b: usize,
_marker: core::marker::PhantomData<(&'a T, &'b T, Arch, Align, Mode)>,
}
unsafe impl<'a, 'b, T: Send, Arch: SimdArch + SimdKernel<T>, Align: Alignment, Mode: ExecutionMode>
Send for ZipChunks<'a, 'b, T, Arch, Align, Mode>
where
T: Scalar,
{
}
unsafe impl<'a, 'b, T: Sync, Arch: SimdArch + SimdKernel<T>, Align: Alignment, Mode: ExecutionMode>
Sync for ZipChunks<'a, 'b, T, Arch, Align, Mode>
where
T: Scalar,
{
}
impl<'a, 'b, T: Scalar, Arch: SimdArch + SimdKernel<T>, Align: Alignment, Mode: ExecutionMode>
ZipChunks<'a, 'b, T, Arch, Align, Mode>
{
#[inline]
pub(crate) unsafe fn from_raw_parts(
base_a: *const T,
total_a: usize,
base_b: *const T,
total_b: usize,
) -> Self {
let lane_count = Arch::LANE_COUNT;
let min_total = total_a.min(total_b);
let simd_end = (min_total / lane_count) * lane_count;
Self {
base_a,
base_b,
pos: 0,
simd_end,
total_a,
total_b,
_marker: core::marker::PhantomData,
}
}
#[inline(always)]
pub fn remainder(&self) -> (&'a [T], &'b [T]) {
unsafe {
(
core::slice::from_raw_parts(
self.base_a.add(self.simd_end),
self.total_a - self.simd_end,
),
core::slice::from_raw_parts(
self.base_b.add(self.simd_end),
self.total_b - self.simd_end,
),
)
}
}
#[inline(always)]
pub fn chunks_remaining(&self) -> usize {
if self.simd_end > self.pos {
(self.simd_end - self.pos) / Arch::LANE_COUNT
} else {
0
}
}
}
impl<'a, 'b, T: Scalar, Arch: SimdArch + SimdKernel<T>, Align: Alignment, Mode: ExecutionMode>
Iterator for ZipChunks<'a, 'b, T, Arch, Align, Mode>
{
type Item = (
SimdView<'a, T, Arch, Align, Mode, &'a [T]>,
SimdView<'b, T, Arch, Align, Mode, &'b [T]>,
);
#[inline(always)]
fn next(&mut self) -> Option<Self::Item> {
if self.pos >= self.simd_end {
return None;
}
let (chunk_a, chunk_b) = unsafe {
(
core::slice::from_raw_parts(self.base_a.add(self.pos), Arch::LANE_COUNT),
core::slice::from_raw_parts(self.base_b.add(self.pos), Arch::LANE_COUNT),
)
};
self.pos += Arch::LANE_COUNT;
Some((
SimdView::new(chunk_a).expect("zip chunk_a alignment invariant violated"),
SimdView::new(chunk_b).expect("zip chunk_b alignment invariant violated"),
))
}
#[inline(always)]
fn size_hint(&self) -> (usize, Option<usize>) {
let r = self.chunks_remaining();
(r, Some(r))
}
}
impl<'a, 'b, T: Scalar, Arch: SimdArch + SimdKernel<T>, Align: Alignment, Mode: ExecutionMode>
ExactSizeIterator for ZipChunks<'a, 'b, T, Arch, Align, Mode>
{
}
impl<'a, 'b, T: Scalar, Arch: SimdArch + SimdKernel<T>, Align: Alignment, Mode: ExecutionMode>
core::iter::FusedIterator for ZipChunks<'a, 'b, T, Arch, Align, Mode>
{
}
pub struct ZipChunksMut<'a, 'b, T: 'a + 'b, Arch: SimdArch, Align: Alignment, Mode: ExecutionMode> {
ptr_a: *mut T,
ptr_b: *const T,
pos: usize,
total: usize,
simd_end: usize,
_marker: core::marker::PhantomData<(&'a mut T, &'b T, Arch, Align, Mode)>,
}
unsafe impl<'a, 'b, T, Arch, Align, Mode> Send for ZipChunksMut<'a, 'b, T, Arch, Align, Mode>
where
T: Scalar + Send + Sync,
Arch: SimdArch + SimdKernel<T>,
Align: Alignment,
Mode: ExecutionMode,
{
}
unsafe impl<'a, 'b, T, Arch, Align, Mode> Sync for ZipChunksMut<'a, 'b, T, Arch, Align, Mode>
where
T: Scalar + Send + Sync,
Arch: SimdArch + SimdKernel<T>,
Align: Alignment,
Mode: ExecutionMode,
{
}
impl<'a, 'b, T: 'a + 'b, Arch: SimdArch + SimdKernel<T>, Align: Alignment, Mode: ExecutionMode>
ZipChunksMut<'a, 'b, T, Arch, Align, Mode>
where
T: Scalar,
{
#[inline]
pub(crate) unsafe fn from_raw_parts(
ptr_a: *mut T,
total_a: usize,
ptr_b: *const T,
total_b: usize,
) -> Self {
let lane_count = Arch::LANE_COUNT;
let total = total_a.min(total_b);
let simd_end = (total / lane_count) * lane_count;
Self {
ptr_a,
ptr_b,
pos: 0,
total,
simd_end,
_marker: core::marker::PhantomData,
}
}
#[inline(always)]
pub fn into_remainder(self) -> (&'a mut [T], &'b [T]) {
let len = self.total - self.simd_end;
unsafe {
(
core::slice::from_raw_parts_mut(self.ptr_a.add(self.simd_end), len),
core::slice::from_raw_parts(self.ptr_b.add(self.simd_end), len),
)
}
}
#[inline(always)]
pub fn chunks_remaining(&self) -> usize {
if self.simd_end > self.pos {
(self.simd_end - self.pos) / Arch::LANE_COUNT
} else {
0
}
}
}
impl<'a, 'b, T: Scalar, Arch: SimdArch + SimdKernel<T>, Align: Alignment, Mode: ExecutionMode>
Iterator for ZipChunksMut<'a, 'b, T, Arch, Align, Mode>
{
type Item = (
SimdView<'a, T, Arch, Align, Mode, &'a mut [T]>,
SimdView<'b, T, Arch, Align, Mode, &'b [T]>,
);
#[inline(always)]
fn next(&mut self) -> Option<Self::Item> {
if self.pos >= self.simd_end {
return None;
}
let lane = Arch::LANE_COUNT;
let (chunk_a, chunk_b) = unsafe {
(
core::slice::from_raw_parts_mut(self.ptr_a.add(self.pos), lane),
core::slice::from_raw_parts(self.ptr_b.add(self.pos), lane),
)
};
self.pos += lane;
Some((
SimdView::new_mut(chunk_a).expect("ZipChunksMut chunk_a alignment violated"),
SimdView::new(chunk_b).expect("ZipChunksMut chunk_b alignment violated"),
))
}
#[inline(always)]
fn size_hint(&self) -> (usize, Option<usize>) {
let r = self.chunks_remaining();
(r, Some(r))
}
}
impl<'a, 'b, T: Scalar, Arch: SimdArch + SimdKernel<T>, Align: Alignment, Mode: ExecutionMode>
ExactSizeIterator for ZipChunksMut<'a, 'b, T, Arch, Align, Mode>
{
}
impl<'a, 'b, T: Scalar, Arch: SimdArch + SimdKernel<T>, Align: Alignment, Mode: ExecutionMode>
core::iter::FusedIterator for ZipChunksMut<'a, 'b, T, Arch, Align, Mode>
{
}