openvm_rv32im_circuit/hintstore/
cuda.rs1use std::sync::Arc;
2
3use derive_new::new;
4use openvm_circuit::{
5 arch::{DenseRecordArena, RecordSeeker},
6 utils::next_power_of_two_or_zero,
7};
8use openvm_circuit_primitives::{
9 bitwise_op_lookup::BitwiseOperationLookupChipGPU, var_range::VariableRangeCheckerChipGPU, Chip,
10};
11use openvm_cuda_backend::{base::DeviceMatrix, prelude::F, GpuBackend};
12use openvm_cuda_common::copy::MemCopyH2D;
13use openvm_instructions::riscv::RV32_CELL_BITS;
14use openvm_stark_backend::prover::AirProvingContext;
15
16use crate::{
17 cuda_abi::hintstore_cuda::tracegen, Rv32HintStoreCols, Rv32HintStoreLayout,
18 Rv32HintStoreRecordMut,
19};
20
21#[derive(new)]
22pub struct Rv32HintStoreChipGpu {
23 pub range_checker: Arc<VariableRangeCheckerChipGPU>,
24 pub bitwise_lookup: Arc<BitwiseOperationLookupChipGPU<RV32_CELL_BITS>>,
25 pub pointer_max_bits: usize,
26 pub timestamp_max_bits: usize,
27}
28
29#[repr(C)]
31#[derive(new)]
32pub struct OffsetInfo {
33 pub record_offset: u32,
34 pub local_idx: u32,
35}
36
37impl Chip<DenseRecordArena, GpuBackend> for Rv32HintStoreChipGpu {
38 fn generate_proving_ctx(&self, mut arena: DenseRecordArena) -> AirProvingContext<GpuBackend> {
39 let width = Rv32HintStoreCols::<u8>::width();
40 let records = arena.allocated_mut();
41 if records.is_empty() {
42 return AirProvingContext::simple_no_pis(DeviceMatrix::dummy());
43 }
44
45 let mut offsets = Vec::<OffsetInfo>::new();
46 let mut offset = 0;
47
48 while offset < records.len() {
49 let prev_offset = offset;
50 let record = RecordSeeker::<
51 DenseRecordArena,
52 Rv32HintStoreRecordMut,
53 Rv32HintStoreLayout,
54 >::get_record_at(&mut offset, records);
55 for idx in 0..record.inner.num_words {
56 offsets.push(OffsetInfo::new(prev_offset as u32, idx));
57 }
58 }
59 let device_ctx = &self.range_checker.device_ctx;
60
61 let d_records = tracing::info_span!("trace_gen.h2d_records")
62 .in_scope(|| records.to_device_on(device_ctx))
63 .unwrap();
64 let d_record_offsets = offsets.to_device_on(device_ctx).unwrap();
65
66 let trace_height = next_power_of_two_or_zero(offsets.len());
67 let d_trace = DeviceMatrix::<F>::with_capacity_on(trace_height, width, device_ctx);
68
69 unsafe {
70 tracegen(
71 d_trace.buffer(),
72 trace_height,
73 &d_records,
74 offsets.len(),
75 &d_record_offsets,
76 self.pointer_max_bits as u32,
77 &self.range_checker.count,
78 &self.bitwise_lookup.count,
79 RV32_CELL_BITS as u32,
80 self.timestamp_max_bits as u32,
81 device_ctx.stream.as_raw(),
82 )
83 .unwrap();
84 }
85
86 AirProvingContext::simple_no_pis(d_trace)
87 }
88}