From c0f64009a790acd27e76f55bcf5ff9038e778d7b Mon Sep 17 00:00:00 2001 From: Henry de Valence Date: Mon, 18 Dec 2017 11:42:53 -0800 Subject: [PATCH] Implement readdition using a CachedPoint type --- src/backend/avx2/edwards.rs | 136 ++++++++++++++++++++++++++++++++---- src/backend/avx2/field.rs | 34 ++++++++- 2 files changed, 156 insertions(+), 14 deletions(-) diff --git a/src/backend/avx2/edwards.rs b/src/backend/avx2/edwards.rs index 2185f6b..cbd0c9a 100644 --- a/src/backend/avx2/edwards.rs +++ b/src/backend/avx2/edwards.rs @@ -29,6 +29,8 @@ use traits::Identity; use backend::avx2::field::FieldElement32x4; use backend::avx2::field::P_TIMES_2_MASKED; +use backend::avx2::field::{A_LANES, B_LANES, C_LANES, D_LANES, ALL_LANES}; + use backend::avx2; /// A point on Curve25519, represented in an AVX2-friendly format. @@ -74,13 +76,69 @@ impl Identity for ExtendedPoint { } } +/// A cached point with some precomputed variables used for readdition. +#[derive(Copy, Clone, Debug)] +pub struct CachedPoint(pub(super) FieldElement32x4); + +impl From for CachedPoint { + fn from(mut P: ExtendedPoint) -> CachedPoint { + let mut x = P.0; + + // x = (S2 S3 Z2 T2) + x.diff_sum(0b00001111); + + // x = (121666*S2 121666*S3 2*121666*Z2 2*121665*T2) + x.scale_by_curve_constants(); + + // x = (121666*S2 121666*S3 2*121666*Z2 -2*121665*T2) + x.negate(D_LANES); + + CachedPoint(x) + } +} + +impl Default for CachedPoint { + fn default() -> CachedPoint { + CachedPoint::identity() + } +} + +impl Identity for CachedPoint { + fn identity() -> CachedPoint { + CachedPoint(FieldElement32x4([ + u32x8::new(121647, 121666, 0, 0, 243332, 67108845, 0, 33554431), + u32x8::new(67108864, 0, 33554431, 0, 0, 67108863, 0, 33554431), + u32x8::new(67108863, 0, 33554431, 0, 0, 67108863, 0, 33554431), + u32x8::new(67108863, 0, 33554431, 0, 0, 67108863, 0, 33554431), + u32x8::new(67108863, 0, 33554431, 0, 0, 67108863, 0, 33554431), + ])) + } +} + +impl ConditionallyAssignable for CachedPoint { + fn conditional_assign(&mut self, other: &CachedPoint, choice: u8) { + self.0.conditional_assign(&other.0, choice); + } +} + +impl<'a> Neg for &'a CachedPoint { + type Output = CachedPoint; + + fn neg(self) -> CachedPoint { + let mut neg = *self; + neg.0.swap_AB(); + neg.0.negate_lazy(D_LANES); + neg + } +} + impl<'a> Neg for &'a ExtendedPoint { type Output = ExtendedPoint; fn neg(self) -> ExtendedPoint { let mut neg = *self; // (X Y Z T) -> (-X Y Z -T) - neg.0.mask_negate(0b10100101); + neg.0.negate(A_LANES | D_LANES); neg } } @@ -206,6 +264,49 @@ impl ExtendedPoint { } } +impl<'a, 'b> Add<&'b CachedPoint> for &'a ExtendedPoint { + type Output = ExtendedPoint; + + /// Uses a slight tweak of the parallel unified formulas of HWCD'08 + fn add(self, other: &'b CachedPoint) -> ExtendedPoint { + unsafe { + use stdsimd::vendor::_mm256_permute2x128_si256; + use stdsimd::vendor::_mm256_permutevar8x32_epi32; + use stdsimd::vendor::_mm256_blend_epi32; + use stdsimd::vendor::_mm256_shuffle_epi32; + + let mut tmp = self.0; + + // tmp = (Y1-X1 Y1+X1 Z1 T1) = (S0 S1 Z1 T1) + tmp.diff_sum(A_LANES | B_LANES); + + // tmp = (S0*S2' S1*S3' Z1*Z2' T1*T2') = (S8 S9 S10 S11) + tmp = &tmp * &other.0; + + // tmp = (S8 S9 S11 S10) + tmp.swap_CD(); + + // tmp = (S9-S8 S9+S8 S10-S11 S10+S11) = (S12 S13 S14 S15) + tmp.diff_sum(ALL_LANES); + + let c0 = u32x8::new(0,5,2,7,5,0,7,2); // (ABCD) -> (ADDA) + let c1 = u32x8::new(4,1,6,3,4,1,6,3); // (ABCD) -> (CBCB) + + // set t0 = (S12 S15 S15 S12) + // set t1 = (S14 S13 S14 S13) + let mut t0 = FieldElement32x4::zero(); + let mut t1 = FieldElement32x4::zero(); + for i in 0..5 { + t0.0[i] = _mm256_permutevar8x32_epi32(tmp.0[i], c0); + t1.0[i] = _mm256_permutevar8x32_epi32(tmp.0[i], c1); + } + + // return (S12*S14 S15*S13 S15*S14 S12*S13) = (X3 Y3 Z3 T3) + ExtendedPoint(&t0 * &t1) + } + } +} + impl<'a, 'b> Add<&'b ExtendedPoint> for &'a ExtendedPoint { type Output = ExtendedPoint; @@ -280,11 +381,11 @@ impl<'a, 'b> Sub<&'b ExtendedPoint> for &'a ExtendedPoint { } } -impl From for LookupTable { +impl From for LookupTable { fn from(P: ExtendedPoint) -> Self { - let mut points = [P; 8]; + let mut points = [CachedPoint::from(P); 8]; for i in 0..7 { - points[i+1] = &P + &points[i]; + points[i+1] = (&P + &points[i]).into(); } LookupTable(points) } @@ -297,7 +398,7 @@ impl<'a, 'b> Mul<&'b Scalar> for &'a ExtendedPoint { /// Uses a window of size 4. fn mul(self, scalar: &'b Scalar) -> ExtendedPoint { // Construct a lookup table of [P,2P,3P,4P,5P,6P,7P,8P] - let lookup_table = LookupTable::from(*self); + let lookup_table = LookupTable::::from(*self); // Setting s = scalar, compute // @@ -325,7 +426,7 @@ impl<'a, 'b> Mul<&'b Scalar> for &'a ExtendedPoint { } #[derive(Clone)] -pub struct EdwardsBasepointTable(pub [LookupTable; 32]); +pub struct EdwardsBasepointTable(pub [LookupTable; 32]); impl<'a, 'b> Mul<&'b Scalar> for &'a EdwardsBasepointTable { type Output = ExtendedPoint; @@ -372,12 +473,6 @@ impl EdwardsBasepointTable { } table } - - /// Get the basepoint for this table as an `ExtendedPoint`. - pub fn basepoint(&self) -> ExtendedPoint { - // self.0[0].select(1) = 1*(16^2)^0*B - self.0[0].select(1) - } } /// Given a vector of (possibly secret) scalars and a vector of @@ -639,16 +734,23 @@ mod test { // Test the vector implementation of the parallel subtraction formulas let S_vector: edwards::ExtendedPoint = (&ExtendedPoint::from(P) - &ExtendedPoint::from(Q)).into(); + // Test the vector implementation of the parallel readdition formulas + let cached_Q = CachedPoint::from(ExtendedPoint::from(Q)); + let T_vector: edwards::ExtendedPoint = (&ExtendedPoint::from(P) + &cached_Q).into(); + println!("Testing point addition:"); println!("P = {:?}", P); println!("Q = {:?}", Q); + println!("cached Q = {:?}", cached_Q); println!("R = P + Q = {:?}", &P + &Q); println!("R_serial = {:?}", R_serial); println!("R_vector = {:?}", R_vector); + println!("T_vector = {:?}", T_vector); println!("S = P - Q = {:?}", &P - &Q); println!("S_vector = {:?}", S_vector); assert_eq!(R_serial.compress(), (&P + &Q).compress()); assert_eq!(R_vector.compress(), (&P + &Q).compress()); + assert_eq!(T_vector.compress(), (&P + &Q).compress()); assert_eq!(S_vector.compress(), (&P - &Q).compress()); println!("OK!\n"); } @@ -874,6 +976,16 @@ mod bench { b.iter(|| edwards::ExtendedPoint::from(B_avx2)); } + #[bench] + fn point_readdition(b: &mut Bencher) { + let B = &constants::ED25519_BASEPOINT_TABLE; + let P = ExtendedPoint::from(B * &Scalar::from_u64(83973422)); + let Q = ExtendedPoint::from(B * &Scalar::from_u64(98932328)); + let Q_cached = CachedPoint::from(Q); + + b.iter(|| &P + &Q_cached ); + } + #[bench] fn point_addition(b: &mut Bencher) { let B = &constants::ED25519_BASEPOINT_TABLE; diff --git a/src/backend/avx2/field.rs b/src/backend/avx2/field.rs index 0c52739..12e36c1 100644 --- a/src/backend/avx2/field.rs +++ b/src/backend/avx2/field.rs @@ -13,6 +13,13 @@ #![allow(bad_style)] +pub const A_LANES: u8 = 0b0000_0101; +pub const B_LANES: u8 = 0b0000_1010; +pub const C_LANES: u8 = 0b0101_0000; +pub const D_LANES: u8 = 0b1010_0000; + +pub const ALL_LANES: u8 = A_LANES | B_LANES | C_LANES | D_LANES; + use std::ops::Mul; use stdsimd::simd::{u32x8, i32x8, u64x4}; @@ -111,9 +118,21 @@ impl FieldElement32x4 { return out; } + pub fn negate_lazy(&mut self, mask: u8) { + let mask = mask as i32; + unsafe { + use stdsimd::vendor::_mm256_blend_epi32; + self.0[0] = _mm256_blend_epi32(self.0[0].into(), (P_TIMES_2_LO - self.0[0]).into(), mask).into(); + self.0[1] = _mm256_blend_epi32(self.0[1].into(), (P_TIMES_2_HI - self.0[1]).into(), mask).into(); + self.0[2] = _mm256_blend_epi32(self.0[2].into(), (P_TIMES_2_HI - self.0[2]).into(), mask).into(); + self.0[3] = _mm256_blend_epi32(self.0[3].into(), (P_TIMES_2_HI - self.0[3]).into(), mask).into(); + self.0[4] = _mm256_blend_epi32(self.0[4].into(), (P_TIMES_2_HI - self.0[4]).into(), mask).into(); + } + } + /// Negate variables in lanes where mask is set /// XXX fix up api - pub fn mask_negate(&mut self, mask: u8) { + pub fn negate(&mut self, mask: u8) { let mask = mask as i32; unsafe { use stdsimd::vendor::_mm256_blend_epi32; @@ -126,6 +145,18 @@ impl FieldElement32x4 { self.reduce32(); } + /// Given `self = (A,B,C,D)`, set `self = (B,A,C,D)` + pub fn swap_AB(&mut self) { + unsafe { + use stdsimd::vendor::_mm256_shuffle_epi32; + use stdsimd::vendor::_mm256_blend_epi32; + for i in 0..5 { + let swapped = _mm256_shuffle_epi32(self.0[i].into(), 0b10_11_00_01); + self.0[i] = _mm256_blend_epi32(self.0[i].into(), swapped, 0b00001111).into(); + } + } + } + /// Given `self = (A,B,C,D)`, set `self = (A,B,D,C)` pub fn swap_CD(&mut self) { unsafe { @@ -192,7 +223,6 @@ impl FieldElement32x4 { unsafe { use stdsimd::vendor::_mm256_mul_epu32; - use stdsimd::vendor::_mm256_blend_epi32; let (b0, b1) = unpack_pair(self.0[0]); b[0] = _mm256_mul_epu32(b0, consts);