use std::marker::PhantomData;
use openvm_cpu_backend::CpuBackend;
use openvm_cuda_backend::{
base::DeviceMatrix, data_transporter::transport_matrix_h2d_col_major,
hash_scheme::GpuHashScheme, prelude::SC, GenericGpuBackend, GpuBackend,
};
use openvm_cuda_common::stream::GpuDeviceCtx;
use openvm_stark_backend::prover::{AirProvingContext, ColMajorMatrix};
use crate::Chip;
pub fn get_empty_air_proving_ctx<HS: GpuHashScheme>() -> AirProvingContext<GenericGpuBackend<HS>> {
AirProvingContext {
cached_mains: vec![],
common_main: DeviceMatrix::dummy(),
public_values: vec![],
}
}
pub struct HybridChip<RA, C: Chip<RA, CpuBackend<SC>>> {
pub cpu_chip: C,
pub device_ctx: GpuDeviceCtx,
_marker: PhantomData<RA>,
}
impl<RA, C: Chip<RA, CpuBackend<SC>>> HybridChip<RA, C> {
pub fn new(cpu_chip: C, device_ctx: GpuDeviceCtx) -> Self {
Self {
cpu_chip,
device_ctx,
_marker: PhantomData,
}
}
}
impl<RA, C: Chip<RA, CpuBackend<SC>>> Chip<RA, GpuBackend> for HybridChip<RA, C> {
fn generate_proving_ctx(&self, arena: RA) -> AirProvingContext<GpuBackend> {
let ctx = self.cpu_chip.generate_proving_ctx(arena);
cpu_proving_ctx_to_gpu(ctx, &self.device_ctx)
}
}
pub fn cpu_proving_ctx_to_gpu<HS: GpuHashScheme>(
cpu_ctx: AirProvingContext<CpuBackend<SC>>,
device_ctx: &GpuDeviceCtx,
) -> AirProvingContext<GenericGpuBackend<HS>> {
assert!(
cpu_ctx.cached_mains.is_empty(),
"CPU to GPU transfer of cached traces not supported"
);
let cm = ColMajorMatrix::from_row_major(&cpu_ctx.common_main);
let trace = transport_matrix_h2d_col_major(&cm, device_ctx).unwrap();
AirProvingContext {
cached_mains: vec![],
common_main: trace,
public_values: cpu_ctx.public_values,
}
}