openvm_rv32im_circuit/mul/
cuda.rs

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