use std::sync::Arc;
use openvm_circuit::{
arch::{
testing::memory::air::{MemoryDummyAir, MemoryDummyChip},
MemoryConfig, DEFAULT_BLOCK_SIZE,
},
system::memory::{
offline_checker::{MemoryBridge, MemoryBus},
online::TracingMemory,
},
};
use openvm_circuit_primitives::{
var_range::{VariableRangeCheckerBus, VariableRangeCheckerChipGPU},
Chip,
};
use openvm_cuda_backend::{base::DeviceMatrix, prelude::F, GpuBackend};
use openvm_cuda_common::{copy::MemCopyH2D, stream::GpuDeviceCtx};
use openvm_instructions::DEFERRAL_AS;
use openvm_stark_backend::{
p3_air::BaseAir,
p3_field::{PrimeCharacteristicRing, PrimeField32},
prover::AirProvingContext,
};
use crate::{
cuda_abi::memory_testing,
system::cuda::{memory::MemoryInventoryGPU, poseidon2::Poseidon2PeripheryChipGPU},
};
pub struct DeviceMemoryTester {
pub(crate) chip: FixedSizeMemoryTester,
pub memory: TracingMemory,
pub inventory: MemoryInventoryGPU,
pub hasher_chip: Option<Arc<Poseidon2PeripheryChipGPU>>,
pub config: MemoryConfig,
pub mem_bus: MemoryBus,
pub range_bus: VariableRangeCheckerBus,
}
impl DeviceMemoryTester {
pub fn new(
memory: TracingMemory,
mem_bus: MemoryBus,
mem_config: MemoryConfig,
range_checker: Arc<VariableRangeCheckerChipGPU>,
device_ctx: GpuDeviceCtx,
) -> Self {
let range_bus = range_checker.cpu_chip.as_ref().unwrap().bus();
let sbox_regs = 1;
let poseidon2_periphery = Arc::new(Poseidon2PeripheryChipGPU::new(
sbox_regs,
device_ctx.clone(),
));
let mut inventory = MemoryInventoryGPU::new(
mem_config.clone(),
poseidon2_periphery.clone(),
device_ctx.clone(),
);
inventory.set_initial_memory(&memory.data.memory);
Self {
chip: FixedSizeMemoryTester::new(mem_bus, device_ctx),
memory,
inventory,
hasher_chip: Some(poseidon2_periphery),
config: mem_config,
mem_bus,
range_bus,
}
}
pub fn memory_bridge(&self) -> MemoryBridge {
MemoryBridge::new(self.mem_bus, self.config.timestamp_max_bits, self.range_bus)
}
pub fn read<const N: usize>(&mut self, addr_space: usize, ptr: usize) -> [F; N] {
const { assert!(N == DEFAULT_BLOCK_SIZE) };
let t = self.memory.timestamp();
let (t_prev, data) = if addr_space as u32 == DEFERRAL_AS {
unsafe { self.memory.read::<F, N>(addr_space as u32, ptr as u32) }
} else {
let (t_prev, data) =
unsafe { self.memory.read::<u8, N>(addr_space as u32, ptr as u32) };
(t_prev, data.map(F::from_u8))
};
self.chip
.receive(addr_space as u32, ptr as u32, &data, t_prev);
self.chip.send(addr_space as u32, ptr as u32, &data, t);
data
}
pub fn write<const N: usize>(&mut self, addr_space: usize, ptr: usize, data: [F; N]) {
const { assert!(N == DEFAULT_BLOCK_SIZE) };
let t = self.memory.timestamp();
let (t_prev, data_prev) = unsafe {
self.memory.write::<u8, N>(
addr_space as u32,
ptr as u32,
data.map(|x| x.as_canonical_u32() as u8),
)
};
let data_prev = data_prev.map(F::from_u8);
self.chip
.receive(addr_space as u32, ptr as u32, &data_prev, t_prev);
self.chip.send(addr_space as u32, ptr as u32, &data, t);
}
}
pub struct FixedSizeMemoryTester(pub(crate) MemoryDummyChip<F>, GpuDeviceCtx);
impl FixedSizeMemoryTester {
pub fn new(bus: MemoryBus, device_ctx: GpuDeviceCtx) -> Self {
Self(MemoryDummyChip::new(MemoryDummyAir::new(bus)), device_ctx)
}
pub fn send(&mut self, addr_space: u32, ptr: u32, data: &[F], timestamp: u32) {
self.0.send(addr_space, ptr, data, timestamp);
}
pub fn receive(&mut self, addr_space: u32, ptr: u32, data: &[F], timestamp: u32) {
self.0.receive(addr_space, ptr, data, timestamp);
}
pub fn push(&mut self, addr_space: u32, ptr: u32, data: &[F], timestamp: u32, count: F) {
self.0.push(addr_space, ptr, data, timestamp, count);
}
}
impl<RA> Chip<RA, GpuBackend> for FixedSizeMemoryTester {
fn generate_proving_ctx(&self, _: RA) -> AirProvingContext<GpuBackend> {
let width = BaseAir::<F>::width(&self.0.air);
let height = (self.0.trace.len() / width).next_power_of_two();
let mut records = self.0.trace.clone();
records.resize(height * width, F::ZERO);
let num_records = height;
let trace = DeviceMatrix::<F>::with_capacity_on(height, width, &self.1);
trace.buffer().fill_zero_on(&self.1).unwrap();
unsafe {
memory_testing::tracegen(
trace.buffer(),
height,
width,
&records.to_device_on(&self.1).unwrap(),
num_records,
DEFAULT_BLOCK_SIZE,
self.1.stream.as_raw(),
)
.unwrap();
}
AirProvingContext::simple_no_pis(trace)
}
}