openvm_bigint_circuit/extension/
cuda.rs1use openvm_circuit::{
2 arch::DenseRecordArena,
3 system::cuda::{
4 extensions::{
5 get_inventory_range_checker, get_or_create_bitwise_op_lookup, SystemGpuBuilder,
6 },
7 SystemChipInventoryGPU,
8 },
9};
10use openvm_circuit_primitives::range_tuple::RangeTupleCheckerChipGPU;
11use openvm_cuda_backend::{BabyBearPoseidon2GpuEngine as GpuBabyBearPoseidon2Engine, GpuBackend};
12use openvm_rv32im_circuit::Rv32ImGpuProverExt;
13use openvm_stark_sdk::config::baby_bear_poseidon2::BabyBearPoseidon2Config;
14
15use super::*;
16
17pub struct Int256GpuProverExt;
18
19impl VmProverExtension<GpuBabyBearPoseidon2Engine, DenseRecordArena, Int256>
22 for Int256GpuProverExt
23{
24 fn extend_prover(
25 &self,
26 extension: &Int256,
27 inventory: &mut ChipInventory<BabyBearPoseidon2Config, DenseRecordArena, GpuBackend>,
28 ) -> Result<(), ChipInventoryError> {
29 let pointer_max_bits = inventory.airs().pointer_max_bits();
30 let timestamp_max_bits = inventory.timestamp_max_bits();
31
32 let range_checker = get_inventory_range_checker(inventory);
33 let bitwise_lu = get_or_create_bitwise_op_lookup(inventory)?;
34
35 let range_tuple_checker = {
36 let existing_chip = inventory
37 .find_chip::<Arc<RangeTupleCheckerChipGPU<2>>>()
38 .find(|c| {
39 c.sizes[0] >= extension.range_tuple_checker_sizes[0]
40 && c.sizes[1] >= extension.range_tuple_checker_sizes[1]
41 });
42 if let Some(chip) = existing_chip {
43 chip.clone()
44 } else {
45 inventory.next_air::<RangeTupleCheckerAir<2>>()?;
46 let chip = Arc::new(RangeTupleCheckerChipGPU::new(
47 extension.range_tuple_checker_sizes,
48 range_checker.device_ctx.clone(),
49 ));
50 inventory.add_periphery_chip(chip.clone());
51 chip
52 }
53 };
54
55 inventory.next_air::<Rv32BaseAlu256Air>()?;
56 let base_alu = BaseAlu256ChipGpu::new(
57 range_checker.clone(),
58 bitwise_lu.clone(),
59 pointer_max_bits,
60 timestamp_max_bits,
61 );
62 inventory.add_executor_chip(base_alu);
63
64 inventory.next_air::<Rv32LessThan256Air>()?;
65 let lt = LessThan256ChipGpu::new(
66 range_checker.clone(),
67 bitwise_lu.clone(),
68 pointer_max_bits,
69 timestamp_max_bits,
70 );
71 inventory.add_executor_chip(lt);
72
73 inventory.next_air::<Rv32BranchEqual256Air>()?;
74 let beq = BranchEqual256ChipGpu::new(
75 range_checker.clone(),
76 bitwise_lu.clone(),
77 pointer_max_bits,
78 timestamp_max_bits,
79 );
80 inventory.add_executor_chip(beq);
81
82 inventory.next_air::<Rv32BranchLessThan256Air>()?;
83 let blt = BranchLessThan256ChipGpu::new(
84 range_checker.clone(),
85 bitwise_lu.clone(),
86 pointer_max_bits,
87 timestamp_max_bits,
88 );
89 inventory.add_executor_chip(blt);
90
91 inventory.next_air::<Rv32Multiplication256Air>()?;
92 let mult = Multiplication256ChipGpu::new(
93 range_checker.clone(),
94 bitwise_lu.clone(),
95 range_tuple_checker.clone(),
96 pointer_max_bits,
97 timestamp_max_bits,
98 );
99 inventory.add_executor_chip(mult);
100
101 inventory.next_air::<Rv32Shift256Air>()?;
102 let shift = Shift256ChipGpu::new(
103 range_checker.clone(),
104 bitwise_lu.clone(),
105 pointer_max_bits,
106 timestamp_max_bits,
107 );
108 inventory.add_executor_chip(shift);
109
110 Ok(())
111 }
112}
113
114#[derive(Clone)]
115pub struct Int256Rv32GpuBuilder;
116
117type E = GpuBabyBearPoseidon2Engine;
118
119impl VmBuilder<E> for Int256Rv32GpuBuilder {
120 type VmConfig = Int256Rv32Config;
121 type SystemChipInventory = SystemChipInventoryGPU;
122 type RecordArena = DenseRecordArena;
123
124 fn create_chip_complex(
125 &self,
126 config: &Int256Rv32Config,
127 circuit: AirInventory<<E as StarkEngine>::SC>,
128 device_ctx: &openvm_stark_backend::EngineDeviceCtx<E>,
129 ) -> Result<
130 VmChipComplex<
131 <E as StarkEngine>::SC,
132 Self::RecordArena,
133 <E as StarkEngine>::PB,
134 Self::SystemChipInventory,
135 >,
136 ChipInventoryError,
137 > {
138 let mut chip_complex = VmBuilder::<E>::create_chip_complex(
139 &SystemGpuBuilder,
140 &config.system,
141 circuit,
142 device_ctx,
143 )?;
144 let inventory = &mut chip_complex.inventory;
145 VmProverExtension::<E, _, _>::extend_prover(&Rv32ImGpuProverExt, &config.rv32i, inventory)?;
146 VmProverExtension::<E, _, _>::extend_prover(&Rv32ImGpuProverExt, &config.rv32m, inventory)?;
147 VmProverExtension::<E, _, _>::extend_prover(&Rv32ImGpuProverExt, &config.io, inventory)?;
148 VmProverExtension::<E, _, _>::extend_prover(
149 &Int256GpuProverExt,
150 &config.bigint,
151 inventory,
152 )?;
153 Ok(chip_complex)
154 }
155}