openvm_circuit/system/cuda/
memory.rs

1use std::sync::Arc;
2
3use openvm_circuit::{
4    arch::{AddressSpaceHostLayout, MemoryConfig, ADDR_SPACE_OFFSET},
5    system::{memory::AddressMap, TouchedMemory},
6};
7use openvm_circuit_primitives::Chip;
8use openvm_cuda_backend::{prelude::F, GpuBackend};
9use openvm_cuda_common::{
10    copy::{cuda_memcpy_on, MemCopyD2H, MemCopyH2D},
11    d_buffer::DeviceBuffer,
12    memory_manager::MemTracker,
13    stream::GpuDeviceCtx,
14};
15use openvm_stark_backend::{p3_field::PrimeCharacteristicRing, prover::AirProvingContext};
16use tracing::instrument;
17
18use super::{
19    boundary::BoundaryChipGPU,
20    merkle_tree::{MemoryMerkleTree, MERKLE_TOUCHED_BLOCK_WIDTH},
21    Poseidon2PeripheryChipGPU, DIGEST_WIDTH,
22};
23use crate::{cuda_abi::inventory, system::memory::online::LinearMemory};
24
25pub struct MemoryInventoryGPU {
26    pub device_ctx: GpuDeviceCtx,
27    pub boundary: BoundaryChipGPU,
28    pub merkle_tree: MemoryMerkleTree,
29    pub hasher_chip: Arc<Poseidon2PeripheryChipGPU>,
30    pub initial_memory: Vec<DeviceBuffer<u8>>,
31    pub merkle_records: Option<DeviceBuffer<u32>>,
32    #[cfg(feature = "metrics")]
33    pub(super) unpadded_merkle_height: usize,
34}
35
36#[repr(C)]
37#[derive(Clone, Copy)]
38struct MemoryInventoryRecord<const CHUNK: usize, const BLOCKS: usize> {
39    address_space: u32,
40    ptr: u32,
41    timestamps: [u32; BLOCKS],
42    values: [u32; CHUNK],
43}
44
45#[repr(C)]
46#[derive(Clone, Copy)]
47struct MemoryMerkleRecord {
48    address_space: u32,
49    ptr: u32,
50    timestamp: u32,
51    values: [u32; DIGEST_WIDTH],
52}
53
54impl MemoryInventoryGPU {
55    #[inline]
56    fn field_to_raw_u32(value: F) -> u32 {
57        unsafe { std::mem::transmute::<F, u32>(value) }
58    }
59
60    pub fn new(
61        config: MemoryConfig,
62        hasher_chip: Arc<Poseidon2PeripheryChipGPU>,
63        device_ctx: GpuDeviceCtx,
64    ) -> Self {
65        Self {
66            device_ctx: device_ctx.clone(),
67            boundary: BoundaryChipGPU::new(hasher_chip.shared_buffer(), device_ctx.clone()),
68            merkle_tree: MemoryMerkleTree::new(config.clone(), hasher_chip.clone(), device_ctx),
69            hasher_chip,
70            initial_memory: Vec::new(),
71            merkle_records: None,
72            #[cfg(feature = "metrics")]
73            unpadded_merkle_height: 0,
74        }
75    }
76
77    #[instrument(name = "set_initial_memory", skip_all)]
78    pub fn set_initial_memory(&mut self, initial_memory: &AddressMap) {
79        let mem = MemTracker::start("set initial memory");
80        for (addr_sp, raw_mem) in initial_memory
81            .get_memory()
82            .iter()
83            .map(|mem| mem.as_slice())
84            .enumerate()
85        {
86            tracing::debug!(
87                "Setting initial memory for address space {}: {} bytes",
88                addr_sp,
89                raw_mem.len()
90            );
91            self.initial_memory.push(if raw_mem.is_empty() {
92                DeviceBuffer::new()
93            } else {
94                raw_mem
95                    .to_device_on(&self.device_ctx)
96                    .expect("failed to copy memory to device")
97            });
98            self.merkle_tree
99                .build_async(&self.initial_memory[addr_sp], addr_sp);
100        }
101        self.boundary.initial_leaves = self
102            .initial_memory
103            .iter()
104            .skip(1)
105            .map(|per_as| per_as.as_raw_ptr())
106            .collect();
107        mem.emit_metrics();
108    }
109
110    #[instrument(name = "generate_proving_ctxs", skip_all)]
111    pub fn generate_proving_ctxs(
112        &mut self,
113        touched_memory: TouchedMemory<F>,
114    ) -> Vec<AirProvingContext<GpuBackend>> {
115        let mem = MemTracker::start("generate mem proving ctxs");
116        let partition = touched_memory;
117        let merkle_proof_ctx = if partition.is_empty() {
118            let leftmost_values = 'left: {
119                let mut res = [F::ZERO; DIGEST_WIDTH];
120                if self.initial_memory[ADDR_SPACE_OFFSET as usize].is_empty() {
121                    break 'left res;
122                }
123                let layout =
124                    &self.merkle_tree.mem_config().addr_spaces[ADDR_SPACE_OFFSET as usize].layout;
125                let one_cell_size = layout.size();
126                let mut values = vec![0u8; one_cell_size * DIGEST_WIDTH];
127                unsafe {
128                    cuda_memcpy_on::<true, false>(
129                        values.as_mut_ptr() as *mut std::ffi::c_void,
130                        self.initial_memory[ADDR_SPACE_OFFSET as usize].as_ptr()
131                            as *const std::ffi::c_void,
132                        values.len(),
133                        &self.device_ctx,
134                    )
135                    .unwrap();
136                    for i in 0..DIGEST_WIDTH {
137                        res[i] = layout.to_field::<F>(&values[i * one_cell_size..]);
138                    }
139                }
140                res
141            };
142
143            let values_u32 = leftmost_values.map(Self::field_to_raw_u32);
144            let merkle_record = MemoryMerkleRecord {
145                address_space: ADDR_SPACE_OFFSET,
146                ptr: 0,
147                timestamp: 0,
148                values: values_u32,
149            };
150            let merkle_records = [merkle_record];
151            let merkle_words: &[u32] = unsafe {
152                std::slice::from_raw_parts(
153                    merkle_records.as_ptr() as *const u32,
154                    MERKLE_TOUCHED_BLOCK_WIDTH,
155                )
156            };
157            let d_merkle_touched_memory = merkle_words.to_device_on(&self.device_ctx).unwrap();
158
159            let unpadded_merkle_height = self.merkle_tree.calculate_unpadded_height(&partition);
160            #[cfg(feature = "metrics")]
161            {
162                self.unpadded_merkle_height = unpadded_merkle_height;
163            }
164
165            self.boundary.finalize_records::<DIGEST_WIDTH>(Vec::new());
166            self.prepare_poseidon2_records(0, unpadded_merkle_height);
167            mem.tracing_info("merkle update");
168            self.merkle_tree.finalize();
169            self.merkle_tree.update_with_touched_blocks(
170                unpadded_merkle_height,
171                &d_merkle_touched_memory,
172                true,
173            )
174        } else {
175            // Convert MemoryInventoryRecord<4, 1> to MemoryInventoryRecord<8, 2>
176            let in_records: Vec<MemoryInventoryRecord<4, 1>> = partition
177                .iter()
178                .map(|&((addr_space, ptr), ts_values)| MemoryInventoryRecord {
179                    address_space: addr_space,
180                    ptr,
181                    timestamps: [ts_values.timestamp],
182                    values: ts_values.values.map(Self::field_to_raw_u32),
183                })
184                .collect();
185            let in_num_records = in_records.len();
186            let out_words = in_num_records
187                * (std::mem::size_of::<MemoryInventoryRecord<8, 2>>() / std::mem::size_of::<u32>());
188            let d_in_records = in_records
189                .to_device_on(&self.device_ctx)
190                .unwrap()
191                .as_buffer::<u32>();
192            let d_tmp_records = DeviceBuffer::<u32>::with_capacity_on(out_words, &self.device_ctx);
193            let d_out_records = DeviceBuffer::<u32>::with_capacity_on(out_words, &self.device_ctx);
194            let d_out_num_records = DeviceBuffer::<usize>::with_capacity_on(1, &self.device_ctx);
195            let d_flags = DeviceBuffer::<u32>::with_capacity_on(in_num_records, &self.device_ctx);
196            let d_positions =
197                DeviceBuffer::<u32>::with_capacity_on(in_num_records, &self.device_ctx);
198            let d_initial_mem = self
199                .boundary
200                .initial_leaves
201                .to_device_on(&self.device_ctx)
202                .unwrap();
203            let mut temp_bytes = 0usize;
204            unsafe {
205                inventory::merge_records_get_temp_bytes(
206                    &d_flags,
207                    in_num_records,
208                    &mut temp_bytes,
209                    self.device_ctx.stream.as_raw(),
210                )
211                .expect("merge_records_get_temp_bytes failed");
212            }
213            let d_temp_storage = if temp_bytes == 0 {
214                DeviceBuffer::<u8>::new()
215            } else {
216                DeviceBuffer::<u8>::with_capacity_on(temp_bytes, &self.device_ctx)
217            };
218            unsafe {
219                inventory::merge_records(
220                    &d_in_records,
221                    in_num_records,
222                    &d_initial_mem,
223                    &d_tmp_records,
224                    &d_out_records,
225                    &d_flags,
226                    &d_positions,
227                    &d_temp_storage,
228                    temp_bytes,
229                    &d_out_num_records,
230                    self.device_ctx.stream.as_raw(),
231                )
232                .expect("merge_records failed");
233            }
234
235            // Send records to boundary chip
236            let out_num_records = d_out_num_records.to_host_on(&self.device_ctx).unwrap()[0];
237            self.boundary
238                .finalize_records_device::<DIGEST_WIDTH>(d_out_records, out_num_records);
239
240            // Send records to memory merkle tree
241            let out_records = self
242                .boundary
243                .records()
244                .to_host_on(&self.device_ctx)
245                .unwrap();
246            let record_words = 4 + DIGEST_WIDTH;
247            let mut merkle_records = Vec::with_capacity(out_num_records);
248            for i in 0..out_num_records {
249                let base = i * record_words;
250                let mut values = [0u32; DIGEST_WIDTH];
251                values.copy_from_slice(&out_records[base + 4..base + 4 + DIGEST_WIDTH]);
252                let record = MemoryMerkleRecord {
253                    address_space: out_records[base],
254                    ptr: out_records[base + 1],
255                    timestamp: out_records[base + 2].max(out_records[base + 3]),
256                    values,
257                };
258                merkle_records.push(record);
259            }
260            let merkle_words: &[u32] = unsafe {
261                std::slice::from_raw_parts(
262                    merkle_records.as_ptr() as *const u32,
263                    merkle_records.len() * MERKLE_TOUCHED_BLOCK_WIDTH,
264                )
265            };
266            self.merkle_records = Some(merkle_words.to_device_on(&self.device_ctx).unwrap());
267
268            let unpadded_merkle_height = self.merkle_tree.calculate_unpadded_height(&partition);
269            #[cfg(feature = "metrics")]
270            {
271                self.unpadded_merkle_height = unpadded_merkle_height;
272            }
273
274            self.prepare_poseidon2_records(out_num_records, unpadded_merkle_height);
275            mem.tracing_info("merkle update");
276            self.merkle_tree.finalize();
277            self.merkle_tree.update_with_touched_blocks(
278                unpadded_merkle_height,
279                self.merkle_records
280                    .as_ref()
281                    .expect("missing merkle records"),
282                false,
283            )
284        };
285        mem.tracing_info("boundary tracegen");
286        let ret = vec![self.boundary.generate_proving_ctx(()), merkle_proof_ctx];
287        mem.tracing_info("dropping merkle tree");
288        self.merkle_tree.drop_subtrees();
289        self.initial_memory = Vec::new();
290        mem.emit_metrics();
291        ret
292    }
293
294    fn prepare_poseidon2_records(&self, boundary_records: usize, merkle_height: usize) {
295        let num_records = boundary_records
296            .checked_mul(2)
297            .and_then(|n| n.checked_add(merkle_height))
298            .expect("Poseidon2 records count overflow");
299        self.hasher_chip.prepare_records(num_records);
300    }
301}
302
303impl Drop for MemoryInventoryGPU {
304    fn drop(&mut self) {
305        // WARNING: The merkle subtree events must be completed before dropping the initial memory
306        // buffers. This prevents buffers from dropping before build_async completes.
307        self.merkle_tree.drop_subtrees();
308        self.initial_memory.clear();
309    }
310}
311
312#[cfg(test)]
313mod tests {
314    use std::sync::Arc;
315
316    use openvm_circuit::{
317        arch::{vm_poseidon2_config, MemoryConfig},
318        system::{
319            memory::{merkle::MerkleTree, online::GuestMemory, AddressMap, TimestampedValues},
320            poseidon2::Poseidon2PeripheryChip,
321        },
322    };
323    use openvm_cuda_backend::prelude::F;
324    use openvm_cuda_common::{
325        common::get_device,
326        stream::{CudaStream, GpuDeviceCtx, StreamGuard},
327    };
328    use openvm_instructions::riscv::{RV32_MEMORY_AS, RV32_REGISTER_AS};
329    use openvm_stark_backend::prover::MatrixDimensions;
330
331    use super::*;
332    #[test]
333    fn test_empty_touched_memory_uses_full_chunk_values() {
334        let mut addr_spaces = MemoryConfig::empty_address_space_configs(5);
335        for addr_space in [RV32_REGISTER_AS, RV32_MEMORY_AS] {
336            addr_spaces[addr_space as usize].num_cells = 2 * DIGEST_WIDTH;
337        }
338        let mem_config = MemoryConfig::new(2, addr_spaces, 4, 29, 17);
339
340        let mut memory = GuestMemory::new(AddressMap::from_mem_config(&mem_config));
341        unsafe {
342            memory.write::<u8, DIGEST_WIDTH>(RV32_REGISTER_AS, 0, [1, 2, 3, 4, 5, 6, 7, 8]);
343            memory.write::<u8, { DIGEST_WIDTH / 2 }>(RV32_MEMORY_AS, 0, [9, 10, 11, 12]);
344        }
345
346        let cpu_hasher = Poseidon2PeripheryChip::new(vm_poseidon2_config(), 3);
347        let cpu_merkle_tree = MerkleTree::<F, DIGEST_WIDTH>::from_memory(
348            &memory.memory,
349            &mem_config.memory_dimensions(),
350            &cpu_hasher,
351        );
352        let expected_root = cpu_merkle_tree.root();
353
354        let device_ctx = GpuDeviceCtx {
355            device_id: get_device().unwrap() as u32,
356            stream: StreamGuard::new(CudaStream::new_non_blocking().unwrap()),
357        };
358        let hasher_chip = Arc::new(Poseidon2PeripheryChipGPU::new(1, device_ctx.clone()));
359        let mut inventory =
360            MemoryInventoryGPU::new(mem_config.clone(), hasher_chip, device_ctx.clone());
361        inventory.set_initial_memory(&memory.memory);
362
363        let ctxs = inventory.generate_proving_ctxs(Vec::new());
364        let boundary_ctx = ctxs.first().expect("missing boundary ctx");
365        assert_eq!(
366            boundary_ctx.common_main.height(),
367            1,
368            "boundary trace should be a single padding row for empty touched memory"
369        );
370        assert!(
371            boundary_ctx.public_values.is_empty(),
372            "boundary chip should not emit public values"
373        );
374
375        let merkle_ctx = ctxs
376            .iter()
377            .find(|ctx| ctx.public_values.len() >= 2 * DIGEST_WIDTH)
378            .expect("missing merkle ctx");
379        let gpu_root_slice =
380            &merkle_ctx.public_values[merkle_ctx.public_values.len() - DIGEST_WIDTH..];
381        let gpu_root: [F; DIGEST_WIDTH] = gpu_root_slice.try_into().unwrap();
382
383        assert_eq!(expected_root, gpu_root);
384    }
385
386    #[test]
387    fn test_touched_memory_updates_memory_address_space() {
388        let mut addr_spaces = MemoryConfig::empty_address_space_configs(5);
389        for addr_space in [RV32_REGISTER_AS, RV32_MEMORY_AS] {
390            addr_spaces[addr_space as usize].num_cells = 2 * DIGEST_WIDTH;
391        }
392        let mem_config = MemoryConfig::new(2, addr_spaces, 4, 29, 17);
393
394        let mut memory = GuestMemory::new(AddressMap::from_mem_config(&mem_config));
395        unsafe {
396            memory.write::<u8, DIGEST_WIDTH>(RV32_REGISTER_AS, 0, [1, 2, 3, 4, 5, 6, 7, 8]);
397            memory.write::<u8, { DIGEST_WIDTH / 2 }>(RV32_MEMORY_AS, 0, [9, 10, 11, 12]);
398        }
399
400        let mut final_memory = memory.clone();
401        let touched_bytes = [101u8, 102, 103, 104];
402        let touched_bytes_late = [111u8, 112, 113, 114];
403        unsafe {
404            final_memory.write::<u8, { crate::arch::DEFAULT_BLOCK_SIZE }>(
405                RV32_MEMORY_AS,
406                0,
407                touched_bytes,
408            );
409            final_memory.write::<u8, { crate::arch::DEFAULT_BLOCK_SIZE }>(
410                RV32_MEMORY_AS,
411                crate::arch::DEFAULT_BLOCK_SIZE as u32,
412                touched_bytes_late,
413            );
414        }
415
416        let cpu_hasher = Poseidon2PeripheryChip::new(vm_poseidon2_config(), 3);
417        let cpu_merkle_tree = MerkleTree::<F, DIGEST_WIDTH>::from_memory(
418            &final_memory.memory,
419            &mem_config.memory_dimensions(),
420            &cpu_hasher,
421        );
422        let expected_root = cpu_merkle_tree.root();
423
424        let device_ctx = GpuDeviceCtx {
425            device_id: get_device().unwrap() as u32,
426            stream: StreamGuard::new(CudaStream::new_non_blocking().unwrap()),
427        };
428        let hasher_chip = Arc::new(Poseidon2PeripheryChipGPU::new(1, device_ctx.clone()));
429        let mut inventory =
430            MemoryInventoryGPU::new(mem_config.clone(), hasher_chip, device_ctx.clone());
431        inventory.set_initial_memory(&memory.memory);
432
433        let touched_memory = vec![
434            (
435                (RV32_MEMORY_AS, 0),
436                TimestampedValues {
437                    timestamp: 1,
438                    values: touched_bytes.map(F::from_u8),
439                },
440            ),
441            (
442                (RV32_MEMORY_AS, crate::arch::DEFAULT_BLOCK_SIZE as u32),
443                TimestampedValues {
444                    timestamp: 3,
445                    values: touched_bytes_late.map(F::from_u8),
446                },
447            ),
448        ];
449        let ctxs = inventory.generate_proving_ctxs(touched_memory);
450        let boundary_ctx = ctxs.first().expect("missing boundary ctx");
451        assert!(
452            boundary_ctx.common_main.height() > 0,
453            "boundary trace should be present when touched memory is non-empty"
454        );
455        assert!(
456            boundary_ctx.public_values.is_empty(),
457            "boundary chip should not emit public values"
458        );
459
460        let merkle_ctx = ctxs
461            .iter()
462            .find(|ctx| ctx.public_values.len() >= 2 * DIGEST_WIDTH)
463            .expect("missing merkle ctx");
464        let gpu_root_slice =
465            &merkle_ctx.public_values[merkle_ctx.public_values.len() - DIGEST_WIDTH..];
466        let gpu_root: [F; DIGEST_WIDTH] = gpu_root_slice.try_into().unwrap();
467
468        assert_eq!(expected_root, gpu_root);
469    }
470}