Add a shuffling abstraction for FieldElement32x4

This commit is contained in:
Henry de Valence 2018-06-14 14:07:33 -07:00
parent c8dc2a6418
commit 14ce6d3da6
2 changed files with 81 additions and 60 deletions

View file

@ -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)
}
}

View file

@ -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])
}