use crate::num::Bytes;
#[derive(Debug, Clone, Copy)]
pub(crate) struct Checked<P>(P);
impl<P> Checked<P>
where
P: Prefetch,
{
pub(crate) fn new(prefetcher: P, len: Bytes) -> Result<Self, InvalidPrefetch> {
prefetcher.check(len)?;
Ok(Self(prefetcher))
}
pub(crate) unsafe fn prefetch(self, ptr: *const u8, len: Bytes) {
debug_assert!(self.check(len).is_ok());
unsafe { self.0.prefetch(ptr, len) }
}
pub(crate) fn check(self, len: Bytes) -> Result<(), InvalidPrefetch> {
self.0.check(len)
}
#[cfg(test)]
fn safe_prefetch(self, x: &[u8]) -> Result<(), InvalidPrefetch> {
let bytes = Bytes::new(x.len());
self.check(bytes)?;
unsafe { self.prefetch(x.as_ptr(), bytes) };
Ok(())
}
}
pub(crate) unsafe trait Prefetch:
std::fmt::Debug + Send + Sync + 'static + Copy
{
fn check(self, len: Bytes) -> Result<(), InvalidPrefetch>;
unsafe fn prefetch(self, ptr: *const u8, len: Bytes);
}
#[derive(Debug)]
pub(crate) struct InvalidPrefetch(());
impl InvalidPrefetch {
const fn new() -> Self {
Self(())
}
}
impl std::fmt::Display for InvalidPrefetch {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "invalid prefetch")
}
}
impl std::error::Error for InvalidPrefetch {}
diskann::convert_error!(InvalidPrefetch);
#[derive(Debug, Clone, Copy)]
pub(crate) struct Loop(());
impl Loop {
pub(crate) const fn new() -> Self {
Self(())
}
}
unsafe impl Prefetch for Loop {
fn check(self, _len: Bytes) -> Result<(), InvalidPrefetch> {
Ok(())
}
#[inline(always)]
unsafe fn prefetch(self, ptr: *const u8, len: Bytes) {
unsafe { prefetch(ptr, len.value()) }
}
}
#[derive(Debug, Clone, Copy)]
pub(crate) struct Unrolled<const BYTES: usize>(());
impl<const BYTES: usize> Unrolled<BYTES> {
pub(crate) const fn new() -> Self {
Self(())
}
}
unsafe impl<const BYTES: usize> Prefetch for Unrolled<BYTES> {
fn check(self, bytes: Bytes) -> Result<(), InvalidPrefetch> {
if bytes == Bytes::new(BYTES) {
Ok(())
} else {
Err(InvalidPrefetch::new())
}
}
#[inline(always)]
unsafe fn prefetch(self, ptr: *const u8, _len: Bytes) {
debug_assert!(self.check(_len).is_ok());
unsafe { prefetch(ptr, BYTES) }
}
}
#[cfg(all(target_arch = "x86_64", target_feature = "avx2"))]
#[inline(always)]
pub(crate) unsafe fn prefetch(ptr: *const u8, len: usize) {
use std::arch::x86_64::*;
let stride = Bytes::CACHELINE.value();
let ptr = ptr.cast::<i8>();
let lines = len.div_ceil(stride);
if lines == 0 {
return;
}
unsafe { _mm_prefetch(ptr.add(stride * (lines - 1)), _MM_HINT_T0) };
for i in 0..(lines - 1) {
unsafe {
_mm_prefetch(ptr.add(stride * i), _MM_HINT_T0);
}
}
}
#[cfg(not(all(target_arch = "x86_64", target_feature = "avx2")))]
pub(crate) unsafe fn prefetch(_ptr: *const u8, _len: usize) {}
#[cfg(test)]
mod test {
use super::*;
#[test]
fn test_loop() {
let p = Loop::new();
for i in 0..=1024 {
let v = vec![0u8; i];
let checked = Checked::new(p, Bytes::new(v.len())).unwrap();
checked.safe_prefetch(&v).unwrap();
}
}
fn unrolled_prefetch<const BYTES: usize>() {
let p = Unrolled::<BYTES>::new();
{
let c = Checked::new(p, Bytes::new(BYTES)).unwrap();
let v = vec![0u8; BYTES];
c.safe_prefetch(&v).unwrap();
}
if let Some(under) = BYTES.checked_sub(1) {
assert!(Checked::new(p, Bytes::new(under)).is_err());
let v = vec![0u8; under];
let c = Checked::new(p, Bytes::new(BYTES)).unwrap();
assert!(c.safe_prefetch(&v).is_err());
}
if let Some(over) = BYTES.checked_add(1) {
assert!(Checked::new(p, Bytes::new(over)).is_err());
let v = vec![0u8; over];
let c = Checked::new(p, Bytes::new(BYTES)).unwrap();
assert!(c.safe_prefetch(&v).is_err());
}
}
#[test]
fn test_unrolled() {
unrolled_prefetch::<0>();
unrolled_prefetch::<63>();
unrolled_prefetch::<64>();
unrolled_prefetch::<65>();
unrolled_prefetch::<127>();
unrolled_prefetch::<128>();
unrolled_prefetch::<129>();
}
}