use crate::error::{DriverError, IntoResult};
use crate::simt::launch::DeviceLaunchLimits;
use crate::simt::stream::CudaStream;
use std::ffi::c_int;
use std::mem::MaybeUninit;
use std::sync::atomic::{AtomicBool, AtomicU32, AtomicUsize, Ordering};
use std::sync::Arc;
#[derive(Debug)]
pub struct CudaContext {
pub(crate) cu_device: cuda_bindings::CUdevice,
pub(crate) cu_ctx: cuda_bindings::CUcontext,
pub(crate) ordinal: usize,
pub(crate) num_streams: AtomicUsize,
pub(crate) event_tracking: AtomicBool,
pub(crate) error_state: AtomicU32,
}
unsafe impl Send for CudaContext {}
unsafe impl Sync for CudaContext {}
impl Drop for CudaContext {
fn drop(&mut self) {
self.record_err(self.bind_to_thread());
let ctx = std::mem::replace(&mut self.cu_ctx, std::ptr::null_mut());
if !ctx.is_null() {
self.record_err(unsafe {
cuda_bindings::cuDevicePrimaryCtxRelease_v2(self.cu_device).result()
});
}
}
}
impl PartialEq for CudaContext {
fn eq(&self, other: &Self) -> bool {
self.cu_device == other.cu_device
&& self.cu_ctx == other.cu_ctx
&& self.ordinal == other.ordinal
}
}
impl Eq for CudaContext {}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct StreamPriorityRange {
least: i32,
greatest: i32,
}
impl StreamPriorityRange {
pub fn least(&self) -> i32 {
self.least
}
pub fn greatest(&self) -> i32 {
self.greatest
}
pub fn is_supported(&self) -> bool {
self.least != self.greatest
}
pub fn contains(&self, priority: i32) -> bool {
(self.greatest..=self.least).contains(&priority)
}
pub fn clamp(&self, priority: i32) -> i32 {
priority.clamp(self.greatest, self.least)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ContextLimit {
StackSize,
PrintfFifoSize,
MallocHeapSize,
DevRuntimeSyncDepth,
DevRuntimePendingLaunchCount,
MaxL2FetchGranularity,
PersistingL2CacheSize,
}
impl ContextLimit {
fn to_raw(self) -> cuda_bindings::CUlimit {
match self {
ContextLimit::StackSize => cuda_bindings::CUlimit_enum_CU_LIMIT_STACK_SIZE,
ContextLimit::PrintfFifoSize => cuda_bindings::CUlimit_enum_CU_LIMIT_PRINTF_FIFO_SIZE,
ContextLimit::MallocHeapSize => cuda_bindings::CUlimit_enum_CU_LIMIT_MALLOC_HEAP_SIZE,
ContextLimit::DevRuntimeSyncDepth => {
cuda_bindings::CUlimit_enum_CU_LIMIT_DEV_RUNTIME_SYNC_DEPTH
}
ContextLimit::DevRuntimePendingLaunchCount => {
cuda_bindings::CUlimit_enum_CU_LIMIT_DEV_RUNTIME_PENDING_LAUNCH_COUNT
}
ContextLimit::MaxL2FetchGranularity => {
cuda_bindings::CUlimit_enum_CU_LIMIT_MAX_L2_FETCH_GRANULARITY
}
ContextLimit::PersistingL2CacheSize => {
cuda_bindings::CUlimit_enum_CU_LIMIT_PERSISTING_L2_CACHE_SIZE
}
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum SyncPolicy {
Auto,
Spin,
Yield,
BlockingSync,
}
impl SyncPolicy {
fn to_raw(self) -> cuda_bindings::CUctx_flags_enum {
match self {
SyncPolicy::Auto => cuda_bindings::CUctx_flags_enum_CU_CTX_SCHED_AUTO,
SyncPolicy::Spin => cuda_bindings::CUctx_flags_enum_CU_CTX_SCHED_SPIN,
SyncPolicy::Yield => cuda_bindings::CUctx_flags_enum_CU_CTX_SCHED_YIELD,
SyncPolicy::BlockingSync => cuda_bindings::CUctx_flags_enum_CU_CTX_SCHED_BLOCKING_SYNC,
}
}
fn from_raw(raw: cuda_bindings::CUctx_flags_enum) -> Option<Self> {
match raw & cuda_bindings::CUctx_flags_enum_CU_CTX_SCHED_MASK {
cuda_bindings::CUctx_flags_enum_CU_CTX_SCHED_AUTO => Some(SyncPolicy::Auto),
cuda_bindings::CUctx_flags_enum_CU_CTX_SCHED_SPIN => Some(SyncPolicy::Spin),
cuda_bindings::CUctx_flags_enum_CU_CTX_SCHED_YIELD => Some(SyncPolicy::Yield),
cuda_bindings::CUctx_flags_enum_CU_CTX_SCHED_BLOCKING_SYNC => {
Some(SyncPolicy::BlockingSync)
}
_ => None,
}
}
}
impl CudaContext {
fn device_attribute(
&self,
attribute: cuda_bindings::CUdevice_attribute,
) -> Result<u32, DriverError> {
self.bind_to_thread()?;
let mut value = MaybeUninit::uninit();
unsafe {
cuda_bindings::cuDeviceGetAttribute(value.as_mut_ptr(), attribute, self.cu_device)
.result()?;
u32::try_from(value.assume_init())
.map_err(|_| DriverError(cuda_bindings::cudaError_enum_CUDA_ERROR_INVALID_VALUE))
}
}
pub fn new(ordinal: usize) -> Result<Arc<Self>, DriverError> {
unsafe { crate::init(0)? };
let cu_device = unsafe {
let mut cu_device = MaybeUninit::uninit();
cuda_bindings::cuDeviceGet(cu_device.as_mut_ptr(), ordinal as c_int).result()?;
cu_device.assume_init()
};
let cu_ctx = unsafe {
let mut cu_ctx = MaybeUninit::uninit();
cuda_bindings::cuDevicePrimaryCtxRetain(cu_ctx.as_mut_ptr(), cu_device).result()?;
cu_ctx.assume_init()
};
let ctx = Arc::new(CudaContext {
cu_device,
cu_ctx,
ordinal,
num_streams: AtomicUsize::new(0),
event_tracking: AtomicBool::new(true),
error_state: AtomicU32::new(0),
});
ctx.bind_to_thread()?;
Ok(ctx)
}
pub fn ordinal(&self) -> usize {
self.ordinal
}
pub fn cu_device(&self) -> cuda_bindings::CUdevice {
self.cu_device
}
pub fn cu_ctx(&self) -> cuda_bindings::CUcontext {
self.cu_ctx
}
pub fn bind_to_thread(&self) -> Result<(), DriverError> {
self.check_err()?;
let mut current = MaybeUninit::uninit();
unsafe {
cuda_bindings::cuCtxGetCurrent(current.as_mut_ptr()).result()?;
let current = current.assume_init();
if current.is_null() || current != self.cu_ctx {
cuda_bindings::cuCtxSetCurrent(self.cu_ctx).result()?;
}
}
Ok(())
}
pub fn synchronize(&self) -> Result<(), DriverError> {
self.bind_to_thread()?;
unsafe { cuda_bindings::cuCtxSynchronize() }.result()
}
pub fn default_stream(self: &Arc<Self>) -> Arc<CudaStream> {
Arc::new(CudaStream {
cu_stream: std::ptr::null_mut(),
ctx: self.clone(),
})
}
pub fn new_stream(self: &Arc<Self>) -> Result<Arc<CudaStream>, DriverError> {
self.create_stream(None)
}
pub fn new_stream_with_priority(
self: &Arc<Self>,
priority: i32,
) -> Result<Arc<CudaStream>, DriverError> {
self.create_stream(Some(priority))
}
fn create_stream(
self: &Arc<Self>,
priority: Option<i32>,
) -> Result<Arc<CudaStream>, DriverError> {
self.bind_to_thread()?;
let prev = self.num_streams.fetch_add(1, Ordering::Relaxed);
if prev == 0 && self.event_tracking.load(Ordering::Relaxed) {
self.synchronize()?;
}
let flags = cuda_bindings::CUstream_flags_enum_CU_STREAM_NON_BLOCKING;
let mut cu_stream = MaybeUninit::uninit();
let cu_stream = unsafe {
match priority {
Some(priority) => cuda_bindings::cuStreamCreateWithPriority(
cu_stream.as_mut_ptr(),
flags,
priority,
),
None => cuda_bindings::cuStreamCreate(cu_stream.as_mut_ptr(), flags),
}
.result()?;
cu_stream.assume_init()
};
Ok(Arc::new(CudaStream {
cu_stream,
ctx: self.clone(),
}))
}
pub fn stream_priority_range(&self) -> Result<StreamPriorityRange, DriverError> {
self.bind_to_thread()?;
let mut least = MaybeUninit::uninit();
let mut greatest = MaybeUninit::uninit();
unsafe {
cuda_bindings::cuCtxGetStreamPriorityRange(least.as_mut_ptr(), greatest.as_mut_ptr())
.result()?;
Ok(StreamPriorityRange {
least: least.assume_init(),
greatest: greatest.assume_init(),
})
}
}
pub fn device_name(&self) -> Result<String, DriverError> {
self.bind_to_thread()?;
let mut buf = [0; 256];
unsafe {
cuda_bindings::cuDeviceGetName(buf.as_mut_ptr(), buf.len() as c_int, self.cu_device)
.result()?;
}
let bytes: Vec<u8> = buf
.iter()
.take_while(|&&c| c != 0)
.map(|&c| c as u8)
.collect();
Ok(String::from_utf8_lossy(&bytes).into_owned())
}
pub fn compute_capability(&self) -> Result<(i32, i32), DriverError> {
self.bind_to_thread()?;
let mut major = MaybeUninit::uninit();
let mut minor = MaybeUninit::uninit();
unsafe {
cuda_bindings::cuDeviceGetAttribute(
major.as_mut_ptr(),
cuda_bindings::CUdevice_attribute_enum_CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MAJOR,
self.cu_device,
)
.result()?;
cuda_bindings::cuDeviceGetAttribute(
minor.as_mut_ptr(),
cuda_bindings::CUdevice_attribute_enum_CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MINOR,
self.cu_device,
)
.result()?;
Ok((major.assume_init(), minor.assume_init()))
}
}
pub fn launch_limits(&self) -> Result<DeviceLaunchLimits, DriverError> {
Ok(DeviceLaunchLimits {
max_threads_per_block: self.device_attribute(
cuda_bindings::CUdevice_attribute_enum_CU_DEVICE_ATTRIBUTE_MAX_THREADS_PER_BLOCK,
)?,
max_block_dim: (
self.device_attribute(
cuda_bindings::CUdevice_attribute_enum_CU_DEVICE_ATTRIBUTE_MAX_BLOCK_DIM_X,
)?,
self.device_attribute(
cuda_bindings::CUdevice_attribute_enum_CU_DEVICE_ATTRIBUTE_MAX_BLOCK_DIM_Y,
)?,
self.device_attribute(
cuda_bindings::CUdevice_attribute_enum_CU_DEVICE_ATTRIBUTE_MAX_BLOCK_DIM_Z,
)?,
),
max_grid_dim: (
self.device_attribute(
cuda_bindings::CUdevice_attribute_enum_CU_DEVICE_ATTRIBUTE_MAX_GRID_DIM_X,
)?,
self.device_attribute(
cuda_bindings::CUdevice_attribute_enum_CU_DEVICE_ATTRIBUTE_MAX_GRID_DIM_Y,
)?,
self.device_attribute(
cuda_bindings::CUdevice_attribute_enum_CU_DEVICE_ATTRIBUTE_MAX_GRID_DIM_Z,
)?,
),
max_shared_memory_per_block: self.device_attribute(
cuda_bindings::CUdevice_attribute_enum_CU_DEVICE_ATTRIBUTE_MAX_SHARED_MEMORY_PER_BLOCK,
)?,
})
}
pub fn max_opt_in_shared_memory_per_block(&self) -> Result<u32, DriverError> {
self.device_attribute(
cuda_bindings::CUdevice_attribute_enum_CU_DEVICE_ATTRIBUTE_MAX_SHARED_MEMORY_PER_BLOCK_OPTIN,
)
}
pub fn supports_cooperative_launch(&self) -> Result<bool, DriverError> {
self.device_attribute(
cuda_bindings::CUdevice_attribute_enum_CU_DEVICE_ATTRIBUTE_COOPERATIVE_LAUNCH,
)
.map(|value| value != 0)
}
pub fn supports_cluster_launch(&self) -> Result<bool, DriverError> {
self.device_attribute(
cuda_bindings::CUdevice_attribute_enum_CU_DEVICE_ATTRIBUTE_CLUSTER_LAUNCH,
)
.map(|value| value != 0)
}
pub fn multiprocessor_count(&self) -> Result<u32, DriverError> {
self.device_attribute(
cuda_bindings::CUdevice_attribute_enum_CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT,
)
}
pub fn max_threads_per_multiprocessor(&self) -> Result<u32, DriverError> {
self.device_attribute(
cuda_bindings::CUdevice_attribute_enum_CU_DEVICE_ATTRIBUTE_MAX_THREADS_PER_MULTIPROCESSOR,
)
}
pub fn limit(&self, limit: ContextLimit) -> Result<usize, DriverError> {
self.bind_to_thread()?;
let mut value = MaybeUninit::uninit();
unsafe {
cuda_bindings::cuCtxGetLimit(value.as_mut_ptr(), limit.to_raw()).result()?;
Ok(value.assume_init())
}
}
pub fn set_limit(&self, limit: ContextLimit, value: usize) -> Result<(), DriverError> {
self.bind_to_thread()?;
unsafe { cuda_bindings::cuCtxSetLimit(limit.to_raw(), value) }.result()
}
pub fn stack_size(&self) -> Result<usize, DriverError> {
self.limit(ContextLimit::StackSize)
}
pub fn set_stack_size(&self, bytes: usize) -> Result<(), DriverError> {
self.set_limit(ContextLimit::StackSize, bytes)
}
pub fn set_sync_policy(&self, policy: SyncPolicy) -> Result<(), DriverError> {
self.bind_to_thread()?;
unsafe {
let mut current = MaybeUninit::uninit();
let mut active = MaybeUninit::uninit();
cuda_bindings::cuDevicePrimaryCtxGetState(
self.cu_device,
current.as_mut_ptr(),
active.as_mut_ptr(),
)
.result()?;
let current = current.assume_init();
let new_flags =
(current & !cuda_bindings::CUctx_flags_enum_CU_CTX_SCHED_MASK) | policy.to_raw();
cuda_bindings::cuDevicePrimaryCtxSetFlags_v2(self.cu_device, new_flags).result()
}
}
pub fn sync_policy(&self) -> Result<Option<SyncPolicy>, DriverError> {
self.bind_to_thread()?;
unsafe {
let mut flags = MaybeUninit::uninit();
let mut active = MaybeUninit::uninit();
cuda_bindings::cuDevicePrimaryCtxGetState(
self.cu_device,
flags.as_mut_ptr(),
active.as_mut_ptr(),
)
.result()?;
Ok(SyncPolicy::from_raw(flags.assume_init()))
}
}
pub fn check_err(&self) -> Result<(), DriverError> {
let error_state = self.error_state.swap(0, Ordering::Relaxed);
if error_state == 0 {
Ok(())
} else {
Err(DriverError(error_state))
}
}
pub fn record_err<T>(&self, result: Result<T, DriverError>) {
if let Err(err) = result {
self.error_state.store(err.0, Ordering::Relaxed)
}
}
}