Skip to main content

sp1_gpu_commit/
commit.rs

1use std::{iter::once, sync::Arc};
2
3use slop_algebra::AbstractField;
4use slop_challenger::IopCtx;
5use slop_jagged::JaggedProverData;
6use slop_symmetric::{CryptographicHasher, PseudoCompressionFunction as _};
7use sp1_gpu_basefold::{CudaStackedPcsProverData, FriCudaProver};
8use sp1_gpu_cudart::TaskScope;
9use sp1_gpu_merkle_tree::{CudaTcsProver, SingleLayerMerkleTreeProverError};
10use sp1_gpu_utils::{traces::JaggedTraceMle, Ext, Felt};
11
12/// TODO: document
13#[allow(clippy::type_complexity)]
14pub fn commit_multilinears<GC: IopCtx<F = Felt, EF = Ext>, P: CudaTcsProver<GC>>(
15    jagged_trace_mle: &JaggedTraceMle<Felt, TaskScope>,
16    max_log_row_count: u32,
17    use_preprocessed: bool,
18    drop_main_traces: bool,
19    basefold_prover: &FriCudaProver<GC, P, Felt>,
20) -> Result<
21    (GC::Digest, JaggedProverData<GC, CudaStackedPcsProverData<GC>>),
22    SingleLayerMerkleTreeProverError,
23> {
24    let (index, padding) = if use_preprocessed {
25        (
26            &jagged_trace_mle.dense().preprocessed_table_index,
27            jagged_trace_mle.dense().preprocessed_padding,
28        )
29    } else {
30        (&jagged_trace_mle.dense().main_table_index, jagged_trace_mle.dense().main_padding)
31    };
32    let (mut row_counts, mut column_counts) = (
33        index.values().map(|x| x.poly_size).collect::<Vec<_>>(),
34        index.values().map(|x| x.num_polys).collect::<Vec<_>>(),
35    );
36
37    let drop_traces = drop_main_traces && !use_preprocessed;
38
39    let (commitment, data) =
40        basefold_prover.encode_and_commit(use_preprocessed, drop_traces, jagged_trace_mle)?;
41
42    let num_added_cols = padding.div_ceil(1 << max_log_row_count).max(1);
43
44    row_counts.push(1 << max_log_row_count);
45    row_counts.push(padding - (num_added_cols - 1) * (1 << max_log_row_count));
46    column_counts.push(num_added_cols - 1);
47    column_counts.push(1);
48
49    let (hasher, compressor) = GC::default_hasher_and_compressor();
50
51    let hash = hasher.hash_iter(
52        once(Felt::from_canonical_u32(row_counts.len() as u32))
53            .chain(row_counts.clone().into_iter().map(|x| Felt::from_canonical_u32(x as u32)))
54            .chain(column_counts.clone().into_iter().map(|x| Felt::from_canonical_u32(x as u32))),
55    );
56
57    let final_commitment = compressor.compress([commitment, hash]);
58
59    let jagged_prover_data = JaggedProverData {
60        pcs_prover_data: data,
61        row_counts: Arc::new(row_counts),
62        column_counts: Arc::new(column_counts),
63        padding_column_count: num_added_cols,
64        original_commitment: commitment,
65    };
66
67    Ok((final_commitment, jagged_prover_data))
68}
69
70#[cfg(test)]
71mod tests {
72    use std::sync::Arc;
73
74    use serial_test::serial;
75    use slop_alloc::{CpuBackend, ToHost};
76    use slop_challenger::IopCtx;
77    use slop_futures::queue::WorkerQueue;
78    use slop_jagged::{JaggedPcsVerifier, JaggedProver};
79    use slop_merkle_tree::Poseidon2KoalaBear16Prover;
80    use slop_stacked::StackedPcsProver;
81    use sp1_core_machine::io::SP1Stdin;
82    use sp1_gpu_basefold::FriCudaProver;
83    use sp1_gpu_cudart::{run_in_place, PinnedBuffer};
84    use sp1_gpu_jagged_tracegen::test_utils::tracegen_setup::{
85        self, CORE_MAX_LOG_ROW_COUNT, LOG_STACKING_HEIGHT,
86    };
87    use sp1_gpu_jagged_tracegen::{full_tracegen, CORE_MAX_TRACE_SIZE};
88    use sp1_gpu_merkle_tree::{CudaTcsProver, Poseidon2SP1Field16CudaProver};
89    use sp1_gpu_utils::{Felt, TestGC};
90    use sp1_hypercube::prover::{DefaultTraceGenerator, ProverSemaphore, TraceGenerator};
91    use sp1_hypercube::{SP1InnerPcs, SP1PcsProofInner};
92    use sp1_primitives::fri_params::core_fri_config;
93
94    use crate::commit::commit_multilinears;
95    #[serial]
96    #[tokio::test]
97    async fn test_commit_matches() {
98        let (machine, record, program) =
99            tracegen_setup::setup(&test_artifacts::FIBONACCI_ELF, SP1Stdin::new()).await;
100
101        type JC = SP1InnerPcs;
102        type Prover = JaggedProver<
103            TestGC,
104            SP1PcsProofInner,
105            StackedPcsProver<Poseidon2KoalaBear16Prover, TestGC>,
106        >;
107
108        run_in_place(|scope| async move {
109            let semaphore = ProverSemaphore::new(1);
110            // Generate traces using the host tracegen.
111            let trace_generator = DefaultTraceGenerator::new_in(machine.clone(), CpuBackend);
112            let old_traces = trace_generator
113                .generate_traces(
114                    program.clone(),
115                    record.clone(),
116                    CORE_MAX_LOG_ROW_COUNT as usize,
117                    semaphore.clone(),
118                )
119                .await;
120
121            tracing::info!(
122                "warmup traces generated: {:?}",
123                old_traces.main_trace_data.shard_chips.len()
124            );
125
126            let num_rounds = 2;
127
128            let jagged_verifier = JaggedPcsVerifier::<_, JC>::new_from_basefold_params(
129                core_fri_config(),
130                LOG_STACKING_HEIGHT,
131                CORE_MAX_LOG_ROW_COUNT as usize,
132                num_rounds,
133            );
134
135            // Commit to preprocessed and main using the old prover.
136            let jagged_prover = Prover::from_verifier(&jagged_verifier);
137
138            let mut preprocessed_host_values = Vec::new();
139            for mle in old_traces.preprocessed_traces.values() {
140                let mle_host = mle.to_host().unwrap();
141                preprocessed_host_values.push(mle_host);
142            }
143
144            let mut main_host_values = Vec::new();
145            for mle in old_traces.main_trace_data.traces.values() {
146                let mle_host = mle.to_host().unwrap();
147                main_host_values.push(mle_host);
148            }
149
150            let preprocessed_message = preprocessed_host_values.into_iter().collect();
151            let main_message = main_host_values.into_iter().collect();
152
153            let (old_preprocessed_commitment, old_preprocessed_data) =
154                jagged_prover.commit_multilinears(preprocessed_message).ok().unwrap();
155            let (old_main_commitment, old_main_data) =
156                jagged_prover.commit_multilinears(main_message).ok().unwrap();
157
158            // Commit to preprocessed and main using the new prover.
159            // Do tracegen with the new setup.
160            let record = Arc::new(record);
161            let capacity = CORE_MAX_TRACE_SIZE as usize;
162            let buffer = PinnedBuffer::<Felt>::with_capacity(capacity);
163            let queue = Arc::new(WorkerQueue::new(vec![buffer]));
164            let buffer = queue.pop().await.unwrap();
165            let (_public_values, jagged_trace_data, _chip_set, _permit) = full_tracegen(
166                &machine,
167                program.clone(),
168                record.clone(),
169                &buffer,
170                CORE_MAX_TRACE_SIZE as usize,
171                LOG_STACKING_HEIGHT,
172                CORE_MAX_LOG_ROW_COUNT,
173                &scope,
174                ProverSemaphore::new(1),
175                false,
176            )
177            .await;
178
179            let tcs_prover = Poseidon2SP1Field16CudaProver::new(&scope);
180
181            let basefold_prover = FriCudaProver::<TestGC, _, <TestGC as IopCtx>::F>::new(
182                tcs_prover,
183                jagged_verifier.pcs_verifier.basefold_verifier.fri_config,
184                LOG_STACKING_HEIGHT,
185            );
186
187            let (new_preprocessed_commitment, new_preprocessed_data) =
188                commit_multilinears::<TestGC, _>(
189                    &jagged_trace_data,
190                    CORE_MAX_LOG_ROW_COUNT,
191                    true,
192                    false,
193                    &basefold_prover,
194                )
195                .unwrap();
196
197            let (new_main_commitment, new_main_data) = commit_multilinears::<TestGC, _>(
198                &jagged_trace_data,
199                CORE_MAX_LOG_ROW_COUNT,
200                false,
201                false,
202                &basefold_prover,
203            )
204            .unwrap();
205
206            assert_eq!(old_preprocessed_data.row_counts, new_preprocessed_data.row_counts);
207            assert_eq!(old_preprocessed_data.column_counts, new_preprocessed_data.column_counts);
208            assert_eq!(
209                old_preprocessed_data.padding_column_count,
210                new_preprocessed_data.padding_column_count
211            );
212            assert_eq!(old_main_data.row_counts, new_main_data.row_counts);
213            assert_eq!(old_main_data.column_counts, new_main_data.column_counts);
214            assert_eq!(old_main_data.padding_column_count, new_main_data.padding_column_count);
215            assert_eq!(
216                old_preprocessed_data.original_commitment,
217                new_preprocessed_data.original_commitment
218            );
219            assert_eq!(old_main_data.original_commitment, new_main_data.original_commitment);
220            assert_eq!(old_preprocessed_commitment, new_preprocessed_commitment);
221            assert_eq!(old_main_commitment, new_main_commitment);
222        })
223        .await;
224    }
225}