use alloc::vec::Vec;
use core::arch::wasm32::{
i32x4_shuffle, i64x2_add, i64x2_extmul_low_u32x4, i64x2_gt, i64x2_shl, i64x2_shuffle,
i64x2_sub, u64x2_shr, u64x2_splat, v128, v128_and, v128_andnot, v128_or, v128_xor,
};
use core::fmt::Debug;
use core::iter::{Product, Sum};
use core::mem::transmute;
use core::ops::{Add, AddAssign, Div, DivAssign, Mul, MulAssign, Neg, Sub, SubAssign};
use p3_field::exponentiation::exp_10540996611094048183;
use p3_field::op_assign_macros::{
impl_add_assign, impl_add_base_field, impl_div_methods, impl_mul_base_field, impl_mul_methods,
impl_packed_field_div, impl_packed_value, impl_rng, impl_sub_assign, impl_sub_base_field,
impl_sum_prod_base_field, ring_sum,
};
use p3_field::{
Algebra, Field, InjectiveMonomial, PackedField, PackedFieldPow2, PackedValue,
PermutationMonomial, PrimeCharacteristicRing, PrimeField64,
};
use p3_util::reconstitute_from_base;
use rand::distr::{Distribution, StandardUniform};
use rand::{Rng, RngExt};
use crate::{Goldilocks, P};
const WIDTH: usize = 2;
const EPSILON: u64 = Goldilocks::ORDER_U64.wrapping_neg();
const _LAYOUT_INVARIANTS: () = {
assert!(size_of::<[Goldilocks; WIDTH]>() == size_of::<v128>());
assert!(size_of::<Goldilocks>() == size_of::<u64>());
};
#[derive(Copy, Clone, Debug, Default, PartialEq, Eq)]
#[repr(transparent)]
#[must_use]
pub struct PackedGoldilocksWasmSimd128(pub [Goldilocks; WIDTH]);
impl PackedGoldilocksWasmSimd128 {
#[inline]
#[must_use]
pub(crate) fn to_vector(self) -> v128 {
unsafe { transmute(self) }
}
#[inline]
pub(crate) fn from_vector(vector: v128) -> Self {
unsafe { transmute(vector) }
}
#[inline]
const fn broadcast(value: Goldilocks) -> Self {
Self([value; WIDTH])
}
}
impl From<Goldilocks> for PackedGoldilocksWasmSimd128 {
fn from(x: Goldilocks) -> Self {
Self::broadcast(x)
}
}
impl Add for PackedGoldilocksWasmSimd128 {
type Output = Self;
#[inline]
fn add(self, rhs: Self) -> Self {
Self::from_vector(add(self.to_vector(), rhs.to_vector()))
}
}
impl Sub for PackedGoldilocksWasmSimd128 {
type Output = Self;
#[inline]
fn sub(self, rhs: Self) -> Self {
Self::from_vector(sub(self.to_vector(), rhs.to_vector()))
}
}
impl Neg for PackedGoldilocksWasmSimd128 {
type Output = Self;
#[inline]
fn neg(self) -> Self {
Self::from_vector(neg(self.to_vector()))
}
}
impl Mul for PackedGoldilocksWasmSimd128 {
type Output = Self;
#[inline]
fn mul(self, rhs: Self) -> Self {
Self::from_vector(mul(self.to_vector(), rhs.to_vector()))
}
}
impl_add_assign!(PackedGoldilocksWasmSimd128);
impl_sub_assign!(PackedGoldilocksWasmSimd128);
impl_mul_methods!(PackedGoldilocksWasmSimd128);
ring_sum!(PackedGoldilocksWasmSimd128);
impl_rng!(PackedGoldilocksWasmSimd128);
impl PrimeCharacteristicRing for PackedGoldilocksWasmSimd128 {
type PrimeSubfield = Goldilocks;
const ZERO: Self = Self::broadcast(Goldilocks::ZERO);
const ONE: Self = Self::broadcast(Goldilocks::ONE);
const TWO: Self = Self::broadcast(Goldilocks::TWO);
const NEG_ONE: Self = Self::broadcast(Goldilocks::NEG_ONE);
#[inline]
fn from_prime_subfield(f: Self::PrimeSubfield) -> Self {
f.into()
}
#[inline]
fn halve(&self) -> Self {
Self::from_vector(halve(self.to_vector()))
}
#[inline]
fn double(&self) -> Self {
Self::from_vector(double(self.to_vector()))
}
#[inline]
fn square(&self) -> Self {
Self::from_vector(square(self.to_vector()))
}
#[inline]
fn zero_vec(len: usize) -> Vec<Self> {
unsafe { reconstitute_from_base(Goldilocks::zero_vec(len * WIDTH)) }
}
#[inline]
fn sum_array<const N: usize>(input: &[Self]) -> Self {
assert_eq!(N, input.len());
match N {
0 => Self::ZERO,
1 => input[0],
2 => input[0] + input[1],
_ => {
let vectors: [v128; N] = core::array::from_fn(|i| input[i].to_vector());
Self::from_vector(sum_delayed_reduce::<N>(&vectors))
}
}
}
#[inline]
fn dot_product<const N: usize>(lhs: &[Self; N], rhs: &[Self; N]) -> Self {
match N {
0 => Self::ZERO,
1 => lhs[0] * rhs[0],
_ => Self::from_vector(dot_pairs::<N>(|i| (lhs[i].to_vector(), rhs[i].to_vector()))),
}
}
}
impl InjectiveMonomial<7> for PackedGoldilocksWasmSimd128 {}
impl PermutationMonomial<7> for PackedGoldilocksWasmSimd128 {
fn injective_exp_root_n(&self) -> Self {
exp_10540996611094048183(*self)
}
}
impl_add_base_field!(PackedGoldilocksWasmSimd128, Goldilocks);
impl_sub_base_field!(PackedGoldilocksWasmSimd128, Goldilocks);
impl_mul_base_field!(PackedGoldilocksWasmSimd128, Goldilocks);
impl_div_methods!(PackedGoldilocksWasmSimd128, Goldilocks);
impl_packed_field_div!(PackedGoldilocksWasmSimd128);
impl_sum_prod_base_field!(PackedGoldilocksWasmSimd128, Goldilocks);
impl Algebra<Goldilocks> for PackedGoldilocksWasmSimd128 {
const BATCHED_LC_CHUNK: usize = 4;
#[inline]
fn mixed_dot_product<const N: usize>(a: &[Self; N], f: &[Goldilocks; N]) -> Self {
match N {
0 => Self::ZERO,
1 => a[0] * f[0],
_ => Self::from_vector(dot_pairs::<N>(|i| {
(a[i].to_vector(), Self::from(f[i]).to_vector())
})),
}
}
}
impl_packed_value!(PackedGoldilocksWasmSimd128, Goldilocks, WIDTH);
unsafe impl PackedField for PackedGoldilocksWasmSimd128 {
type Scalar = Goldilocks;
}
#[inline]
pub fn interleave_u64(v0: v128, v1: v128) -> (v128, v128) {
let r0 = i64x2_shuffle::<0, 2>(v0, v1);
let r1 = i64x2_shuffle::<1, 3>(v0, v1);
(r0, r1)
}
unsafe impl PackedFieldPow2 for PackedGoldilocksWasmSimd128 {
fn interleave(&self, other: Self, block_len: usize) -> (Self, Self) {
let (v0, v1) = (self.to_vector(), other.to_vector());
let (res0, res1) = match block_len {
1 => interleave_u64(v0, v1),
2 => (v0, v1),
_ => panic!("unsupported block length"),
};
(Self::from_vector(res0), Self::from_vector(res1))
}
}
const SIGN_BIT: v128 =
unsafe { transmute::<[u64; WIDTH], v128>([0x8000_0000_0000_0000u64; WIDTH]) };
const SHIFTED_FIELD_ORDER: v128 = unsafe {
transmute::<[u64; WIDTH], v128>([Goldilocks::ORDER_U64 ^ 0x8000_0000_0000_0000u64; WIDTH])
};
const EPSILON_VEC: v128 = unsafe { transmute::<[u64; WIDTH], v128>([EPSILON; WIDTH]) };
#[inline(always)]
fn shift(x: v128) -> v128 {
v128_xor(x, SIGN_BIT)
}
#[inline(always)]
fn canonicalize_s(x_s: v128) -> v128 {
let mask = i64x2_gt(SHIFTED_FIELD_ORDER, x_s);
let wrapback_amt = v128_andnot(EPSILON_VEC, mask);
i64x2_add(x_s, wrapback_amt)
}
#[inline(always)]
fn add_no_double_overflow_64_64s_s(x: v128, y_s: v128) -> v128 {
let res_wrapped_s = i64x2_add(x, y_s);
let mask = i64x2_gt(y_s, res_wrapped_s);
let wrapback_amt = u64x2_shr(mask, 32);
i64x2_add(res_wrapped_s, wrapback_amt)
}
#[inline]
fn add(x: v128, y: v128) -> v128 {
let y_s = shift(y);
let res_s = add_no_double_overflow_64_64s_s(x, canonicalize_s(y_s));
shift(res_s)
}
#[inline]
fn sub(x: v128, y: v128) -> v128 {
let y_s = canonicalize_s(shift(y));
let x_s = shift(x);
let mask = i64x2_gt(y_s, x_s);
let wrapback_amt = u64x2_shr(mask, 32);
let res_wrapped = i64x2_sub(x_s, y_s);
i64x2_sub(res_wrapped, wrapback_amt)
}
#[inline]
fn neg(y: v128) -> v128 {
let y_s = shift(y);
i64x2_sub(SHIFTED_FIELD_ORDER, canonicalize_s(y_s))
}
#[inline(always)]
pub(crate) fn halve(input: v128) -> v128 {
let one = u64x2_splat(1);
let zero = u64x2_splat(0);
let half_v = u64x2_splat(P.div_ceil(2));
let least_bit = v128_and(input, one);
let t = u64x2_shr(input, 1);
let neg_least_bit = i64x2_sub(zero, least_bit);
let maybe_half = v128_and(half_v, neg_least_bit);
i64x2_add(t, maybe_half)
}
#[inline(always)]
fn lo32(a: v128) -> v128 {
i32x4_shuffle::<0, 2, 0, 0>(a, a)
}
#[inline(always)]
fn hi32(a: v128) -> v128 {
i32x4_shuffle::<1, 3, 0, 0>(a, a)
}
#[inline(always)]
fn mul_u32_lanes(a_packed: v128, b_packed: v128) -> v128 {
i64x2_extmul_low_u32x4(a_packed, b_packed)
}
#[inline]
fn mul64_64(x: v128, y: v128) -> (v128, v128) {
let x_lo = lo32(x);
let x_hi = hi32(x);
let y_lo = lo32(y);
let y_hi = hi32(y);
let ll = mul_u32_lanes(x_lo, y_lo); let lh = mul_u32_lanes(x_lo, y_hi); let hl = mul_u32_lanes(x_hi, y_lo);
let hh = mul_u32_lanes(x_hi, y_hi);
let ll_hi = u64x2_shr(ll, 32);
let t0 = i64x2_add(hl, ll_hi);
let t0_lo = v128_and(t0, EPSILON_VEC);
let t0_hi = u64x2_shr(t0, 32);
let t1 = i64x2_add(lh, t0_lo);
let t2 = i64x2_add(hh, t0_hi);
let t1_hi = u64x2_shr(t1, 32);
let res_hi = i64x2_add(t2, t1_hi);
let ll_lo32 = v128_and(ll, EPSILON_VEC);
let t1_lo32 = v128_and(t1, EPSILON_VEC);
let t1_shifted = i64x2_shl(t1_lo32, 32);
let res_lo = v128_or(ll_lo32, t1_shifted);
(res_hi, res_lo)
}
#[inline(always)]
fn add_small_64s_64_s(x_s: v128, y: v128) -> v128 {
let res_wrapped_s = i64x2_add(x_s, y);
let mask = i64x2_gt(x_s, res_wrapped_s); let wrapback_amt = u64x2_shr(mask, 32); i64x2_add(res_wrapped_s, wrapback_amt)
}
#[inline(always)]
fn sub_small_64s_64_s(x_s: v128, y: v128) -> v128 {
let res_wrapped_s = i64x2_sub(x_s, y);
let mask = i64x2_gt(res_wrapped_s, x_s); let wrapback_amt = u64x2_shr(mask, 32);
i64x2_sub(res_wrapped_s, wrapback_amt)
}
#[inline]
fn reduce128(hi: v128, lo: v128) -> v128 {
let lo_s = shift(lo);
let hi_hi = u64x2_shr(hi, 32);
let lo1_s = sub_small_64s_64_s(lo_s, hi_hi);
let hi_lo32 = v128_and(hi, EPSILON_VEC);
let hi_lo32_shifted = i64x2_shl(hi_lo32, 32);
let t1 = i64x2_sub(hi_lo32_shifted, hi_lo32);
let lo2_s = add_small_64s_64_s(lo1_s, t1);
shift(lo2_s)
}
#[inline(always)]
fn unsigned_lt_as_carry(a: v128, b: v128) -> v128 {
let mask = i64x2_gt(shift(b), shift(a));
u64x2_shr(mask, 63)
}
#[inline]
fn dot_pairs<const N: usize>(get: impl Fn(usize) -> (v128, v128)) -> v128 {
const {
assert!((N as u32) <= (1 << 31));
}
let mut acc_lo_hi = u64x2_splat(0);
let mut acc_lo_lo = u64x2_splat(0);
let mut acc_hi96 = u64x2_splat(0);
for i in 0..N {
let (lhs, rhs) = get(i);
let (term_hi, term_lo) = mul64_64(lhs, rhs);
let term_hi96 = u64x2_shr(term_hi, 32);
let new_lo_lo = i64x2_add(acc_lo_lo, term_lo);
let carry = unsigned_lt_as_carry(new_lo_lo, acc_lo_lo);
acc_lo_hi = i64x2_add(i64x2_add(acc_lo_hi, term_hi), carry);
acc_lo_lo = new_lo_lo;
acc_hi96 = i64x2_add(acc_hi96, term_hi96);
}
let hi96_shifted = i64x2_shl(acc_hi96, 32);
let lo_hi = i64x2_sub(acc_lo_hi, hi96_shifted);
let lo_lo = acc_lo_lo;
let p_minus_hi = i64x2_sub(u64x2_splat(P), acc_hi96);
let sum_lo = i64x2_add(lo_lo, p_minus_hi);
let carry2 = unsigned_lt_as_carry(sum_lo, lo_lo);
let sum_hi = i64x2_add(lo_hi, carry2);
reduce128(sum_hi, sum_lo)
}
#[inline]
fn sum_delayed_reduce<const N: usize>(terms: &[v128; N]) -> v128 {
let mut acc_hi = u64x2_splat(0);
let mut acc_lo = u64x2_splat(0);
for &term in terms {
let new_lo = i64x2_add(acc_lo, term);
let carry = unsigned_lt_as_carry(new_lo, acc_lo);
acc_hi = i64x2_add(acc_hi, carry);
acc_lo = new_lo;
}
reduce128(acc_hi, acc_lo)
}
#[inline]
fn mul(x: v128, y: v128) -> v128 {
let (hi, lo) = mul64_64(x, y);
reduce128(hi, lo)
}
#[inline]
fn square64(x: v128) -> (v128, v128) {
let x_lo = lo32(x);
let x_hi = hi32(x);
let ll = mul_u32_lanes(x_lo, x_lo);
let lh = mul_u32_lanes(x_lo, x_hi);
let hh = mul_u32_lanes(x_hi, x_hi);
let ll_hi = u64x2_shr(ll, 33);
let t0 = i64x2_add(lh, ll_hi);
let t0_hi = u64x2_shr(t0, 31);
let res_hi = i64x2_add(hh, t0_hi);
let lh_shifted = i64x2_shl(lh, 33);
let res_lo = i64x2_add(ll, lh_shifted);
(res_hi, res_lo)
}
#[inline]
fn square(x: v128) -> v128 {
let (hi, lo) = square64(x);
reduce128(hi, lo)
}
#[inline(always)]
fn double(x: v128) -> v128 {
add(x, x)
}
#[cfg(test)]
mod tests {
use p3_field_testing::test_packed_field;
use super::{Goldilocks, PackedGoldilocksWasmSimd128, WIDTH};
const SPECIAL_VALS: [Goldilocks; WIDTH] =
Goldilocks::new_array([0xFFFF_FFFF_0000_0000, 0xFFFF_FFFF_FFFF_FFFF]);
const ZEROS: PackedGoldilocksWasmSimd128 =
PackedGoldilocksWasmSimd128(Goldilocks::new_array([
0x0000_0000_0000_0000,
0xFFFF_FFFF_0000_0001, ]));
const ONES: PackedGoldilocksWasmSimd128 = PackedGoldilocksWasmSimd128(Goldilocks::new_array([
0x0000_0000_0000_0001,
0xFFFF_FFFF_0000_0002, ]));
test_packed_field!(
crate::PackedGoldilocksWasmSimd128,
&[super::ZEROS],
&[super::ONES],
crate::PackedGoldilocksWasmSimd128(super::SPECIAL_VALS)
);
#[test]
fn sum_array_delayed_reduction_matches_scalar() {
use p3_field::{PackedValue, PrimeCharacteristicRing, PrimeField64};
use rand::rngs::SmallRng;
use rand::{RngExt, SeedableRng};
fn check<const N: usize>(terms0: [Goldilocks; N], terms1: [Goldilocks; N]) {
let packed: [PackedGoldilocksWasmSimd128; N] =
core::array::from_fn(|i| PackedGoldilocksWasmSimd128([terms0[i], terms1[i]]));
let expected0 = Goldilocks::sum_array::<N>(&terms0);
let expected1 = Goldilocks::sum_array::<N>(&terms1);
let actual = PackedGoldilocksWasmSimd128::sum_array::<N>(&packed);
assert_eq!(
actual.as_slice()[0].as_canonical_u64(),
expected0.as_canonical_u64(),
"N={N} mismatch at lane 0: terms={terms0:?}"
);
assert_eq!(
actual.as_slice()[1].as_canonical_u64(),
expected1.as_canonical_u64(),
"N={N} mismatch at lane 1: terms={terms1:?}"
);
}
macro_rules! check_edge_n {
($n:literal) => {
check::<$n>([Goldilocks::new(u64::MAX); $n], [Goldilocks::ZERO; $n]);
};
}
check::<2>([Goldilocks::new(u64::MAX); 2], [Goldilocks::ZERO; 2]);
check_edge_n!(3);
check_edge_n!(4);
check_edge_n!(5);
check_edge_n!(7);
check_edge_n!(8);
check_edge_n!(11);
check_edge_n!(12);
check_edge_n!(15);
check_edge_n!(16);
check_edge_n!(32);
let mut rng = SmallRng::seed_from_u64(0x005A_A0D1_CA7E);
macro_rules! check_random_n {
($n:literal, $count:literal) => {
for _ in 0..$count {
let terms0: [Goldilocks; $n] = core::array::from_fn(|_| rng.random());
let terms1: [Goldilocks; $n] = core::array::from_fn(|_| rng.random());
check::<$n>(terms0, terms1);
}
};
}
check_random_n!(3, 32);
check_random_n!(7, 32);
check_random_n!(11, 16);
check_random_n!(15, 16);
check_random_n!(64, 8);
}
#[test]
fn dot_product_delayed_reduction_matches_scalar() {
use p3_field::{PackedValue, PrimeCharacteristicRing, PrimeField64};
use rand::rngs::SmallRng;
use rand::{RngExt, SeedableRng};
const EDGE_VALUES: [u64; 5] = [
0,
1,
Goldilocks::ORDER_U64 - 1,
0xFFFF_FFFF_0000_0000, u64::MAX, ];
fn check<const N: usize>(
lhs0: [Goldilocks; N],
rhs0: [Goldilocks; N],
lhs1: [Goldilocks; N],
rhs1: [Goldilocks; N],
) {
let packed_lhs: [PackedGoldilocksWasmSimd128; N] =
core::array::from_fn(|i| PackedGoldilocksWasmSimd128([lhs0[i], lhs1[i]]));
let packed_rhs: [PackedGoldilocksWasmSimd128; N] =
core::array::from_fn(|i| PackedGoldilocksWasmSimd128([rhs0[i], rhs1[i]]));
let expected0 = Goldilocks::dot_product(&lhs0, &rhs0);
let expected1 = Goldilocks::dot_product(&lhs1, &rhs1);
let actual = PackedGoldilocksWasmSimd128::dot_product(&packed_lhs, &packed_rhs);
assert_eq!(
actual.as_slice()[0].as_canonical_u64(),
expected0.as_canonical_u64(),
"N={N} mismatch at lane 0: lhs={lhs0:?} rhs={rhs0:?}"
);
assert_eq!(
actual.as_slice()[1].as_canonical_u64(),
expected1.as_canonical_u64(),
"N={N} mismatch at lane 1: lhs={lhs1:?} rhs={rhs1:?}"
);
}
macro_rules! check_edge_n {
($n:literal) => {
check::<$n>(
[Goldilocks::new(u64::MAX); $n],
[Goldilocks::new(u64::MAX); $n],
[Goldilocks::ZERO; $n],
[Goldilocks::new(u64::MAX); $n],
);
};
}
check_edge_n!(2);
check_edge_n!(3);
check_edge_n!(4);
check_edge_n!(5);
check_edge_n!(8);
check_edge_n!(12);
check_edge_n!(16);
check_edge_n!(32);
for &a in &EDGE_VALUES {
for &b in &EDGE_VALUES {
for &c in &EDGE_VALUES {
check::<3>(
[Goldilocks::new(a), Goldilocks::new(b), Goldilocks::new(c)],
[Goldilocks::new(c), Goldilocks::new(b), Goldilocks::new(a)],
[Goldilocks::new(c), Goldilocks::new(b), Goldilocks::new(a)],
[Goldilocks::new(a), Goldilocks::new(b), Goldilocks::new(c)],
);
}
}
}
let mut rng = SmallRng::seed_from_u64(0x00D0_79A0_D7CE);
macro_rules! check_random_n {
($n:literal, $count:literal) => {
for _ in 0..$count {
let lhs0: [Goldilocks; $n] = core::array::from_fn(|_| rng.random());
let rhs0: [Goldilocks; $n] = core::array::from_fn(|_| rng.random());
let lhs1: [Goldilocks; $n] = core::array::from_fn(|_| rng.random());
let rhs1: [Goldilocks; $n] = core::array::from_fn(|_| rng.random());
check::<$n>(lhs0, rhs0, lhs1, rhs1);
}
};
}
check_random_n!(2, 32);
check_random_n!(3, 32);
check_random_n!(4, 32);
check_random_n!(7, 32);
check_random_n!(16, 16);
check_random_n!(64, 8);
}
#[test]
fn mixed_dot_product_delayed_reduction_matches_scalar() {
use p3_field::{Algebra, PackedValue, PrimeCharacteristicRing, PrimeField64};
use rand::rngs::SmallRng;
use rand::{RngExt, SeedableRng};
fn check<const N: usize>(a0: [Goldilocks; N], a1: [Goldilocks; N], f: [Goldilocks; N]) {
let packed_a: [PackedGoldilocksWasmSimd128; N] =
core::array::from_fn(|i| PackedGoldilocksWasmSimd128([a0[i], a1[i]]));
let expected0 = Goldilocks::dot_product(&a0, &f);
let expected1 = Goldilocks::dot_product(&a1, &f);
let actual = PackedGoldilocksWasmSimd128::mixed_dot_product(&packed_a, &f);
assert_eq!(
actual.as_slice()[0].as_canonical_u64(),
expected0.as_canonical_u64(),
"N={N} mismatch at lane 0"
);
assert_eq!(
actual.as_slice()[1].as_canonical_u64(),
expected1.as_canonical_u64(),
"N={N} mismatch at lane 1"
);
}
macro_rules! check_edge_n {
($n:literal) => {
check::<$n>(
[Goldilocks::new(u64::MAX); $n],
[Goldilocks::ZERO; $n],
[Goldilocks::new(u64::MAX); $n],
);
};
}
check_edge_n!(2);
check_edge_n!(5);
check_edge_n!(8);
check_edge_n!(16);
check_edge_n!(32);
let mut rng = SmallRng::seed_from_u64(0x011E_DD07_9A0D);
macro_rules! check_random_n {
($n:literal, $count:literal) => {
for _ in 0..$count {
let a0: [Goldilocks; $n] = core::array::from_fn(|_| rng.random());
let a1: [Goldilocks; $n] = core::array::from_fn(|_| rng.random());
let f: [Goldilocks; $n] = core::array::from_fn(|_| rng.random());
check::<$n>(a0, a1, f);
}
};
}
check_random_n!(2, 16);
check_random_n!(3, 16);
check_random_n!(8, 16);
check_random_n!(16, 8);
}
}