pub use oxiblas_core::memory::{CACHE_LINE_SIZE, PrefetchLocality, prefetch_read, prefetch_write};
pub const PREFETCH_DISTANCE_LINES: usize = 8;
pub const PREFETCH_DISTANCE_BYTES: usize = PREFETCH_DISTANCE_LINES * CACHE_LINE_SIZE;
#[inline]
pub fn prefetch_range_read<T>(ptr: *const T, len: usize, locality: PrefetchLocality) {
oxiblas_core::memory::prefetch_read_range(ptr, len, locality);
}
#[inline]
pub fn prefetch_range_write<T>(ptr: *mut T, len: usize, locality: PrefetchLocality) {
oxiblas_core::memory::prefetch_write_range(ptr, len, locality);
}
#[inline]
fn strided_prefetch_row_indices(nrows: usize) -> impl Iterator<Item = usize> {
0..nrows
}
#[allow(clippy::not_unsafe_ptr_arg_deref)]
#[inline]
pub fn prefetch_column<T>(
ptr: *const T,
nrows: usize,
row_stride: usize,
locality: PrefetchLocality,
) {
let elem_size = core::mem::size_of::<T>();
if row_stride == 1 || (row_stride * elem_size) <= CACHE_LINE_SIZE {
prefetch_range_read(ptr, nrows, locality);
} else {
for row in strided_prefetch_row_indices(nrows) {
let addr = ptr.wrapping_add(row.wrapping_mul(row_stride));
prefetch_read(addr, locality);
}
}
}
#[allow(clippy::not_unsafe_ptr_arg_deref)]
#[inline]
pub fn prefetch_block<T>(
ptr: *const T,
block_rows: usize,
block_cols: usize,
row_stride: usize,
locality: PrefetchLocality,
) {
for j in 0..block_cols {
let col_ptr = ptr.wrapping_add(j.wrapping_mul(row_stride));
prefetch_column(col_ptr, block_rows, 1, locality);
}
}
pub struct MatrixPrefetcher<T> {
ptr: *const T,
nrows: usize,
ncols: usize,
row_stride: usize,
current_col: usize,
distance: usize,
locality: PrefetchLocality,
}
impl<T> MatrixPrefetcher<T> {
#[allow(clippy::not_unsafe_ptr_arg_deref)]
#[inline]
pub fn new(
ptr: *const T,
nrows: usize,
ncols: usize,
row_stride: usize,
distance: usize,
locality: PrefetchLocality,
) -> Self {
let prefetcher = MatrixPrefetcher {
ptr,
nrows,
ncols,
row_stride,
current_col: 0,
distance,
locality,
};
for j in 0..distance.min(ncols) {
let col_ptr = ptr.wrapping_add(j.wrapping_mul(row_stride));
prefetch_column(col_ptr, nrows, 1, locality);
}
prefetcher
}
#[inline]
pub fn advance(&mut self) {
self.current_col += 1;
let prefetch_col = self.current_col.saturating_add(self.distance);
if prefetch_col < self.ncols {
let col_ptr = self
.ptr
.wrapping_add(prefetch_col.wrapping_mul(self.row_stride));
prefetch_column(col_ptr, self.nrows, 1, self.locality);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_prefetch_locality() {
assert_ne!(PrefetchLocality::High, PrefetchLocality::Low);
assert_eq!(PrefetchLocality::Medium, PrefetchLocality::Medium);
}
#[test]
#[cfg_attr(miri, ignore)]
fn test_prefetch_read_safety() {
let data = [1.0f64; 1024];
prefetch_read(data.as_ptr(), PrefetchLocality::High);
prefetch_read(data.as_ptr().wrapping_add(100), PrefetchLocality::Medium);
prefetch_read(data.as_ptr().wrapping_add(500), PrefetchLocality::Low);
prefetch_read(
data.as_ptr().wrapping_add(900),
PrefetchLocality::NonTemporal,
);
}
#[test]
#[cfg_attr(miri, ignore)]
fn test_prefetch_write_safety() {
let mut data = [1.0f64; 1024];
prefetch_write(data.as_mut_ptr(), PrefetchLocality::High);
prefetch_write(
data.as_mut_ptr().wrapping_add(100),
PrefetchLocality::Medium,
);
}
#[test]
#[cfg_attr(miri, ignore)]
fn test_prefetch_range() {
let data = vec![1.0f64; 4096];
prefetch_range_read(data.as_ptr(), data.len(), PrefetchLocality::Medium);
prefetch_range_read(data.as_ptr(), 0, PrefetchLocality::High); prefetch_range_read(data.as_ptr(), 1, PrefetchLocality::Low); }
#[test]
#[cfg_attr(miri, ignore)]
fn test_prefetch_column() {
let data = vec![1.0f64; 1000];
prefetch_column(data.as_ptr(), 100, 1, PrefetchLocality::High);
prefetch_column(data.as_ptr(), 10, 100, PrefetchLocality::Medium);
}
#[test]
#[cfg_attr(miri, ignore)]
fn test_prefetch_block() {
let data = vec![1.0f64; 10000];
prefetch_block(data.as_ptr(), 64, 64, 100, PrefetchLocality::High);
}
#[test]
#[cfg_attr(miri, ignore)]
fn test_matrix_prefetcher() {
let data = vec![1.0f64; 10000];
let mut prefetcher = MatrixPrefetcher::new(
data.as_ptr(),
100, 100, 100, 8, PrefetchLocality::Medium,
);
for _ in 0..100 {
prefetcher.advance();
}
}
#[test]
fn test_cache_constants() {
assert!(CACHE_LINE_SIZE.is_power_of_two());
assert_eq!(CACHE_LINE_SIZE, oxiblas_core::memory::CACHE_LINE_SIZE);
const { assert!(PREFETCH_DISTANCE_LINES > 0) };
assert_eq!(
PREFETCH_DISTANCE_BYTES,
PREFETCH_DISTANCE_LINES * CACHE_LINE_SIZE
);
}
#[test]
fn test_strided_column_prefetch_covers_all_rows() {
for &nrows in &[0usize, 1, 7, 8, 9, 64, 100, 137, 1000] {
let rows: Vec<usize> = strided_prefetch_row_indices(nrows).collect();
assert_eq!(
rows.len(),
nrows,
"strided prefetch must visit every row, not a fraction of them (nrows = {nrows})"
);
assert_eq!(
rows,
(0..nrows).collect::<Vec<_>>(),
"strided prefetch must visit rows 0..nrows in order (nrows = {nrows})"
);
}
}
#[test]
#[cfg_attr(miri, ignore)]
fn test_prefetch_column_strided_large() {
let row_stride = CACHE_LINE_SIZE + 1;
let nrows = 200;
let data = vec![1.0f64; nrows * row_stride];
prefetch_column(data.as_ptr(), nrows, row_stride, PrefetchLocality::Medium);
}
}