openvm_rv32im_circuit/divrem/
cuda.rs

1use std::sync::Arc;
2
3use derive_new::new;
4use openvm_circuit::{arch::DenseRecordArena, utils::next_power_of_two_or_zero};
5use openvm_circuit_primitives::{
6    bitwise_op_lookup::BitwiseOperationLookupChipGPU, range_tuple::RangeTupleCheckerChipGPU,
7    var_range::VariableRangeCheckerChipGPU, Chip,
8};
9use openvm_cuda_backend::{base::DeviceMatrix, prelude::F, GpuBackend};
10use openvm_cuda_common::copy::MemCopyH2D;
11use openvm_instructions::riscv::{RV32_CELL_BITS, RV32_REGISTER_NUM_LIMBS};
12use openvm_stark_backend::prover::AirProvingContext;
13
14use crate::{
15    adapters::{Rv32MultAdapterCols, Rv32MultAdapterRecord},
16    cuda_abi::{divrem_cuda::tracegen, UInt2},
17    DivRemCoreCols, DivRemCoreRecord,
18};
19
20#[derive(new)]
21pub struct Rv32DivRemChipGpu {
22    pub range_checker: Arc<VariableRangeCheckerChipGPU>,
23    pub bitwise_lookup: Arc<BitwiseOperationLookupChipGPU<RV32_CELL_BITS>>,
24    pub range_tuple_checker: Arc<RangeTupleCheckerChipGPU<2>>,
25    pub pointer_max_bits: usize,
26    pub timestamp_max_bits: usize,
27}
28
29impl Chip<DenseRecordArena, GpuBackend> for Rv32DivRemChipGpu {
30    fn generate_proving_ctx(&self, arena: DenseRecordArena) -> AirProvingContext<GpuBackend> {
31        const RECORD_SIZE: usize = size_of::<(
32            Rv32MultAdapterRecord,
33            DivRemCoreRecord<RV32_REGISTER_NUM_LIMBS>,
34        )>();
35        let records = arena.allocated();
36        if records.is_empty() {
37            return AirProvingContext::simple_no_pis(DeviceMatrix::dummy());
38        }
39        debug_assert_eq!(records.len() % RECORD_SIZE, 0);
40
41        let trace_width = DivRemCoreCols::<F, RV32_REGISTER_NUM_LIMBS, RV32_CELL_BITS>::width()
42            + Rv32MultAdapterCols::<F>::width();
43        let height = records.len() / RECORD_SIZE;
44        let padded_height = next_power_of_two_or_zero(height);
45
46        let tuple_checker_sizes = self.range_tuple_checker.sizes;
47        let tuple_checker_sizes = UInt2::new(tuple_checker_sizes[0], tuple_checker_sizes[1]);
48        let device_ctx = &self.range_checker.device_ctx;
49
50        let d_records = tracing::info_span!("trace_gen.h2d_records")
51            .in_scope(|| records.to_device_on(device_ctx))
52            .unwrap();
53        let d_trace = DeviceMatrix::<F>::with_capacity_on(padded_height, trace_width, device_ctx);
54        unsafe {
55            tracegen(
56                d_trace.buffer(),
57                padded_height,
58                trace_width,
59                &d_records,
60                &self.range_checker.count,
61                &self.bitwise_lookup.count,
62                RV32_CELL_BITS as u32,
63                &self.range_tuple_checker.count,
64                tuple_checker_sizes,
65                self.timestamp_max_bits as u32,
66                device_ctx.stream.as_raw(),
67            )
68            .unwrap();
69        }
70
71        AirProvingContext::simple_no_pis(d_trace)
72    }
73}