use oxicuda_driver::error::{CudaError, CudaResult};
use oxicuda_driver::stream::Stream;
use crate::kernel::{Kernel, KernelArgs};
use crate::params::LaunchParams;
#[derive(Debug)]
pub struct CooperativeLaunch;
impl CooperativeLaunch {
pub fn launch<A: KernelArgs>(
kernel: &Kernel,
params: &LaunchParams,
stream: &Stream,
args: &A,
) -> CudaResult<()> {
if params.grid.x == 0
|| params.grid.y == 0
|| params.grid.z == 0
|| params.block.x == 0
|| params.block.y == 0
|| params.block.z == 0
{
return Err(CudaError::InvalidValue);
}
let config = oxicuda_driver::cooperative_launch::CooperativeLaunchConfig {
grid_dim: (params.grid.x, params.grid.y, params.grid.z),
block_dim: (params.block.x, params.block.y, params.block.z),
shared_mem_bytes: params.shared_mem_bytes,
stream: Some(stream.raw()),
};
let param_ptrs = args.as_param_ptrs();
oxicuda_driver::cooperative_launch::cooperative_launch(
kernel.function(),
&config,
¶m_ptrs,
)
}
pub fn max_active_blocks(
kernel: &Kernel,
block_size: u32,
dynamic_smem: usize,
) -> CudaResult<u32> {
let result = Self::max_active_blocks_inner(kernel, block_size as i32, dynamic_smem)?;
Ok(result as u32)
}
fn max_active_blocks_inner(
kernel: &Kernel,
block_size: i32,
dynamic_smem: usize,
) -> CudaResult<i32> {
kernel
.function()
.max_active_blocks_per_sm(block_size, dynamic_smem)
}
pub fn optimal_block_size(kernel: &Kernel, dynamic_smem: usize) -> CudaResult<(i32, i32)> {
kernel.function().optimal_block_size(dynamic_smem)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::grid::Dim3;
#[test]
fn cooperative_launch_struct_exists() {
let _: fn(&Kernel, &LaunchParams, &Stream, &(u64, u32)) -> CudaResult<()> =
CooperativeLaunch::launch;
}
#[test]
fn max_active_blocks_signature_compiles() {
let _: fn(&Kernel, u32, usize) -> CudaResult<u32> = CooperativeLaunch::max_active_blocks;
}
#[test]
fn optimal_block_size_signature_compiles() {
let _: fn(&Kernel, usize) -> CudaResult<(i32, i32)> = CooperativeLaunch::optimal_block_size;
}
#[test]
fn launch_rejects_zero_grid_x() {
let params = LaunchParams {
grid: Dim3::new(0, 1, 1),
block: Dim3::x(256),
shared_mem_bytes: 0,
};
assert_eq!(params.grid.x, 0);
}
#[test]
fn launch_rejects_zero_block_y() {
let params = LaunchParams {
grid: Dim3::x(4),
block: Dim3::new(256, 0, 1),
shared_mem_bytes: 0,
};
assert_eq!(params.block.y, 0);
}
#[test]
fn cooperative_launch_is_send() {
fn assert_send<T: Send>() {}
assert_send::<CooperativeLaunch>();
}
#[test]
fn cooperative_launch_is_sync() {
fn assert_sync<T: Sync>() {}
assert_sync::<CooperativeLaunch>();
}
#[test]
fn cooperative_dim3_total_nonzero() {
let d = Dim3::new(4, 2, 1);
assert_eq!(d.total(), 8);
assert!(d.total() > 0);
}
#[test]
fn cooperative_config_valid_fields() {
let params = LaunchParams {
grid: Dim3::new(4, 1, 1),
block: Dim3::new(256, 1, 1),
shared_mem_bytes: 1024,
};
assert_eq!(params.grid.x, 4);
assert_eq!(params.block.x, 256);
assert_eq!(params.shared_mem_bytes, 1024);
}
#[test]
fn cooperative_max_blocks_constraint_signature() {
let _: fn(&Kernel, &LaunchParams, &Stream, &(u64, u32)) -> CudaResult<()> =
CooperativeLaunch::launch;
}
#[test]
fn cooperative_debug_display() {
let coop = CooperativeLaunch;
let dbg = format!("{coop:?}");
assert!(
dbg.contains("CooperativeLaunch"),
"Debug output must contain type name, got: {dbg}"
);
}
#[cfg(feature = "gpu-tests")]
const NOOP_PTX: &str = "\
.version 7.0
.target sm_70
.address_size 64
.visible .entry noop_kernel()
{
ret;
}
";
#[cfg(feature = "gpu-tests")]
#[test]
fn cooperative_launch_succeeds_across_all_sms() {
use std::sync::Arc;
let Ok(dev) = oxicuda_driver::device::Device::get(0) else {
return;
};
let ctx = match oxicuda_driver::context::Context::new(&dev) {
Ok(c) => Arc::new(c),
Err(_) => return,
};
let stream = match Stream::new(&ctx) {
Ok(s) => s,
Err(_) => return,
};
let module = match oxicuda_driver::module::Module::from_ptx(NOOP_PTX) {
Ok(m) => Arc::new(m),
Err(_) => return,
};
let kernel = match Kernel::from_module(module, "noop_kernel") {
Ok(k) => k,
Err(_) => return,
};
let block_size: u32 = 128;
let max_per_sm = match CooperativeLaunch::max_active_blocks(&kernel, block_size, 0) {
Ok(m) if m > 0 => m,
_ => return,
};
let sm_count = match dev.multiprocessor_count() {
Ok(c) if c > 0 => c as u32,
_ => return,
};
let grid_x = max_per_sm * sm_count;
let params = LaunchParams::new(grid_x, block_size);
let result = CooperativeLaunch::launch(&kernel, ¶ms, &stream, &());
assert!(
result.is_ok(),
"cooperative launch spanning all {sm_count} SMs \
({max_per_sm} blocks/SM) must succeed, got {result:?}"
);
stream
.synchronize()
.expect("stream sync after cooperative launch");
}
}