use super::{
TestParams, scalar_znx_backend_ref, upload_mat_znx, upload_scalar_znx,
vec_znx_dft::{dft_of_uploaded_vec_znx, idft_apply_to_host},
};
use crate::layouts::SvpPPolToBackendMut;
use crate::layouts::VmpPMatToBackendMut;
use crate::{
api::{
ScratchOwnedAlloc, SvpPPolAlloc, SvpPrepare, VecZnxBigAlloc, VecZnxBigNormalize, VecZnxBigNormalizeTmpBytes,
VecZnxDftAlloc, VecZnxDftApply, VecZnxIdftApply, VmpPMatAlloc, VmpPrepare, VmpPrepareTmpBytes,
},
layouts::{
DataView, FillUniform, HostBytesBackend, MatZnx, MatZnxToBackendRef, Module, ScratchOwned, SvpPPolLayoutCompatible,
SvpPPolOwned, VecZnxDftLayoutCompatible, VecZnxDftOwned, VmpPMatLayoutCompatible, VmpPMatOwned,
},
source::Source,
};
pub fn test_word_compat_dft_bytes<BA, BB>(
params: &TestParams,
module_host: &Module<HostBytesBackend>,
module_a: &Module<BA>,
module_b: &Module<BB>,
) where
BA: crate::test_suite::TestBackend + VecZnxDftLayoutCompatible<BB>,
BB: crate::test_suite::TestBackend,
Module<BA>: VecZnxDftAlloc<BA> + VecZnxDftApply<BA>,
Module<BB>: VecZnxDftAlloc<BB> + VecZnxDftApply<BB>,
{
let base2k = params.base2k;
assert_eq!(module_a.n(), module_b.n());
let cols = 2;
let mut source = Source::new([0u8; 32]);
for size in [1, 2, 3, 4] {
let mut a = module_host.vec_znx_alloc(cols, size);
a.fill_uniform(base2k, &mut source);
let dft_a = dft_of_uploaded_vec_znx(module_a, &a, 1, 0);
let dft_b = dft_of_uploaded_vec_znx(module_b, &a, 1, 0);
assert!(
BA::to_host_bytes(&dft_a.data) == BB::to_host_bytes(&dft_b.data),
"shared DftWord but different DFT buffer bytes (size={size}): one backend violates the word contract"
);
}
}
pub fn test_word_compat_svp_prepare_bytes<BA, BB>(
params: &TestParams,
module_host: &Module<HostBytesBackend>,
module_a: &Module<BA>,
module_b: &Module<BB>,
) where
BA: crate::test_suite::TestBackend + SvpPPolLayoutCompatible<BB>,
BB: crate::test_suite::TestBackend,
Module<BA>: SvpPPolAlloc<BA> + SvpPrepare<BA>,
Module<BB>: SvpPPolAlloc<BB> + SvpPrepare<BB>,
{
let base2k = params.base2k;
assert_eq!(module_a.n(), module_b.n());
let cols = 2;
let mut source = Source::new([0u8; 32]);
let mut scalar = module_host.scalar_znx_alloc(cols);
scalar.fill_uniform(base2k, &mut source);
let scalar_a = upload_scalar_znx::<BA>(&scalar);
let scalar_b = upload_scalar_znx::<BB>(&scalar);
let mut svp_a: SvpPPolOwned<BA> = module_a.svp_ppol_alloc(cols);
let mut svp_b: SvpPPolOwned<BB> = module_b.svp_ppol_alloc(cols);
for j in 0..cols {
module_a.svp_prepare(&mut svp_a.to_backend_mut(), j, &scalar_znx_backend_ref::<BA>(&scalar_a), j);
module_b.svp_prepare(&mut svp_b.to_backend_mut(), j, &scalar_znx_backend_ref::<BB>(&scalar_b), j);
}
assert!(
BA::to_host_bytes(&svp_a.data) == BB::to_host_bytes(&svp_b.data),
"shared DftWord but different SvpPPol buffer bytes: one backend violates the word contract"
);
}
pub fn test_word_compat_vmp_prepare_bytes<BA, BB>(
params: &TestParams,
module_host: &Module<HostBytesBackend>,
module_a: &Module<BA>,
module_b: &Module<BB>,
) where
BA: crate::test_suite::TestBackend + VmpPMatLayoutCompatible<BB>,
BB: crate::test_suite::TestBackend,
Module<BA>: VmpPMatAlloc<BA> + VmpPrepare<BA> + VmpPrepareTmpBytes,
Module<BB>: VmpPMatAlloc<BB> + VmpPrepare<BB> + VmpPrepareTmpBytes,
ScratchOwned<BA>: ScratchOwnedAlloc<BA>,
ScratchOwned<BB>: ScratchOwnedAlloc<BB>,
{
let base2k = params.base2k;
assert_eq!(module_a.n(), module_b.n());
let (rows, cols_in, cols_out, size) = (2, 2, 2, 3);
let mut source = Source::new([0u8; 32]);
let mut scratch_a: ScratchOwned<BA> = ScratchOwned::alloc(module_a.vmp_prepare_tmp_bytes(rows, cols_in, cols_out, size));
let mut scratch_b: ScratchOwned<BB> = ScratchOwned::alloc(module_b.vmp_prepare_tmp_bytes(rows, cols_in, cols_out, size));
let mut mat = module_host.mat_znx_alloc(rows, cols_in, cols_out, size);
mat.fill_uniform(base2k, &mut source);
let mat_a = upload_mat_znx::<BA>(&mat);
let mat_b = upload_mat_znx::<BB>(&mat);
let mut pmat_a: VmpPMatOwned<BA> = module_a.vmp_pmat_alloc(rows, cols_in, cols_out, size);
let mut pmat_b: VmpPMatOwned<BB> = module_b.vmp_pmat_alloc(rows, cols_in, cols_out, size);
module_a.vmp_prepare(
&mut pmat_a.to_backend_mut(),
&<MatZnx<BA::OwnedBuf, BA::ZnxWord> as MatZnxToBackendRef<BA>>::to_backend_ref(&mat_a),
&mut scratch_a.arena(),
);
module_b.vmp_prepare(
&mut pmat_b.to_backend_mut(),
&<MatZnx<BB::OwnedBuf, BB::ZnxWord> as MatZnxToBackendRef<BB>>::to_backend_ref(&mat_b),
&mut scratch_b.arena(),
);
assert!(
BA::to_host_bytes(pmat_a.data()) == BB::to_host_bytes(pmat_b.data()),
"shared DftWord but different VmpPMat buffer bytes: one backend violates the word contract"
);
}
pub fn test_word_compat_dft_cross_idft<BA, BB>(
params: &TestParams,
module_host: &Module<HostBytesBackend>,
module_a: &Module<BA>,
module_b: &Module<BB>,
) where
BA: crate::test_suite::TestBackend + VecZnxDftLayoutCompatible<BB>,
BB: crate::test_suite::TestBackend<OwnedBuf = BA::OwnedBuf, DftWord = BA::DftWord, ZnxWord = BA::ZnxWord>
+ VecZnxDftLayoutCompatible<BA>,
Module<BA>: VecZnxDftAlloc<BA>
+ VecZnxDftApply<BA>
+ VecZnxBigAlloc<BA>
+ VecZnxIdftApply<BA>
+ VecZnxBigNormalize<BA>
+ VecZnxBigNormalizeTmpBytes,
Module<BB>: VecZnxDftAlloc<BB>
+ VecZnxDftApply<BB>
+ VecZnxBigAlloc<BB>
+ VecZnxIdftApply<BB>
+ VecZnxBigNormalize<BB>
+ VecZnxBigNormalizeTmpBytes,
ScratchOwned<BA>: ScratchOwnedAlloc<BA>,
ScratchOwned<BB>: ScratchOwnedAlloc<BB>,
{
let base2k = params.base2k;
assert_eq!(module_a.n(), module_b.n());
let cols = 2;
let mut source = Source::new([0u8; 32]);
let mut scratch_a: ScratchOwned<BA> = ScratchOwned::alloc(module_a.vec_znx_big_normalize_tmp_bytes());
let mut scratch_b: ScratchOwned<BB> = ScratchOwned::alloc(module_b.vec_znx_big_normalize_tmp_bytes());
for size in [1, 2, 3, 4] {
let mut a = module_host.vec_znx_alloc(cols, size);
a.fill_uniform(base2k, &mut source);
let dft_a = dft_of_uploaded_vec_znx(module_a, &a, 1, 0);
let dft_b = dft_of_uploaded_vec_znx(module_b, &a, 1, 0);
let res_aa = idft_apply_to_host(module_a, base2k, &dft_a, size, &mut scratch_a);
let res_bb = idft_apply_to_host(module_b, base2k, &dft_b, size, &mut scratch_b);
let dft_ab: VecZnxDftOwned<BB> = dft_a.into_backend::<BB>();
let dft_ba: VecZnxDftOwned<BA> = dft_b.into_backend::<BA>();
let res_ab = idft_apply_to_host(module_b, base2k, &dft_ab, size, &mut scratch_b);
let res_ba = idft_apply_to_host(module_a, base2k, &dft_ba, size, &mut scratch_a);
assert_eq!(res_aa, res_ab, "consuming A's DFT buffer on B diverges (size={size})");
assert_eq!(res_bb, res_ba, "consuming B's DFT buffer on A diverges (size={size})");
}
}