use crate::error::{DriverError, IntoResult};
use cuda_bindings::CUdeviceptr;
use std::mem::MaybeUninit;
unsafe fn set_mem_location_device(
loc: &mut cuda_bindings::CUmemLocation_st,
device: cuda_bindings::CUdevice,
) {
loc.type_ = cuda_bindings::CUmemLocationType_enum_CU_MEM_LOCATION_TYPE_DEVICE;
unsafe {
let base = loc as *mut _ as *mut u8;
(base.add(4) as *mut i32).write(device);
}
}
pub struct PhysicalAllocation {
handle: cuda_bindings::CUmemGenericAllocationHandle,
size: usize,
device: cuda_bindings::CUdevice,
}
impl PhysicalAllocation {
pub fn new(device: cuda_bindings::CUdevice, size: usize) -> Result<Self, DriverError> {
let mut prop: cuda_bindings::CUmemAllocationProp_st = unsafe { std::mem::zeroed() };
prop.type_ = cuda_bindings::CUmemAllocationType_enum_CU_MEM_ALLOCATION_TYPE_PINNED;
unsafe { set_mem_location_device(&mut prop.location, device) };
let mut handle = MaybeUninit::uninit();
unsafe {
cuda_bindings::cuMemCreate(handle.as_mut_ptr(), size, &prop, 0).result()?;
Ok(Self {
handle: handle.assume_init(),
size,
device,
})
}
}
pub fn handle(&self) -> cuda_bindings::CUmemGenericAllocationHandle {
self.handle
}
pub fn size(&self) -> usize {
self.size
}
pub fn device(&self) -> cuda_bindings::CUdevice {
self.device
}
}
impl Drop for PhysicalAllocation {
fn drop(&mut self) {
unsafe {
let _ = cuda_bindings::cuMemRelease(self.handle).result();
}
}
}
pub struct VirtualReservation {
base: CUdeviceptr,
size: usize,
}
impl VirtualReservation {
pub fn new(size: usize, alignment: usize) -> Result<Self, DriverError> {
let mut base = MaybeUninit::uninit();
unsafe {
cuda_bindings::cuMemAddressReserve(base.as_mut_ptr(), size, alignment, 0, 0)
.result()?;
Ok(Self {
base: base.assume_init(),
size,
})
}
}
pub fn base(&self) -> CUdeviceptr {
self.base
}
pub fn size(&self) -> usize {
self.size
}
}
impl Drop for VirtualReservation {
fn drop(&mut self) {
unsafe {
let _ = cuda_bindings::cuMemAddressFree(self.base, self.size).result();
}
}
}
pub struct Mapping {
va: CUdeviceptr,
size: usize,
}
impl Mapping {
pub fn new(
va: CUdeviceptr,
size: usize,
phys: &PhysicalAllocation,
offset: usize,
) -> Result<Self, DriverError> {
unsafe {
cuda_bindings::cuMemMap(va, size, offset, phys.handle(), 0).result()?;
}
Ok(Self { va, size })
}
#[cfg(cuda_has_multicast)]
pub fn new_multicast(
va: CUdeviceptr,
size: usize,
multicast: &MulticastObject,
offset: usize,
) -> Result<Self, DriverError> {
unsafe {
cuda_bindings::cuMemMap(va, size, offset, multicast.handle(), 0).result()?;
}
Ok(Self { va, size })
}
pub fn va(&self) -> CUdeviceptr {
self.va
}
pub fn size(&self) -> usize {
self.size
}
}
impl Drop for Mapping {
fn drop(&mut self) {
unsafe {
let _ = cuda_bindings::cuMemUnmap(self.va, self.size).result();
}
}
}
pub fn set_access(
va: CUdeviceptr,
size: usize,
devices: &[cuda_bindings::CUdevice],
) -> Result<(), DriverError> {
let descs: Vec<cuda_bindings::CUmemAccessDesc_st> = devices
.iter()
.map(|&dev| {
let mut desc: cuda_bindings::CUmemAccessDesc_st = unsafe { std::mem::zeroed() };
unsafe { set_mem_location_device(&mut desc.location, dev) };
desc.flags = cuda_bindings::CUmemAccess_flags_enum_CU_MEM_ACCESS_FLAGS_PROT_READWRITE;
desc
})
.collect();
unsafe { cuda_bindings::cuMemSetAccess(va, size, descs.as_ptr(), descs.len()) }.result()
}
pub fn allocation_granularity(device: cuda_bindings::CUdevice) -> Result<usize, DriverError> {
let mut prop: cuda_bindings::CUmemAllocationProp_st = unsafe { std::mem::zeroed() };
prop.type_ = cuda_bindings::CUmemAllocationType_enum_CU_MEM_ALLOCATION_TYPE_PINNED;
unsafe { set_mem_location_device(&mut prop.location, device) };
let mut granularity = MaybeUninit::uninit();
unsafe {
cuda_bindings::cuMemGetAllocationGranularity(
granularity.as_mut_ptr(),
&prop,
cuda_bindings::CUmemAllocationGranularity_flags_enum_CU_MEM_ALLOC_GRANULARITY_MINIMUM,
)
.result()?;
Ok(granularity.assume_init())
}
}
pub fn align_size(size: usize, granularity: usize) -> usize {
(size + granularity - 1) & !(granularity - 1)
}
#[cfg(cuda_has_multicast)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum MulticastGranularity {
Minimum,
Recommended,
}
#[cfg(cuda_has_multicast)]
impl MulticastGranularity {
fn to_flag(self) -> cuda_bindings::CUmulticastGranularity_flags {
match self {
MulticastGranularity::Minimum => {
cuda_bindings::CUmulticastGranularity_flags_enum_CU_MULTICAST_GRANULARITY_MINIMUM
}
MulticastGranularity::Recommended => {
cuda_bindings::CUmulticastGranularity_flags_enum_CU_MULTICAST_GRANULARITY_RECOMMENDED
}
}
}
}
#[cfg(cuda_has_multicast)]
fn multicast_prop(num_devices: u32, size: usize) -> cuda_bindings::CUmulticastObjectProp_st {
let mut prop: cuda_bindings::CUmulticastObjectProp_st = unsafe { std::mem::zeroed() };
prop.numDevices = num_devices;
prop.size = size;
prop
}
#[cfg(cuda_has_multicast)]
pub fn multicast_supported(device: cuda_bindings::CUdevice) -> Result<bool, DriverError> {
let mut value = MaybeUninit::uninit();
let status = unsafe {
cuda_bindings::cuDeviceGetAttribute(
value.as_mut_ptr(),
cuda_bindings::CUdevice_attribute_enum_CU_DEVICE_ATTRIBUTE_MULTICAST_SUPPORTED,
device,
)
};
if status == cuda_bindings::cudaError_enum_CUDA_ERROR_INVALID_VALUE {
return Ok(false);
}
status.result()?;
Ok(unsafe { value.assume_init() } != 0)
}
#[cfg(not(cuda_has_multicast))]
pub fn multicast_supported(_device: cuda_bindings::CUdevice) -> Result<bool, DriverError> {
Ok(false)
}
#[cfg(cuda_has_multicast)]
pub fn multicast_granularity(
num_devices: u32,
size: usize,
granularity: MulticastGranularity,
) -> Result<usize, DriverError> {
let prop = multicast_prop(num_devices, size);
let mut value = MaybeUninit::uninit();
unsafe {
cuda_bindings::cuMulticastGetGranularity(value.as_mut_ptr(), &prop, granularity.to_flag())
.result()?;
Ok(value.assume_init())
}
}
#[cfg(cuda_has_multicast)]
pub struct MulticastObject {
handle: cuda_bindings::CUmemGenericAllocationHandle,
size: usize,
num_devices: u32,
}
#[cfg(cuda_has_multicast)]
impl MulticastObject {
pub fn new(num_devices: u32, size: usize) -> Result<Self, DriverError> {
let prop = multicast_prop(num_devices, size);
let mut handle = MaybeUninit::uninit();
unsafe {
cuda_bindings::cuMulticastCreate(handle.as_mut_ptr(), &prop).result()?;
Ok(Self {
handle: handle.assume_init(),
size,
num_devices,
})
}
}
pub fn add_device(&self, device: cuda_bindings::CUdevice) -> Result<(), DriverError> {
unsafe { cuda_bindings::cuMulticastAddDevice(self.handle, device).result() }
}
pub fn bind_mem(
&self,
mc_offset: usize,
phys: &PhysicalAllocation,
mem_offset: usize,
size: usize,
) -> Result<MulticastBinding, DriverError> {
unsafe {
cuda_bindings::cuMulticastBindMem(
self.handle,
mc_offset,
phys.handle(),
mem_offset,
size,
0,
)
.result()?;
}
Ok(MulticastBinding {
mc_handle: self.handle,
device: phys.device(),
mc_offset,
size,
})
}
pub fn handle(&self) -> cuda_bindings::CUmemGenericAllocationHandle {
self.handle
}
pub fn size(&self) -> usize {
self.size
}
pub fn num_devices(&self) -> u32 {
self.num_devices
}
}
#[cfg(cuda_has_multicast)]
impl Drop for MulticastObject {
fn drop(&mut self) {
unsafe {
let _ = cuda_bindings::cuMemRelease(self.handle).result();
}
}
}
#[cfg(cuda_has_multicast)]
pub struct MulticastBinding {
mc_handle: cuda_bindings::CUmemGenericAllocationHandle,
device: cuda_bindings::CUdevice,
mc_offset: usize,
size: usize,
}
#[cfg(cuda_has_multicast)]
impl MulticastBinding {
pub fn device(&self) -> cuda_bindings::CUdevice {
self.device
}
}
#[cfg(cuda_has_multicast)]
impl Drop for MulticastBinding {
fn drop(&mut self) {
unsafe {
let _ = cuda_bindings::cuMulticastUnbind(
self.mc_handle,
self.device,
self.mc_offset,
self.size,
)
.result();
}
}
}