mod common;
use common::u32_bytes;
use vyre_driver::{DispatchConfig, VyreBackend};
use vyre_driver_wgpu::WgpuBackend;
use vyre_foundation::ir::{BufferAccess, BufferDecl, DataType, Expr, Node, Program};
fn pairs() -> Vec<(i32, i32)> {
vec![
(-7, 3),
(7, 3),
(-8, 3),
(100, 7),
(-100, 7),
(-1, 2),
(5, -3),
(-2147483648, 3),
]
}
fn rem_program(n: u32) -> Program {
let mut body = Vec::new();
for i in 0..n {
body.push(Node::store(
"out",
Expr::u32(i),
Expr::rem(Expr::load("a", Expr::u32(i)), Expr::load("b", Expr::u32(i))),
));
}
Program::wrapped(
vec![
BufferDecl::storage("out", 0, BufferAccess::ReadWrite, DataType::I32).with_count(n),
BufferDecl::storage("a", 1, BufferAccess::ReadOnly, DataType::I32).with_count(n),
BufferDecl::storage("b", 2, BufferAccess::ReadOnly, DataType::I32).with_count(n),
],
[1, 1, 1],
body,
)
}
fn div_program(n: u32) -> Program {
let mut body = Vec::new();
for i in 0..n {
body.push(Node::store(
"out",
Expr::u32(i),
Expr::div(Expr::load("a", Expr::u32(i)), Expr::load("b", Expr::u32(i))),
));
}
Program::wrapped(
vec![
BufferDecl::storage("out", 0, BufferAccess::ReadWrite, DataType::I32).with_count(n),
BufferDecl::storage("a", 1, BufferAccess::ReadOnly, DataType::I32).with_count(n),
BufferDecl::storage("b", 2, BufferAccess::ReadOnly, DataType::I32).with_count(n),
],
[1, 1, 1],
body,
)
}
fn run(backend: &WgpuBackend, program: &Program, ps: &[(i32, i32)]) -> Vec<i32> {
let a = u32_bytes(&ps.iter().map(|&(a, _)| a as u32).collect::<Vec<_>>());
let b = u32_bytes(&ps.iter().map(|&(_, b)| b as u32).collect::<Vec<_>>());
let out_init = u32_bytes(&vec![0u32; ps.len()]);
let outputs = backend
.dispatch_borrowed(
program,
&[out_init.as_slice(), a.as_slice(), b.as_slice()],
&DispatchConfig::default(),
)
.expect("Fix: WGPU must dispatch the signed modulo contract.");
outputs[0]
.chunks_exact(4)
.map(|c| i32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect()
}
#[test]
fn signed_modulo_matches_rust_on_gpu() {
let backend =
WgpuBackend::acquire().expect("Fix: signed modulo parity requires a live GPU backend.");
let ps = pairs();
let gpu = run(&backend, &rem_program(ps.len() as u32), &ps);
let expected: Vec<i32> = ps.iter().map(|&(a, b)| a % b).collect();
assert_eq!(
expected,
vec![-1, 1, -2, 2, -2, -1, 2, -2],
"Rust signed-remainder reference drifted"
);
assert_eq!(
gpu, expected,
"GPU signed modulo diverged from Rust `%` (the naga unsigned-Modulo bug regressed).\n \
pairs: {ps:?}\n expected: {expected:?}\n gpu: {gpu:?}"
);
}
#[test]
fn signed_division_matches_rust_on_gpu() {
let backend =
WgpuBackend::acquire().expect("Fix: signed division parity requires a live GPU backend.");
let ps = pairs();
let gpu = run(&backend, &div_program(ps.len() as u32), &ps);
let expected: Vec<i32> = ps.iter().map(|&(a, b)| a / b).collect();
assert_eq!(
expected,
vec![-2, 2, -2, 14, -14, 0, -1, -715827882],
"Rust signed-division reference drifted"
);
assert_eq!(
gpu, expected,
"GPU signed division diverged from Rust `/`.\n pairs: {ps:?}\n expected: {expected:?}\n gpu: {gpu:?}"
);
}