openvm_deferral_circuit/extension/
cuda.rs

1use 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}