1use std::sync::Arc;
17
18use baracuda_cuda_sys::runtime::{
19 cudaGraph_t, cudaGraphExec_t, cudaGraphNode_t, runtime, types::cudaStreamCaptureStatus,
20};
21
22use crate::error::{Result, check};
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 {
104 stream: self,
105 armed: true,
106 };
107 let inner_result = f(self);
108 guard.armed = false;
111 let end_result = self.end_capture();
112 match (inner_result, end_result) {
113 (Ok(()), Ok(graph)) => Ok(graph),
114 (Err(e), _) => Err(e),
115 (Ok(()), Err(e)) => Err(e),
116 }
117 }
118
119 pub fn is_capturing(&self) -> Result<bool> {
121 let r = runtime()?;
122 let cu = r.cuda_stream_is_capturing()?;
123 let mut status: core::ffi::c_int = 0;
124 check(unsafe { cu(self.as_raw(), &mut status) })?;
125 Ok(status == cudaStreamCaptureStatus::ACTIVE)
126 }
127}
128
129#[derive(Clone)]
131pub struct Graph {
132 inner: Arc<GraphInner>,
133}
134
135struct GraphInner {
136 handle: cudaGraph_t,
137}
138
139unsafe impl Send for GraphInner {}
140unsafe impl Sync for GraphInner {}
141
142impl core::fmt::Debug for GraphInner {
143 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
144 f.debug_struct("Graph")
145 .field("handle", &self.handle)
146 .finish_non_exhaustive()
147 }
148}
149
150impl core::fmt::Debug for Graph {
151 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
152 self.inner.fmt(f)
153 }
154}
155
156impl Graph {
157 pub fn new() -> Result<Self> {
159 let r = runtime()?;
160 let cu = r.cuda_graph_create()?;
161 let mut graph: cudaGraph_t = core::ptr::null_mut();
162 check(unsafe { cu(&mut graph, 0) })?;
163 Ok(Self {
164 inner: Arc::new(GraphInner { handle: graph }),
165 })
166 }
167
168 pub fn instantiate(&self) -> Result<GraphExec> {
170 let r = runtime()?;
171 let cu = r.cuda_graph_instantiate()?;
172 let mut exec: cudaGraphExec_t = core::ptr::null_mut();
173 check(unsafe { cu(&mut exec, self.inner.handle, 0) })?;
174 Ok(GraphExec {
175 inner: Arc::new(GraphExecInner { handle: exec }),
176 })
177 }
178
179 pub fn node_count(&self) -> Result<usize> {
181 let r = runtime()?;
182 let cu = r.cuda_graph_get_nodes()?;
183 let mut count: usize = 0;
184 check(unsafe { cu(self.inner.handle, core::ptr::null_mut(), &mut count) })?;
185 Ok(count)
186 }
187
188 #[inline]
190 pub fn as_raw(&self) -> cudaGraph_t {
191 self.inner.handle
192 }
193
194 pub fn add_empty_node(&self, dependencies: &[GraphNode]) -> Result<GraphNode> {
196 let r = runtime()?;
197 let cu = r.cuda_graph_add_empty_node()?;
198 let deps: Vec<_> = dependencies.iter().map(|n| n.raw).collect();
199 let (dp, dl) = deps_raw(&deps);
200 let mut node: cudaGraphNode_t = core::ptr::null_mut();
201 check(unsafe { cu(&mut node, self.inner.handle, dp, dl) })?;
202 Ok(GraphNode { raw: node })
203 }
204
205 pub unsafe fn add_kernel_node(
212 &self,
213 dependencies: &[GraphNode],
214 kernel: &crate::Kernel,
215 grid: crate::Dim3,
216 block: crate::Dim3,
217 shared_mem_bytes: u32,
218 args: &mut [*mut core::ffi::c_void],
219 ) -> Result<GraphNode> {
220 unsafe {
221 use baracuda_cuda_sys::runtime::types::{cudaKernelNodeParams, dim3};
222 let r = runtime()?;
223 let cu = r.cuda_graph_add_kernel_node()?;
224 let params = cudaKernelNodeParams {
225 func: kernel.as_launch_ptr() as *mut core::ffi::c_void,
226 grid_dim: dim3::new(grid.x, grid.y, grid.z),
227 block_dim: dim3::new(block.x, block.y, block.z),
228 shared_mem_bytes,
229 kernel_params: if args.is_empty() {
230 core::ptr::null_mut()
231 } else {
232 args.as_mut_ptr()
233 },
234 extra: core::ptr::null_mut(),
235 };
236 let deps: Vec<_> = dependencies.iter().map(|n| n.raw).collect();
237 let (dp, dl) = deps_raw(&deps);
238 let mut node: cudaGraphNode_t = core::ptr::null_mut();
239 check(cu(&mut node, self.inner.handle, dp, dl, ¶ms))?;
240 Ok(GraphNode { raw: node })
241 }
242 }
243
244 pub fn add_memset_u32_node(
246 &self,
247 dependencies: &[GraphNode],
248 dst: *mut core::ffi::c_void,
249 value: u32,
250 count: usize,
251 ) -> Result<GraphNode> {
252 use baracuda_cuda_sys::runtime::types::cudaMemsetParams;
253 let r = runtime()?;
254 let cu = r.cuda_graph_add_memset_node()?;
255 let params = cudaMemsetParams {
256 dst,
257 pitch: 0,
258 value,
259 element_size: 4,
260 width: count,
261 height: 1,
262 };
263 let deps: Vec<_> = dependencies.iter().map(|n| n.raw).collect();
264 let (dp, dl) = deps_raw(&deps);
265 let mut node: cudaGraphNode_t = core::ptr::null_mut();
266 check(unsafe { cu(&mut node, self.inner.handle, dp, dl, ¶ms) })?;
267 Ok(GraphNode { raw: node })
268 }
269
270 pub unsafe fn add_host_node(
278 &self,
279 dependencies: &[GraphNode],
280 fn_: unsafe extern "C" fn(*mut core::ffi::c_void),
281 user_data: *mut core::ffi::c_void,
282 ) -> Result<GraphNode> {
283 unsafe {
284 use baracuda_cuda_sys::runtime::types::cudaHostNodeParams;
285 let r = runtime()?;
286 let cu = r.cuda_graph_add_host_node()?;
287 let params = cudaHostNodeParams {
288 fn_: Some(fn_),
289 user_data,
290 };
291 let deps: Vec<_> = dependencies.iter().map(|n| n.raw).collect();
292 let (dp, dl) = deps_raw(&deps);
293 let mut node: cudaGraphNode_t = core::ptr::null_mut();
294 check(cu(&mut node, self.inner.handle, dp, dl, ¶ms))?;
295 Ok(GraphNode { raw: node })
296 }
297 }
298
299 pub fn add_child_graph_node(
301 &self,
302 dependencies: &[GraphNode],
303 child: &Graph,
304 ) -> Result<GraphNode> {
305 let r = runtime()?;
306 let cu = r.cuda_graph_add_child_graph_node()?;
307 let deps: Vec<_> = dependencies.iter().map(|n| n.raw).collect();
308 let (dp, dl) = deps_raw(&deps);
309 let mut node: cudaGraphNode_t = core::ptr::null_mut();
310 check(unsafe { cu(&mut node, self.inner.handle, dp, dl, child.as_raw()) })?;
311 Ok(GraphNode { raw: node })
312 }
313
314 pub fn add_event_record_node(
316 &self,
317 dependencies: &[GraphNode],
318 event: &crate::Event,
319 ) -> Result<GraphNode> {
320 let r = runtime()?;
321 let cu = r.cuda_graph_add_event_record_node()?;
322 let deps: Vec<_> = dependencies.iter().map(|n| n.raw).collect();
323 let (dp, dl) = deps_raw(&deps);
324 let mut node: cudaGraphNode_t = core::ptr::null_mut();
325 check(unsafe { cu(&mut node, self.inner.handle, dp, dl, event.as_raw()) })?;
326 Ok(GraphNode { raw: node })
327 }
328
329 pub fn add_event_wait_node(
331 &self,
332 dependencies: &[GraphNode],
333 event: &crate::Event,
334 ) -> Result<GraphNode> {
335 let r = runtime()?;
336 let cu = r.cuda_graph_add_event_wait_node()?;
337 let deps: Vec<_> = dependencies.iter().map(|n| n.raw).collect();
338 let (dp, dl) = deps_raw(&deps);
339 let mut node: cudaGraphNode_t = core::ptr::null_mut();
340 check(unsafe { cu(&mut node, self.inner.handle, dp, dl, event.as_raw()) })?;
341 Ok(GraphNode { raw: node })
342 }
343
344 pub fn add_mem_alloc_node(
347 &self,
348 dependencies: &[GraphNode],
349 device: &crate::Device,
350 bytesize: usize,
351 ) -> Result<(GraphNode, *mut core::ffi::c_void)> {
352 use baracuda_cuda_sys::runtime::types::{
353 cudaMemAllocNodeParams, cudaMemAllocationHandleType, cudaMemAllocationType,
354 cudaMemLocation, cudaMemLocationType, cudaMemPoolProps,
355 };
356 let r = runtime()?;
357 let cu = r.cuda_graph_add_mem_alloc_node()?;
358 let mut params = cudaMemAllocNodeParams {
359 pool_props: cudaMemPoolProps {
360 alloc_type: cudaMemAllocationType::PINNED,
361 handle_types: cudaMemAllocationHandleType::NONE,
362 location: cudaMemLocation {
363 type_: cudaMemLocationType::DEVICE,
364 id: device.ordinal(),
365 },
366 ..Default::default()
367 },
368 access_descs: core::ptr::null(),
369 access_desc_count: 0,
370 bytesize,
371 dptr: core::ptr::null_mut(),
372 };
373 let deps: Vec<_> = dependencies.iter().map(|n| n.raw).collect();
374 let (dp, dl) = deps_raw(&deps);
375 let mut node: cudaGraphNode_t = core::ptr::null_mut();
376 check(unsafe { cu(&mut node, self.inner.handle, dp, dl, &mut params) })?;
377 Ok((GraphNode { raw: node }, params.dptr))
378 }
379
380 pub unsafe fn add_mem_free_node(
387 &self,
388 dependencies: &[GraphNode],
389 dptr: *mut core::ffi::c_void,
390 ) -> Result<GraphNode> {
391 unsafe {
392 let r = runtime()?;
393 let cu = r.cuda_graph_add_mem_free_node()?;
394 let deps: Vec<_> = dependencies.iter().map(|n| n.raw).collect();
395 let (dp, dl) = deps_raw(&deps);
396 let mut node: cudaGraphNode_t = core::ptr::null_mut();
397 check(cu(&mut node, self.inner.handle, dp, dl, dptr))?;
398 Ok(GraphNode { raw: node })
399 }
400 }
401
402 pub fn conditional_handle_create(&self, default_launch_value: u32, flags: u32) -> Result<u64> {
410 use baracuda_types::{Feature, supports};
411 let installed = crate::init::driver_version()?;
412 if !supports(installed, Feature::GraphConditionalNodes) {
413 return Err(crate::error::Error::FeatureNotSupported {
414 api: "cudaGraphConditionalHandleCreate",
415 since: Feature::GraphConditionalNodes.required_version(),
416 });
417 }
418 let r = runtime()?;
419 let cu = r.cuda_graph_conditional_handle_create()?;
420 let mut handle: u64 = 0;
421 check(unsafe { cu(&mut handle, self.inner.handle, default_launch_value, flags) })?;
422 Ok(handle)
423 }
424
425 pub unsafe fn add_node_raw(
437 &self,
438 dependencies: &[GraphNode],
439 node_params: *mut core::ffi::c_void,
440 ) -> Result<GraphNode> {
441 unsafe {
442 let r = runtime()?;
443 let cu = r.cuda_graph_add_node()?;
444 let deps: Vec<_> = dependencies.iter().map(|n| n.raw).collect();
445 let (dp, dl) = deps_raw(&deps);
446 let mut node: cudaGraphNode_t = core::ptr::null_mut();
447 check(cu(&mut node, self.inner.handle, dp, dl, node_params))?;
448 Ok(GraphNode { raw: node })
449 }
450 }
451
452 pub fn add_dependencies(&self, from: &[GraphNode], to: &[GraphNode]) -> Result<()> {
454 assert_eq!(from.len(), to.len());
455 if from.is_empty() {
456 return Ok(());
457 }
458 let r = runtime()?;
459 let cu = r.cuda_graph_add_dependencies()?;
460 let f: Vec<_> = from.iter().map(|n| n.raw).collect();
461 let t: Vec<_> = to.iter().map(|n| n.raw).collect();
462 check(unsafe { cu(self.inner.handle, f.as_ptr(), t.as_ptr(), f.len()) })
463 }
464}
465
466fn deps_raw(deps: &[cudaGraphNode_t]) -> (*const cudaGraphNode_t, usize) {
467 if deps.is_empty() {
468 (core::ptr::null(), 0)
469 } else {
470 (deps.as_ptr(), deps.len())
471 }
472}
473
474#[derive(Copy, Clone, Debug)]
477pub struct GraphNode {
478 raw: cudaGraphNode_t,
479}
480
481impl GraphNode {
482 #[inline]
485 pub fn as_raw(&self) -> cudaGraphNode_t {
486 self.raw
487 }
488
489 pub fn node_type(&self) -> Result<i32> {
494 let r = runtime()?;
495 let cu = r.cuda_graph_node_get_type()?;
496 let mut t: core::ffi::c_int = 0;
497 check(unsafe { cu(self.raw, &mut t) })?;
498 Ok(t)
499 }
500
501 pub fn mem_free_ptr(&self) -> Result<*mut core::ffi::c_void> {
503 let r = runtime()?;
504 let cu = r.cuda_graph_mem_free_node_get_params()?;
505 let mut p: *mut core::ffi::c_void = core::ptr::null_mut();
506 check(unsafe { cu(self.raw, &mut p) })?;
507 Ok(p)
508 }
509}
510
511impl Drop for GraphInner {
512 fn drop(&mut self) {
513 if let Ok(r) = runtime() {
514 if let Ok(cu) = r.cuda_graph_destroy() {
515 let _ = unsafe { cu(self.handle) };
516 }
517 }
518 }
519}
520
521#[derive(Clone)]
523pub struct GraphExec {
524 inner: Arc<GraphExecInner>,
525}
526
527struct GraphExecInner {
528 handle: cudaGraphExec_t,
529}
530
531unsafe impl Send for GraphExecInner {}
532unsafe impl Sync for GraphExecInner {}
533
534impl core::fmt::Debug for GraphExecInner {
535 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
536 f.debug_struct("GraphExec")
537 .field("handle", &self.handle)
538 .finish_non_exhaustive()
539 }
540}
541
542impl core::fmt::Debug for GraphExec {
543 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
544 self.inner.fmt(f)
545 }
546}
547
548impl GraphExec {
549 pub fn launch(&self, stream: &Stream) -> Result<()> {
551 let r = runtime()?;
552 let cu = r.cuda_graph_launch()?;
553 check(unsafe { cu(self.inner.handle, stream.as_raw()) })
554 }
555
556 pub fn update(&self, new_template: &Graph) -> Result<UpdateResult> {
560 let r = runtime()?;
561 let cu = r.cuda_graph_exec_update()?;
562 let mut error_node: cudaGraphNode_t = core::ptr::null_mut();
563 let mut result: core::ffi::c_int = 0;
564 let rc = unsafe {
569 cu(
570 self.inner.handle,
571 new_template.as_raw(),
572 &mut error_node,
573 &mut result,
574 )
575 };
576 if rc != baracuda_cuda_sys::runtime::cudaError_t::Success
577 && result == baracuda_cuda_sys::runtime::types::cudaGraphExecUpdateResult::SUCCESS
578 {
579 return Err(crate::error::Error::Status { status: rc });
580 }
581 Ok(UpdateResult {
582 result,
583 error_node: if error_node.is_null() {
584 None
585 } else {
586 Some(GraphNode { raw: error_node })
587 },
588 })
589 }
590
591 #[inline]
593 pub fn as_raw(&self) -> cudaGraphExec_t {
594 self.inner.handle
595 }
596}
597
598#[derive(Clone, Debug)]
602pub struct UpdateResult {
603 pub result: core::ffi::c_int,
607 pub error_node: Option<GraphNode>,
609}
610
611impl UpdateResult {
612 pub fn is_success(&self) -> bool {
614 self.result == baracuda_cuda_sys::runtime::types::cudaGraphExecUpdateResult::SUCCESS
615 }
616}
617
618impl Drop for GraphExecInner {
619 fn drop(&mut self) {
620 if let Ok(r) = runtime() {
621 if let Ok(cu) = r.cuda_graph_exec_destroy() {
622 let _ = unsafe { cu(self.handle) };
623 }
624 }
625 }
626}