openvm_sha2_circuit/extension/
cuda.rs

1use openvm_circuit::{
2    arch::{
3        AirInventory, ChipInventory, ChipInventoryError, DenseRecordArena, VmBuilder,
4        VmChipComplex, VmProverExtension,
5    },
6    system::cuda::{
7        extensions::{
8            get_inventory_range_checker, get_or_create_bitwise_op_lookup, SystemGpuBuilder,
9        },
10        SystemChipInventoryGPU,
11    },
12};
13use openvm_cuda_backend::{BabyBearPoseidon2GpuEngine as GpuBabyBearPoseidon2Engine, GpuBackend};
14use openvm_rv32im_circuit::Rv32ImGpuProverExt;
15use openvm_sha2_air::{Sha256Config, Sha512Config};
16use openvm_stark_sdk::config::baby_bear_poseidon2::BabyBearPoseidon2Config;
17
18use super::*;
19use crate::{
20    cuda::{Sha2BlockHasherChipGpu, Sha2MainChipGpu},
21    Sha2BlockHasherVmAir, Sha2MainAir,
22};
23
24pub struct Sha2GpuProverExt;
25
26impl VmProverExtension<GpuBabyBearPoseidon2Engine, DenseRecordArena, Sha2> for Sha2GpuProverExt {
27    fn extend_prover(
28        &self,
29        _: &Sha2,
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_gpu = get_inventory_range_checker(inventory);
36        let bitwise_gpu = get_or_create_bitwise_op_lookup(inventory)?;
37
38        // SHA-256
39        inventory.next_air::<Sha2BlockHasherVmAir<Sha256Config>>()?;
40        let sha256_shared_records = Arc::new(Mutex::new(None));
41        let sha256_block_gpu = Sha2BlockHasherChipGpu::<Sha256Config>::new(
42            sha256_shared_records.clone(),
43            bitwise_gpu.clone(),
44        );
45        inventory.add_periphery_chip(sha256_block_gpu);
46
47        inventory.next_air::<Sha2MainAir<Sha256Config>>()?;
48        let sha256_main_gpu = Sha2MainChipGpu::<Sha256Config>::new(
49            sha256_shared_records,
50            range_checker_gpu.clone(),
51            bitwise_gpu.clone(),
52            pointer_max_bits as u32,
53            timestamp_max_bits as u32,
54        );
55        inventory.add_executor_chip(sha256_main_gpu);
56
57        // SHA-512 (also covers SHA-384 constraints)
58        inventory.next_air::<Sha2BlockHasherVmAir<Sha512Config>>()?;
59        let sha512_shared_records = Arc::new(Mutex::new(None));
60        let sha512_block_gpu = Sha2BlockHasherChipGpu::<Sha512Config>::new(
61            sha512_shared_records.clone(),
62            bitwise_gpu.clone(),
63        );
64        inventory.add_periphery_chip(sha512_block_gpu);
65
66        inventory.next_air::<Sha2MainAir<Sha512Config>>()?;
67        let sha512_main_gpu = Sha2MainChipGpu::<Sha512Config>::new(
68            sha512_shared_records,
69            range_checker_gpu,
70            bitwise_gpu,
71            pointer_max_bits as u32,
72            timestamp_max_bits as u32,
73        );
74        inventory.add_executor_chip(sha512_main_gpu);
75
76        Ok(())
77    }
78}
79
80pub struct Sha2Rv32GpuBuilder;
81
82type E = GpuBabyBearPoseidon2Engine;
83
84impl VmBuilder<E> for Sha2Rv32GpuBuilder {
85    type VmConfig = Sha2Rv32Config;
86    type SystemChipInventory = SystemChipInventoryGPU;
87    type RecordArena = DenseRecordArena;
88
89    fn create_chip_complex(
90        &self,
91        config: &Sha2Rv32Config,
92        circuit: AirInventory<<E as StarkEngine>::SC>,
93        device_ctx: &openvm_stark_backend::EngineDeviceCtx<E>,
94    ) -> Result<
95        VmChipComplex<
96            <E as StarkEngine>::SC,
97            Self::RecordArena,
98            <E as StarkEngine>::PB,
99            Self::SystemChipInventory,
100        >,
101        ChipInventoryError,
102    > {
103        let mut chip_complex = VmBuilder::<E>::create_chip_complex(
104            &SystemGpuBuilder,
105            &config.system,
106            circuit,
107            device_ctx,
108        )?;
109        let inventory = &mut chip_complex.inventory;
110        VmProverExtension::<E, _, _>::extend_prover(&Rv32ImGpuProverExt, &config.rv32i, inventory)?;
111        VmProverExtension::<E, _, _>::extend_prover(&Rv32ImGpuProverExt, &config.rv32m, inventory)?;
112        VmProverExtension::<E, _, _>::extend_prover(&Rv32ImGpuProverExt, &config.io, inventory)?;
113        VmProverExtension::<E, _, _>::extend_prover(&Sha2GpuProverExt, &config.sha2, inventory)?;
114        Ok(chip_complex)
115    }
116}