From 42648aa4601c4c5300fbe04a553295fb70f2713e Mon Sep 17 00:00:00 2001 From: Oleg Andreev Date: Tue, 21 May 2019 13:31:50 -0700 Subject: [PATCH] cgs --- src/backend/serial/scalar_mul/mod.rs | 3 + src/backend/serial/scalar_mul/pippenger.rs | 207 +++++++++++++++++++++ src/backend/vector/scalar_mul/mod.rs | 3 + src/backend/vector/scalar_mul/pippenger.rs | 172 +++++++++++++++++ src/scalar.rs | 32 ++-- 5 files changed, 401 insertions(+), 16 deletions(-) create mode 100644 src/backend/serial/scalar_mul/pippenger.rs create mode 100644 src/backend/vector/scalar_mul/pippenger.rs diff --git a/src/backend/serial/scalar_mul/mod.rs b/src/backend/serial/scalar_mul/mod.rs index bec874b..8d859eb 100644 --- a/src/backend/serial/scalar_mul/mod.rs +++ b/src/backend/serial/scalar_mul/mod.rs @@ -26,3 +26,6 @@ pub mod straus; #[cfg(feature = "alloc")] pub mod precomputed_straus; + +#[cfg(feature = "alloc")] +pub mod pippenger; diff --git a/src/backend/serial/scalar_mul/pippenger.rs b/src/backend/serial/scalar_mul/pippenger.rs new file mode 100644 index 0000000..82cf1bd --- /dev/null +++ b/src/backend/serial/scalar_mul/pippenger.rs @@ -0,0 +1,207 @@ +// -*- mode: rust; -*- +// +// This file is part of curve25519-dalek. +// Copyright (c) 2016-2018 Isis Lovecruft, Henry de Valence +// See LICENSE for licensing information. +// +// Authors: +// - Isis Agora Lovecruft +// - Henry de Valence +// - Oleg Andreev + +//! Implementation of a variant of Pippenger's algorithm. + +#![allow(non_snake_case)] + +use core::borrow::Borrow; + +use edwards::EdwardsPoint; +use scalar::Scalar; +use traits::VartimeMultiscalarMul; + +#[allow(unused_imports)] +use prelude::*; + +/// Implements a version of Pippenger's algorithm. +/// +/// The algorithm works as follows: +/// +/// Let `n` be a number of point-scalar pairs. +/// Let `w` be a window of bits (6..8, chosen based on `n`, see cost factor). +/// +/// 1. Prepare `2^(w-1) - 1` buckets with indices `[1..2^(w-1))` initialized with identity points. +/// Bucket 0 is not needed as it would contain points multiplied by 0. +/// 2. Convert scalars to a radix-`2^w` representation with signed digits in `[-2^w/2, 2^w/2]`. +/// Note: only the last digit may equal `2^w/2`. +/// 3. Starting with the last window, for each point `i=[0..n)` add it to a a bucket indexed by +/// the point's scalar's value in the window. +/// 4. Once all points in a window are sorted into buckets, add buckets by multiplying each +/// by their index. Efficient way of doing it is to start with the last bucket and compute two sums: +/// intermediate sum from the last to the first, and the full sum made of all intermediate sums. +/// 5. Shift the resulting sum of buckets by `w` bits by using `w` doublings. +/// 6. Add to the return value. +/// 7. Repeat the loop. +/// +/// Approximate cost w/o wNAF optimizations (A = addition, D = doubling): +/// +/// ```ascii +/// cost = (n*A + 2*(2^w/2)*A + w*D + A)*256/w +/// | | | | | +/// | | | | looping over 256/w windows +/// | | | adding to the result +/// sorting points | shifting the sum by w bits (to the next window, starting from last window) +/// one by one | +/// into buckets adding/subtracting all buckets +/// multiplied by their indexes +/// using a sum of intermediate sums +/// ``` +/// +/// For large `n`, dominant factor is (n*256/w) additions. +/// However, if `w` is too big and `n` is not too big, then `(2^w/2)*A` could dominate. +/// Therefore, the optimal choice of `w` grows slowly as `n` grows. +/// +pub struct Pippenger; + +#[cfg(any(feature = "alloc", feature = "std"))] +impl VartimeMultiscalarMul for Pippenger { + type Point = EdwardsPoint; + + fn optional_multiscalar_mul(scalars: I, points: J) -> Option + where + I: IntoIterator, + I::Item: Borrow, + J: IntoIterator>, + { + use backend::serial::curve_models::{ProjectiveNielsPoint}; + use traits::Identity; + + let mut scalars = scalars.into_iter(); + let size = scalars.by_ref().size_hint().0; + + // Digit width in bits. As digit width grows, + // number of point additions goes down, but amount of + // buckets and bucket additions grows exponentially. + let w = if size < 500 { + 6 + } else if size < 800 { + 7 + } else { + 8 + }; + + let max_digit: usize = 1 << w; + let digits_count: usize = (256 + w - 1) / w; // == ceil(256/w) + let buckets_count: usize = max_digit / 2; // digits are signed+centered hence 2^w/2, excluding 0-th bucket + + // Collect optimized scalars and points in buffers for repeated access + // (scanning the whole set per digit position). + let scalars = scalars.into_iter() + .map(|s| s.borrow().to_pippenger_radix(w).0 ) + .collect::>(); + let points: Vec = match points + .into_iter() + .map(|p| p.map(|P| P.to_projective_niels())) + .collect::>>() { + Some(x) => x, + None => return None, + }; + + // Prepare 2^w/2 buckets. + // buckets[i] corresponds to a multiplication factor (i+1). + let mut buckets: Vec<_> = (0..buckets_count) + .map(|_| EdwardsPoint::identity()) + .collect(); + + let columns: Vec<_> = (0..digits_count).map(|digit_index| { + + // Clear the buckets when processing another digit. + for i in 0..buckets_count { + buckets[i] = EdwardsPoint::identity(); + } + + // Iterate over pairs of (point, scalar) + // and add/sub the point to the corresponding bucket. + // Note: if we add support for precomputed lookup tables, + // we'll be adding/subtractiong point premultiplied by `digits[i]` to buckets[0]. + for (digits, pt) in scalars.iter().zip(points.iter()) { + let digit = digits[digit_index]; + if digit > 0 { + let b = (digit - 1) as usize; + buckets[b] = (&buckets[b] + pt).to_extended(); + } else if digit < 0 { + let b = (-digit - 1) as usize; + buckets[b] = (&buckets[b] - pt).to_extended(); + } + } + + // Add the buckets applying the multiplication factor to each bucket. + // The most efficient way to do that is to have a single sum with two running sums: + // an intermediate sum from last bucket to the first, and a sum of intermediate sums. + // + // For example, to add buckets 1*A, 2*B, 3*C we need to add these points: + // C + // C B + // C B A Sum = C + (C+B) + (C+B+A) + let mut buckets_intermediate_sum = buckets[buckets_count - 1]; + let mut buckets_sum = buckets[buckets_count - 1]; + for i in (0..(buckets_count - 1)).rev() { + buckets_intermediate_sum += buckets[i]; + buckets_sum += buckets_intermediate_sum; + } + + buckets_sum + }) + .collect(); + // ^ Note: we collect points because if we chain .rev().fold() + // then the .map() will run in reversed order, producing incorrect digit values + // (they can only be produced in lo->hi order). + + // Add the intermediate per-digit results in hi->lo order + // so that we can minimize doublings. + Some(columns[0..(digits_count - 1)].iter().rev().fold( + columns[digits_count - 1], + |total, &p| total.mul_by_pow_2(w as u32) + p, + )) + } +} + +#[cfg(test)] +mod test { + use super::*; + use constants; + use scalar::Scalar; + + #[test] + fn test_vartime_pippenger() { + // Reuse points across different tests + let mut n = 512; + let x = Scalar::from(2128506u64).invert(); + let y = Scalar::from(4443282u64).invert(); + let points: Vec<_> = (0..n) + .map(|i| { + constants::ED25519_BASEPOINT_POINT * Scalar::from(1 + i as u64) + }) + .collect(); + let scalars: Vec<_> = (0..n) + .map(|i| x + (Scalar::from(i as u64)*y)) // fast way to make ~random but deterministic scalars + .collect(); + + let premultiplied: Vec = scalars + .iter() + .zip(points.iter()) + .map(|(sc, pt)| sc * pt) + .collect(); + + while n > 0 { + let scalars = &scalars[0..n].to_vec(); + let points = &points[0..n].to_vec(); + let control: EdwardsPoint = premultiplied[0..n].iter().sum(); + + let subject = Pippenger::vartime_multiscalar_mul(scalars.clone(), points.clone()); + + assert_eq!(subject.compress(), control.compress()); + + n = n / 2; + } + } +} diff --git a/src/backend/vector/scalar_mul/mod.rs b/src/backend/vector/scalar_mul/mod.rs index 5c8734d..bdd6c4a 100644 --- a/src/backend/vector/scalar_mul/mod.rs +++ b/src/backend/vector/scalar_mul/mod.rs @@ -17,3 +17,6 @@ pub mod straus; #[cfg(feature = "alloc")] pub mod precomputed_straus; + +#[cfg(feature = "alloc")] +pub mod pippenger; diff --git a/src/backend/vector/scalar_mul/pippenger.rs b/src/backend/vector/scalar_mul/pippenger.rs new file mode 100644 index 0000000..4ec72f6 --- /dev/null +++ b/src/backend/vector/scalar_mul/pippenger.rs @@ -0,0 +1,172 @@ +// -*- mode: rust; -*- +// +// This file is part of curve25519-dalek. +// Copyright (c) 2016-2018 Isis Lovecruft, Henry de Valence +// See LICENSE for licensing information. +// +// Authors: +// - Isis Agora Lovecruft +// - Henry de Valence +// - Oleg Andreev + +#![allow(non_snake_case)] + +use core::borrow::Borrow; + +use clear_on_drop::ClearOnDrop; + +use backend::vector::{CachedPoint, ExtendedPoint}; +use edwards::EdwardsPoint; +use scalar::Scalar; +use window::{LookupTable, NafLookupTable5}; +use traits::{Identity, MultiscalarMul, VartimeMultiscalarMul}; + +#[allow(unused_imports)] +use prelude::*; + +/// Implements a version of Pippenger's algorithm. +/// +/// See the documentation in the serial `scalar_mul::pippenger` module for details. +pub struct Pippenger; + +#[cfg(any(feature = "alloc", feature = "std"))] +impl VartimeMultiscalarMul for Pippenger { + type Point = EdwardsPoint; + + fn optional_multiscalar_mul(scalars: I, points: J) -> Option + where + I: IntoIterator, + I::Item: Borrow, + J: IntoIterator>, + { + let mut scalars = scalars.into_iter(); + let size = scalars.by_ref().size_hint().0; + let w = if size < 500 { + 6 + } else if size < 800 { + 7 + } else { + 8 + }; + + let max_digit: usize = 1 << w; + let digits_count: usize = (256 + w - 1) / w; // == ceil(256/w) + let buckets_count: usize = max_digit / 2; // digits are signed+centered hence 2^w/2, excluding 0-th bucket + + // Collect optimized scalars and points in buffers for repeated access + // (scanning the whole set per digit position). + let scalars = scalars.into_iter() + .map(|s| s.borrow().to_pippenger_radix(w).0 ) + .collect::>(); + let points: Vec = match points + .into_iter() + .map(|p| p.map(|P| CachedPoint::from(ExtendedPoint::from(P)))) + .collect::>>() { + Some(x) => x, + None => return None, + }; + + // Prepare 2^w/2 buckets. + // buckets[i] corresponds to a multiplication factor (i+1). + let mut buckets: Vec<_> = (0..buckets_count) + .map(|_| ExtendedPoint::identity()) + .collect(); + + let columns: Vec = (0..digits_count).map(|digit_index| { + + // Clear the buckets when processing another digit. + for i in 0..buckets_count { + buckets[i] = ExtendedPoint::identity(); + } + + // Iterate over pairs of (point, scalar) + // and add/sub the point to the corresponding bucket. + // Note: if we add support for precomputed lookup tables, + // we'll be adding/subtractiong point premultiplied by `digits[i]` to buckets[0]. + for (digits, pt) in scalars.iter().zip(points.iter()) { + let digit = digits[digit_index]; + if digit > 0 { + let b = (digit - 1) as usize; + buckets[b] = &buckets[b] + pt; + } else if digit < 0 { + let b = (-digit - 1) as usize; + buckets[b] = &buckets[b] - pt; + } + } + + // Add the buckets applying the multiplication factor to each bucket. + // The most efficient way to do that is to have a single sum with two running sums: + // an intermediate sum from last bucket to the first, and a sum of intermediate sums. + // + // For example, to add buckets 1*A, 2*B, 3*C we need to add these points: + // C + // C B + // C B A Sum = C + (C+B) + (C+B+A) + let mut buckets_intermediate_sum = buckets[buckets_count - 1]; + let mut buckets_sum = buckets[buckets_count - 1]; + for i in (0..(buckets_count - 1)).rev() { + buckets_intermediate_sum = &buckets_intermediate_sum + &buckets[i]; + buckets_sum = &buckets_sum + &buckets_intermediate_sum; + } + + buckets_sum + }) + .collect(); + // ^ Note: we collect points because if we chain .rev().fold() + // then the .map() will run in reversed order, producing incorrect digit values + // (they can only be produced in lo->hi order). + + // Add the intermediate per-digit results in hi->lo order + // so that we can minimize doublings. + Some( + columns[0..(digits_count - 1)] + .iter() + .rev() + .fold(columns[digits_count - 1], |total, &p| { + &total.mul_by_pow_2(w as u32) + &p + }) + .into(), + ) + } +} + +#[cfg(test)] +mod test { + use super::*; + use constants; + use scalar::Scalar; + + #[test] + fn test_vartime_pippenger() { + // Reuse points across different tests + let mut n = 512; + let x = Scalar::from(2128506u64).invert(); + let y = Scalar::from(4443282u64).invert(); + let points: Vec<_> = (0..n) + .map(|i| { + constants::ED25519_BASEPOINT_POINT * Scalar::from(1 + i as u64) + }) + .collect(); + let scalars: Vec<_> = (0..n) + .map(|i| x + (Scalar::from(i as u64)*y)) // fast way to make ~random but deterministic scalars + .collect(); + + let premultiplied: Vec = scalars + .iter() + .zip(points.iter()) + .map(|(sc, pt)| sc * pt) + .collect(); + + while n > 0 { + let scalars = &scalars[0..n].to_vec(); + let points = &points[0..n].to_vec(); + let control: EdwardsPoint = premultiplied[0..n].iter().sum(); + + let subject = Pippenger::vartime_multiscalar_mul(scalars.clone(), points.clone()); + + assert_eq!(subject.compress(), control.compress()); + + n = n / 2; + } + } +} \ No newline at end of file diff --git a/src/scalar.rs b/src/scalar.rs index ae2344e..73b6d82 100644 --- a/src/scalar.rs +++ b/src/scalar.rs @@ -972,18 +972,18 @@ impl Scalar { /// /// ## Scalar representation /// - /// Radix \\(2\^r\\), with \\(n = ceil(256/r)\\) coefficients in \\([-(2\^r)/2,(2\^r)/2)\\), + /// Radix \\(2\^w\\), with \\(n = ceil(256/w)\\) coefficients in \\([-(2\^w)/2,(2\^w)/2)\\), /// i.e., scalar is represented using digits \\(a\_i\\) such that /// $$ - /// a = a\_0 + a\_1 2\^1r + \cdots + a_{n-1} 2\^{r*(n-1)}, + /// a = a\_0 + a\_1 2\^1w + \cdots + a_{n-1} 2\^{w*(n-1)}, /// $$ - /// with \\(-2\^r/2 \leq a_i < 2\^r/2\\) for \\(0 \leq i < (n-1)\\) and \\(-2\^r/2 \leq a_{n-1} \leq 2\^r/2\\). + /// with \\(-2\^w/2 \leq a_i < 2\^w/2\\) for \\(0 \leq i < (n-1)\\) and \\(-2\^w/2 \leq a_{n-1} \leq 2\^w/2\\). /// - pub(crate) fn to_pippenger_radix(&self, r: usize) -> ([i8; 43], usize) { - debug_assert!(r >= 6); - debug_assert!(r <= 8); + pub(crate) fn to_pippenger_radix(&self, w: usize) -> ([i8; 43], usize) { + debug_assert!(w >= 6); + debug_assert!(w <= 8); - let digits_count = (256 + r - 1)/r as usize; + let digits_count = (256 + w - 1)/w as usize; debug_assert!(digits_count <= 43); use byteorder::{ByteOrder, LittleEndian}; @@ -992,20 +992,20 @@ impl Scalar { let mut scalar64x4 = [0u64; 4]; LittleEndian::read_u64_into(&self.bytes, &mut scalar64x4[0..4]); - let radix: u64 = 1 << r; + let radix: u64 = 1 << w; let window_mask: u64 = radix - 1; let mut carry = 0u64; let mut digits = [0i8; 43]; for i in 0..digits_count { // Construct a buffer of bits of the scalar, starting at `bit_offset`. - let bit_offset = i*r; + let bit_offset = i*w; let u64_idx = bit_offset / 64; let bit_idx = bit_offset % 64; // Read the bits from the scalar let bit_buf: u64; - if bit_idx < 64 - r || u64_idx == 3 { + if bit_idx < 64 - w || u64_idx == 3 { // This window's bits are contained in a single u64, // or it's the last u64 anyway. bit_buf = scalar64x4[u64_idx] >> bit_idx; @@ -1018,8 +1018,8 @@ impl Scalar { let coef = carry + (bit_buf & window_mask); // coef = [0, 2^r) // Recenter coefficients from [0,2^r) to [-2^r/2, 2^r/2) - carry = (coef + (radix/2) as u64) >> r; - digits[i] = ((coef as i64) - (carry << r) as i64) as i8; + carry = (coef + (radix/2) as u64) >> w; + digits[i] = ((coef as i64) - (carry << w) as i64) as i8; } // Apply the resulting carry to the last digit @@ -1030,7 +1030,7 @@ impl Scalar { // we allow the last word to touch the value 2^r/2. // XXX: make sure tests cover this case, so the carry is non-zero and this line matters. // Maybe it never happens to be non-zero for r=6/7/8?... - digits[digits_count-1] += (carry << r) as i8; + digits[digits_count-1] += (carry << w) as i8; (digits, digits_count) } @@ -1517,11 +1517,11 @@ mod test { use std::iter; // For each valid radix it tests that 1000 random-ish scalars can be restored // from the produced representation precisely. - for r in 6..9 { + for w in 6..9 { for scalar in (2..100).map(|s| Scalar::from(s as u64).invert() ).chain(iter::once(-Scalar::one())) { - let (digits, digits_count) = scalar.to_pippenger_radix(r); + let (digits, digits_count) = scalar.to_pippenger_radix(w); - let radix = Scalar::from((1<