openvm_circuit/system/poseidon2/
trace.rs1use std::borrow::BorrowMut;
2
3use openvm_circuit_primitives::{utils::next_power_of_two_or_zero, Chip};
4use openvm_cpu_backend::CpuBackend;
5use openvm_stark_backend::{
6 p3_air::BaseAir, p3_field::PrimeCharacteristicRing, p3_matrix::dense::RowMajorMatrix,
7 p3_maybe_rayon::prelude::*, prover::AirProvingContext, StarkProtocolConfig, Val,
8};
9
10use super::{columns::*, Poseidon2PeripheryBaseChip, PERIPHERY_POSEIDON2_WIDTH};
11use crate::arch::VmField;
12
13impl<RA, SC: StarkProtocolConfig, const SBOX_REGISTERS: usize> Chip<RA, CpuBackend<SC>>
14 for Poseidon2PeripheryBaseChip<Val<SC>, SBOX_REGISTERS>
15where
16 Val<SC>: VmField,
17{
18 fn generate_proving_ctx(&self, _: RA) -> AirProvingContext<CpuBackend<SC>> {
20 let width = Poseidon2PeripheryCols::<Val<SC>, SBOX_REGISTERS>::width();
21 if !self.nonempty.load(std::sync::atomic::Ordering::Relaxed) {
22 let trace = RowMajorMatrix::new(vec![], width);
23 return AirProvingContext::simple_no_pis(trace);
24 }
25 let height = next_power_of_two_or_zero(self.records.len());
26
27 let mut inputs = Vec::with_capacity(height);
28 let mut multiplicities = Vec::with_capacity(height);
29 #[cfg(feature = "parallel")]
30 let records_iter = self.records.par_iter();
31 #[cfg(not(feature = "parallel"))]
32 let records_iter = self.records.iter();
33 let (actual_inputs, actual_multiplicities): (Vec<_>, Vec<_>) = records_iter
34 .map(|r| {
35 let (input, mult) = r.pair();
36 (*input, mult.load(std::sync::atomic::Ordering::Relaxed))
37 })
38 .unzip();
39 inputs.extend(actual_inputs);
40 multiplicities.extend(actual_multiplicities);
41 inputs.resize(height, [Val::<SC>::ZERO; PERIPHERY_POSEIDON2_WIDTH]);
42 multiplicities.resize(height, 0);
43
44 let inner_trace = self.subchip.generate_trace(inputs);
46 let inner_width = self.subchip.air.width();
47
48 let mut values = Val::<SC>::zero_vec(height * width);
49 values
50 .par_chunks_mut(width)
51 .zip(inner_trace.values.par_chunks(inner_width))
52 .zip(multiplicities)
53 .for_each(|((row, inner_row), mult)| {
54 row[..inner_width].copy_from_slice(inner_row);
56 let cols: &mut Poseidon2PeripheryCols<Val<SC>, SBOX_REGISTERS> = row.borrow_mut();
57 cols.mult = Val::<SC>::from_u32(mult);
58 });
59 self.records.clear();
60 self.nonempty
61 .store(false, std::sync::atomic::Ordering::Relaxed);
62
63 let trace = RowMajorMatrix::new(values, width);
64 AirProvingContext::simple_no_pis(trace)
65 }
66}