use crate::cache::{prefetch_read_data, prefetch_write_data};
use std::mem;
pub const PREFETCH_DISTANCE: usize = 4;
pub struct PrefetchIterator<I: Iterator> {
iter: I,
prefetch_distance: usize,
}
impl<I: Iterator> PrefetchIterator<I> {
pub fn new(iter: I) -> Self {
Self {
iter,
prefetch_distance: PREFETCH_DISTANCE,
}
}
pub fn with_distance(iter: I, distance: usize) -> Self {
Self {
iter,
prefetch_distance: distance,
}
}
}
impl<I> Iterator for PrefetchIterator<I>
where
I: Iterator,
I::Item: Sized,
{
type Item = I::Item;
fn next(&mut self) -> Option<Self::Item> {
let item = self.iter.next()?;
if let Some(size_hint) = self.iter.size_hint().1
&& size_hint > self.prefetch_distance
{
}
Some(item)
}
fn size_hint(&self) -> (usize, Option<usize>) {
self.iter.size_hint()
}
}
pub trait PrefetchExt: Iterator + Sized {
fn prefetch(self) -> PrefetchIterator<Self> {
PrefetchIterator::new(self)
}
fn prefetch_distance(self, distance: usize) -> PrefetchIterator<Self> {
PrefetchIterator::with_distance(self, distance)
}
}
impl<I: Iterator> PrefetchExt for I {}
pub struct PrefetchSliceIter<'a, T> {
slice: &'a [T],
prefetch_distance: usize,
}
impl<'a, T> PrefetchSliceIter<'a, T> {
pub fn new(slice: &'a [T]) -> Self {
let prefetch_distance = PREFETCH_DISTANCE;
if !slice.is_empty() && prefetch_distance < slice.len() {
unsafe {
let future_ptr = slice.as_ptr().add(prefetch_distance);
prefetch_read_data(future_ptr as *const u8, 0);
}
}
Self {
slice,
prefetch_distance,
}
}
}
impl<'a, T> Iterator for PrefetchSliceIter<'a, T> {
type Item = &'a T;
#[inline]
fn next(&mut self) -> Option<Self::Item> {
if self.slice.is_empty() {
return None;
}
if self.prefetch_distance < self.slice.len() {
unsafe {
let future_ptr = self.slice.as_ptr().add(self.prefetch_distance);
prefetch_read_data(future_ptr as *const u8, 0);
}
}
let (first, rest) = self.slice.split_at(1);
self.slice = rest;
Some(&first[0])
}
fn size_hint(&self) -> (usize, Option<usize>) {
(self.slice.len(), Some(self.slice.len()))
}
}
pub struct PrefetchSliceIterMut<'a, T> {
slice: &'a mut [T],
prefetch_distance: usize,
}
impl<'a, T> PrefetchSliceIterMut<'a, T> {
pub fn new(slice: &'a mut [T]) -> Self {
let prefetch_distance = PREFETCH_DISTANCE;
if !slice.is_empty() && prefetch_distance < slice.len() {
unsafe {
let future_ptr = slice.as_mut_ptr().add(prefetch_distance);
prefetch_write_data(future_ptr as *mut u8, 0);
}
}
Self {
slice,
prefetch_distance,
}
}
}
impl<'a, T> Iterator for PrefetchSliceIterMut<'a, T> {
type Item = &'a mut T;
#[inline]
fn next(&mut self) -> Option<Self::Item> {
if self.slice.is_empty() {
return None;
}
if self.prefetch_distance < self.slice.len() {
unsafe {
let future_ptr = self.slice.as_mut_ptr().add(self.prefetch_distance);
prefetch_write_data(future_ptr as *mut u8, 0);
}
}
let slice = std::mem::take(&mut self.slice);
let (first, rest) = slice.split_at_mut(1);
self.slice = rest;
Some(&mut first[0])
}
fn size_hint(&self) -> (usize, Option<usize>) {
(self.slice.len(), Some(self.slice.len()))
}
}
pub trait SlicePrefetchExt<T> {
fn prefetch_iter(&self) -> PrefetchSliceIter<'_, T>;
fn prefetch_iter_mut(&mut self) -> PrefetchSliceIterMut<'_, T>;
}
impl<T> SlicePrefetchExt<T> for [T] {
fn prefetch_iter(&self) -> PrefetchSliceIter<'_, T> {
PrefetchSliceIter::new(self)
}
fn prefetch_iter_mut(&mut self) -> PrefetchSliceIterMut<'_, T> {
PrefetchSliceIterMut::new(self)
}
}
pub struct PrefetchChunks<'a, T> {
slice: &'a [T],
chunk_size: usize,
position: usize,
}
impl<'a, T> PrefetchChunks<'a, T> {
pub fn new(slice: &'a [T], chunk_size: usize) -> Self {
assert!(chunk_size > 0, "Chunk size must be positive");
let iter = Self {
slice,
chunk_size,
position: 0,
};
if !slice.is_empty() {
unsafe {
let chunk_end = chunk_size.min(slice.len());
for i in (0..chunk_end).step_by(64 / mem::size_of::<T>().max(1)) {
prefetch_read_data(slice.as_ptr().add(i) as *const u8, 0);
}
}
}
iter
}
}
impl<'a, T> Iterator for PrefetchChunks<'a, T> {
type Item = &'a [T];
fn next(&mut self) -> Option<Self::Item> {
if self.position >= self.slice.len() {
return None;
}
let chunk_end = (self.position + self.chunk_size).min(self.slice.len());
let chunk = &self.slice[self.position..chunk_end];
let next_position = self.position + self.chunk_size;
if next_position < self.slice.len() {
unsafe {
let next_chunk_end = (next_position + self.chunk_size).min(self.slice.len());
for i in (next_position..next_chunk_end).step_by(64 / mem::size_of::<T>().max(1)) {
prefetch_read_data(self.slice.as_ptr().add(i) as *const u8, 1);
}
}
}
self.position = chunk_end;
Some(chunk)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_prefetch_slice_iter() {
let data: Vec<i32> = (0..1000).collect();
let sum: i32 = data.prefetch_iter().sum();
assert_eq!(sum, (0..1000).sum::<i32>());
}
#[test]
fn test_prefetch_chunks() {
let data: Vec<i32> = (0..1000).collect();
let chunks = PrefetchChunks::new(&data, 100);
let count = chunks.count();
assert_eq!(count, 10);
}
#[test]
fn test_prefetch_mut_iter() {
let mut data: Vec<i32> = vec![0; 1000];
for (i, val) in data.prefetch_iter_mut().enumerate() {
*val = i as i32;
}
for (i, &val) in data.iter().enumerate() {
assert_eq!(val, i as i32);
}
}
}