diff --git a/src/backend/vector/ifma/field.rs b/src/backend/vector/ifma/field.rs index 1b67ea8..d263eec 100644 --- a/src/backend/vector/ifma/field.rs +++ b/src/backend/vector/ifma/field.rs @@ -20,24 +20,28 @@ extern "C" { fn madd52hi(z: u64x4, x: u64x4, y: u64x4) -> u64x4; } -#[derive(Copy, Clone)] -pub struct FieldElement51x4([u64x4; 5]); +/// A vector of four field elements in radix 2^51, with unreduced coefficients. +#[derive(Copy, Clone, Debug)] +pub struct F51x4Unreduced(pub(crate) [u64x4; 5]); -impl FieldElement51x4 { +/// A vector of four field elements in radix 2^51, with reduced coefficients. +#[derive(Copy, Clone, Debug)] +pub struct F51x4Reduced(pub(crate) [u64x4; 5]); + +impl F51x4Unreduced { pub fn new( x0: &FieldElement51, x1: &FieldElement51, x2: &FieldElement51, x3: &FieldElement51, - ) -> FieldElement51x4 { - FieldElement51x4([ + ) -> F51x4Unreduced { + F51x4Unreduced([ u64x4::new(x0.0[0], x1.0[0], x2.0[0], x3.0[0]), u64x4::new(x0.0[1], x1.0[1], x2.0[1], x3.0[1]), u64x4::new(x0.0[2], x1.0[2], x2.0[2], x3.0[2]), u64x4::new(x0.0[3], x1.0[3], x2.0[3], x3.0[3]), u64x4::new(x0.0[4], x1.0[4], x2.0[4], x3.0[4]), ]) - .reduce() } pub fn split(&self) -> [FieldElement51; 4] { @@ -73,35 +77,44 @@ impl FieldElement51x4 { ]), ] } +} +impl From for F51x4Unreduced { #[inline] - pub fn reduce(&self) -> FieldElement51x4 { + fn from(x: F51x4Reduced) -> F51x4Unreduced { + F51x4Unreduced(x.0) + } +} + +impl From for F51x4Reduced { + #[inline] + fn from(x: F51x4Unreduced) -> F51x4Reduced { let mask = u64x4::splat((1 << 51) - 1); let r19 = u64x4::splat(19); // Compute carryouts in parallel - let c0 = self.0[0] >> 51; - let c1 = self.0[1] >> 51; - let c2 = self.0[2] >> 51; - let c3 = self.0[3] >> 51; - let c4 = self.0[4] >> 51; + let c0 = x.0[0] >> 51; + let c1 = x.0[1] >> 51; + let c2 = x.0[2] >> 51; + let c3 = x.0[3] >> 51; + let c4 = x.0[4] >> 51; unsafe { - FieldElement51x4([ - madd52lo(self.0[0] & mask, c4, r19), - (self.0[1] & mask) + c0, - (self.0[2] & mask) + c1, - (self.0[3] & mask) + c2, - (self.0[4] & mask) + c3, + F51x4Reduced([ + madd52lo(x.0[0] & mask, c4, r19), + (x.0[1] & mask) + c0, + (x.0[2] & mask) + c1, + (x.0[3] & mask) + c2, + (x.0[4] & mask) + c3, ]) } } } -impl<'a> Mul<(u32, u32, u32, u32)> for &'a FieldElement51x4 { - type Output = FieldElement51x4; +impl<'a> Mul<(u32, u32, u32, u32)> for &'a F51x4Reduced { + type Output = F51x4Unreduced; #[inline] - fn mul(self, scalars: (u32, u32, u32, u32)) -> FieldElement51x4 { + fn mul(self, scalars: (u32, u32, u32, u32)) -> F51x4Unreduced { unsafe { let x = &self.0; let y = u64x4::new( @@ -137,29 +150,23 @@ impl<'a> Mul<(u32, u32, u32, u32)> for &'a FieldElement51x4 { // Wave 2 z2hi = madd52hi(z2hi, y, x[1]); z1hi = madd52hi(z1hi, y, x[0]); + z0lo = madd52lo(z0lo, z5hi + z5hi, r19); - let z4 = z4hi + z4hi + z4lo; - let c4 = z4 >> 51; - let z5 = z5hi + z5hi + c4; - let z3 = z3hi + z3hi + z3lo; - let z2 = z2hi + z2hi + z2lo; - let z1 = z1hi + z1hi + z1lo; - - FieldElement51x4([ - madd52lo(z0lo & mask, z5, r19), - (z1 & mask) + (z0lo >> 51), - (z2 & mask) + (z1lo >> 51), - (z3 & mask) + (z2lo >> 51), - (z4 & mask) + (z3lo >> 51), + F51x4Unreduced([ + z0lo, + z1hi + z1hi + z1lo, + z2hi + z2hi + z2lo, + z3hi + z3hi + z3lo, + z4hi + z4hi + z4lo, ]) } } } -impl<'a, 'b> Mul<&'b FieldElement51x4> for &'a FieldElement51x4 { - type Output = FieldElement51x4; +impl<'a, 'b> Mul<&'b F51x4Reduced> for &'a F51x4Reduced { + type Output = F51x4Unreduced; #[inline] - fn mul(self, rhs: &'b FieldElement51x4) -> FieldElement51x4 { + fn mul(self, rhs: &'b F51x4Reduced) -> F51x4Unreduced { unsafe { // Inputs let x = &self.0; @@ -285,14 +292,13 @@ impl<'a, 'b> Mul<&'b FieldElement51x4> for &'a FieldElement51x4 { z3lo = madd52lo(z3lo, r38, z7 >> 52); z0lo = madd52lo(z0lo, r1938, z9 >> 52); - FieldElement51x4([ + F51x4Unreduced([ z0lo, z1lo + z1hi + z1hi, z2lo + z2hi + z2hi, z3lo + z3hi + z3hi, z4lo + z4hi + z4hi, ]) - .reduce() } } } @@ -317,7 +323,7 @@ mod test { // Invert a small field element to get a big one let a = FieldElement51([2438, 24, 243, 0, 0]).invert(); - let ax4 = FieldElement51x4::new(&a, &a, &a, &a); + let ax4 = F51x4Unreduced::new(&a, &a, &a, &a); let splits = ax4.split(); for i in 0..4 { @@ -338,7 +344,7 @@ mod test { a.0[4] << 4, ]); - let a16x4 = FieldElement51x4::new(&a16, &a16, &a16, &a16); + let a16x4 = F51x4Unreduced::new(&a16, &a16, &a16, &a16); let splits = a16x4.split(); for i in 0..4 { @@ -346,6 +352,29 @@ mod test { } } + #[test] + fn test_reduction() { + // Invert a small field element to get a big one + let a = FieldElement51([2438, 24, 243, 0, 0]).invert(); + // ... but now multiply it by 128 without reducing coeffs + let abig = FieldElement51([ + a.0[0] << 8, + a.0[1] << 8, + a.0[2] << 8, + a.0[3] << 8, + a.0[4] << 8, + ]); + + let abigx4: F51x4Reduced = F51x4Unreduced::new(&abig, &abig, &abig, &abig).into(); + + let splits = F51x4Unreduced::from(abigx4).split(); + let c = &a * &FieldElement51([128, 0, 0, 0, 0]); + + for i in 0..4 { + assert_eq!(c, splits[i]); + } + } + #[test] fn mul_matches_serial() { // Invert a small field element to get a big one @@ -353,8 +382,8 @@ mod test { let b = FieldElement51([98098, 87987897, 0, 1, 0]).invert(); let c = &a * &b; - let ax4 = FieldElement51x4::new(&a, &a, &a, &a); - let bx4 = FieldElement51x4::new(&b, &b, &b, &b); + let ax4: F51x4Reduced = F51x4Unreduced::new(&a, &a, &a, &a).into(); + let bx4: F51x4Reduced = F51x4Unreduced::new(&b, &b, &b, &b).into(); let cx4 = &ax4 * &bx4; let splits = cx4.split(); @@ -375,12 +404,12 @@ mod test { c = &b * &c; } - let ax4 = FieldElement51x4::new(&a, &a, &a, &a); - let bx4 = FieldElement51x4::new(&b, &b, &b, &b); + let ax4: F51x4Reduced = F51x4Unreduced::new(&a, &a, &a, &a).into(); + let bx4: F51x4Reduced = F51x4Unreduced::new(&b, &b, &b, &b).into(); let mut cx4 = &ax4 * &bx4; for i in 0..1024 { - cx4 = &ax4 * &cx4; - cx4 = &bx4 * &cx4; + cx4 = &ax4 * &F51x4Reduced::from(cx4); + cx4 = &bx4 * &F51x4Reduced::from(cx4); } let splits = cx4.split(); @@ -400,11 +429,11 @@ mod test { c = &b * &c; } - let ax4 = FieldElement51x4::new(&a, &a, &a, &a); + let ax4 = F51x4Unreduced::new(&a, &a, &a, &a); let bx4 = (121665u32, 121665u32, 121665u32, 121665u32); - let mut cx4 = &ax4 * bx4; + let mut cx4 = &F51x4Reduced::from(ax4) * bx4; for i in 0..1024 { - cx4 = &cx4 * bx4; + cx4 = &F51x4Reduced::from(cx4) * bx4; } let splits = cx4.split();