pub mod context;
pub mod device_buffer;
pub mod embedded;
pub mod event;
pub mod launch;
pub mod memory;
pub mod module;
pub mod peer;
pub mod pinned_host_buffer;
pub mod stream;
pub mod vmm;
pub use context::{ContextLimit, CudaContext, StreamPriorityRange, SyncPolicy};
pub use cuda_core_derive::DeviceCopy;
pub use device_buffer::{DeviceBuffer, DeviceCopy};
pub use embedded::{EmbeddedModule, EmbeddedModuleError};
pub use event::CudaEvent;
pub use launch::{
BlockRequirement, DeviceLaunchLimits, DynamicSharedMemoryRequirement, KernelLaunchConfig,
KernelLaunchContract, LaunchAxis, LaunchConfig, LaunchConfig1D, LaunchConfig2D, LaunchConfig3D,
LaunchContractError, LaunchContractSpec, LaunchDimension, PreparedLaunch,
};
pub use module::{ConstantHandle, CudaFunction, CudaModule};
pub use pinned_host_buffer::PinnedHostBuffer;
pub use stream::CudaStream;
pub use crate::error::{DriverError, IntoResult};
pub use crate::init;
pub use crate::launch_kernel;
pub use cuda_bindings as sys;
pub use oxide_artifacts as artifacts;
#[inline]
pub unsafe fn launch_kernel_on_stream(
func: &CudaFunction,
grid_dim: (u32, u32, u32),
block_dim: (u32, u32, u32),
shared_mem_bytes: u32,
stream: &CudaStream,
kernel_params: &mut [*mut std::ffi::c_void],
) -> Result<(), DriverError> {
stream.context().bind_to_thread()?;
unsafe {
launch_kernel(
func.cu_function(),
grid_dim,
block_dim,
shared_mem_bytes,
stream.cu_stream(),
kernel_params,
)
}
}
#[inline]
pub unsafe fn launch_kernel_ex(
func: cuda_bindings::CUfunction,
grid_dim: (u32, u32, u32),
block_dim: (u32, u32, u32),
shared_mem_bytes: u32,
cluster_dim: (u32, u32, u32),
stream: cuda_bindings::CUstream,
kernel_params: &mut [*mut std::ffi::c_void],
) -> Result<(), DriverError> {
let mut cluster_attr: cuda_bindings::CUlaunchAttribute_st = unsafe { std::mem::zeroed() };
unsafe {
let base = &mut cluster_attr as *mut _ as *mut u8;
(base as *mut u32)
.write(cuda_bindings::CUlaunchAttributeID_enum_CU_LAUNCH_ATTRIBUTE_CLUSTER_DIMENSION);
let dim_ptr = base.add(8) as *mut u32;
dim_ptr.write(cluster_dim.0);
dim_ptr.add(1).write(cluster_dim.1);
dim_ptr.add(2).write(cluster_dim.2);
}
let config = cuda_bindings::CUlaunchConfig_st {
gridDimX: grid_dim.0,
gridDimY: grid_dim.1,
gridDimZ: grid_dim.2,
blockDimX: block_dim.0,
blockDimY: block_dim.1,
blockDimZ: block_dim.2,
sharedMemBytes: shared_mem_bytes,
hStream: stream,
attrs: &mut cluster_attr,
numAttrs: 1,
};
unsafe {
cuda_bindings::cuLaunchKernelEx(
&config,
func,
kernel_params.as_mut_ptr(),
std::ptr::null_mut(),
)
}
.result()
}
#[inline]
pub unsafe fn launch_kernel_ex_on_stream(
func: &CudaFunction,
grid_dim: (u32, u32, u32),
block_dim: (u32, u32, u32),
shared_mem_bytes: u32,
cluster_dim: (u32, u32, u32),
stream: &CudaStream,
kernel_params: &mut [*mut std::ffi::c_void],
) -> Result<(), DriverError> {
stream.context().bind_to_thread()?;
unsafe {
launch_kernel_ex(
func.cu_function(),
grid_dim,
block_dim,
shared_mem_bytes,
cluster_dim,
stream.cu_stream(),
kernel_params,
)
}
}
#[inline]
pub unsafe fn launch_kernel_cooperative(
func: cuda_bindings::CUfunction,
grid_dim: (u32, u32, u32),
block_dim: (u32, u32, u32),
shared_mem_bytes: u32,
stream: cuda_bindings::CUstream,
kernel_params: &mut [*mut std::ffi::c_void],
) -> Result<(), DriverError> {
let mut coop_attr: cuda_bindings::CUlaunchAttribute_st = unsafe { std::mem::zeroed() };
unsafe {
let base = &mut coop_attr as *mut _ as *mut u8;
(base as *mut u32)
.write(cuda_bindings::CUlaunchAttributeID_enum_CU_LAUNCH_ATTRIBUTE_COOPERATIVE);
let val_ptr = base.add(8) as *mut i32;
val_ptr.write(1);
}
let config = cuda_bindings::CUlaunchConfig_st {
gridDimX: grid_dim.0,
gridDimY: grid_dim.1,
gridDimZ: grid_dim.2,
blockDimX: block_dim.0,
blockDimY: block_dim.1,
blockDimZ: block_dim.2,
sharedMemBytes: shared_mem_bytes,
hStream: stream,
attrs: &mut coop_attr,
numAttrs: 1,
};
unsafe {
cuda_bindings::cuLaunchKernelEx(
&config,
func,
kernel_params.as_mut_ptr(),
std::ptr::null_mut(),
)
}
.result()
}
#[inline]
pub unsafe fn launch_kernel_cooperative_on_stream(
func: &CudaFunction,
grid_dim: (u32, u32, u32),
block_dim: (u32, u32, u32),
shared_mem_bytes: u32,
stream: &CudaStream,
kernel_params: &mut [*mut std::ffi::c_void],
) -> Result<(), DriverError> {
stream.context().bind_to_thread()?;
unsafe {
launch_kernel_cooperative(
func.cu_function(),
grid_dim,
block_dim,
shared_mem_bytes,
stream.cu_stream(),
kernel_params,
)
}
}
#[inline]
pub unsafe fn launch_kernel_ex_cooperative(
func: cuda_bindings::CUfunction,
grid_dim: (u32, u32, u32),
block_dim: (u32, u32, u32),
shared_mem_bytes: u32,
cluster_dim: (u32, u32, u32),
stream: cuda_bindings::CUstream,
kernel_params: &mut [*mut std::ffi::c_void],
) -> Result<(), DriverError> {
let mut attrs: [cuda_bindings::CUlaunchAttribute_st; 2] = unsafe { std::mem::zeroed() };
unsafe {
let base = &mut attrs[0] as *mut _ as *mut u8;
(base as *mut u32)
.write(cuda_bindings::CUlaunchAttributeID_enum_CU_LAUNCH_ATTRIBUTE_CLUSTER_DIMENSION);
let dim_ptr = base.add(8) as *mut u32;
dim_ptr.write(cluster_dim.0);
dim_ptr.add(1).write(cluster_dim.1);
dim_ptr.add(2).write(cluster_dim.2);
let base = &mut attrs[1] as *mut _ as *mut u8;
(base as *mut u32)
.write(cuda_bindings::CUlaunchAttributeID_enum_CU_LAUNCH_ATTRIBUTE_COOPERATIVE);
let val_ptr = base.add(8) as *mut i32;
val_ptr.write(1);
}
let config = cuda_bindings::CUlaunchConfig_st {
gridDimX: grid_dim.0,
gridDimY: grid_dim.1,
gridDimZ: grid_dim.2,
blockDimX: block_dim.0,
blockDimY: block_dim.1,
blockDimZ: block_dim.2,
sharedMemBytes: shared_mem_bytes,
hStream: stream,
attrs: attrs.as_mut_ptr(),
numAttrs: 2,
};
unsafe {
cuda_bindings::cuLaunchKernelEx(
&config,
func,
kernel_params.as_mut_ptr(),
std::ptr::null_mut(),
)
}
.result()
}
#[inline]
pub unsafe fn launch_kernel_ex_cooperative_on_stream(
func: &CudaFunction,
grid_dim: (u32, u32, u32),
block_dim: (u32, u32, u32),
shared_mem_bytes: u32,
cluster_dim: (u32, u32, u32),
stream: &CudaStream,
kernel_params: &mut [*mut std::ffi::c_void],
) -> Result<(), DriverError> {
stream.context().bind_to_thread()?;
unsafe {
launch_kernel_ex_cooperative(
func.cu_function(),
grid_dim,
block_dim,
shared_mem_bytes,
cluster_dim,
stream.cu_stream(),
kernel_params,
)
}
}