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