openvm_rv32im_circuit/branch_eq/
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::{var_range::VariableRangeCheckerChipGPU, Chip};
6use openvm_cuda_backend::{base::DeviceMatrix, prelude::F, GpuBackend};
7use openvm_cuda_common::copy::MemCopyH2D;
8use openvm_stark_backend::prover::AirProvingContext;
9
10use crate::{
11    adapters::{Rv32BranchAdapterCols, Rv32BranchAdapterRecord, RV32_REGISTER_NUM_LIMBS},
12    cuda_abi::beq_cuda::tracegen,
13    BranchEqualCoreCols, BranchEqualCoreRecord,
14};
15
16#[derive(new)]
17pub struct Rv32BranchEqualChipGpu {
18    pub range_checker: Arc<VariableRangeCheckerChipGPU>,
19    pub timestamp_max_bits: usize,
20}
21
22impl Chip<DenseRecordArena, GpuBackend> for Rv32BranchEqualChipGpu {
23    fn generate_proving_ctx(&self, arena: DenseRecordArena) -> AirProvingContext<GpuBackend> {
24        const RECORD_SIZE: usize = size_of::<(
25            Rv32BranchAdapterRecord,
26            BranchEqualCoreRecord<RV32_REGISTER_NUM_LIMBS>,
27        )>();
28        let records = arena.allocated();
29        if records.is_empty() {
30            return AirProvingContext::simple_no_pis(DeviceMatrix::dummy());
31        }
32        debug_assert_eq!(records.len() % RECORD_SIZE, 0);
33
34        let trace_width = BranchEqualCoreCols::<F, RV32_REGISTER_NUM_LIMBS>::width()
35            + Rv32BranchAdapterCols::<F>::width();
36        let trace_height = next_power_of_two_or_zero(records.len() / RECORD_SIZE);
37        let device_ctx = &self.range_checker.device_ctx;
38
39        let d_records = tracing::info_span!("trace_gen.h2d_records")
40            .in_scope(|| records.to_device_on(device_ctx))
41            .unwrap();
42        let d_trace = DeviceMatrix::<F>::with_capacity_on(trace_height, trace_width, device_ctx);
43
44        unsafe {
45            tracegen(
46                d_trace.buffer(),
47                trace_height,
48                &d_records,
49                &self.range_checker.count,
50                self.timestamp_max_bits as u32,
51                device_ctx.stream.as_raw(),
52            )
53            .unwrap();
54        }
55        AirProvingContext::simple_no_pis(d_trace)
56    }
57}