#![allow(dead_code)]
use gemmkit::Parallelism;
use proptest::prelude::*;
#[path = "../oracle_common/mod.rs"]
mod oracle_common;
pub use oracle_common::*;
pub fn cases(default: u32) -> u32 {
if let Some(n) = std::env::var("PROPTEST_CASES")
.ok()
.and_then(|s| s.trim().parse().ok())
{
return n;
}
if fast_test() {
return (default / 8).max(8);
}
default
}
#[derive(Copy, Clone, Debug)]
pub enum PLayout {
Row { pad: usize },
Col { pad: usize },
General { cs: usize, pad: usize },
}
pub fn layout() -> impl Strategy<Value = PLayout> {
prop_oneof![
3 => (0usize..=7).prop_map(|pad| PLayout::Row { pad }),
3 => (0usize..=7).prop_map(|pad| PLayout::Col { pad }),
2 => (2usize..=4, 0usize..=5).prop_map(|(cs, pad)| PLayout::General { cs, pad }),
]
}
fn strides_for(rows: usize, cols: usize, l: PLayout) -> (usize, usize) {
match l {
PLayout::Row { pad } => (cols + pad, 1),
PLayout::Col { pad } => (1, rows + pad),
PLayout::General { cs, pad } => (cols * cs + pad, cs),
}
}
pub fn build_view_rowmajor<T: Copy>(
vals: &[T],
rows: usize,
cols: usize,
zero: T,
l: PLayout,
) -> (Vec<T>, isize, isize) {
let (rs, cs) = strides_for(rows, cols, l);
let need = if rows == 0 || cols == 0 {
0
} else {
(rows - 1) * rs + (cols - 1) * cs + 1
};
let mut buf = vec![zero; need];
for i in 0..rows {
for j in 0..cols {
buf[i * rs + j * cs] = vals[i * cols + j];
}
}
(buf, rs as isize, cs as isize)
}
pub fn build_view<T: Elem>(m: &Mat<T>, l: PLayout) -> (Vec<T>, isize, isize) {
build_view_rowmajor(&m.v, m.rows, m.cols, T::from_f64(0.0), l)
}
pub fn dim() -> impl Strategy<Value = usize> {
prop_oneof![
2 => Just(0usize),
3 => Just(1usize),
5 => proptest::sample::select(
&[2usize, 4, 5, 6, 11, 12, 13, 15, 16, 17, 24, 31, 32, 33, 47, 48, 49, 63, 64, 65][..]),
8 => 2usize..=96,
]
}
pub fn pos_dim() -> impl Strategy<Value = usize> {
prop_oneof![
3 => Just(1usize),
5 => proptest::sample::select(
&[2usize, 4, 5, 6, 11, 12, 13, 15, 16, 17, 24, 31, 32, 33, 47, 48, 49, 63, 64, 65][..]),
8 => 2usize..=96,
]
}
pub fn kdim() -> impl Strategy<Value = usize> {
prop_oneof![
8 => dim(),
1 => proptest::sample::select(&[200usize, 511, 512, 513][..]),
]
}
pub fn kdim_pos() -> impl Strategy<Value = usize> {
prop_oneof![
3 => Just(1usize),
5 => proptest::sample::select(
&[2usize, 4, 5, 6, 11, 12, 13, 15, 16, 17, 24, 31, 32, 33, 47, 48, 49, 63, 64, 65][..]),
8 => 2usize..=96,
1 => proptest::sample::select(&[200usize, 511, 512, 513][..]),
]
}
pub fn coeff() -> impl Strategy<Value = f64> {
proptest::sample::select(&[0.0f64, 1.0, -1.0, 0.5, -1.5, 2.0, 2.5, 1e-3][..])
}
pub fn par() -> impl Strategy<Value = Parallelism> {
proptest::sample::select(
&[
Parallelism::Serial,
Parallelism::Rayon(0),
Parallelism::Rayon(3),
][..],
)
}
pub fn frob_norm<T: Elem>(m: &Mat<T>) -> f64 {
m.v.iter().map(|x| x.to_f64().powi(2)).sum::<f64>().sqrt()
}
pub fn bits_identical<T: Elem>(x: &[T], y: &[T]) -> bool {
x.len() == y.len()
&& x.iter()
.zip(y)
.all(|(a, b)| a.to_bits_u64() == b.to_bits_u64())
}
#[cfg(feature = "int8")]
pub fn fill_i8(n: usize, seed: u64) -> Vec<i8> {
let mut s = seed.wrapping_add(0x9E3779B97F4A7C15);
(0..n)
.map(|_| {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
(s >> 24) as i8
})
.collect()
}
#[cfg(feature = "int8")]
#[allow(clippy::too_many_arguments)]
pub fn ref_i8_wrapping(
a: &[i8],
b: &[i8],
c0: &[i32],
m: usize,
k: usize,
n: usize,
alpha: i32,
beta: i32,
) -> Vec<i32> {
let mut out = vec![0i32; m * n];
for i in 0..m {
for j in 0..n {
let mut acc: i32 = 0;
for p in 0..k {
let prod = a[i * k + p] as i32 * b[p * n + j] as i32;
acc = acc.wrapping_add(prod);
}
out[i * n + j] = beta
.wrapping_mul(c0[i * n + j])
.wrapping_add(alpha.wrapping_mul(acc));
}
}
out
}
#[cfg(feature = "complex")]
pub fn cplx_bits_identical<T: CElem>(x: &[T], y: &[T]) -> bool {
x.len() == y.len()
&& x.iter().zip(y).all(|(a, b)| {
let (ar, ai) = a.parts();
let (br, bi) = b.parts();
ar.to_bits() == br.to_bits() && ai.to_bits() == bi.to_bits()
})
}