openvm_bigint_circuit/extension/
cuda.rs

1use 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
19// This implementation is specific to GpuBackend because the lookup chips
20// (VariableRangeCheckerChipGPU, BitwiseOperationLookupChipGPU) are specific to GpuBackend.
21impl 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}