use crate::threading::run_serial_and_parallel;
use crate::*;
fn strides(dims: &[usize]) -> Vec<isize> {
let mut stride = 1isize;
dims.iter()
.map(|&dim| {
let current = stride;
stride *= dim as isize;
current
})
.collect()
}
fn data(len: usize) -> Vec<f64> {
(0..len)
.map(|i| ((i * 7919) % 1009) as f64 / 97.0 - 5.0)
.collect()
}
const CASES: [([usize; 2], usize); 3] = [([4096, 16], 0), ([130, 512], 1), ([600, 70], 1)];
#[test]
fn scan_parallel_matches_serial() {
for (dims, axis) in CASES {
let st = strides(&dims);
let src = data(dims.iter().product());
let plan = ErasedScanPlan::compile(
KernelDType::F64,
ScanOp::Sum,
&dims,
&st,
&st,
axis,
ScanOptions::new().reverse(true),
)
.unwrap();
let (serial, parallel) = run_serial_and_parallel(|| {
let mut out = vec![0.0f64; src.len()];
let src = ErasedRawStridedRef::from_slice(&src, &dims, &st, 0).unwrap();
let mut dest = ErasedRawStridedMut::from_slice_mut(&mut out, &dims, &st, 0).unwrap();
plan.execute(&ExecContext::ambient(), &mut dest, &src)
.unwrap();
out
});
assert_eq!(serial, parallel);
}
}
#[test]
fn argreduce_parallel_matches_serial() {
for (dims, axis) in CASES {
let st = strides(&dims);
let src = data(dims.iter().product());
let out_dims = [dims[1 - axis]];
let plan = ErasedArgReducePlan::compile(
KernelDType::F64,
KernelDType::I64,
ArgReduceOp::MaxAbs,
&dims,
&st,
&out_dims,
&[1],
axis,
)
.unwrap();
let (serial, parallel) = run_serial_and_parallel(|| {
let mut out = vec![-1i64; out_dims[0]];
let src = ErasedRawStridedRef::from_slice(&src, &dims, &st, 0).unwrap();
let mut dest =
ErasedRawStridedMut::from_slice_mut(&mut out, &out_dims, &[1], 0).unwrap();
plan.execute(&ExecContext::ambient(), &mut dest, &src)
.unwrap();
out
});
assert_eq!(serial, parallel);
assert!(serial.iter().all(|&i| i >= 0));
}
}
#[test]
fn norm_parallel_matches_serial() {
for (dims, axis) in CASES {
let st = strides(&dims);
let src = data(dims.iter().product());
let weight = data(dims[axis]);
let spec = NormSpec::layer_norm(1e-5).with_weight(1).with_bias(1);
let plan = ErasedNormPlan::compile(KernelDType::F64, spec, &dims, &st, &st, axis).unwrap();
let (serial, parallel) = run_serial_and_parallel(|| {
let mut out = vec![0.0f64; src.len()];
let src = ErasedRawStridedRef::from_slice(&src, &dims, &st, 0).unwrap();
let w = ErasedRawStridedRef::from_slice(&weight, &dims[axis..=axis], &[1], 0).unwrap();
let mut dest = ErasedRawStridedMut::from_slice_mut(&mut out, &dims, &st, 0).unwrap();
plan.execute(&ExecContext::ambient(), &mut dest, &src, Some(&w), Some(&w))
.unwrap();
out
});
assert_eq!(
serial.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
parallel.iter().map(|v| v.to_bits()).collect::<Vec<_>>()
);
}
}