openvm_rv32im_circuit/hintstore/
cuda.rs

1use std::sync::Arc;
2
3use derive_new::new;
4use openvm_circuit::{
5    arch::{DenseRecordArena, RecordSeeker},
6    utils::next_power_of_two_or_zero,
7};
8use openvm_circuit_primitives::{
9    bitwise_op_lookup::BitwiseOperationLookupChipGPU, var_range::VariableRangeCheckerChipGPU, Chip,
10};
11use openvm_cuda_backend::{base::DeviceMatrix, prelude::F, GpuBackend};
12use openvm_cuda_common::copy::MemCopyH2D;
13use openvm_instructions::riscv::RV32_CELL_BITS;
14use openvm_stark_backend::prover::AirProvingContext;
15
16use crate::{
17    cuda_abi::hintstore_cuda::tracegen, Rv32HintStoreCols, Rv32HintStoreLayout,
18    Rv32HintStoreRecordMut,
19};
20
21#[derive(new)]
22pub struct Rv32HintStoreChipGpu {
23    pub range_checker: Arc<VariableRangeCheckerChipGPU>,
24    pub bitwise_lookup: Arc<BitwiseOperationLookupChipGPU<RV32_CELL_BITS>>,
25    pub pointer_max_bits: usize,
26    pub timestamp_max_bits: usize,
27}
28
29// This is the info needed by each row to do parallel tracegen
30#[repr(C)]
31#[derive(new)]
32pub struct OffsetInfo {
33    pub record_offset: u32,
34    pub local_idx: u32,
35}
36
37impl Chip<DenseRecordArena, GpuBackend> for Rv32HintStoreChipGpu {
38    fn generate_proving_ctx(&self, mut arena: DenseRecordArena) -> AirProvingContext<GpuBackend> {
39        let width = Rv32HintStoreCols::<u8>::width();
40        let records = arena.allocated_mut();
41        if records.is_empty() {
42            return AirProvingContext::simple_no_pis(DeviceMatrix::dummy());
43        }
44
45        let mut offsets = Vec::<OffsetInfo>::new();
46        let mut offset = 0;
47
48        while offset < records.len() {
49            let prev_offset = offset;
50            let record = RecordSeeker::<
51                DenseRecordArena,
52                Rv32HintStoreRecordMut,
53                Rv32HintStoreLayout,
54            >::get_record_at(&mut offset, records);
55            for idx in 0..record.inner.num_words {
56                offsets.push(OffsetInfo::new(prev_offset as u32, idx));
57            }
58        }
59        let device_ctx = &self.range_checker.device_ctx;
60
61        let d_records = tracing::info_span!("trace_gen.h2d_records")
62            .in_scope(|| records.to_device_on(device_ctx))
63            .unwrap();
64        let d_record_offsets = offsets.to_device_on(device_ctx).unwrap();
65
66        let trace_height = next_power_of_two_or_zero(offsets.len());
67        let d_trace = DeviceMatrix::<F>::with_capacity_on(trace_height, width, device_ctx);
68
69        unsafe {
70            tracegen(
71                d_trace.buffer(),
72                trace_height,
73                &d_records,
74                offsets.len(),
75                &d_record_offsets,
76                self.pointer_max_bits as u32,
77                &self.range_checker.count,
78                &self.bitwise_lookup.count,
79                RV32_CELL_BITS as u32,
80                self.timestamp_max_bits as u32,
81                device_ctx.stream.as_raw(),
82            )
83            .unwrap();
84        }
85
86        AirProvingContext::simple_no_pis(d_trace)
87    }
88}