openvm_rv32im_circuit/jalr/
cuda.rs1use 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 bitwise_op_lookup::BitwiseOperationLookupChipGPU, 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::{Rv32JalrAdapterCols, Rv32JalrAdapterRecord, RV32_CELL_BITS},
14 cuda_abi::jalr_cuda::tracegen,
15 Rv32JalrCoreCols, Rv32JalrCoreRecord,
16};
17#[derive(new)]
18pub struct Rv32JalrChipGpu {
19 pub range_checker: Arc<VariableRangeCheckerChipGPU>,
20 pub bitwise_lookup: Arc<BitwiseOperationLookupChipGPU<RV32_CELL_BITS>>,
21 pub timestamp_max_bits: usize,
22}
23
24impl Chip<DenseRecordArena, GpuBackend> for Rv32JalrChipGpu {
25 fn generate_proving_ctx(&self, arena: DenseRecordArena) -> AirProvingContext<GpuBackend> {
26 const RECORD_SIZE: usize = size_of::<(Rv32JalrAdapterRecord, Rv32JalrCoreRecord)>();
27 let records = arena.allocated();
28 if records.is_empty() {
29 return AirProvingContext::simple_no_pis(DeviceMatrix::dummy());
30 }
31 debug_assert_eq!(records.len() % RECORD_SIZE, 0);
32
33 let trace_width = Rv32JalrCoreCols::<F>::width() + Rv32JalrAdapterCols::<F>::width();
34 let trace_height = next_power_of_two_or_zero(records.len() / RECORD_SIZE);
35 let device_ctx = &self.range_checker.device_ctx;
36
37 let d_records = tracing::info_span!("trace_gen.h2d_records")
38 .in_scope(|| records.to_device_on(device_ctx))
39 .unwrap();
40 let d_trace = DeviceMatrix::<F>::with_capacity_on(trace_height, trace_width, device_ctx);
41
42 unsafe {
43 tracegen(
44 d_trace.buffer(),
45 trace_height,
46 &d_records,
47 &self.range_checker.count,
48 &self.bitwise_lookup.count,
49 RV32_CELL_BITS,
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}