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