openvm_keccak256_circuit/extension/
cuda.rs

1use std::sync::{Arc, Mutex};
2
3use openvm_circuit::{
4    arch::DenseRecordArena,
5    system::cuda::{
6        extensions::{
7            get_inventory_range_checker, get_or_create_bitwise_op_lookup, SystemGpuBuilder,
8        },
9        SystemChipInventoryGPU,
10    },
11};
12use openvm_cuda_backend::{BabyBearPoseidon2GpuEngine as GpuBabyBearPoseidon2Engine, GpuBackend};
13use openvm_rv32im_circuit::Rv32ImGpuProverExt;
14use openvm_stark_sdk::config::baby_bear_poseidon2::BabyBearPoseidon2Config;
15
16use super::*;
17use crate::{
18    cuda::{KeccakfOpChipGpu, KeccakfPermChipGpu, SharedKeccakfRecords, XorinVmChipGpu},
19    keccakf_perm::KeccakfPermAir,
20};
21
22pub struct Keccak256GpuProverExt;
23
24impl VmProverExtension<GpuBabyBearPoseidon2Engine, DenseRecordArena, Keccak256>
25    for Keccak256GpuProverExt
26{
27    fn extend_prover(
28        &self,
29        _extension: &Keccak256,
30        inventory: &mut ChipInventory<BabyBearPoseidon2Config, DenseRecordArena, GpuBackend>,
31    ) -> Result<(), ChipInventoryError> {
32        let pointer_max_bits = inventory.airs().pointer_max_bits();
33        let timestamp_max_bits = inventory.timestamp_max_bits();
34
35        let range_checker = get_inventory_range_checker(inventory);
36        let bitwise_lu = get_or_create_bitwise_op_lookup(inventory)?;
37
38        // XorinVmChip
39        inventory.next_air::<XorinVmAir>()?;
40        let xorin_chip = XorinVmChipGpu::new(
41            range_checker.clone(),
42            bitwise_lu.clone(),
43            pointer_max_bits,
44            timestamp_max_bits as u32,
45        );
46        inventory.add_executor_chip(xorin_chip);
47
48        // Create shared state for passing records between Op and Perm chips
49        let shared_records = Arc::new(Mutex::new(SharedKeccakfRecords::default()));
50
51        // NOTE: AIRs are added in extend_circuit in this order: XorinVmAir, KeccakfPermAir,
52        // KeccakfOpAir The prover extension must consume AIRs in the same order.
53
54        // Register KeccakfPermChip (periphery chip - added BEFORE OpChip to ensure OpChip tracegen
55        // runs first)
56        inventory.next_air::<KeccakfPermAir>()?;
57        let perm_chip =
58            KeccakfPermChipGpu::new(shared_records.clone(), range_checker.device_ctx.clone());
59        inventory.add_periphery_chip(perm_chip);
60
61        // Register KeccakfOpChip (executor chip - generates first due to executor vs periphery
62        // ordering)
63        inventory.next_air::<KeccakfOpAir>()?;
64        let op_chip = KeccakfOpChipGpu::new(
65            range_checker,
66            bitwise_lu,
67            pointer_max_bits,
68            timestamp_max_bits as u32,
69            shared_records,
70        );
71        inventory.add_executor_chip(op_chip);
72
73        Ok(())
74    }
75}
76
77#[derive(Clone)]
78pub struct Keccak256Rv32GpuBuilder;
79
80type E = GpuBabyBearPoseidon2Engine;
81
82impl VmBuilder<E> for Keccak256Rv32GpuBuilder {
83    type VmConfig = Keccak256Rv32Config;
84    type SystemChipInventory = SystemChipInventoryGPU;
85    type RecordArena = DenseRecordArena;
86
87    fn create_chip_complex(
88        &self,
89        config: &Keccak256Rv32Config,
90        circuit: AirInventory<<E as StarkEngine>::SC>,
91        device_ctx: &openvm_stark_backend::EngineDeviceCtx<E>,
92    ) -> Result<
93        VmChipComplex<
94            <E as StarkEngine>::SC,
95            Self::RecordArena,
96            <E as StarkEngine>::PB,
97            Self::SystemChipInventory,
98        >,
99        ChipInventoryError,
100    > {
101        let mut chip_complex = VmBuilder::<E>::create_chip_complex(
102            &SystemGpuBuilder,
103            &config.system,
104            circuit,
105            device_ctx,
106        )?;
107        let inventory = &mut chip_complex.inventory;
108        VmProverExtension::<E, _, _>::extend_prover(&Rv32ImGpuProverExt, &config.rv32i, inventory)?;
109        VmProverExtension::<E, _, _>::extend_prover(&Rv32ImGpuProverExt, &config.rv32m, inventory)?;
110        VmProverExtension::<E, _, _>::extend_prover(&Rv32ImGpuProverExt, &config.io, inventory)?;
111        VmProverExtension::<E, _, _>::extend_prover(
112            &Keccak256GpuProverExt,
113            &config.keccak,
114            inventory,
115        )?;
116        Ok(chip_complex)
117    }
118}