From 34ae1b15e00096c22376ee7b54ac5a647a364c23 Mon Sep 17 00:00:00 2001 From: Henry de Valence Date: Sun, 12 Nov 2017 17:30:01 -0800 Subject: [PATCH] Add a squaring implementation --- src/avx2/edwards.rs | 5 +++- src/avx2/field.rs | 72 +++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 76 insertions(+), 1 deletion(-) diff --git a/src/avx2/edwards.rs b/src/avx2/edwards.rs index 2c7d10a..ca73b87 100644 --- a/src/avx2/edwards.rs +++ b/src/avx2/edwards.rs @@ -126,7 +126,7 @@ impl ExtendedPoint { t0.0[3] = _mm256_blend_epi32(t0.0[3].into(), P.0[3].into(), 0b01011111).into(); t0.0[4] = _mm256_blend_epi32(t0.0[4].into(), P.0[4].into(), 0b01011111).into(); - t1 = &t0 * &t0; // replace with .square() + t1 = t0.square(); // Now t1 = (S1 S2 S3 S4) @@ -167,6 +167,9 @@ impl ExtendedPoint { // // Can we tighten these bounds to avoid a reduction? Alternately, can we do better than // the 64-bit reduction that reduce32() calls internally? + // + // Or, could we do the arithmetic on the intermediate [u64x4;10], then do + // the reduction we'd need to do for the squaring? // // Also, can we do better than the mess below? for i in 0..5 { diff --git a/src/avx2/field.rs b/src/avx2/field.rs index cca421d..53c6405 100644 --- a/src/avx2/field.rs +++ b/src/avx2/field.rs @@ -341,6 +341,60 @@ pub fn repack_pair(x: u32x8, y: u32x8) -> u32x8 { } } +impl FieldElement32x4 { + pub fn square(&self) -> FieldElement32x4 { + #[inline(always)] + fn m(x: u32x8, y: u32x8) -> u64x4 { + use stdsimd::vendor::_mm256_mul_epu32; + unsafe { _mm256_mul_epu32(x,y) } + } + + #[inline(always)] + fn m_lo(x: u32x8, y: u32x8) -> u32x8 { + use stdsimd::vendor::_mm256_mul_epu32; + unsafe { u32x8::from(_mm256_mul_epu32(x,y)) } + } + + let v19 = u32x8::new(19,0,19,0,19,0,19,0); + + let mut z = [u64x4::splat(0); 10]; + + let (x0, x1) = unpack_pair(self.0[0]); + let (x2, x3) = unpack_pair(self.0[1]); + let (x4, x5) = unpack_pair(self.0[2]); + let (x6, x7) = unpack_pair(self.0[3]); + let (x8, x9) = unpack_pair(self.0[4]); + + let x0_2 = x0 << 1; + let x1_2 = x1 << 1; + let x2_2 = x2 << 1; + let x3_2 = x3 << 1; + let x4_2 = x4 << 1; + let x5_2 = x5 << 1; + let x6_2 = x6 << 1; + let x7_2 = x7 << 1; + + let x5_19 = m_lo(v19, x5); + let x6_19 = m_lo(v19, x6); + let x7_19 = m_lo(v19, x7); + let x8_19 = m_lo(v19, x8); + let x9_19 = m_lo(v19, x9); + + z[0] = m(x0, x0) + m(x2_2,x8_19) + m(x4_2,x6_19) + ((m(x1_2,x9_19) + m(x3_2,x7_19) + m(x5,x5_19)) << 1); + z[1] = m(x0_2,x1) + m(x3_2,x8_19) + m(x5_2,x6_19) + ((m(x2,x9_19) + m(x4,x7_19)) << 1); + z[2] = m(x0_2,x2) + m(x1_2,x1) + m(x4_2,x8_19) + m(x6,x6_19) + ((m(x3_2,x9_19) + m(x5_2,x7_19)) << 1); + z[3] = m(x0_2,x3) + m(x1_2,x2) + m(x5_2,x8_19) + ((m(x4,x9_19) + m(x6,x7_19)) << 1); + z[4] = m(x0_2,x4) + m(x1_2,x3_2) + m(x2, x2) + m(x6_2,x8_19) + ((m(x5_2,x9_19) + m(x7,x7_19)) << 1); + z[5] = m(x0_2,x5) + m(x1_2,x4) + m(x2_2,x3) + m(x7_2,x8_19) + ((m(x6,x9_19)) << 1); + z[6] = m(x0_2,x6) + m(x1_2,x5_2) + m(x2_2,x4) + m(x3_2,x3) + m(x8,x8_19) + ((m(x7_2,x9_19)) << 1); + z[7] = m(x0_2,x7) + m(x1_2,x6) + m(x2_2,x5) + m(x3_2,x4) + ((m(x8,x9_19)) << 1); + z[8] = m(x0_2,x8) + m(x1_2,x7_2) + m(x2_2,x6) + m(x3_2,x5_2) + m(x4,x4) + ((m(x9,x9_19)) << 1); + z[9] = m(x0_2,x9) + m(x1_2,x8) + m(x2_2,x7) + m(x3_2,x6) + m(x4_2,x5); + + return FieldElement32x4::reduce64(z); + } +} + impl<'a, 'b> Mul<&'b FieldElement32x4> for &'a FieldElement32x4 { type Output = FieldElement32x4; fn mul(self, _rhs: &'b FieldElement32x4) -> FieldElement32x4 { @@ -457,6 +511,24 @@ mod test { assert_eq!(result[3], &x3 + &x2); } + #[test] + fn square_vs_serial() { + let x0 = FieldElement32([10000, 10001, 10002, 10003, 10004, 10005, 10006, 10007, 10008, 10009]); + let x1 = FieldElement32([10100, 10101, 10102, 10103, 10104, 10105, 10106, 10107, 10108, 10109]); + let x2 = FieldElement32([10200, 10201, 10202, 10203, 10204, 10205, 10206, 10207, 10208, 10209]); + let x3 = FieldElement32([10300, 10301, 10302, 10303, 10304, 10305, 10306, 10307, 10308, 10309]); + + let vec = FieldElement32x4::new(&x0, &x1, &x2, &x3); + + let result = vec.square().split(); + + assert_eq!(result[0], &x0 * &x0); + assert_eq!(result[1], &x1 * &x1); + assert_eq!(result[2], &x2 * &x2); + assert_eq!(result[3], &x3 * &x3); + } + + #[test] fn multiply_vs_serial() { let x0 = FieldElement32([10000, 10001, 10002, 10003, 10004, 10005, 10006, 10007, 10008, 10009]);