Skip to main content

baracuda_runtime/
multicast.rs

1//! Multicast objects (CUDA 12.0+).
2//!
3//! A multicast object is a single VMM handle bound to multiple devices;
4//! writes through the handle are implicitly replicated across them.
5//! Useful for NVLink-connected GPUs (A100/H100) on all-reduce-like
6//! workloads when you don't want to go through NCCL.
7//!
8//! Not supported on older drivers — returns
9//! [`crate::Error::FeatureNotSupported`].
10
11use core::ffi::c_void;
12
13use baracuda_cuda_sys::runtime::runtime;
14use baracuda_cuda_sys::runtime::types::cudaMemGenericAllocationHandle_t;
15use baracuda_types::{Feature, supports};
16
17use crate::device::Device;
18use crate::error::{Error, Result, check};
19
20/// Properties passed to [`MulticastObject::new`]. Layout matches
21/// `cudaMulticastObjectProp`.
22#[repr(C)]
23#[derive(Copy, Clone, Debug, Default)]
24#[allow(non_camel_case_types)]
25pub struct MulticastProp {
26    /// Number of devices that will be added via
27    /// [`MulticastObject::add_device`].
28    pub num_devices: core::ffi::c_uint,
29    /// Size of the multicast window, in bytes.
30    pub size: usize,
31    /// `cudaMemAllocationHandleType` value (use `NONE` unless exporting).
32    pub handle_types: core::ffi::c_int,
33    /// Reserved — pass 0.
34    pub flags: u64,
35}
36
37fn require_multicast() -> Result<()> {
38    let installed = crate::init::driver_version()?;
39    if supports(installed, Feature::MulticastObjects) {
40        Ok(())
41    } else {
42        Err(Error::FeatureNotSupported {
43            api: "cudaMulticast*",
44            since: Feature::MulticastObjects.required_version(),
45        })
46    }
47}
48
49/// A multicast object. Drop releases it via `cudaMemRelease`.
50#[derive(Debug)]
51pub struct MulticastObject {
52    handle: cudaMemGenericAllocationHandle_t,
53}
54
55impl MulticastObject {
56    /// Create a multicast object with the given props.
57    pub fn new(prop: &MulticastProp) -> Result<Self> {
58        require_multicast()?;
59        let r = runtime()?;
60        let cu = r.cuda_multicast_create()?;
61        let mut h: cudaMemGenericAllocationHandle_t = 0;
62        check(unsafe { cu(&mut h, prop as *const MulticastProp as *const c_void) })?;
63        Ok(Self { handle: h })
64    }
65
66    /// Raw `cudaMemGenericAllocationHandle_t` for the multicast object.
67    /// Use with care — released on drop.
68    #[inline]
69    pub fn as_raw(&self) -> cudaMemGenericAllocationHandle_t {
70        self.handle
71    }
72
73    /// Add a participating device to this object.
74    pub fn add_device(&self, device: &Device) -> Result<()> {
75        let r = runtime()?;
76        let cu = r.cuda_multicast_add_device()?;
77        check(unsafe { cu(self.handle, device.ordinal()) })
78    }
79
80    /// Bind a physical-memory handle (from [`crate::vmm::MemHandle`]) at
81    /// `mc_offset` within this object, `size` bytes.
82    ///
83    /// # Safety
84    ///
85    /// `mem_handle` must be a live VMM allocation on a device that was
86    /// already added via [`Self::add_device`].
87    pub unsafe fn bind_mem(
88        &self,
89        mc_offset: usize,
90        mem_handle: cudaMemGenericAllocationHandle_t,
91        mem_offset: usize,
92        size: usize,
93        flags: u64,
94    ) -> Result<()> {
95        unsafe {
96            let r = runtime()?;
97            let cu = r.cuda_multicast_bind_mem()?;
98            check(cu(
99                self.handle,
100                mc_offset,
101                mem_handle,
102                mem_offset,
103                size,
104                flags,
105            ))
106        }
107    }
108
109    /// Bind a device address (instead of a handle).
110    ///
111    /// # Safety
112    ///
113    /// `mem_ptr` must be a mapped VMM address on a registered device.
114    pub unsafe fn bind_addr(
115        &self,
116        mc_offset: usize,
117        mem_ptr: *mut c_void,
118        size: usize,
119        flags: u64,
120    ) -> Result<()> {
121        unsafe {
122            let r = runtime()?;
123            let cu = r.cuda_multicast_bind_addr()?;
124            check(cu(self.handle, mc_offset, mem_ptr, size, flags))
125        }
126    }
127
128    /// Unbind the region `[mc_offset, mc_offset + size)` from `device`.
129    pub fn unbind(&self, device: &Device, mc_offset: usize, size: usize) -> Result<()> {
130        let r = runtime()?;
131        let cu = r.cuda_multicast_unbind()?;
132        check(unsafe { cu(self.handle, device.ordinal(), mc_offset, size) })
133    }
134}
135
136impl Drop for MulticastObject {
137    fn drop(&mut self) {
138        if let Ok(r) = runtime() {
139            if let Ok(cu) = r.cuda_mem_release() {
140                let _ = unsafe { cu(self.handle) };
141            }
142        }
143    }
144}
145
146/// Report the granularity (min alignment / min size) for multicast
147/// objects with the given props. `option`: 0 = minimum, 1 = recommended.
148pub fn multicast_granularity(prop: &MulticastProp, option: i32) -> Result<usize> {
149    require_multicast()?;
150    let r = runtime()?;
151    let cu = r.cuda_multicast_get_granularity()?;
152    let mut g: usize = 0;
153    check(unsafe {
154        cu(
155            &mut g,
156            prop as *const MulticastProp as *const c_void,
157            option,
158        )
159    })?;
160    Ok(g)
161}