use super::{
limbs::{WideLimbs, mul_4_by_4},
montgomery::{MontgomeryLimbs, montgomery_reduce_9},
};
use ff::PrimeField;
use num_traits::Zero;
use std::ops::AddAssign;
pub trait DelayedReduction<Value>: Sized {
type Accumulator: Copy + Clone + Default + AddAssign + Send + Sync + Zero;
fn unreduced_multiply_accumulate(acc: &mut Self::Accumulator, field: &Self, value: &Value);
fn reduce(acc: &Self::Accumulator) -> Self;
}
impl<F: MontgomeryLimbs + PrimeField + Copy> DelayedReduction<F> for F {
type Accumulator = WideLimbs<9>;
#[inline(always)]
fn unreduced_multiply_accumulate(acc: &mut Self::Accumulator, field_a: &Self, field_b: &F) {
let product = mul_4_by_4(field_a.to_limbs(), field_b.to_limbs());
let mut carry = 0u128;
for (acc_limb, &prod_limb) in acc.0.iter_mut().take(8).zip(product.iter()) {
let sum = (*acc_limb as u128) + (prod_limb as u128) + carry;
*acc_limb = sum as u64;
carry = sum >> 64;
}
let old_limb8 = acc.0[8];
acc.0[8] = acc.0[8].wrapping_add(carry as u64);
debug_assert!(
acc.0[8] >= old_limb8,
"DelayedReduction accumulator overflow: limb 8 wrapped from {} to {} (carry={}). \
Too many products accumulated without reduction.",
old_limb8,
acc.0[8],
carry
);
}
#[inline(always)]
fn reduce(acc: &Self::Accumulator) -> Self {
F::from_limbs(montgomery_reduce_9::<F>(&acc.0))
}
}
#[cfg(test)]
pub(crate) fn test_delayed_reduction_sum_impl<F: MontgomeryLimbs + PrimeField + Copy>() {
use rand::{SeedableRng, rngs::StdRng};
let mut rng = StdRng::seed_from_u64(54321);
let n = 1000;
let a_vec: Vec<F> = (0..n).map(|_| F::random(&mut rng)).collect();
let b_vec: Vec<F> = (0..n).map(|_| F::random(&mut rng)).collect();
let expected: F = a_vec.iter().zip(b_vec.iter()).map(|(a, b)| *a * *b).sum();
let mut acc = WideLimbs::<9>::default();
for (a, b) in a_vec.iter().zip(b_vec.iter()) {
<F as DelayedReduction<F>>::unreduced_multiply_accumulate(&mut acc, a, b);
}
let result = <F as DelayedReduction<F>>::reduce(&acc);
assert_eq!(
result, expected,
"Delayed reduction sum failed: accumulated result != direct sum"
);
}
#[cfg(test)]
#[macro_export]
macro_rules! test_delayed_reduction {
($mod_name:ident, $field:ty) => {
mod $mod_name {
#[test]
fn delayed_reduction_sum() {
$crate::big_num::delayed_reduction::test_delayed_reduction_sum_impl::<$field>();
}
}
};
}