use alloc::vec::Vec;
use p3_field::{InjectiveMonomial, PrimeCharacteristicRing};
use p3_poseidon2::{
ExternalLayer, ExternalLayerConstants, ExternalLayerConstructor, InternalLayer,
InternalLayerConstructor, MDSMat4, mds_light_permutation,
};
use crate::poseidon1::GOLDILOCKS_S_BOX_DEGREE;
use crate::x86_64_avx512::packing::PackedGoldilocksAVX512;
use crate::{Goldilocks, Poseidon2ExternalLayerGoldilocks, Poseidon2InternalLayerGoldilocks};
#[inline(always)]
fn add_rc_and_sbox(val: &mut PackedGoldilocksAVX512, rc: PackedGoldilocksAVX512) {
*val = (*val + rc).injective_exp_n();
}
#[inline(always)]
fn sbox_array<const WIDTH: usize>(state: &mut [PackedGoldilocksAVX512; WIDTH]) {
let x2: [PackedGoldilocksAVX512; WIDTH] = core::array::from_fn(|i| state[i].square());
let mut x3 = x2;
let mut x4 = x2;
for i in 0..WIDTH {
x3[i] *= state[i];
x4[i] = x4[i].square();
}
for i in 0..WIDTH {
state[i] = x3[i] * x4[i];
}
}
#[inline(always)]
fn external_round<const WIDTH: usize>(
state: &mut [PackedGoldilocksAVX512; WIDTH],
rc: &[PackedGoldilocksAVX512; WIDTH],
) {
for i in 0..WIDTH {
state[i] += rc[i];
}
sbox_array(state);
mds_light_permutation(state, &MDSMat4);
}
#[inline(always)]
fn internal_round_goldilocks_8(
state: &mut [PackedGoldilocksAVX512; 8],
rc: PackedGoldilocksAVX512,
) {
let s1 = state[1];
let s2 = state[2];
let s3 = state[3];
let s4 = state[4];
let s5 = state[5];
let s6 = state[6];
let s7 = state[7];
let sum_tail = s1 + s2 + s3 + s4 + s5 + s6 + s7;
add_rc_and_sbox(&mut state[0], rc);
let s0 = state[0];
let sum = sum_tail + s0;
state[0] = sum - (s0 + s0);
state[1] = sum + s1;
state[2] = sum + (s2 + s2);
state[3] = sum + s3.halve();
state[4] = sum + (s4 + s4 + s4);
state[5] = sum - s5.halve();
state[6] = sum - (s6 + s6 + s6);
let two_s7 = s7 + s7;
state[7] = sum - (two_s7 + two_s7);
}
#[inline(always)]
fn internal_round_goldilocks_12(
state: &mut [PackedGoldilocksAVX512; 12],
rc: PackedGoldilocksAVX512,
) {
let s1 = state[1];
let s2 = state[2];
let s3 = state[3];
let s4 = state[4];
let s5 = state[5];
let s6 = state[6];
let s7 = state[7];
let s8 = state[8];
let s9 = state[9];
let s10 = state[10];
let s11 = state[11];
let sum_tail = s1 + s2 + s3 + s4 + s5 + s6 + s7 + s8 + s9 + s10 + s11;
add_rc_and_sbox(&mut state[0], rc);
let s0 = state[0];
let sum = sum_tail + s0;
state[0] = sum - (s0 + s0);
state[1] = sum + s1;
state[2] = sum + (s2 + s2);
state[3] = sum + s3.halve();
state[4] = sum + (s4 + s4 + s4);
let two_s5 = s5 + s5;
state[5] = sum + (two_s5 + two_s5);
state[6] = sum - s6.halve();
state[7] = sum - (s7 + s7 + s7);
let two_s8 = s8 + s8;
state[8] = sum - (two_s8 + two_s8);
state[9] = sum + s9.halve().halve();
state[10] = sum - s10.halve().halve();
state[11] = sum + s11.halve().halve().halve();
}
#[inline(always)]
fn internal_round_goldilocks_16(
state: &mut [PackedGoldilocksAVX512; 16],
rc: PackedGoldilocksAVX512,
) {
let s1 = state[1];
let s2 = state[2];
let s3 = state[3];
let s4 = state[4];
let s5 = state[5];
let s6 = state[6];
let s7 = state[7];
let s8 = state[8];
let s9 = state[9];
let s10 = state[10];
let s11 = state[11];
let s12 = state[12];
let s13 = state[13];
let s14 = state[14];
let s15 = state[15];
let sum_tail = s1 + s2 + s3 + s4 + s5 + s6 + s7 + s8 + s9 + s10 + s11 + s12 + s13 + s14 + s15;
add_rc_and_sbox(&mut state[0], rc);
let s0 = state[0];
let sum = sum_tail + s0;
state[0] = sum - (s0 + s0);
state[1] = sum + s1;
state[2] = sum + (s2 + s2);
state[3] = sum + s3.halve();
state[4] = sum + (s4 + s4 + s4);
let two_s5 = s5 + s5;
state[5] = sum + (two_s5 + two_s5);
state[6] = sum - s6.halve();
state[7] = sum - (s7 + s7 + s7);
let two_s8 = s8 + s8;
state[8] = sum - (two_s8 + two_s8);
state[9] = sum + s9.halve().halve().halve();
state[10] = sum + s10.halve().halve().halve().halve();
state[11] = sum + s11.halve().halve().halve().halve().halve();
state[12] = sum - s12.halve().halve().halve();
state[13] = sum - s13.halve().halve().halve().halve();
state[14] = sum - s14.halve().halve().halve().halve().halve();
let inv_2_32 = crate::MATRIX_DIAG_16_GOLDILOCKS[15];
state[15] = sum + s15 * inv_2_32;
}
#[derive(Clone, Debug, Default)]
pub struct Poseidon2InternalLayerGoldilocksAVX512 {
inner: Poseidon2InternalLayerGoldilocks,
packed_internal_constants: Vec<PackedGoldilocksAVX512>,
}
impl InternalLayerConstructor<Goldilocks> for Poseidon2InternalLayerGoldilocksAVX512 {
fn new_from_constants(internal_constants: Vec<Goldilocks>) -> Self {
let packed_internal_constants = internal_constants
.iter()
.copied()
.map(PackedGoldilocksAVX512::from)
.collect();
let inner = Poseidon2InternalLayerGoldilocks::new_from_constants(internal_constants);
Self {
inner,
packed_internal_constants,
}
}
}
impl InternalLayer<Goldilocks, 8, GOLDILOCKS_S_BOX_DEGREE>
for Poseidon2InternalLayerGoldilocksAVX512
{
fn permute_state(&self, state: &mut [Goldilocks; 8]) {
InternalLayer::<Goldilocks, 8, GOLDILOCKS_S_BOX_DEGREE>::permute_state(&self.inner, state);
}
}
impl InternalLayer<Goldilocks, 12, GOLDILOCKS_S_BOX_DEGREE>
for Poseidon2InternalLayerGoldilocksAVX512
{
fn permute_state(&self, state: &mut [Goldilocks; 12]) {
InternalLayer::<Goldilocks, 12, GOLDILOCKS_S_BOX_DEGREE>::permute_state(&self.inner, state);
}
}
impl InternalLayer<Goldilocks, 16, GOLDILOCKS_S_BOX_DEGREE>
for Poseidon2InternalLayerGoldilocksAVX512
{
fn permute_state(&self, state: &mut [Goldilocks; 16]) {
InternalLayer::<Goldilocks, 16, GOLDILOCKS_S_BOX_DEGREE>::permute_state(&self.inner, state);
}
}
impl InternalLayer<Goldilocks, 20, GOLDILOCKS_S_BOX_DEGREE>
for Poseidon2InternalLayerGoldilocksAVX512
{
fn permute_state(&self, state: &mut [Goldilocks; 20]) {
InternalLayer::<Goldilocks, 20, GOLDILOCKS_S_BOX_DEGREE>::permute_state(&self.inner, state);
}
}
impl InternalLayer<PackedGoldilocksAVX512, 20, GOLDILOCKS_S_BOX_DEGREE>
for Poseidon2InternalLayerGoldilocksAVX512
{
fn permute_state(&self, state: &mut [PackedGoldilocksAVX512; 20]) {
InternalLayer::<PackedGoldilocksAVX512, 20, GOLDILOCKS_S_BOX_DEGREE>::permute_state(
&self.inner,
state,
);
}
}
impl InternalLayer<PackedGoldilocksAVX512, 8, GOLDILOCKS_S_BOX_DEGREE>
for Poseidon2InternalLayerGoldilocksAVX512
{
fn permute_state(&self, state: &mut [PackedGoldilocksAVX512; 8]) {
for &rc in &self.packed_internal_constants {
internal_round_goldilocks_8(state, rc);
}
}
}
impl InternalLayer<PackedGoldilocksAVX512, 12, GOLDILOCKS_S_BOX_DEGREE>
for Poseidon2InternalLayerGoldilocksAVX512
{
fn permute_state(&self, state: &mut [PackedGoldilocksAVX512; 12]) {
for &rc in &self.packed_internal_constants {
internal_round_goldilocks_12(state, rc);
}
}
}
impl InternalLayer<PackedGoldilocksAVX512, 16, GOLDILOCKS_S_BOX_DEGREE>
for Poseidon2InternalLayerGoldilocksAVX512
{
fn permute_state(&self, state: &mut [PackedGoldilocksAVX512; 16]) {
for &rc in &self.packed_internal_constants {
internal_round_goldilocks_16(state, rc);
}
}
}
#[derive(Clone)]
pub struct Poseidon2ExternalLayerGoldilocksAVX512<const WIDTH: usize> {
inner: Poseidon2ExternalLayerGoldilocks<WIDTH>,
packed_initial_external_constants: Vec<[PackedGoldilocksAVX512; WIDTH]>,
packed_terminal_external_constants: Vec<[PackedGoldilocksAVX512; WIDTH]>,
}
impl<const WIDTH: usize> ExternalLayerConstructor<Goldilocks, WIDTH>
for Poseidon2ExternalLayerGoldilocksAVX512<WIDTH>
{
fn new_from_constants(external_constants: ExternalLayerConstants<Goldilocks, WIDTH>) -> Self {
let pack_round = |rc: &[Goldilocks; WIDTH]| rc.map(PackedGoldilocksAVX512::from);
let packed_initial_external_constants = external_constants
.get_initial_constants()
.iter()
.map(pack_round)
.collect();
let packed_terminal_external_constants = external_constants
.get_terminal_constants()
.iter()
.map(pack_round)
.collect();
let inner = Poseidon2ExternalLayerGoldilocks::new_from_constants(external_constants);
Self {
inner,
packed_initial_external_constants,
packed_terminal_external_constants,
}
}
}
impl<const WIDTH: usize> ExternalLayer<Goldilocks, WIDTH, GOLDILOCKS_S_BOX_DEGREE>
for Poseidon2ExternalLayerGoldilocksAVX512<WIDTH>
{
fn permute_state_initial(&self, state: &mut [Goldilocks; WIDTH]) {
ExternalLayer::<Goldilocks, WIDTH, GOLDILOCKS_S_BOX_DEGREE>::permute_state_initial(
&self.inner,
state,
);
}
fn permute_state_terminal(&self, state: &mut [Goldilocks; WIDTH]) {
ExternalLayer::<Goldilocks, WIDTH, GOLDILOCKS_S_BOX_DEGREE>::permute_state_terminal(
&self.inner,
state,
);
}
}
impl<const WIDTH: usize> ExternalLayer<PackedGoldilocksAVX512, WIDTH, GOLDILOCKS_S_BOX_DEGREE>
for Poseidon2ExternalLayerGoldilocksAVX512<WIDTH>
{
fn permute_state_initial(&self, state: &mut [PackedGoldilocksAVX512; WIDTH]) {
mds_light_permutation(state, &MDSMat4);
for rc in &self.packed_initial_external_constants {
external_round(state, rc);
}
}
fn permute_state_terminal(&self, state: &mut [PackedGoldilocksAVX512; WIDTH]) {
for rc in &self.packed_terminal_external_constants {
external_round(state, rc);
}
}
}
#[cfg(test)]
mod tests {
use p3_field::PackedValue;
use p3_symmetric::Permutation;
use rand::rngs::SmallRng;
use rand::{RngExt, SeedableRng};
use super::*;
use crate::{
default_goldilocks_poseidon2_8, default_goldilocks_poseidon2_12,
default_goldilocks_poseidon2_16,
};
const PACKING_WIDTH: usize = <PackedGoldilocksAVX512 as PackedValue>::WIDTH;
fn assert_packed_matches_scalar<const WIDTH: usize>(
scalar_perm: &impl Permutation<[Goldilocks; WIDTH]>,
packed_perm: &impl Permutation<[PackedGoldilocksAVX512; WIDTH]>,
rng: &mut SmallRng,
) {
let lanes: [[Goldilocks; WIDTH]; PACKING_WIDTH] = core::array::from_fn(|_| rng.random());
let mut packed_state: [PackedGoldilocksAVX512; WIDTH] =
core::array::from_fn(|i| PackedGoldilocksAVX512::from_fn(|l| lanes[l][i]));
packed_perm.permute_mut(&mut packed_state);
for (l, lane) in lanes.into_iter().enumerate() {
let mut scalar_state = lane;
scalar_perm.permute_mut(&mut scalar_state);
for i in 0..WIDTH {
assert_eq!(
scalar_state[i],
packed_state[i].as_slice()[l],
"width {WIDTH}, lane {l}, element {i}"
);
}
}
}
#[test]
fn packed_matches_scalar_width_8() {
let mut rng = SmallRng::seed_from_u64(1);
let perm = default_goldilocks_poseidon2_8();
assert_packed_matches_scalar(&perm, &perm, &mut rng);
}
#[test]
fn packed_matches_scalar_width_12() {
let mut rng = SmallRng::seed_from_u64(2);
let perm = default_goldilocks_poseidon2_12();
assert_packed_matches_scalar(&perm, &perm, &mut rng);
}
#[test]
fn packed_matches_scalar_width_16() {
let mut rng = SmallRng::seed_from_u64(3);
let perm = default_goldilocks_poseidon2_16();
assert_packed_matches_scalar(&perm, &perm, &mut rng);
}
#[test]
fn packed_matches_scalar_width_20() {
let mut rng = SmallRng::seed_from_u64(4);
let perm: p3_poseidon2::Poseidon2<
Goldilocks,
Poseidon2ExternalLayerGoldilocksAVX512<20>,
Poseidon2InternalLayerGoldilocksAVX512,
20,
GOLDILOCKS_S_BOX_DEGREE,
> = p3_poseidon2::Poseidon2::new_from_rng(8, 22, &mut rng);
assert_packed_matches_scalar(&perm, &perm, &mut rng);
}
}