use std::borrow::BorrowMut;
use openvm_circuit_primitives::{utils::next_power_of_two_or_zero, Chip};
use openvm_cpu_backend::CpuBackend;
use openvm_stark_backend::{
p3_air::BaseAir, p3_field::PrimeCharacteristicRing, p3_matrix::dense::RowMajorMatrix,
p3_maybe_rayon::prelude::*, prover::AirProvingContext, StarkProtocolConfig, Val,
};
use super::{columns::*, Poseidon2PeripheryBaseChip, PERIPHERY_POSEIDON2_WIDTH};
use crate::arch::VmField;
impl<RA, SC: StarkProtocolConfig, const SBOX_REGISTERS: usize> Chip<RA, CpuBackend<SC>>
for Poseidon2PeripheryBaseChip<Val<SC>, SBOX_REGISTERS>
where
Val<SC>: VmField,
{
fn generate_proving_ctx(&self, _: RA) -> AirProvingContext<CpuBackend<SC>> {
let width = Poseidon2PeripheryCols::<Val<SC>, SBOX_REGISTERS>::width();
if !self.nonempty.load(std::sync::atomic::Ordering::Relaxed) {
let trace = RowMajorMatrix::new(vec![], width);
return AirProvingContext::simple_no_pis(trace);
}
let height = next_power_of_two_or_zero(self.records.len());
let mut inputs = Vec::with_capacity(height);
let mut multiplicities = Vec::with_capacity(height);
#[cfg(feature = "parallel")]
let records_iter = self.records.par_iter();
#[cfg(not(feature = "parallel"))]
let records_iter = self.records.iter();
let (actual_inputs, actual_multiplicities): (Vec<_>, Vec<_>) = records_iter
.map(|r| {
let (input, mult) = r.pair();
(*input, mult.load(std::sync::atomic::Ordering::Relaxed))
})
.unzip();
inputs.extend(actual_inputs);
multiplicities.extend(actual_multiplicities);
inputs.resize(height, [Val::<SC>::ZERO; PERIPHERY_POSEIDON2_WIDTH]);
multiplicities.resize(height, 0);
let inner_trace = self.subchip.generate_trace(inputs);
let inner_width = self.subchip.air.width();
let mut values = Val::<SC>::zero_vec(height * width);
values
.par_chunks_mut(width)
.zip(inner_trace.values.par_chunks(inner_width))
.zip(multiplicities)
.for_each(|((row, inner_row), mult)| {
row[..inner_width].copy_from_slice(inner_row);
let cols: &mut Poseidon2PeripheryCols<Val<SC>, SBOX_REGISTERS> = row.borrow_mut();
cols.mult = Val::<SC>::from_u32(mult);
});
self.records.clear();
self.nonempty
.store(false, std::sync::atomic::Ordering::Relaxed);
let trace = RowMajorMatrix::new(values, width);
AirProvingContext::simple_no_pis(trace)
}
}