1#![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 #[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 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 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, 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}