openvm_sha2_circuit/extension/
cuda.rs1use 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 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 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}