openvm_deferral_circuit/extension/
cuda.rs1use std::sync::Arc;
2
3use openvm_circuit::{
4 arch::{
5 AirInventory, ChipInventory, ChipInventoryError, DenseRecordArena, VmBuilder,
6 VmChipComplex, VmProverExtension,
7 },
8 system::cuda::{
9 extensions::{
10 get_inventory_range_checker, get_or_create_bitwise_op_lookup, SystemGpuBuilder,
11 },
12 SystemChipInventoryGPU,
13 },
14};
15use openvm_cuda_backend::{
16 prelude::F as CudaF, BabyBearPoseidon2GpuEngine as GpuBabyBearPoseidon2Engine, GpuBackend,
17};
18use openvm_cuda_common::d_buffer::DeviceBuffer;
19use openvm_rv32im_circuit::Rv32ImGpuProverExt;
20use openvm_stark_sdk::config::baby_bear_poseidon2::BabyBearPoseidon2Config;
21
22use crate::{
23 call::{DeferralCallAir, DeferralCallChipGpu},
24 count::{DeferralCircuitCountAir, DeferralCircuitCountChipGpu},
25 output::{DeferralOutputAir, DeferralOutputChipGpu},
26 poseidon2::{DeferralPoseidon2Air, DeferralPoseidon2ChipGpu},
27 DeferralExtension, Rv32DeferralConfig,
28};
29
30pub struct DeferralGpuProverExt;
31
32const DEFAULT_DEFERRAL_POSEIDON2_MAX_TRACE_HEIGHT: usize = 1 << 24;
33
34impl VmProverExtension<GpuBabyBearPoseidon2Engine, DenseRecordArena, DeferralExtension>
35 for DeferralGpuProverExt
36{
37 fn extend_prover(
38 &self,
39 extension: &DeferralExtension,
40 inventory: &mut ChipInventory<BabyBearPoseidon2Config, DenseRecordArena, GpuBackend>,
41 ) -> Result<(), ChipInventoryError> {
42 let num_deferral_circuits = extension.fns.len();
43 let address_bits = inventory.airs().pointer_max_bits();
44 let timestamp_max_bits = inventory.timestamp_max_bits();
45
46 let range_checker = get_inventory_range_checker(inventory);
47 let bitwise_lu = get_or_create_bitwise_op_lookup(inventory)?;
48
49 let count = Arc::new(if num_deferral_circuits == 0 {
50 DeviceBuffer::<u32>::new()
51 } else {
52 DeviceBuffer::<u32>::with_capacity_on(num_deferral_circuits, &range_checker.device_ctx)
53 });
54 if num_deferral_circuits > 0 {
55 count.fill_zero_on(&range_checker.device_ctx).unwrap();
56 }
57
58 inventory.next_air::<DeferralCircuitCountAir>()?;
59 let count_chip = Arc::new(DeferralCircuitCountChipGpu::new(
60 count.clone(),
61 num_deferral_circuits,
62 range_checker.device_ctx.clone(),
63 ));
64 inventory.add_periphery_chip(count_chip);
65
66 inventory.next_air::<DeferralPoseidon2Air<CudaF>>()?;
67 let poseidon2_chip = Arc::new(DeferralPoseidon2ChipGpu::new(
68 DEFAULT_DEFERRAL_POSEIDON2_MAX_TRACE_HEIGHT,
69 1,
70 range_checker.device_ctx.clone(),
71 ));
72 let poseidon2_shared = poseidon2_chip.shared_buffer();
73 inventory.add_periphery_chip(poseidon2_chip);
74
75 inventory.next_air::<DeferralCallAir>()?;
76 let call_chip = DeferralCallChipGpu::new(
77 range_checker.clone(),
78 bitwise_lu.clone(),
79 address_bits,
80 timestamp_max_bits,
81 count.clone(),
82 num_deferral_circuits,
83 poseidon2_shared.clone(),
84 );
85 inventory.add_executor_chip(call_chip);
86
87 inventory.next_air::<DeferralOutputAir>()?;
88 let output_chip = DeferralOutputChipGpu::new(
89 range_checker,
90 bitwise_lu,
91 address_bits,
92 timestamp_max_bits,
93 count,
94 num_deferral_circuits,
95 poseidon2_shared,
96 );
97 inventory.add_executor_chip(output_chip);
98
99 Ok(())
100 }
101}
102
103#[derive(Clone)]
104pub struct Rv32DeferralGpuBuilder;
105
106impl VmBuilder<GpuBabyBearPoseidon2Engine> for Rv32DeferralGpuBuilder {
107 type VmConfig = Rv32DeferralConfig;
108 type SystemChipInventory = SystemChipInventoryGPU;
109 type RecordArena = DenseRecordArena;
110
111 fn create_chip_complex(
112 &self,
113 config: &Self::VmConfig,
114 circuit: AirInventory<BabyBearPoseidon2Config>,
115 device_ctx: &openvm_stark_backend::EngineDeviceCtx<GpuBabyBearPoseidon2Engine>,
116 ) -> Result<
117 VmChipComplex<
118 BabyBearPoseidon2Config,
119 Self::RecordArena,
120 GpuBackend,
121 Self::SystemChipInventory,
122 >,
123 ChipInventoryError,
124 > {
125 let mut chip_complex = VmBuilder::<GpuBabyBearPoseidon2Engine>::create_chip_complex(
126 &SystemGpuBuilder,
127 &config.system,
128 circuit,
129 device_ctx,
130 )?;
131 let inventory = &mut chip_complex.inventory;
132 VmProverExtension::<GpuBabyBearPoseidon2Engine, _, _>::extend_prover(
133 &Rv32ImGpuProverExt,
134 &config.rv32i,
135 inventory,
136 )?;
137 VmProverExtension::<GpuBabyBearPoseidon2Engine, _, _>::extend_prover(
138 &Rv32ImGpuProverExt,
139 &config.rv32m,
140 inventory,
141 )?;
142 VmProverExtension::<GpuBabyBearPoseidon2Engine, _, _>::extend_prover(
143 &Rv32ImGpuProverExt,
144 &config.io,
145 inventory,
146 )?;
147 VmProverExtension::<GpuBabyBearPoseidon2Engine, _, _>::extend_prover(
148 &DeferralGpuProverExt,
149 &config.deferral,
150 inventory,
151 )?;
152 Ok(chip_complex)
153 }
154}