#[derive(Clone, Debug, Default, PartialEq)]
pub struct NvrtcOptions {
pub arch: Option<String>,
pub std: Option<String>,
pub extra: Vec<String>,
}
impl NvrtcOptions {
fn to_cudarc(&self) -> cudarc::nvrtc::CompileOptions {
let mut options = cudarc::nvrtc::CompileOptions::default();
if let Some(arch) = &self.arch {
options.options.push(format!("--gpu-architecture={arch}"));
}
if let Some(std_flag) = &self.std {
options.options.push(format!("--std={std_flag}"));
}
options.options.extend(self.extra.iter().cloned());
options
}
}
pub fn compile_nvrtc(src: &str, opts: &NvrtcOptions) -> crate::Result<cudarc::nvrtc::Ptx> {
if src.as_bytes().contains(&0) {
return Err(crate::Error::invalid_argument(
"nvrtc.compile",
"source",
"CUDA source cannot contain NUL bytes",
));
}
cudarc::nvrtc::compile_ptx_with_opts(src, opts.to_cudarc())
.map_err(|err| crate::Error::backend_source("nvrtc.compile", err))
}