use super::gpu::{shared_device, CudaError};
use std::sync::Mutex;
pub fn cuda_graph_enabled() -> bool {
use std::sync::OnceLock;
static ENABLED: OnceLock<bool> = OnceLock::new();
*ENABLED.get_or_init(|| {
matches!(
std::env::var("FERROX_CUDA_GRAPH").ok().as_deref(),
Some("1") | Some("true") | Some("on")
)
})
}
pub struct CudaDecodeGraph {
inner: Mutex<Option<GraphExec>>,
}
struct GraphExec {
exec: usize,
}
unsafe impl Send for GraphExec {}
impl CudaDecodeGraph {
pub fn new() -> Self {
Self {
inner: Mutex::new(None),
}
}
pub fn has_exec(&self) -> bool {
self.inner.lock().unwrap().is_some()
}
pub fn begin_capture(&self) -> Result<(), CudaError> {
if !cuda_graph_enabled() {
return Err(CudaError::Launch("FERROX_CUDA_GRAPH not enabled".into()));
}
let _dev = shared_device()?;
unsafe {
use cudarc::driver::sys::{self, CUstreamCaptureMode};
let err = sys::lib().cuStreamBeginCapture_v2(
std::ptr::null_mut(),
CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_GLOBAL,
);
if err != sys::CUresult::CUDA_SUCCESS {
return Err(CudaError::Launch(format!(
"cuStreamBeginCapture failed: {err:?}"
)));
}
}
Ok(())
}
pub fn end_capture(&self) -> Result<(), CudaError> {
let _dev = shared_device()?;
unsafe {
use cudarc::driver::sys;
let mut graph: sys::CUgraph = std::ptr::null_mut();
let err = sys::lib().cuStreamEndCapture(std::ptr::null_mut(), &mut graph);
if err != sys::CUresult::CUDA_SUCCESS {
return Err(CudaError::Launch(format!(
"cuStreamEndCapture failed: {err:?}"
)));
}
let mut exec: sys::CUgraphExec = std::ptr::null_mut();
let err = sys::lib().cuGraphInstantiate_v2(
&mut exec,
graph,
std::ptr::null_mut(),
std::ptr::null_mut(),
0,
);
let _ = sys::lib().cuGraphDestroy(graph);
if err != sys::CUresult::CUDA_SUCCESS {
return Err(CudaError::Launch(format!(
"cuGraphInstantiate failed: {err:?}"
)));
}
*self.inner.lock().unwrap() = Some(GraphExec {
exec: exec as usize,
});
}
Ok(())
}
pub fn launch(&self) -> Result<(), CudaError> {
let guard = self.inner.lock().unwrap();
let Some(exec) = guard.as_ref() else {
return Err(CudaError::Launch("no CUDA graph instantiated".into()));
};
unsafe {
use cudarc::driver::sys;
let err = sys::lib().cuGraphLaunch(exec.exec as sys::CUgraphExec, std::ptr::null_mut());
if err != sys::CUresult::CUDA_SUCCESS {
return Err(CudaError::Launch(format!("cuGraphLaunch failed: {err:?}")));
}
}
Ok(())
}
}
impl Default for CudaDecodeGraph {
fn default() -> Self {
Self::new()
}
}
impl Drop for CudaDecodeGraph {
fn drop(&mut self) {
if let Some(exec) = self.inner.lock().unwrap().take() {
unsafe {
use cudarc::driver::sys;
let _ = sys::lib().cuGraphExecDestroy(exec.exec as sys::CUgraphExec);
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn graph_disabled_reports_no_exec() {
let g = CudaDecodeGraph::new();
assert!(!g.has_exec());
}
}