#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub enum SimdWidth {
Scalar, Sse, Avx2, Neon, Sve, }
impl SimdWidth {
pub fn from_name(name: &str) -> Self {
match name {
"SSE" | "NEON" => SimdWidth::Sse,
"AVX2" | "SVE" => SimdWidth::Avx2,
_ => SimdWidth::Scalar,
}
}
pub fn lane_count_f64(self) -> usize {
match self {
SimdWidth::Scalar => 1,
SimdWidth::Sse | SimdWidth::Neon => 2,
SimdWidth::Avx2 | SimdWidth::Sve => 4,
}
}
pub fn lane_count_f32(self) -> usize {
match self {
SimdWidth::Scalar => 1,
SimdWidth::Sse | SimdWidth::Neon => 4,
SimdWidth::Avx2 | SimdWidth::Sve => 8,
}
}
}
pub fn padded_len(n: usize, alignment: usize) -> usize {
if alignment == 0 { return n; }
let rem = n % alignment;
if rem == 0 { n } else { n + alignment - rem }
}
pub struct PaddedVec<T: Clone + Default> {
data: Vec<T>,
#[allow(dead_code)]
alignment: usize,
original_len: usize,
}
impl<T: Clone + Default> PaddedVec<T> {
pub fn new(src: &[T], alignment: usize) -> Self {
let original_len = src.len();
let total = padded_len(original_len, alignment);
let mut data = src.to_vec();
data.resize(total, T::default());
Self { data, alignment, original_len }
}
pub fn as_slice(&self) -> &[T] {
&self.data
}
pub fn as_mut_slice(&mut self) -> &mut [T] {
&mut self.data
}
pub fn original_len(&self) -> usize {
self.original_len
}
pub fn into_result(mut self) -> Vec<T> {
self.data.truncate(self.original_len);
self.data
}
}
pub fn pad_pair<T: Clone + Default>(a: &[T], b: &[T], alignment: usize) -> (PaddedVec<T>, PaddedVec<T>) {
let len = a.len().min(b.len());
(PaddedVec::new(&a[..len], alignment), PaddedVec::new(&b[..len], alignment))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::shape::{SizeAwareDispatch, ThresholdTuner};
#[test]
fn test_padded_len() {
assert_eq!(padded_len(0, 4), 0);
assert_eq!(padded_len(3, 4), 4);
assert_eq!(padded_len(4, 4), 4);
assert_eq!(padded_len(7, 8), 8);
}
#[test]
fn test_padded_vec() {
let v = PaddedVec::new(&[1, 2, 3], 4);
assert_eq!(v.as_slice(), &[1, 2, 3, 0]);
assert_eq!(v.original_len(), 3);
assert_eq!(v.into_result(), vec![1, 2, 3]);
}
#[test]
fn test_pad_pair() {
let (a, b) = pad_pair(&[1.0, 2.0, 3.0], &[4.0, 5.0, 6.0, 7.0], 4);
assert_eq!(a.as_slice(), &[1.0, 2.0, 3.0, 0.0]);
assert_eq!(b.as_slice(), &[4.0, 5.0, 6.0, 0.0]);
}
#[test]
fn test_size_aware_dispatch() {
let mut d = SizeAwareDispatch::new("test", 0);
d.add_rule(1000, 1);
d.add_rule(100000, 2);
assert_eq!(d.select(10), 0);
assert_eq!(d.select(5000), 1);
assert_eq!(d.select(500000), 2);
}
#[test]
fn test_threshold_tuner() {
use std::time::Duration;
let mut t = ThresholdTuner::new(1024);
assert_eq!(t.threshold(), 1024);
for _ in 0..10 {
t.observe_cpu(Duration::from_micros(200));
t.observe_gpu(Duration::from_micros(50));
}
t.adapt();
assert!(t.threshold() < 1024);
}
}