Skip to main content

launchbound_bench/
cuda.rs

1//! Minimal CUDA driver API surface, loaded at runtime with dlopen so the
2//! crate builds (and its tests run) on machines with no CUDA at all —
3//! including CI and the Mac. Only the box can actually produce a timing.
4
5#![allow(clippy::missing_safety_doc)]
6
7use libloading::Library;
8use std::ffi::c_void;
9
10pub type CUresult = i32;
11type CUdeviceptr = u64;
12
13macro_rules! driver_api {
14    ($( $name:ident : fn( $($arg:ty),* ) ; )*) => {
15        // Fields carry the C symbol names verbatim.
16        #[allow(non_snake_case)]
17        pub struct Cuda {
18            _lib: Library,
19            $( $name: unsafe extern "C" fn($($arg),*) -> CUresult, )*
20        }
21
22        impl Cuda {
23            /// dlopen libcuda and resolve the surface. Errors on machines
24            /// without a driver — that is the honest answer there.
25            pub fn load() -> Result<Self, String> {
26                let lib = ["libcuda.so.1", "libcuda.so"]
27                    .iter()
28                    .find_map(|n| unsafe { Library::new(n).ok() })
29                    .ok_or_else(|| {
30                        "libcuda not found: benchmarks need an NVIDIA driver".to_string()
31                    })?;
32                unsafe {
33                    Ok(Cuda {
34                        $( $name: *lib
35                            .get(concat!(stringify!($name), "\0").as_bytes())
36                            .map_err(|e| format!("missing {}: {e}", stringify!($name)))?, )*
37                        _lib: lib,
38                    })
39                }
40            }
41        }
42    };
43}
44
45driver_api! {
46    cuInit: fn(u32);
47    cuDriverGetVersion: fn(*mut i32);
48    cuDeviceGet: fn(*mut i32, i32);
49    cuDeviceGetName: fn(*mut u8, i32, i32);
50    cuDeviceGetAttribute: fn(*mut i32, i32, i32);
51    cuCtxCreate_v2: fn(*mut *mut c_void, u32, i32);
52    cuCtxDestroy_v2: fn(*mut c_void);
53    cuCtxSynchronize: fn();
54    cuModuleLoadData: fn(*mut *mut c_void, *const c_void);
55    cuModuleUnload: fn(*mut c_void);
56    cuModuleGetFunction: fn(*mut *mut c_void, *mut c_void, *const u8);
57    cuMemAlloc_v2: fn(*mut CUdeviceptr, usize);
58    cuMemFree_v2: fn(CUdeviceptr);
59    cuMemcpyHtoD_v2: fn(CUdeviceptr, *const c_void, usize);
60    cuMemcpyDtoH_v2: fn(*mut c_void, CUdeviceptr, usize);
61    cuLaunchKernel: fn(*mut c_void, u32, u32, u32, u32, u32, u32, u32, *mut c_void, *mut *mut c_void, *mut *mut c_void);
62    cuEventCreate: fn(*mut *mut c_void, u32);
63    cuEventDestroy_v2: fn(*mut c_void);
64    cuEventRecord: fn(*mut c_void, *mut c_void);
65    cuEventSynchronize: fn(*mut c_void);
66    cuEventElapsedTime: fn(*mut f32, *mut c_void, *mut c_void);
67}
68
69const CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MAJOR: i32 = 75;
70const CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MINOR: i32 = 76;
71
72fn check(what: &str, code: CUresult) -> Result<(), String> {
73    if code == 0 {
74        Ok(())
75    } else {
76        Err(format!("{what} failed: CUresult {code}"))
77    }
78}
79
80pub struct Device {
81    cuda: Cuda,
82    ctx: *mut c_void,
83    pub name: String,
84    pub cc: String,
85    pub driver_version: String,
86}
87
88pub struct Module<'d> {
89    device: &'d Device,
90    module: *mut c_void,
91    pub function: *mut c_void,
92}
93
94pub struct Buffer<'d> {
95    device: &'d Device,
96    pub ptr: CUdeviceptr,
97    pub bytes: usize,
98}
99
100impl Device {
101    pub fn open() -> Result<Self, String> {
102        let cuda = Cuda::load()?;
103        unsafe {
104            check("cuInit", (cuda.cuInit)(0))?;
105            let mut version = 0i32;
106            check(
107                "cuDriverGetVersion",
108                (cuda.cuDriverGetVersion)(&mut version),
109            )?;
110            let mut dev = 0i32;
111            check("cuDeviceGet", (cuda.cuDeviceGet)(&mut dev, 0))?;
112            let mut name = [0u8; 128];
113            check(
114                "cuDeviceGetName",
115                (cuda.cuDeviceGetName)(name.as_mut_ptr(), name.len() as i32, dev),
116            )?;
117            let (mut major, mut minor) = (0i32, 0i32);
118            check(
119                "cc major",
120                (cuda.cuDeviceGetAttribute)(
121                    &mut major,
122                    CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MAJOR,
123                    dev,
124                ),
125            )?;
126            check(
127                "cc minor",
128                (cuda.cuDeviceGetAttribute)(
129                    &mut minor,
130                    CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MINOR,
131                    dev,
132                ),
133            )?;
134            let mut ctx = std::ptr::null_mut();
135            check("cuCtxCreate", (cuda.cuCtxCreate_v2)(&mut ctx, 0, dev))?;
136            let name = String::from_utf8_lossy(
137                &name[..name.iter().position(|&b| b == 0).unwrap_or(name.len())],
138            )
139            .to_string();
140            Ok(Device {
141                cuda,
142                ctx,
143                name,
144                cc: format!("{major}.{minor}"),
145                driver_version: format!("{}.{}", version / 1000, (version % 1000) / 10),
146            })
147        }
148    }
149
150    pub fn load_module(&self, ptx: &str, entry: &str) -> Result<Module<'_>, String> {
151        let mut ptx_z = ptx.as_bytes().to_vec();
152        ptx_z.push(0);
153        let mut entry_z = entry.as_bytes().to_vec();
154        entry_z.push(0);
155        unsafe {
156            let mut module = std::ptr::null_mut();
157            check(
158                "cuModuleLoadData",
159                (self.cuda.cuModuleLoadData)(&mut module, ptx_z.as_ptr().cast()),
160            )?;
161            let mut function = std::ptr::null_mut();
162            let got = (self.cuda.cuModuleGetFunction)(&mut function, module, entry_z.as_ptr());
163            if got != 0 {
164                (self.cuda.cuModuleUnload)(module);
165                return Err(format!(
166                    "cuModuleGetFunction({entry}) failed: CUresult {got}"
167                ));
168            }
169            Ok(Module {
170                device: self,
171                module,
172                function,
173            })
174        }
175    }
176
177    pub fn alloc(&self, bytes: usize) -> Result<Buffer<'_>, String> {
178        let mut ptr = 0u64;
179        unsafe { check("cuMemAlloc", (self.cuda.cuMemAlloc_v2)(&mut ptr, bytes))? };
180        Ok(Buffer {
181            device: self,
182            ptr,
183            bytes,
184        })
185    }
186
187    pub fn copy_in(&self, buffer: &Buffer<'_>, data: &[u8]) -> Result<(), String> {
188        assert!(data.len() <= buffer.bytes);
189        unsafe {
190            check(
191                "cuMemcpyHtoD",
192                (self.cuda.cuMemcpyHtoD_v2)(buffer.ptr, data.as_ptr().cast(), data.len()),
193            )
194        }
195    }
196
197    pub fn copy_out(&self, buffer: &Buffer<'_>, out: &mut [u8]) -> Result<(), String> {
198        assert!(out.len() <= buffer.bytes);
199        unsafe {
200            check(
201                "cuMemcpyDtoH",
202                (self.cuda.cuMemcpyDtoH_v2)(out.as_mut_ptr().cast(), buffer.ptr, out.len()),
203            )
204        }
205    }
206
207    pub fn synchronize(&self) -> Result<(), String> {
208        unsafe { check("cuCtxSynchronize", (self.cuda.cuCtxSynchronize)()) }
209    }
210
211    /// Launch once and return the kernel-only elapsed milliseconds,
212    /// measured with a cuEvent pair on the default stream.
213    pub fn timed_launch(
214        &self,
215        module: &Module<'_>,
216        grid: [u32; 3],
217        block: [u32; 3],
218        params: &mut [*mut c_void],
219    ) -> Result<f64, String> {
220        unsafe {
221            let mut ev0 = std::ptr::null_mut();
222            let mut ev1 = std::ptr::null_mut();
223            check("cuEventCreate", (self.cuda.cuEventCreate)(&mut ev0, 0))?;
224            check("cuEventCreate", (self.cuda.cuEventCreate)(&mut ev1, 0))?;
225            let stream = std::ptr::null_mut();
226            check("cuEventRecord", (self.cuda.cuEventRecord)(ev0, stream))?;
227            let launched = (self.cuda.cuLaunchKernel)(
228                module.function,
229                grid[0],
230                grid[1],
231                grid[2],
232                block[0],
233                block[1],
234                block[2],
235                0, // static SharedArray only: no dynamic smem
236                stream,
237                params.as_mut_ptr(),
238                std::ptr::null_mut(),
239            );
240            if launched != 0 {
241                (self.cuda.cuEventDestroy_v2)(ev0);
242                (self.cuda.cuEventDestroy_v2)(ev1);
243                return Err(format!("cuLaunchKernel failed: CUresult {launched}"));
244            }
245            check("cuEventRecord", (self.cuda.cuEventRecord)(ev1, stream))?;
246            check("cuEventSynchronize", (self.cuda.cuEventSynchronize)(ev1))?;
247            let mut ms = 0f32;
248            check(
249                "cuEventElapsedTime",
250                (self.cuda.cuEventElapsedTime)(&mut ms, ev0, ev1),
251            )?;
252            (self.cuda.cuEventDestroy_v2)(ev0);
253            (self.cuda.cuEventDestroy_v2)(ev1);
254            Ok(ms as f64)
255        }
256    }
257}
258
259impl Drop for Module<'_> {
260    fn drop(&mut self) {
261        unsafe {
262            (self.device.cuda.cuModuleUnload)(self.module);
263        }
264    }
265}
266
267impl Drop for Buffer<'_> {
268    fn drop(&mut self) {
269        unsafe {
270            (self.device.cuda.cuMemFree_v2)(self.ptr);
271        }
272    }
273}
274
275impl Drop for Device {
276    fn drop(&mut self) {
277        unsafe {
278            (self.cuda.cuCtxDestroy_v2)(self.ctx);
279        }
280    }
281}