use std::slice::from_raw_parts;
use openvm_circuit::{
arch::{
instructions::instruction::Instruction, testing::program::ProgramTester, ExecutionState,
},
primitives::Chip,
system::program::{ProgramBus, ProgramExecutionCols},
utils::next_power_of_two_or_zero,
};
use openvm_cuda_backend::{base::DeviceMatrix, prelude::F, GpuBackend};
use openvm_cuda_common::{copy::MemCopyH2D, stream::GpuDeviceCtx};
use openvm_stark_backend::prover::AirProvingContext;
use crate::cuda_abi::program_testing;
pub struct DeviceProgramTester(ProgramTester<F>, GpuDeviceCtx);
impl DeviceProgramTester {
pub fn new(bus: ProgramBus, device_ctx: GpuDeviceCtx) -> Self {
Self(ProgramTester::new(bus), device_ctx)
}
pub fn bus(&self) -> ProgramBus {
self.0.bus
}
pub fn execute(&mut self, instruction: &Instruction<F>, initial_state: &ExecutionState<u32>) {
self.0.execute(instruction, initial_state);
}
}
impl<RA> Chip<RA, GpuBackend> for DeviceProgramTester {
fn generate_proving_ctx(&self, _: RA) -> AirProvingContext<GpuBackend> {
let height = next_power_of_two_or_zero(self.0.records.len());
let width = ProgramTester::<F>::width();
if height == 0 {
return AirProvingContext::simple_no_pis(DeviceMatrix::dummy());
}
let trace = DeviceMatrix::<F>::with_capacity_on(height, width, &self.1);
trace.buffer().fill_zero_on(&self.1).unwrap();
let records = &self.0.records;
let num_records = records.len();
unsafe {
let bytes_size = num_records * size_of::<ProgramExecutionCols<F>>();
let records_bytes = from_raw_parts(records.as_ptr() as *const u8, bytes_size);
let records = records_bytes.to_device_on(&self.1).unwrap();
program_testing::tracegen(
trace.buffer(),
height,
width,
&records,
num_records,
self.1.stream.as_raw(),
)
.unwrap();
}
AirProvingContext::simple_no_pis(trace)
}
}