mirror of
https://github.com/saymrwulf/risc0-curve25519-dalek-source.git
synced 2026-09-08 20:40:32 +00:00
Add impl Mul<FieldElement51x4> for FieldElement51x4.
This commit is contained in:
parent
aa73d7b1bc
commit
70199d6094
1 changed files with 184 additions and 0 deletions
|
|
@ -20,6 +20,7 @@ extern "C" {
|
||||||
fn madd52hi(z: u64x4, x: u64x4, y: u64x4) -> u64x4;
|
fn madd52hi(z: u64x4, x: u64x4, y: u64x4) -> u64x4;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[derive(Copy, Clone)]
|
||||||
pub struct FieldElement51x4([u64x4; 5]);
|
pub struct FieldElement51x4([u64x4; 5]);
|
||||||
|
|
||||||
impl FieldElement51x4 {
|
impl FieldElement51x4 {
|
||||||
|
|
@ -97,6 +98,146 @@ impl FieldElement51x4 {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl<'a, 'b> Mul<&'b FieldElement51x4> for &'a FieldElement51x4 {
|
||||||
|
type Output = FieldElement51x4;
|
||||||
|
fn mul(self, rhs: &'b FieldElement51x4) -> FieldElement51x4 {
|
||||||
|
unsafe {
|
||||||
|
// Inputs
|
||||||
|
let x = &self.0;
|
||||||
|
let y = &rhs.0;
|
||||||
|
|
||||||
|
// Accumulators for lo-sourced terms
|
||||||
|
let mut z0lo = u64x4::splat(0);
|
||||||
|
let mut z1lo = u64x4::splat(0);
|
||||||
|
let mut z2lo = u64x4::splat(0);
|
||||||
|
let mut z3lo = u64x4::splat(0);
|
||||||
|
let mut z4lo = u64x4::splat(0);
|
||||||
|
let mut z5lo = u64x4::splat(0);
|
||||||
|
let mut z6lo = u64x4::splat(0);
|
||||||
|
let mut z7lo = u64x4::splat(0);
|
||||||
|
let mut z8lo = u64x4::splat(0);
|
||||||
|
|
||||||
|
// Accumulators for hi-sourced terms
|
||||||
|
// Need to be doubled before adding
|
||||||
|
let mut z1hi = u64x4::splat(0);
|
||||||
|
let mut z2hi = u64x4::splat(0);
|
||||||
|
let mut z3hi = u64x4::splat(0);
|
||||||
|
let mut z4hi = u64x4::splat(0);
|
||||||
|
let mut z5hi = u64x4::splat(0);
|
||||||
|
let mut z6hi = u64x4::splat(0);
|
||||||
|
let mut z7hi = u64x4::splat(0);
|
||||||
|
let mut z8hi = u64x4::splat(0);
|
||||||
|
let mut z9hi = u64x4::splat(0);
|
||||||
|
|
||||||
|
// Wave 0
|
||||||
|
z4lo = madd52lo(z4lo, x[4], y[0]);
|
||||||
|
z5hi = madd52hi(z5hi, x[4], y[0]);
|
||||||
|
z5lo = madd52lo(z5lo, x[4], y[1]);
|
||||||
|
z6hi = madd52hi(z6hi, x[4], y[1]);
|
||||||
|
z6lo = madd52lo(z6lo, x[4], y[2]);
|
||||||
|
z7hi = madd52hi(z7hi, x[4], y[2]);
|
||||||
|
z7lo = madd52lo(z7lo, x[4], y[3]);
|
||||||
|
z8hi = madd52hi(z8hi, x[4], y[3]);
|
||||||
|
|
||||||
|
// Wave 1
|
||||||
|
z4lo = madd52lo(z4lo, x[3], y[1]);
|
||||||
|
z5hi = madd52hi(z5hi, x[3], y[1]);
|
||||||
|
z5lo = madd52lo(z5lo, x[3], y[2]);
|
||||||
|
z6hi = madd52hi(z6hi, x[3], y[2]);
|
||||||
|
z6lo = madd52lo(z6lo, x[3], y[3]);
|
||||||
|
z7hi = madd52hi(z7hi, x[3], y[3]);
|
||||||
|
z7lo = madd52lo(z7lo, x[3], y[4]);
|
||||||
|
z8hi = madd52hi(z8hi, x[3], y[4]);
|
||||||
|
|
||||||
|
// Wave 2
|
||||||
|
z8lo = madd52lo(z8lo, x[4], y[4]);
|
||||||
|
z9hi = madd52hi(z9hi, x[4], y[4]);
|
||||||
|
z4lo = madd52lo(z4lo, x[2], y[2]);
|
||||||
|
z5hi = madd52hi(z5hi, x[2], y[2]);
|
||||||
|
z5lo = madd52lo(z5lo, x[2], y[3]);
|
||||||
|
z6hi = madd52hi(z6hi, x[2], y[3]);
|
||||||
|
z6lo = madd52lo(z6lo, x[2], y[4]);
|
||||||
|
z7hi = madd52hi(z7hi, x[2], y[4]);
|
||||||
|
|
||||||
|
let z8 = z8lo + z8hi + z8hi;
|
||||||
|
let z9 = z9hi + z9hi;
|
||||||
|
|
||||||
|
// Wave 3
|
||||||
|
z3lo = madd52lo(z3lo, x[3], y[0]);
|
||||||
|
z4hi = madd52hi(z4hi, x[3], y[0]);
|
||||||
|
z4lo = madd52lo(z4lo, x[1], y[3]);
|
||||||
|
z5hi = madd52hi(z5hi, x[1], y[3]);
|
||||||
|
z5lo = madd52lo(z5lo, x[1], y[4]);
|
||||||
|
z6hi = madd52hi(z6hi, x[1], y[4]);
|
||||||
|
z2lo = madd52lo(z2lo, x[2], y[0]);
|
||||||
|
z3hi = madd52hi(z3hi, x[2], y[0]);
|
||||||
|
|
||||||
|
let z6 = z6lo + z6hi + z6hi;
|
||||||
|
let z7 = z7lo + z7hi + z7hi;
|
||||||
|
|
||||||
|
// Wave 4
|
||||||
|
z3lo = madd52lo(z3lo, x[2], y[1]);
|
||||||
|
z4hi = madd52hi(z4hi, x[2], y[1]);
|
||||||
|
z4lo = madd52lo(z4lo, x[0], y[4]);
|
||||||
|
z5hi = madd52hi(z5hi, x[0], y[4]);
|
||||||
|
z1lo = madd52lo(z1lo, x[1], y[0]);
|
||||||
|
z2hi = madd52hi(z2hi, x[1], y[0]);
|
||||||
|
z2lo = madd52lo(z2lo, x[1], y[1]);
|
||||||
|
z3hi = madd52hi(z3hi, x[1], y[1]);
|
||||||
|
|
||||||
|
let z5 = z5lo + z5hi + z5hi;
|
||||||
|
|
||||||
|
// Wave 5
|
||||||
|
z3lo = madd52lo(z3lo, x[1], y[2]);
|
||||||
|
z4hi = madd52hi(z4hi, x[1], y[2]);
|
||||||
|
z0lo = madd52lo(z0lo, x[0], y[0]);
|
||||||
|
z1hi = madd52hi(z1hi, x[0], y[0]);
|
||||||
|
z1lo = madd52lo(z1lo, x[0], y[1]);
|
||||||
|
z2lo = madd52lo(z2lo, x[0], y[2]);
|
||||||
|
z2hi = madd52hi(z2hi, x[0], y[1]);
|
||||||
|
z3hi = madd52hi(z3hi, x[0], y[2]);
|
||||||
|
|
||||||
|
let r19 = u64x4::splat(19);
|
||||||
|
let r38 = u64x4::splat(38);
|
||||||
|
let r1938 = u64x4::splat(19 * 38);
|
||||||
|
let mut z919hi = u64x4::splat(0);
|
||||||
|
|
||||||
|
// Wave 6
|
||||||
|
z3lo = madd52lo(z3lo, x[0], y[3]);
|
||||||
|
z919hi = madd52hi(z919hi, r19, z9);
|
||||||
|
z0lo = madd52lo(z0lo, r19, z5);
|
||||||
|
z4lo = madd52lo(z4lo, r19, z9);
|
||||||
|
z1lo = madd52lo(z1lo, r19, z6);
|
||||||
|
z2lo = madd52lo(z2lo, r19, z7);
|
||||||
|
z1hi = madd52hi(z1hi, r19, z5);
|
||||||
|
z4hi = madd52hi(z4hi, x[0], y[3]);
|
||||||
|
|
||||||
|
// Wave 7
|
||||||
|
z3lo = madd52lo(z3lo, r19, z8);
|
||||||
|
z2hi = madd52hi(z2hi, r19, z6);
|
||||||
|
z0lo = madd52lo(z0lo, r38, z919hi);
|
||||||
|
z4lo = madd52lo(z4lo, r38, z8 >> 52);
|
||||||
|
z1lo = madd52lo(z1lo, r38, z5 >> 52);
|
||||||
|
z2lo = madd52lo(z2lo, r38, z6 >> 52);
|
||||||
|
z3hi = madd52hi(z3hi, r19, z7);
|
||||||
|
z4hi = madd52hi(z4hi, r19, z8);
|
||||||
|
|
||||||
|
// Wave 8
|
||||||
|
z3lo = madd52lo(z3lo, r38, z7 >> 52);
|
||||||
|
z0lo = madd52lo(z0lo, r1938, z9 >> 52);
|
||||||
|
|
||||||
|
FieldElement51x4([
|
||||||
|
z0lo,
|
||||||
|
z1lo + z1hi + z1hi,
|
||||||
|
z2lo + z2hi + z2hi,
|
||||||
|
z3lo + z3hi + z3hi,
|
||||||
|
z4lo + z4hi + z4hi,
|
||||||
|
])
|
||||||
|
.reduce()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod test {
|
mod test {
|
||||||
use super::*;
|
use super::*;
|
||||||
|
|
@ -146,4 +287,47 @@ mod test {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn mul_matches_serial() {
|
||||||
|
// Invert a small field element to get a big one
|
||||||
|
let a = FieldElement51([2438, 24, 243, 0, 0]).invert();
|
||||||
|
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 cx4 = &ax4 * &bx4;
|
||||||
|
|
||||||
|
let splits = cx4.split();
|
||||||
|
|
||||||
|
for i in 0..4 {
|
||||||
|
assert_eq!(c, splits[i]);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn iterated_mul_matches_serial() {
|
||||||
|
// Invert a small field element to get a big one
|
||||||
|
let a = FieldElement51([2438, 24, 243, 0, 0]).invert();
|
||||||
|
let b = FieldElement51([98098, 87987897, 0, 1, 0]).invert();
|
||||||
|
let mut c = &a * &b;
|
||||||
|
for i in 0..1024 {
|
||||||
|
c = &a * &c;
|
||||||
|
c = &b * &c;
|
||||||
|
}
|
||||||
|
|
||||||
|
let ax4 = FieldElement51x4::new(&a, &a, &a, &a);
|
||||||
|
let bx4 = FieldElement51x4::new(&b, &b, &b, &b);
|
||||||
|
let mut cx4 = &ax4 * &bx4;
|
||||||
|
for i in 0..1024 {
|
||||||
|
cx4 = &ax4 * &cx4;
|
||||||
|
cx4 = &bx4 * &cx4;
|
||||||
|
}
|
||||||
|
|
||||||
|
let splits = cx4.split();
|
||||||
|
|
||||||
|
for i in 0..4 {
|
||||||
|
assert_eq!(c, splits[i]);
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue