1use std::sync::Arc;
17
18use baracuda_cuda_sys::runtime::{
19 cudaGraphExec_t, cudaGraphNode_t, cudaGraph_t, runtime, types::cudaStreamCaptureStatus,
20};
21
22use crate::error::{check, Result};
23use crate::stream::Stream;
24
25#[derive(Copy, Clone, Debug, Eq, PartialEq, Default)]
27pub enum CaptureMode {
28 Global,
31 #[default]
34 ThreadLocal,
35 Relaxed,
38}
39
40impl CaptureMode {
41 #[inline]
42 fn raw(self) -> i32 {
43 match self {
44 CaptureMode::Global => 0,
45 CaptureMode::ThreadLocal => 1,
46 CaptureMode::Relaxed => 2,
47 }
48 }
49}
50
51impl Stream {
52 pub fn begin_capture(&self, mode: CaptureMode) -> Result<()> {
54 let r = runtime()?;
55 let cu = r.cuda_stream_begin_capture()?;
56 check(unsafe { cu(self.as_raw(), mode.raw()) })
57 }
58
59 pub fn end_capture(&self) -> Result<Graph> {
61 let r = runtime()?;
62 let cu = r.cuda_stream_end_capture()?;
63 let mut graph: cudaGraph_t = core::ptr::null_mut();
64 check(unsafe { cu(self.as_raw(), &mut graph) })?;
65 Ok(Graph {
66 inner: Arc::new(GraphInner { handle: graph }),
67 })
68 }
69
70 pub fn capture<F>(&self, mode: CaptureMode, f: F) -> Result<Graph>
82 where
83 F: FnOnce(&Stream) -> Result<()>,
84 {
85 struct CaptureGuard<'a> {
91 stream: &'a Stream,
92 armed: bool,
93 }
94 impl Drop for CaptureGuard<'_> {
95 fn drop(&mut self) {
96 if self.armed {
97 let _ = self.stream.end_capture();
98 }
99 }
100 }
101
102 self.begin_capture(mode)?;
103 let mut guard = CaptureGuard { stream: self, armed: true };
104 let inner_result = f(self);
105 guard.armed = false;
108 let end_result = self.end_capture();
109 match (inner_result, end_result) {
110 (Ok(()), Ok(graph)) => Ok(graph),
111 (Err(e), _) => Err(e),
112 (Ok(()), Err(e)) => Err(e),
113 }
114 }
115
116 pub fn is_capturing(&self) -> Result<bool> {
118 let r = runtime()?;
119 let cu = r.cuda_stream_is_capturing()?;
120 let mut status: core::ffi::c_int = 0;
121 check(unsafe { cu(self.as_raw(), &mut status) })?;
122 Ok(status == cudaStreamCaptureStatus::ACTIVE)
123 }
124}
125
126#[derive(Clone)]
128pub struct Graph {
129 inner: Arc<GraphInner>,
130}
131
132struct GraphInner {
133 handle: cudaGraph_t,
134}
135
136unsafe impl Send for GraphInner {}
137unsafe impl Sync for GraphInner {}
138
139impl core::fmt::Debug for GraphInner {
140 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
141 f.debug_struct("Graph")
142 .field("handle", &self.handle)
143 .finish_non_exhaustive()
144 }
145}
146
147impl core::fmt::Debug for Graph {
148 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
149 self.inner.fmt(f)
150 }
151}
152
153impl Graph {
154 pub fn new() -> Result<Self> {
156 let r = runtime()?;
157 let cu = r.cuda_graph_create()?;
158 let mut graph: cudaGraph_t = core::ptr::null_mut();
159 check(unsafe { cu(&mut graph, 0) })?;
160 Ok(Self {
161 inner: Arc::new(GraphInner { handle: graph }),
162 })
163 }
164
165 pub fn instantiate(&self) -> Result<GraphExec> {
167 let r = runtime()?;
168 let cu = r.cuda_graph_instantiate()?;
169 let mut exec: cudaGraphExec_t = core::ptr::null_mut();
170 check(unsafe { cu(&mut exec, self.inner.handle, 0) })?;
171 Ok(GraphExec {
172 inner: Arc::new(GraphExecInner { handle: exec }),
173 })
174 }
175
176 pub fn node_count(&self) -> Result<usize> {
178 let r = runtime()?;
179 let cu = r.cuda_graph_get_nodes()?;
180 let mut count: usize = 0;
181 check(unsafe { cu(self.inner.handle, core::ptr::null_mut(), &mut count) })?;
182 Ok(count)
183 }
184
185 #[inline]
187 pub fn as_raw(&self) -> cudaGraph_t {
188 self.inner.handle
189 }
190
191 pub fn add_empty_node(&self, dependencies: &[GraphNode]) -> Result<GraphNode> {
193 let r = runtime()?;
194 let cu = r.cuda_graph_add_empty_node()?;
195 let deps: Vec<_> = dependencies.iter().map(|n| n.raw).collect();
196 let (dp, dl) = deps_raw(&deps);
197 let mut node: cudaGraphNode_t = core::ptr::null_mut();
198 check(unsafe { cu(&mut node, self.inner.handle, dp, dl) })?;
199 Ok(GraphNode { raw: node })
200 }
201
202 pub unsafe fn add_kernel_node(
209 &self,
210 dependencies: &[GraphNode],
211 kernel: &crate::Kernel,
212 grid: crate::Dim3,
213 block: crate::Dim3,
214 shared_mem_bytes: u32,
215 args: &mut [*mut core::ffi::c_void],
216 ) -> Result<GraphNode> { unsafe {
217 use baracuda_cuda_sys::runtime::types::{cudaKernelNodeParams, dim3};
218 let r = runtime()?;
219 let cu = r.cuda_graph_add_kernel_node()?;
220 let params = cudaKernelNodeParams {
221 func: kernel.as_launch_ptr() as *mut core::ffi::c_void,
222 grid_dim: dim3::new(grid.x, grid.y, grid.z),
223 block_dim: dim3::new(block.x, block.y, block.z),
224 shared_mem_bytes,
225 kernel_params: if args.is_empty() {
226 core::ptr::null_mut()
227 } else {
228 args.as_mut_ptr()
229 },
230 extra: core::ptr::null_mut(),
231 };
232 let deps: Vec<_> = dependencies.iter().map(|n| n.raw).collect();
233 let (dp, dl) = deps_raw(&deps);
234 let mut node: cudaGraphNode_t = core::ptr::null_mut();
235 check(cu(&mut node, self.inner.handle, dp, dl, ¶ms))?;
236 Ok(GraphNode { raw: node })
237 }}
238
239 pub fn add_memset_u32_node(
241 &self,
242 dependencies: &[GraphNode],
243 dst: *mut core::ffi::c_void,
244 value: u32,
245 count: usize,
246 ) -> Result<GraphNode> {
247 use baracuda_cuda_sys::runtime::types::cudaMemsetParams;
248 let r = runtime()?;
249 let cu = r.cuda_graph_add_memset_node()?;
250 let params = cudaMemsetParams {
251 dst,
252 pitch: 0,
253 value,
254 element_size: 4,
255 width: count,
256 height: 1,
257 };
258 let deps: Vec<_> = dependencies.iter().map(|n| n.raw).collect();
259 let (dp, dl) = deps_raw(&deps);
260 let mut node: cudaGraphNode_t = core::ptr::null_mut();
261 check(unsafe { cu(&mut node, self.inner.handle, dp, dl, ¶ms) })?;
262 Ok(GraphNode { raw: node })
263 }
264
265 pub unsafe fn add_host_node(
273 &self,
274 dependencies: &[GraphNode],
275 fn_: unsafe extern "C" fn(*mut core::ffi::c_void),
276 user_data: *mut core::ffi::c_void,
277 ) -> Result<GraphNode> { unsafe {
278 use baracuda_cuda_sys::runtime::types::cudaHostNodeParams;
279 let r = runtime()?;
280 let cu = r.cuda_graph_add_host_node()?;
281 let params = cudaHostNodeParams {
282 fn_: Some(fn_),
283 user_data,
284 };
285 let deps: Vec<_> = dependencies.iter().map(|n| n.raw).collect();
286 let (dp, dl) = deps_raw(&deps);
287 let mut node: cudaGraphNode_t = core::ptr::null_mut();
288 check(cu(&mut node, self.inner.handle, dp, dl, ¶ms))?;
289 Ok(GraphNode { raw: node })
290 }}
291
292 pub fn add_child_graph_node(
294 &self,
295 dependencies: &[GraphNode],
296 child: &Graph,
297 ) -> Result<GraphNode> {
298 let r = runtime()?;
299 let cu = r.cuda_graph_add_child_graph_node()?;
300 let deps: Vec<_> = dependencies.iter().map(|n| n.raw).collect();
301 let (dp, dl) = deps_raw(&deps);
302 let mut node: cudaGraphNode_t = core::ptr::null_mut();
303 check(unsafe { cu(&mut node, self.inner.handle, dp, dl, child.as_raw()) })?;
304 Ok(GraphNode { raw: node })
305 }
306
307 pub fn add_event_record_node(
309 &self,
310 dependencies: &[GraphNode],
311 event: &crate::Event,
312 ) -> Result<GraphNode> {
313 let r = runtime()?;
314 let cu = r.cuda_graph_add_event_record_node()?;
315 let deps: Vec<_> = dependencies.iter().map(|n| n.raw).collect();
316 let (dp, dl) = deps_raw(&deps);
317 let mut node: cudaGraphNode_t = core::ptr::null_mut();
318 check(unsafe { cu(&mut node, self.inner.handle, dp, dl, event.as_raw()) })?;
319 Ok(GraphNode { raw: node })
320 }
321
322 pub fn add_event_wait_node(
324 &self,
325 dependencies: &[GraphNode],
326 event: &crate::Event,
327 ) -> Result<GraphNode> {
328 let r = runtime()?;
329 let cu = r.cuda_graph_add_event_wait_node()?;
330 let deps: Vec<_> = dependencies.iter().map(|n| n.raw).collect();
331 let (dp, dl) = deps_raw(&deps);
332 let mut node: cudaGraphNode_t = core::ptr::null_mut();
333 check(unsafe { cu(&mut node, self.inner.handle, dp, dl, event.as_raw()) })?;
334 Ok(GraphNode { raw: node })
335 }
336
337 pub fn add_mem_alloc_node(
340 &self,
341 dependencies: &[GraphNode],
342 device: &crate::Device,
343 bytesize: usize,
344 ) -> Result<(GraphNode, *mut core::ffi::c_void)> {
345 use baracuda_cuda_sys::runtime::types::{
346 cudaMemAllocNodeParams, cudaMemAllocationHandleType, cudaMemAllocationType,
347 cudaMemLocation, cudaMemLocationType, cudaMemPoolProps,
348 };
349 let r = runtime()?;
350 let cu = r.cuda_graph_add_mem_alloc_node()?;
351 let mut params = cudaMemAllocNodeParams {
352 pool_props: cudaMemPoolProps {
353 alloc_type: cudaMemAllocationType::PINNED,
354 handle_types: cudaMemAllocationHandleType::NONE,
355 location: cudaMemLocation {
356 type_: cudaMemLocationType::DEVICE,
357 id: device.ordinal(),
358 },
359 ..Default::default()
360 },
361 access_descs: core::ptr::null(),
362 access_desc_count: 0,
363 bytesize,
364 dptr: core::ptr::null_mut(),
365 };
366 let deps: Vec<_> = dependencies.iter().map(|n| n.raw).collect();
367 let (dp, dl) = deps_raw(&deps);
368 let mut node: cudaGraphNode_t = core::ptr::null_mut();
369 check(unsafe { cu(&mut node, self.inner.handle, dp, dl, &mut params) })?;
370 Ok((GraphNode { raw: node }, params.dptr))
371 }
372
373 pub unsafe fn add_mem_free_node(
380 &self,
381 dependencies: &[GraphNode],
382 dptr: *mut core::ffi::c_void,
383 ) -> Result<GraphNode> { unsafe {
384 let r = runtime()?;
385 let cu = r.cuda_graph_add_mem_free_node()?;
386 let deps: Vec<_> = dependencies.iter().map(|n| n.raw).collect();
387 let (dp, dl) = deps_raw(&deps);
388 let mut node: cudaGraphNode_t = core::ptr::null_mut();
389 check(cu(&mut node, self.inner.handle, dp, dl, dptr))?;
390 Ok(GraphNode { raw: node })
391 }}
392
393 pub fn conditional_handle_create(&self, default_launch_value: u32, flags: u32) -> Result<u64> {
401 use baracuda_types::{supports, Feature};
402 let installed = crate::init::driver_version()?;
403 if !supports(installed, Feature::GraphConditionalNodes) {
404 return Err(crate::error::Error::FeatureNotSupported {
405 api: "cudaGraphConditionalHandleCreate",
406 since: Feature::GraphConditionalNodes.required_version(),
407 });
408 }
409 let r = runtime()?;
410 let cu = r.cuda_graph_conditional_handle_create()?;
411 let mut handle: u64 = 0;
412 check(unsafe { cu(&mut handle, self.inner.handle, default_launch_value, flags) })?;
413 Ok(handle)
414 }
415
416 pub unsafe fn add_node_raw(
428 &self,
429 dependencies: &[GraphNode],
430 node_params: *mut core::ffi::c_void,
431 ) -> Result<GraphNode> { unsafe {
432 let r = runtime()?;
433 let cu = r.cuda_graph_add_node()?;
434 let deps: Vec<_> = dependencies.iter().map(|n| n.raw).collect();
435 let (dp, dl) = deps_raw(&deps);
436 let mut node: cudaGraphNode_t = core::ptr::null_mut();
437 check(cu(&mut node, self.inner.handle, dp, dl, node_params))?;
438 Ok(GraphNode { raw: node })
439 }}
440
441 pub fn add_dependencies(&self, from: &[GraphNode], to: &[GraphNode]) -> Result<()> {
443 assert_eq!(from.len(), to.len());
444 if from.is_empty() {
445 return Ok(());
446 }
447 let r = runtime()?;
448 let cu = r.cuda_graph_add_dependencies()?;
449 let f: Vec<_> = from.iter().map(|n| n.raw).collect();
450 let t: Vec<_> = to.iter().map(|n| n.raw).collect();
451 check(unsafe { cu(self.inner.handle, f.as_ptr(), t.as_ptr(), f.len()) })
452 }
453}
454
455fn deps_raw(deps: &[cudaGraphNode_t]) -> (*const cudaGraphNode_t, usize) {
456 if deps.is_empty() {
457 (core::ptr::null(), 0)
458 } else {
459 (deps.as_ptr(), deps.len())
460 }
461}
462
463#[derive(Copy, Clone, Debug)]
466pub struct GraphNode {
467 raw: cudaGraphNode_t,
468}
469
470impl GraphNode {
471 #[inline]
474 pub fn as_raw(&self) -> cudaGraphNode_t {
475 self.raw
476 }
477
478 pub fn node_type(&self) -> Result<i32> {
483 let r = runtime()?;
484 let cu = r.cuda_graph_node_get_type()?;
485 let mut t: core::ffi::c_int = 0;
486 check(unsafe { cu(self.raw, &mut t) })?;
487 Ok(t)
488 }
489
490 pub fn mem_free_ptr(&self) -> Result<*mut core::ffi::c_void> {
492 let r = runtime()?;
493 let cu = r.cuda_graph_mem_free_node_get_params()?;
494 let mut p: *mut core::ffi::c_void = core::ptr::null_mut();
495 check(unsafe { cu(self.raw, &mut p) })?;
496 Ok(p)
497 }
498}
499
500impl Drop for GraphInner {
501 fn drop(&mut self) {
502 if let Ok(r) = runtime() {
503 if let Ok(cu) = r.cuda_graph_destroy() {
504 let _ = unsafe { cu(self.handle) };
505 }
506 }
507 }
508}
509
510#[derive(Clone)]
512pub struct GraphExec {
513 inner: Arc<GraphExecInner>,
514}
515
516struct GraphExecInner {
517 handle: cudaGraphExec_t,
518}
519
520unsafe impl Send for GraphExecInner {}
521unsafe impl Sync for GraphExecInner {}
522
523impl core::fmt::Debug for GraphExecInner {
524 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
525 f.debug_struct("GraphExec")
526 .field("handle", &self.handle)
527 .finish_non_exhaustive()
528 }
529}
530
531impl core::fmt::Debug for GraphExec {
532 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
533 self.inner.fmt(f)
534 }
535}
536
537impl GraphExec {
538 pub fn launch(&self, stream: &Stream) -> Result<()> {
540 let r = runtime()?;
541 let cu = r.cuda_graph_launch()?;
542 check(unsafe { cu(self.inner.handle, stream.as_raw()) })
543 }
544
545 pub fn update(&self, new_template: &Graph) -> Result<UpdateResult> {
549 let r = runtime()?;
550 let cu = r.cuda_graph_exec_update()?;
551 let mut error_node: cudaGraphNode_t = core::ptr::null_mut();
552 let mut result: core::ffi::c_int = 0;
553 let rc = unsafe {
558 cu(
559 self.inner.handle,
560 new_template.as_raw(),
561 &mut error_node,
562 &mut result,
563 )
564 };
565 if rc != baracuda_cuda_sys::runtime::cudaError_t::Success
566 && result == baracuda_cuda_sys::runtime::types::cudaGraphExecUpdateResult::SUCCESS
567 {
568 return Err(crate::error::Error::Status { status: rc });
569 }
570 Ok(UpdateResult {
571 result,
572 error_node: if error_node.is_null() {
573 None
574 } else {
575 Some(GraphNode { raw: error_node })
576 },
577 })
578 }
579
580 #[inline]
582 pub fn as_raw(&self) -> cudaGraphExec_t {
583 self.inner.handle
584 }
585}
586
587#[derive(Clone, Debug)]
591pub struct UpdateResult {
592 pub result: core::ffi::c_int,
596 pub error_node: Option<GraphNode>,
598}
599
600impl UpdateResult {
601 pub fn is_success(&self) -> bool {
603 self.result == baracuda_cuda_sys::runtime::types::cudaGraphExecUpdateResult::SUCCESS
604 }
605}
606
607impl Drop for GraphExecInner {
608 fn drop(&mut self) {
609 if let Ok(r) = runtime() {
610 if let Ok(cu) = r.cuda_graph_exec_destroy() {
611 let _ = unsafe { cu(self.handle) };
612 }
613 }
614 }
615}