openvm_rv32im_circuit/mulh/
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, range_tuple::RangeTupleCheckerChipGPU,
7 var_range::VariableRangeCheckerChipGPU, Chip,
8};
9use openvm_cuda_backend::{base::DeviceMatrix, prelude::F, GpuBackend};
10use openvm_cuda_common::copy::MemCopyH2D;
11use openvm_stark_backend::prover::AirProvingContext;
12
13use crate::{
14 adapters::{
15 Rv32MultAdapterCols, Rv32MultAdapterRecord, RV32_CELL_BITS, RV32_REGISTER_NUM_LIMBS,
16 },
17 cuda_abi::{mulh_cuda::tracegen, UInt2},
18 MulHCoreCols, MulHCoreRecord,
19};
20
21#[derive(new)]
22pub struct Rv32MulHChipGpu {
23 pub range_checker: Arc<VariableRangeCheckerChipGPU>,
24 pub bitwise_lookup: Arc<BitwiseOperationLookupChipGPU<RV32_CELL_BITS>>,
25 pub range_tuple_checker: Arc<RangeTupleCheckerChipGPU<2>>,
26 pub timestamp_max_bits: usize,
27}
28
29impl Chip<DenseRecordArena, GpuBackend> for Rv32MulHChipGpu {
30 fn generate_proving_ctx(&self, arena: DenseRecordArena) -> AirProvingContext<GpuBackend> {
31 const RECORD_SIZE: usize = size_of::<(
32 Rv32MultAdapterRecord,
33 MulHCoreRecord<RV32_REGISTER_NUM_LIMBS, RV32_CELL_BITS>,
34 )>();
35 let records = arena.allocated();
36 if records.is_empty() {
37 return AirProvingContext::simple_no_pis(DeviceMatrix::dummy());
38 }
39 debug_assert_eq!(records.len() % RECORD_SIZE, 0);
40
41 let trace_width = MulHCoreCols::<F, RV32_REGISTER_NUM_LIMBS, RV32_CELL_BITS>::width()
42 + Rv32MultAdapterCols::<F>::width();
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.bitwise_lookup.count,
61 RV32_CELL_BITS,
62 &self.range_tuple_checker.count,
63 tuple_checker_sizes,
64 self.timestamp_max_bits as u32,
65 device_ctx.stream.as_raw(),
66 )
67 .unwrap();
68 }
69
70 AirProvingContext::simple_no_pis(d_trace)
71 }
72}