diff --git a/src/backend/avx2/edwards.rs b/src/backend/avx2/edwards.rs index 702c64a..be154dd 100644 --- a/src/backend/avx2/edwards.rs +++ b/src/backend/avx2/edwards.rs @@ -26,7 +26,7 @@ use scalar_mul::window::{LookupTable, NafLookupTable5, NafLookupTable8}; use traits::Identity; -use backend::avx2::field::{FieldElement32x4, Lanes, D_LANES}; +use backend::avx2::field::{FieldElement32x4, Lanes, Shuffle, D_LANES}; use backend::avx2; @@ -96,41 +96,21 @@ impl ExtendedPoint { // and then adding. // Set t0 = (X1 Y1 X1 Y1) - t0.0[0] = _mm256_permute2x128_si256(P.0[0].into_bits(), P.0[0].into_bits(), 0b0000_0000).into_bits(); - t0.0[1] = _mm256_permute2x128_si256(P.0[1].into_bits(), P.0[1].into_bits(), 0b0000_0000).into_bits(); - t0.0[2] = _mm256_permute2x128_si256(P.0[2].into_bits(), P.0[2].into_bits(), 0b0000_0000).into_bits(); - t0.0[3] = _mm256_permute2x128_si256(P.0[3].into_bits(), P.0[3].into_bits(), 0b0000_0000).into_bits(); - t0.0[4] = _mm256_permute2x128_si256(P.0[4].into_bits(), P.0[4].into_bits(), 0b0000_0000).into_bits(); + let mut t0 = P.shuffle(Shuffle::ABAB); // Set t1 = (Y1 X1 Y1 X1) - t1.0[0] = _mm256_shuffle_epi32(t0.0[0].into_bits(), 0b10_11_00_01).into_bits(); - t1.0[1] = _mm256_shuffle_epi32(t0.0[1].into_bits(), 0b10_11_00_01).into_bits(); - t1.0[2] = _mm256_shuffle_epi32(t0.0[2].into_bits(), 0b10_11_00_01).into_bits(); - t1.0[3] = _mm256_shuffle_epi32(t0.0[3].into_bits(), 0b10_11_00_01).into_bits(); - t1.0[4] = _mm256_shuffle_epi32(t0.0[4].into_bits(), 0b10_11_00_01).into_bits(); - - // Set t0 = (X1+Y1 X1+Y1 X1+Y1 X1+Y1) - t0.0[0] = t0.0[0] + t1.0[0]; - t0.0[1] = t0.0[1] + t1.0[1]; - t0.0[2] = t0.0[2] + t1.0[2]; - t0.0[3] = t0.0[3] + t1.0[3]; - t0.0[4] = t0.0[4] + t1.0[4]; + let mut t1 = t0.shuffle(Shuffle::BADC); // Set t0 = (X1 Y1 Z1 X1+Y1) - // why does this intrinsic take an i32 for the imm8 ??? - t0.0[0] = _mm256_blend_epi32(P.0[0].into_bits(), t0.0[0].into_bits(), D_LANES as i32).into_bits(); - t0.0[1] = _mm256_blend_epi32(P.0[1].into_bits(), t0.0[1].into_bits(), D_LANES as i32).into_bits(); - t0.0[2] = _mm256_blend_epi32(P.0[2].into_bits(), t0.0[2].into_bits(), D_LANES as i32).into_bits(); - t0.0[3] = _mm256_blend_epi32(P.0[3].into_bits(), t0.0[3].into_bits(), D_LANES as i32).into_bits(); - t0.0[4] = _mm256_blend_epi32(P.0[4].into_bits(), t0.0[4].into_bits(), D_LANES as i32).into_bits(); + t0 = P.blend(&t0 + &t1, Lanes::D); // Set t1 = t0^2, negating the D values t1 = t0.square_and_negate_D(); // Now t1 = (S1 S2 S3 -S4) - let c0 = u32x8::new(0,0,2,2,0,0,2,2).into_bits(); // (ABCD) -> (AAAA) - let c1 = u32x8::new(1,1,3,3,1,1,3,3).into_bits(); // (ABCD) -> (BBBB) + let c0 = u32x8::new(0, 0, 2, 2, 0, 0, 2, 2).into_bits(); // (ABCD) -> (AAAA) + let c1 = u32x8::new(1, 1, 3, 3, 1, 1, 3, 3).into_bits(); // (ABCD) -> (BBBB) // See discussion of bounds in the module-level documentation. // @@ -165,14 +145,8 @@ impl ExtendedPoint { t0.0[i] = t0.0[i] - S2_neg; } - let c0 = u32x8::new(4,0,6,2,4,0,6,2).into_bits(); // (ABCD) -> (CACA) - let c1 = u32x8::new(5,1,7,3,1,5,3,7).into_bits(); // (ABCD) -> (DBBD) - - for i in 0..5 { - let tmp = t0.0[i]; - t0.0[i] = _mm256_permutevar8x32_epi32(tmp.into_bits(), c0).into_bits(); - t1.0[i] = _mm256_permutevar8x32_epi32(tmp.into_bits(), c1).into_bits(); - } + t1 = t0.shuffle(Shuffle::DBBD); + t0 = t0.shuffle(Shuffle::CACA); ExtendedPoint(&t0 * &t1) } @@ -248,38 +222,27 @@ impl<'a, 'b> Add<&'b CachedPoint> for &'a ExtendedPoint { /// Uses a slight tweak of the parallel unified formulas of HWCD'08 fn add(self, other: &'b CachedPoint) -> ExtendedPoint { - unsafe { - use core::arch::x86_64::_mm256_permutevar8x32_epi32; + let mut tmp = self.0; - let mut tmp = self.0; + // tmp = (Y1-X1 Y1+X1 Z1 T1) = (S0 S1 Z1 T1) + tmp.diff_sum(Lanes::AB); - // tmp = (Y1-X1 Y1+X1 Z1 T1) = (S0 S1 Z1 T1) - tmp.diff_sum(Lanes::AB); + // tmp = (S0*S2' S1*S3' Z1*Z2' T1*T2') = (S8 S9 S10 S11) + tmp = &tmp * &other.0; - // 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 = (S8 S9 S11 S10) - tmp.swap_CD(); + // tmp = (S9-S8 S9+S8 S10-S11 S10+S11) = (S12 S13 S14 S15) + tmp.diff_sum(Lanes::ALL); - // tmp = (S9-S8 S9+S8 S10-S11 S10+S11) = (S12 S13 S14 S15) - tmp.diff_sum(Lanes::ALL); + // set t0 = (S12 S15 S15 S12) + let t0 = tmp.shuffle(Shuffle::ADDA); + // set t1 = (S14 S13 S14 S13) + let t1 = tmp.shuffle(Shuffle::CBCB); - 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].into_bits(), c0.into_bits()).into_bits(); - t1.0[i] = _mm256_permutevar8x32_epi32(tmp.0[i].into_bits(), c1.into_bits()).into_bits(); - } - - // return (S12*S14 S15*S13 S15*S14 S12*S13) = (X3 Y3 Z3 T3) - ExtendedPoint(&t0 * &t1) - } + // return (S12*S14 S15*S13 S15*S14 S12*S13) = (X3 Y3 Z3 T3) + ExtendedPoint(&t0 * &t1) } } diff --git a/src/backend/avx2/field.rs b/src/backend/avx2/field.rs index 6751b81..d5ac6a8 100644 --- a/src/backend/avx2/field.rs +++ b/src/backend/avx2/field.rs @@ -32,6 +32,8 @@ use backend::u64::field::FieldElement64; #[derive(Copy, Clone)] pub enum Lanes { + C, + D, AB, CD, ALL, @@ -64,6 +66,18 @@ fn blend_lanes(x: u32x8, y: u32x8, control: Lanes) -> u32x8 { } } +#[derive(Copy, Clone)] +pub enum Shuffle { + AAAA, + BBBB, + CACA, + DBBD, + ADDA, + CBCB, + ABAB, + BADC, +} + /// A vector of four `FieldElements`, implemented using AVX2. #[derive(Clone, Copy, Debug)] pub(crate) struct FieldElement32x4(pub(crate) [u32x8; 5]); @@ -103,6 +117,50 @@ impl FieldElement32x4 { out } + #[inline(always)] + pub fn blend(&self, b: FieldElement32x4, control: Lanes) -> FieldElement32x4 { + FieldElement32x4([ + blend_lanes(self.0[0], b.0[0], control), + blend_lanes(self.0[1], b.0[1], control), + blend_lanes(self.0[2], b.0[2], control), + blend_lanes(self.0[3], b.0[3], control), + blend_lanes(self.0[4], b.0[4], control), + ]) + } + + #[inline(always)] + pub fn shuffle(&self, control: Shuffle) -> FieldElement32x4 { + #[inline(always)] + fn shuffle_lanes(x: u32x8, control: Shuffle) -> u32x8 { + unsafe { + use core::arch::x86_64::_mm256_permutevar8x32_epi32; + + let c: u32x8 = match control { + Shuffle::AAAA => u32x8::new(0, 0, 2, 2, 0, 0, 2, 2), + Shuffle::BBBB => u32x8::new(1, 1, 3, 3, 1, 1, 3, 3), + Shuffle::CACA => u32x8::new(4, 0, 6, 2, 4, 0, 6, 2), + Shuffle::DBBD => u32x8::new(5, 1, 7, 3, 1, 5, 3, 7), + Shuffle::ADDA => u32x8::new(0, 5, 2, 7, 5, 0, 7, 2), + Shuffle::CBCB => u32x8::new(4, 1, 6, 3, 4, 1, 6, 3), + Shuffle::ABAB => u32x8::new(0, 1, 2, 3, 0, 1, 2, 3), + Shuffle::BADC => u32x8::new(1, 0, 3, 2, 5, 4, 7, 6), + }; + // Note that this gets turned into a generic LLVM + // shuffle-by-constants, which can be lowered to a simpler + // instruction than a generic permute. + _mm256_permutevar8x32_epi32(x.into_bits(), c.into_bits()).into_bits() + } + } + + FieldElement32x4([ + shuffle_lanes(self.0[0], control), + shuffle_lanes(self.0[1], control), + shuffle_lanes(self.0[2], control), + shuffle_lanes(self.0[3], control), + shuffle_lanes(self.0[4], control), + ]) + } + pub fn zero() -> FieldElement32x4 { FieldElement32x4([u32x8::splat(0); 5]) }