use std::ffi::CString;
use std::rc::Rc;
use super::{Function, Module, ModuleInner};
pub(crate) fn load_module_data(
image: *const std::ffi::c_void,
runtime: crate::cubecl::CudaRuntime,
op: &'static str,
) -> crate::Result<Module> {
let handle = unsafe { cudarc::driver::result::module::load_data(image) }
.map_err(|err| crate::Error::backend_source(op, err))?;
Module::from_handle(op, handle, runtime)
}
pub(crate) fn module_function(
inner: &Rc<ModuleInner>,
name: &str,
op: &'static str,
) -> crate::Result<Function> {
let name_c = CString::new(name)
.map_err(|_| crate::Error::invalid_argument(op, "symbol", "symbol name contains NUL"))?;
let handle = unsafe { cudarc::driver::result::module::get_function(inner.handle, name_c) }
.map_err(|err| crate::Error::backend_source(op, err))?;
Ok(Function {
handle,
module: Rc::clone(inner),
})
}
pub(crate) fn validate_launch_config(config: &super::LaunchConfig) -> crate::Result<()> {
const MAX_BLOCK_THREADS: u32 = 1024;
const MAX_GRID_X: u32 = 2_147_483_647;
if config.grid[0] == 0 || config.grid[1] == 0 || config.grid[2] == 0 {
return Err(crate::Error::invalid_argument(
"launch.validate",
"grid",
"grid dimensions must be non-zero",
));
}
if config.block[0] == 0 || config.block[1] == 0 || config.block[2] == 0 {
return Err(crate::Error::invalid_argument(
"launch.validate",
"block",
"block dimensions must be non-zero",
));
}
let threads = config.block[0]
.checked_mul(config.block[1])
.and_then(|v| v.checked_mul(config.block[2]))
.ok_or_else(|| {
crate::Error::invalid_argument("launch.validate", "block", "block product overflow")
})?;
if threads > MAX_BLOCK_THREADS {
return Err(crate::Error::invalid_argument(
"launch.validate",
"block",
format!("block has {threads} threads; CUDA allows at most {MAX_BLOCK_THREADS}"),
));
}
if config.grid[0] > MAX_GRID_X {
return Err(crate::Error::invalid_argument(
"launch.validate",
"grid",
format!(
"grid.x {} exceeds the CUDA limit {MAX_GRID_X}",
config.grid[0]
),
));
}
Ok(())
}