atlas_core/
kernel.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2
3use std::sync::Arc;
4
5use cudarc::driver::{
6    CudaContext, CudaFunction, CudaModule, CudaStream, LaunchConfig, PushKernelArg,
7};
8use cudarc::nvrtc::Ptx;
9
10use crate::error::{AtlasError, Result};
11
12/// A loaded CUDA kernel module with cached function handles.
13///
14/// Wraps cudarc's CudaModule to provide a clean interface for loading
15/// PTX source (compiled at build time by build.rs) and launching kernels.
16pub struct KernelModule {
17    module: Arc<CudaModule>,
18}
19
20impl KernelModule {
21    /// Load a PTX module from source string into the given CUDA context.
22    ///
23    /// The PTX source should be compiled at build time via nvcc --ptx
24    /// and embedded with include_str!.
25    pub fn from_ptx_src(ctx: &Arc<CudaContext>, ptx_src: &str) -> Result<Self> {
26        let ptx = Ptx::from_src(ptx_src);
27        let module = ctx
28            .load_module(ptx)
29            .map_err(|e| AtlasError::ModuleLoad(format!("PTX load failed: {e}")))?;
30        Ok(Self { module })
31    }
32
33    /// Get a kernel function handle by name.
34    pub fn get_function(&self, name: &str) -> Result<CudaFunction> {
35        self.module
36            .load_function(name)
37            .map_err(|e| AtlasError::ModuleLoad(format!("Function '{name}' not found: {e}")))
38    }
39}
40
41/// Launch configuration helper for SM121.
42pub fn launch_config(n: u32, block_size: u32) -> LaunchConfig {
43    LaunchConfig {
44        grid_dim: (n.div_ceil(block_size), 1, 1),
45        block_dim: (block_size, 1, 1),
46        shared_mem_bytes: 0,
47    }
48}
49
50/// Launch vector_add kernel: `C[i] = A[i] + B[i]`.
51///
52/// This uses the safe cudarc launch_builder API with u64 device pointers.
53/// Since u64 implements DeviceRepr, no FFI conversion is needed.
54///
55/// # Safety
56///
57/// All pointers must be valid CUDA device pointers to f32 arrays of length >= n.
58pub unsafe fn launch_vector_add(
59    stream: &Arc<CudaStream>,
60    func: &CudaFunction,
61    a_ptr: u64,
62    b_ptr: u64,
63    c_ptr: u64,
64    n: u32,
65) -> Result<()> {
66    let cfg = launch_config(n, 256);
67    unsafe {
68        stream
69            .launch_builder(func)
70            .arg(&a_ptr)
71            .arg(&b_ptr)
72            .arg(&c_ptr)
73            .arg(&n)
74            .launch(cfg)
75            .map_err(|e| AtlasError::KernelLaunch(format!("vector_add launch failed: {e}")))?;
76    }
77    Ok(())
78}