Skip to main content

unmtx_gpu/
cuda.rs

1//
2// Copyright (c) 2025-2026 Ɓukasz Szpakowski
3// Copyright (c) 2026 Mateusz Szpakowski
4//
5// This Source Code Form is subject to the terms of the Mozilla Public
6// License, v. 2.0. If a copy of the MPL was not distributed with this
7// file, You can obtain one at https://mozilla.org/MPL/2.0/.
8//
9//! A module that contains a CUDA backend.
10use std::ffi::c_int;
11use std::sync::Arc;
12use std::sync::Mutex;
13use crate::Backend;
14use crate::BackendArray;
15use crate::Error;
16use crate::Result;
17use crate::mutex_lock;
18
19pub use cudarc::cublas::result::CublasError;
20pub use cudarc::driver::DriverError;
21
22use cudarc::cublas::result::sgemm;
23use cudarc::cublas::sys::cublasOperation_t;
24use cudarc::cublas::CudaBlas;
25use cudarc::driver::sys::CUdeviceptr;
26use cudarc::driver::CudaContext;
27use cudarc::driver::CudaModule;
28use cudarc::driver::CudaFunction;
29use cudarc::driver::CudaSlice;
30use cudarc::driver::CudaStream;
31use cudarc::driver::DevicePtr;
32use cudarc::driver::DevicePtrMut;
33use cudarc::driver::LaunchConfig;
34use cudarc::driver::PushKernelArg;
35use cudarc::nvrtc::CompileError;
36use cudarc::nvrtc::Ptx;
37use cudarc::nvrtc::compile_ptx;
38
39const SOURCE: &'static str = include_str!("cuda.cu");
40
41const PTX_SOURCE: &'static str = include_str!("ptx_mul.ptx");
42
43/// A structure of CUDA backend array.
44///
45/// This structure contains the reference to the device memory.
46#[derive(Debug)]
47pub struct CudaBackendArray
48{
49    slice: Arc<Mutex<CudaSlice<f32>>>,
50    len: usize,
51}
52
53struct CudaInnerBackend
54{
55    context: Arc<CudaContext>,
56    stream: Arc<CudaStream>,
57    module: Arc<CudaModule>,
58    ptx_module: Option<Arc<CudaModule>>,
59    cublas: Option<CudaBlas>,
60}
61
62/// A structure of CUDA backend.
63pub struct CudaBackend
64{
65    inner: Mutex<CudaInnerBackend>,
66    has_cublas: bool,
67    has_ptx: bool,
68}
69
70fn preferred_launch_config(n: usize, m: usize, item_row_count: usize, item_col_count: usize, are_swapped_dims: bool, is_mul: bool, is_ptx: bool) -> LaunchConfig
71{
72    if m <= item_col_count && !is_mul {
73        let n2 = (((n + item_row_count - 1) / item_row_count + 1023) / 1024) as u32;
74        if !are_swapped_dims {
75            LaunchConfig {
76                grid_dim: (n2, 1, 1),
77                block_dim: (1024, 1, 1),
78                shared_mem_bytes: 0,
79            }
80        } else {
81            LaunchConfig {
82                grid_dim: (1, n2, 1),
83                block_dim: (1, 1024, 1),
84                shared_mem_bytes: 0,
85            }
86        }
87    } else if n <= item_row_count && !is_mul {
88        let m2 = (((m + item_col_count - 1) / item_col_count + 1023) / 1024) as u32;
89        if !are_swapped_dims {
90            LaunchConfig {
91                grid_dim: (1, m2, 1),
92                block_dim: (1, 1024, 1),
93                shared_mem_bytes: 0,
94            }
95        } else {
96            LaunchConfig {
97                grid_dim: (m2, 1, 1),
98                block_dim: (1024, 1, 1),
99                shared_mem_bytes: 0,
100            }
101        }
102    } else if is_mul {
103        if is_ptx {
104            let n2 = (((n + 3) / 4 + 31) / 32) as u32;
105            let m2 = (((m + 3) / 4 + 31) / 32) as u32;
106            if !are_swapped_dims {
107                LaunchConfig {
108                    grid_dim: (n2, m2, 1),
109                    block_dim: (32, 32, 1),
110                    shared_mem_bytes: 0,
111                }
112            } else {
113                LaunchConfig {
114                    grid_dim: (m2, n2, 1),
115                    block_dim: (32, 32, 1),
116                    shared_mem_bytes: 0,
117                }
118            }
119        } else {
120            let n2 = (((n + 7) / 8 + 15) / 16) as u32;
121            let m2 = (((m + 3) / 4 + 15) / 16) as u32;
122            if !are_swapped_dims {
123                LaunchConfig {
124                    grid_dim: (n2, m2, 1),
125                    block_dim: (16, 16, 1),
126                    shared_mem_bytes: 0,
127                }
128            } else {
129                LaunchConfig {
130                    grid_dim: (m2, n2, 1),
131                    block_dim: (16, 16, 1),
132                    shared_mem_bytes: 0,
133                }
134            }
135        }
136    } else {
137        let n2 = (((n + item_row_count - 1) / item_row_count + 31) / 32) as u32;
138        let m2 = (((m + item_col_count - 1) / item_col_count + 31) / 32) as u32;
139        if !are_swapped_dims {
140            LaunchConfig {
141                grid_dim: (n2, m2, 1),
142                block_dim: (32, 32, 1),
143                shared_mem_bytes: 0,
144            }
145        } else {
146            LaunchConfig {
147                grid_dim: (m2, n2, 1),
148                block_dim: (32, 32, 1),
149                shared_mem_bytes: 0,
150            }
151        }
152    }
153}
154
155impl CudaBackend
156{
157    /// Creates a CUDA backend for a first device.
158    pub fn new() -> Result<CudaBackend>
159    {
160        if cfg!(feature = "default_cublas") {
161            Self::new_with_ordinal_and_cublas_flag(0, true)
162        } else if cfg!(feature = "default_ptx") {
163            Self::new_with_ordinal_and_cublas_flag_and_ptx_flag(0, false, true)
164        } else {
165            Self::new_with_ordinal_and_cublas_flag(0, false)
166        }
167    }
168
169    /// Creates a CUDA backend with the ordinal number and the cuBLAS flag.
170    ///
171    /// See [`new_with_ordinal_and_cublas_flag_and_ptx_flag`](Self::new_with_ordinal_and_cublas_flag_and_ptx_flag).
172    pub fn new_with_ordinal_and_cublas_flag(ordinal: usize, is_cublas: bool) -> Result<CudaBackend>
173    { Self::new_with_ordinal_and_cublas_flag_and_ptx_flag(ordinal, is_cublas, false) }
174
175    /// Creates a CUDA backend with the ordinal number, the cuBLAS flag, and the PTX flag.
176    ///
177    /// This method takes the following flags:
178    ///
179    /// - `is_cublas` - use the cuBLAS library to multiplication of matrices
180    /// - `is_ptx` - use the module in PTX to multiplication of matrices
181    pub fn new_with_ordinal_and_cublas_flag_and_ptx_flag(ordinal: usize, is_cublas: bool, is_ptx: bool) -> Result<CudaBackend>
182    {
183        let context = match CudaContext::new(ordinal) {
184            Ok(tmp_device) => tmp_device,
185            Err(err) => return Err(Error::Cuda(err)),
186        };
187        let ptx = match compile_ptx(SOURCE) {
188            Ok(tmp_ptx) => tmp_ptx,
189            Err(CompileError::CompileError { log, .. }) => return Err(Error::Compilation(log.as_c_str().to_string_lossy().into_owned())),
190            Err(err) => return Err(Error::Compilation(format!("{}", err))),
191        };
192        let module = match context.load_module(ptx) {
193            Ok(tmp_module) => tmp_module,
194            Err(err) => return Err(Error::Cuda(err)),
195        };
196        let is_real_ptx = if !is_cublas {
197            is_ptx
198        } else {
199            false
200        };
201        let ptx_module = if is_real_ptx {
202            match context.load_module(Ptx::from_src(PTX_SOURCE)) {
203                Ok(tmp_ptx_module) => Some(tmp_ptx_module),
204                Err(err) => return Err(Error::Cuda(err)),
205            }
206        } else {
207            None
208        };
209        let stream = context.default_stream();
210        let cublas = if is_cublas {
211            match CudaBlas::new(stream.clone()) {
212                Ok(tmp_cublas) => Some(tmp_cublas),
213                Err(err) => return Err(Error::Cublas(err)),
214            }
215        } else {
216            None
217        };
218        Ok(CudaBackend { inner: Mutex::new(CudaInnerBackend { context, stream, module, ptx_module, cublas, }), has_cublas: is_cublas, has_ptx: is_real_ptx, })
219    }
220        
221    fn check_and_launch2<F, G>(&self, kernel_name: &str, a: &BackendArray, b: &BackendArray, f: F, g: G) -> Result<()>
222        where F: FnOnce(&CudaBackendArray, &CudaBackendArray) -> Result<()>,
223            G: FnOnce(&CudaInnerBackend, CudaFunction, CUdeviceptr, CUdeviceptr) -> Result<()>
224    {
225        #[allow(unreachable_patterns)]
226        match (a, b) {
227            (BackendArray::Cuda(a2), BackendArray::Cuda(b2)) => {
228                f(a2, b2)?;
229                let inner_g = mutex_lock(&self.inner)?;
230                let kernel = match inner_g.module.load_function(kernel_name) {
231                    Ok(tmp_kernel) => tmp_kernel,
232                    Err(_) => return Err(Error::NoKernel(String::from(kernel_name))),
233                };
234                if !Arc::ptr_eq(&a2.slice, &b2.slice) {
235                    let a_slice_g = mutex_lock(&a2.slice)?;
236                    let mut b_slice_g = mutex_lock(&b2.slice)?;
237                    let a_device_ptr = a_slice_g.device_ptr(&inner_g.stream).0;
238                    let b_device_ptr = b_slice_g.device_ptr_mut(&inner_g.stream).0;
239                    g(&*inner_g, kernel, a_device_ptr, b_device_ptr)?;
240                } else {
241                    let mut a_slice_g = mutex_lock(&a2.slice)?;
242                    let a_device_ptr = a_slice_g.device_ptr_mut(&inner_g.stream).0;
243                    g(&*inner_g, kernel, a_device_ptr, a_device_ptr)?;
244                }
245                match inner_g.context.synchronize() {
246                    Ok(()) => (),
247                    Err(err) => return Err(Error::Cuda(err)),
248                }
249                Ok(())
250            },
251            _ => Err(Error::InvalidBackendArray),
252        }
253    }
254
255    fn check_and_launch3<F, G>(&self, kernel_name: &str, a: &BackendArray, b: &BackendArray, c: &BackendArray, f: F, g: G) -> Result<()>
256        where F: FnOnce(&CudaBackendArray, &CudaBackendArray, &CudaBackendArray) -> Result<()>,
257            G: FnOnce(&CudaInnerBackend, CudaFunction, CUdeviceptr, CUdeviceptr, CUdeviceptr) -> Result<()>
258    {
259        #[allow(unreachable_patterns)]
260        match (a, b, c) {
261            (BackendArray::Cuda(a2), BackendArray::Cuda(b2), BackendArray::Cuda(c2)) => {
262                f(a2, b2, c2)?;
263                let inner_g = mutex_lock(&self.inner)?;
264                let kernel = match inner_g.module.load_function(kernel_name) {
265                    Ok(tmp_kernel) => tmp_kernel,
266                    Err(_) => return Err(Error::NoKernel(String::from(kernel_name))),
267                };
268                match (Arc::ptr_eq(&a2.slice, &b2.slice), Arc::ptr_eq(&a2.slice, &c2.slice), Arc::ptr_eq(&b2.slice, &c2.slice)) {
269                    (false, false, false) => {
270                        let a_slice_g = mutex_lock(&a2.slice)?;
271                        let b_slice_g = mutex_lock(&b2.slice)?;
272                        let mut c_slice_g = mutex_lock(&c2.slice)?;
273                        let a_device_ptr = a_slice_g.device_ptr(&inner_g.stream).0;
274                        let b_device_ptr = b_slice_g.device_ptr(&inner_g.stream).0;
275                        let c_device_ptr = c_slice_g.device_ptr_mut(&inner_g.stream).0;
276                        g(&*inner_g, kernel, a_device_ptr, b_device_ptr, c_device_ptr)?
277                    },
278                    (true, false, false) => {
279                        let a_slice_g = mutex_lock(&a2.slice)?;
280                        let mut c_slice_g = mutex_lock(&c2.slice)?;
281                        let a_device_ptr = a_slice_g.device_ptr(&inner_g.stream).0;
282                        let c_device_ptr = c_slice_g.device_ptr_mut(&inner_g.stream).0;
283                        g(&*inner_g, kernel, a_device_ptr, a_device_ptr, c_device_ptr)?
284                    },
285                    (false, true, false) => {
286                        let mut a_slice_g = mutex_lock(&a2.slice)?;
287                        let b_slice_g = mutex_lock(&b2.slice)?;
288                        let a_device_ptr = a_slice_g.device_ptr_mut(&inner_g.stream).0;
289                        let b_device_ptr = b_slice_g.device_ptr(&inner_g.stream).0;
290                        g(&*inner_g, kernel, a_device_ptr, b_device_ptr, a_device_ptr)?
291                    },
292                    (false, false, true) => {
293                        let a_slice_g = mutex_lock(&a2.slice)?;
294                        let mut b_slice_g = mutex_lock(&b2.slice)?;
295                        let a_device_ptr = a_slice_g.device_ptr(&inner_g.stream).0;
296                        let b_device_ptr = b_slice_g.device_ptr_mut(&inner_g.stream).0;
297                        g(&*inner_g, kernel, a_device_ptr, b_device_ptr, b_device_ptr)?
298                    },
299                    _ => {
300                        let mut a_slice_g = mutex_lock(&a2.slice)?;
301                        let a_device_ptr = a_slice_g.device_ptr_mut(&inner_g.stream).0;
302                        g(&*inner_g, kernel, a_device_ptr, a_device_ptr, a_device_ptr)?
303                    },
304                }
305                match inner_g.context.synchronize() {
306                    Ok(()) => (),
307                    Err(err) => return Err(Error::Cuda(err)),
308                }
309                Ok(())
310            },
311            _ => Err(Error::InvalidBackendArray),
312        }
313    }    
314
315    fn check_and_launch_ptx3<F, G>(&self, kernel_name: &str, a: &BackendArray, b: &BackendArray, c: &BackendArray, f: F, g: G) -> Result<()>
316        where F: FnOnce(&CudaBackendArray, &CudaBackendArray, &CudaBackendArray) -> Result<()>,
317            G: FnOnce(&CudaInnerBackend, CudaFunction, CUdeviceptr, CUdeviceptr, CUdeviceptr) -> Result<()>
318    {
319        #[allow(unreachable_patterns)]
320        match (a, b, c) {
321            (BackendArray::Cuda(a2), BackendArray::Cuda(b2), BackendArray::Cuda(c2)) => {
322                f(a2, b2, c2)?;
323                let inner_g = mutex_lock(&self.inner)?;
324                let kernel = match &inner_g.ptx_module {
325                    Some(ptx_module) => {
326                        match ptx_module.load_function(kernel_name) {
327                            Ok(tmp_kernel) => tmp_kernel,
328                            Err(_) => return Err(Error::NoKernel(String::from(kernel_name))),
329                        }
330                    },
331                    None => return Err(Error::NoPtxModule),
332                };
333                match (Arc::ptr_eq(&a2.slice, &b2.slice), Arc::ptr_eq(&a2.slice, &c2.slice), Arc::ptr_eq(&b2.slice, &c2.slice)) {
334                    (false, false, false) => {
335                        let a_slice_g = mutex_lock(&a2.slice)?;
336                        let b_slice_g = mutex_lock(&b2.slice)?;
337                        let mut c_slice_g = mutex_lock(&c2.slice)?;
338                        let a_device_ptr = a_slice_g.device_ptr(&inner_g.stream).0;
339                        let b_device_ptr = b_slice_g.device_ptr(&inner_g.stream).0;
340                        let c_device_ptr = c_slice_g.device_ptr_mut(&inner_g.stream).0;
341                        g(&*inner_g, kernel, a_device_ptr, b_device_ptr, c_device_ptr)?
342                    },
343                    (true, false, false) => {
344                        let a_slice_g = mutex_lock(&a2.slice)?;
345                        let mut c_slice_g = mutex_lock(&c2.slice)?;
346                        let a_device_ptr = a_slice_g.device_ptr(&inner_g.stream).0;
347                        let c_device_ptr = c_slice_g.device_ptr_mut(&inner_g.stream).0;
348                        g(&*inner_g, kernel, a_device_ptr, a_device_ptr, c_device_ptr)?
349                    },
350                    (false, true, false) => {
351                        let mut a_slice_g = mutex_lock(&a2.slice)?;
352                        let b_slice_g = mutex_lock(&b2.slice)?;
353                        let a_device_ptr = a_slice_g.device_ptr_mut(&inner_g.stream).0;
354                        let b_device_ptr = b_slice_g.device_ptr(&inner_g.stream).0;
355                        g(&*inner_g, kernel, a_device_ptr, b_device_ptr, a_device_ptr)?
356                    },
357                    (false, false, true) => {
358                        let a_slice_g = mutex_lock(&a2.slice)?;
359                        let mut b_slice_g = mutex_lock(&b2.slice)?;
360                        let a_device_ptr = a_slice_g.device_ptr(&inner_g.stream).0;
361                        let b_device_ptr = b_slice_g.device_ptr_mut(&inner_g.stream).0;
362                        g(&*inner_g, kernel, a_device_ptr, b_device_ptr, b_device_ptr)?
363                    },
364                    _ => {
365                        let mut a_slice_g = mutex_lock(&a2.slice)?;
366                        let a_device_ptr = a_slice_g.device_ptr_mut(&inner_g.stream).0;
367                        g(&*inner_g, kernel, a_device_ptr, a_device_ptr, a_device_ptr)?
368                    },
369                }
370                match inner_g.context.synchronize() {
371                    Ok(()) => (),
372                    Err(err) => return Err(Error::Cuda(err)),
373                }
374                Ok(())
375            },
376            _ => Err(Error::InvalidBackendArray),
377        }
378    }
379    
380    fn check_and_launch_cublas3<F, G>(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, f: F, g: G) -> Result<()>
381        where F: FnOnce(&CudaBackendArray, &CudaBackendArray, &CudaBackendArray) -> Result<()>,
382            G: FnOnce(&CudaInnerBackend, CUdeviceptr, CUdeviceptr, CUdeviceptr) -> Result<()>
383    {
384        #[allow(unreachable_patterns)]
385        match (a, b, c) {
386            (BackendArray::Cuda(a2), BackendArray::Cuda(b2), BackendArray::Cuda(c2)) => {
387                f(a2, b2, c2)?;
388                let inner_g = mutex_lock(&self.inner)?;
389                match (Arc::ptr_eq(&a2.slice, &b2.slice), Arc::ptr_eq(&a2.slice, &c2.slice), Arc::ptr_eq(&b2.slice, &c2.slice)) {
390                    (false, false, false) => {
391                        let a_slice_g = mutex_lock(&a2.slice)?;
392                        let b_slice_g = mutex_lock(&b2.slice)?;
393                        let mut c_slice_g = mutex_lock(&c2.slice)?;
394                        let a_device_ptr = a_slice_g.device_ptr(&inner_g.stream).0;
395                        let b_device_ptr = b_slice_g.device_ptr(&inner_g.stream).0;
396                        let c_device_ptr = c_slice_g.device_ptr_mut(&inner_g.stream).0;
397                        g(&*inner_g, a_device_ptr, b_device_ptr, c_device_ptr)?
398                    },
399                    (true, false, false) => {
400                        let a_slice_g = mutex_lock(&a2.slice)?;
401                        let mut c_slice_g = mutex_lock(&c2.slice)?;
402                        let a_device_ptr = a_slice_g.device_ptr(&inner_g.stream).0;
403                        let c_device_ptr = c_slice_g.device_ptr_mut(&inner_g.stream).0;
404                        g(&*inner_g, a_device_ptr, a_device_ptr, c_device_ptr)?
405                    },
406                    (false, true, false) => {
407                        let mut a_slice_g = mutex_lock(&a2.slice)?;
408                        let b_slice_g = mutex_lock(&b2.slice)?;
409                        let a_device_ptr = a_slice_g.device_ptr_mut(&inner_g.stream).0;
410                        let b_device_ptr = b_slice_g.device_ptr(&inner_g.stream).0;
411                        g(&*inner_g, a_device_ptr, b_device_ptr, a_device_ptr)?
412                    },
413                    (false, false, true) => {
414                        let a_slice_g = mutex_lock(&a2.slice)?;
415                        let mut b_slice_g = mutex_lock(&b2.slice)?;
416                        let a_device_ptr = a_slice_g.device_ptr(&inner_g.stream).0;
417                        let b_device_ptr = b_slice_g.device_ptr_mut(&inner_g.stream).0;
418                        g(&*inner_g, a_device_ptr, b_device_ptr, b_device_ptr)?
419                    },
420                    _ => {
421                        let mut a_slice_g = mutex_lock(&a2.slice)?;
422                        let a_device_ptr = a_slice_g.device_ptr_mut(&inner_g.stream).0;
423                        g(&*inner_g, a_device_ptr, a_device_ptr, a_device_ptr)?
424                    },
425                }
426                match inner_g.context.synchronize() {
427                    Ok(()) => (),
428                    Err(err) => return Err(Error::Cuda(err)),
429                }
430                Ok(())
431            },
432            _ => Err(Error::InvalidBackendArray),
433        }
434    }
435    
436    fn check_and_launch_for_fun(&self, kernel_name: &str, a: &BackendArray, b: &BackendArray, n: usize, m: usize, item_row_count: usize, item_col_count: usize, are_swapped_dims: bool) -> Result<()>
437    {
438        let is_ptx = self.has_ptx;
439        self.check_and_launch2(kernel_name, a, b, |a2, b2| {
440                if a2.len != n * m {
441                    return Err(Error::BackendArrayElemCount(a2.len, n * m));
442                }
443                if b2.len != n * m {
444                    return Err(Error::BackendArrayElemCount(b2.len, n * m));
445                }
446                Ok(())
447        }, |inner_g, kernel, a_param, b_param| {
448                let config = preferred_launch_config(n, m, item_row_count, item_col_count, are_swapped_dims, false, is_ptx);
449                let mut launch_args = inner_g.stream.launch_builder(&kernel);
450                launch_args.arg(&a_param)
451                    .arg(&b_param)
452                    .arg(&n)
453                    .arg(&m);
454                unsafe {
455                    match launch_args.launch(config) {
456                        Ok(_) => Ok(()),
457                        Err(err) => Err(Error::Cuda(err)),
458                    }
459                }
460        })
461    }
462
463    fn check_and_launch_for_op(&self, kernel_name: &str, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize, item_row_count: usize, item_col_count: usize, are_swapped_dims: bool) -> Result<()>
464    {
465        let is_ptx = self.has_ptx;
466        self.check_and_launch3(kernel_name, a, b, c, |a2, b2, c2| {
467                if a2.len != n * m {
468                    return Err(Error::BackendArrayElemCount(a2.len, n * m));
469                }
470                if b2.len != n * m {
471                    return Err(Error::BackendArrayElemCount(b2.len, n * m));
472                }
473                if c2.len != n * m {
474                    return Err(Error::BackendArrayElemCount(c2.len, n * m));
475                }
476                Ok(())
477        }, |inner_g, kernel, a_param, b_param, c_param| {
478                let config = preferred_launch_config(n, m, item_row_count, item_col_count, are_swapped_dims, false, is_ptx);
479                let mut launch_args = inner_g.stream.launch_builder(&kernel);
480                launch_args.arg(&a_param)
481                    .arg(&b_param)
482                    .arg(&c_param)
483                    .arg(&n)
484                    .arg(&m);
485                unsafe {
486                    match launch_args.launch(config) {
487                        Ok(_) => Ok(()),
488                        Err(err) => Err(Error::Cuda(err)),
489                    }
490                }
491        })
492    }
493
494    fn check_and_launch_for_mul(&self, kernel_name: &str, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize, l: usize, item_row_count: usize, item_col_count: usize, are_swapped_dims: bool) -> Result<()>
495    {
496        let is_ptx = self.has_ptx;
497        self.check_and_launch3(kernel_name, a, b, c, |a2, b2, c2| {
498                if a2.len != n * l {
499                    return Err(Error::BackendArrayElemCount(a2.len, n * l));
500                }
501                if b2.len != l * m {
502                    return Err(Error::BackendArrayElemCount(b2.len, l * m));
503                }
504                if c2.len != n * m {
505                    return Err(Error::BackendArrayElemCount(c2.len, n * m));
506                }
507                Ok(())
508        }, |inner_g, kernel, a_param, b_param, c_param| {
509                let config = preferred_launch_config(n, m, item_row_count, item_col_count, are_swapped_dims, true, is_ptx);
510                let mut launch_args = inner_g.stream.launch_builder(&kernel);
511                launch_args.arg(&a_param)
512                    .arg(&b_param)
513                    .arg(&c_param)
514                    .arg(&n)
515                    .arg(&m)
516                    .arg(&l);
517                unsafe {
518                    match launch_args.launch(config) {
519                        Ok(_) => Ok(()),
520                        Err(err) => Err(Error::Cuda(err)),
521                    }
522                }
523        })
524    }
525
526    fn check_and_launch_for_scalar(&self, kernel_name: &str, a: &BackendArray, b: f32, c: &BackendArray, n: usize, m: usize, item_row_count: usize, item_col_count: usize, are_swapped_dims: bool) -> Result<()>
527    {
528        let is_ptx = self.has_ptx;
529        self.check_and_launch2(kernel_name, a, c, |a2, c2| {
530                if a2.len != n * m  {
531                    return Err(Error::BackendArrayElemCount(a2.len, n * m));
532                }
533                if c2.len != n * m {
534                    return Err(Error::BackendArrayElemCount(c2.len, n * m));
535                }
536                Ok(())
537        }, |inner_g, kernel, a_param, c_param| {
538                let config = preferred_launch_config(n, m, item_row_count, item_col_count, are_swapped_dims, false, is_ptx);
539                let mut launch_args = inner_g.stream.launch_builder(&kernel);
540                launch_args.arg(&a_param)
541                    .arg(&b)
542                    .arg(&c_param)
543                    .arg(&n)
544                    .arg(&m);
545                unsafe {
546                    match launch_args.launch(config) {
547                        Ok(_) => Ok(()),
548                        Err(err) => Err(Error::Cuda(err)),
549                    }
550                }
551        })
552    }
553
554    fn check_and_launch_for_fun_and_tiles(&self, kernel_name: &str, a: &BackendArray, b: &BackendArray, n: usize, m: usize, item_row_count: usize, item_col_count: usize, are_swapped_dims: bool) -> Result<()>
555    {
556        let is_ptx = self.has_ptx;
557        self.check_and_launch2(kernel_name, a, b, |a2, b2| {
558                if a2.len != n * m {
559                    return Err(Error::BackendArrayElemCount(a2.len, n * m));
560                }
561                if b2.len != n * m {
562                    return Err(Error::BackendArrayElemCount(b2.len, n * m));
563                }
564                Ok(())
565        }, |inner_g, kernel, a_param, b_param| {
566                let config = preferred_launch_config(n, m, item_row_count, item_col_count, are_swapped_dims, false, is_ptx);
567                let mut launch_args = inner_g.stream.launch_builder(&kernel);
568                launch_args.arg(&a_param)
569                    .arg(&b_param)
570                    .arg(&n)
571                    .arg(&m);
572                unsafe {
573                    match launch_args.launch(config) {
574                        Ok(_) => Ok(()),
575                        Err(err) => Err(Error::Cuda(err)),
576                    }
577                }
578        })
579    }
580
581    fn check_and_launch_for_repeat_col(&self, kernel_name: &str, a: &BackendArray, b: &BackendArray, n: usize, m: usize, item_row_count: usize, item_col_count: usize, are_swapped_dims: bool) -> Result<()>
582    {
583        let is_ptx = self.has_ptx;
584        self.check_and_launch2(kernel_name, a, b, |a2, b2| {
585                if a2.len != n {
586                    return Err(Error::BackendArrayElemCount(a2.len, n));
587                }
588                if b2.len != n * m {
589                    return Err(Error::BackendArrayElemCount(b2.len, n * m));
590                }
591                Ok(())
592        }, |inner_g, kernel, a_param, b_param| {
593                let config = preferred_launch_config(n, m, item_row_count, item_col_count, are_swapped_dims, false, is_ptx);
594                let mut launch_args = inner_g.stream.launch_builder(&kernel);
595                launch_args.arg(&a_param)
596                    .arg(&b_param)
597                    .arg(&n)
598                    .arg(&m);
599                unsafe {
600                    match launch_args.launch(config) {
601                        Ok(_) => Ok(()),
602                        Err(err) => Err(Error::Cuda(err)),
603                    }
604                }
605        })
606    }
607
608    fn check_and_launch_for_repeat_row(&self, kernel_name: &str, a: &BackendArray, b: &BackendArray, n: usize, m: usize, item_row_count: usize, item_col_count: usize, are_swapped_dims: bool) -> Result<()>
609    {
610        let is_ptx = self.has_ptx;
611        self.check_and_launch2(kernel_name, a, b, |a2, b2| {
612                if a2.len != m {
613                    return Err(Error::BackendArrayElemCount(a2.len, m));
614                }
615                if b2.len != n * m {
616                    return Err(Error::BackendArrayElemCount(b2.len, n * m));
617                }
618                Ok(())
619        }, |inner_g, kernel, a_param, b_param| {
620                let config = preferred_launch_config(n, m, item_row_count, item_col_count, are_swapped_dims, false, is_ptx);
621                let mut launch_args = inner_g.stream.launch_builder(&kernel);
622                launch_args.arg(&a_param)
623                    .arg(&b_param)
624                    .arg(&n)
625                    .arg(&m);
626                unsafe {
627                    match launch_args.launch(config) {
628                        Ok(_) => Ok(()),
629                        Err(err) => Err(Error::Cuda(err)),
630                    }
631                }
632        })
633    }    
634    
635    fn check_and_launch_for_ptx_mul(&self, kernel_name: &str, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize, l: usize, item_row_count: usize, item_col_count: usize, are_swapped_dims: bool) -> Result<()>
636    {
637        let is_ptx = self.has_ptx;
638        self.check_and_launch_ptx3(kernel_name, a, b, c, |a2, b2, c2| {
639                if a2.len != n * l {
640                    return Err(Error::BackendArrayElemCount(a2.len, n * l));
641                }
642                if b2.len != l * m {
643                    return Err(Error::BackendArrayElemCount(b2.len, l * m));
644                }
645                if c2.len != n * m {
646                    return Err(Error::BackendArrayElemCount(c2.len, n * m));
647                }
648                Ok(())
649        }, |inner_g, kernel, a_param, b_param, c_param| {
650                let config = preferred_launch_config(n, m, item_row_count, item_col_count, are_swapped_dims, true, is_ptx);
651                let mut launch_args = inner_g.stream.launch_builder(&kernel);
652                launch_args.arg(&a_param)
653                    .arg(&b_param)
654                    .arg(&c_param)
655                    .arg(&n)
656                    .arg(&m)
657                    .arg(&l);
658                unsafe {
659                    match launch_args.launch(config) {
660                        Ok(_) => Ok(()),
661                        Err(err) => Err(Error::Cuda(err)),
662                    }
663                }
664        })
665    }
666
667    fn check_and_launch_cublas_for_mul(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize, l: usize, is_trans_a: bool, is_trans_b: bool) -> Result<()>
668    {
669        self.check_and_launch_cublas3(a, b, c, |a2, b2, c2| {
670                if a2.len != n * l {
671                    return Err(Error::BackendArrayElemCount(a2.len, n * l));
672                }
673                if b2.len != l * m {
674                    return Err(Error::BackendArrayElemCount(b2.len, l * m));
675                }
676                if c2.len != n * m {
677                    return Err(Error::BackendArrayElemCount(c2.len, n * m));
678                }
679                Ok(())
680        }, |inner, a_device_ptr, b_device_ptr, c_device_ptr| {
681                unsafe {
682                    match &inner.cublas {
683                        Some(cublas) => {
684                            let (transa, lda) = if is_trans_a {
685                                (cublasOperation_t::CUBLAS_OP_T, n as c_int)
686                            } else {
687                                (cublasOperation_t::CUBLAS_OP_N, l as c_int)
688                            };
689                            let (transb, ldb) = if is_trans_b {
690                                (cublasOperation_t::CUBLAS_OP_T, l as c_int)
691                            } else {
692                                (cublasOperation_t::CUBLAS_OP_N, m as c_int)
693                            };
694                            let alpha = 1.0f32;
695                            let beta = 0.0f32;
696                            let res = sgemm(*cublas.handle(),
697                                transb, transa,
698                                m as c_int, n as c_int, l as c_int,
699                                (&alpha) as *const _,
700                                b_device_ptr as *const _, ldb,
701                                a_device_ptr as *const _, lda,
702                                (&beta) as *const _,
703                                c_device_ptr as *mut _, m as c_int);
704                            match res {
705                                Ok(()) => Ok(()),
706                                Err(err) => Err(Error::Cublas(err)),
707                            }
708                        },
709                        None => Err(Error::NoCublas),
710                    }
711                }
712        })
713    }
714}
715
716impl Backend for CudaBackend
717{
718    fn name(&self) -> &'static str
719    {
720        if self.has_cublas {
721            "CUDA(cuBLAS)"
722        } else if self.has_ptx {
723            "CUDA(PTX)"
724        } else {
725            "CUDA"
726        }
727    }
728    
729    fn has_cublas(&self) -> bool
730    { self.has_cublas }
731
732    unsafe fn alloc(&self, n: usize) -> Result<BackendArray>
733    {
734        let inner_g = mutex_lock(&self.inner)?;
735        let slice: CudaSlice<f32> = match inner_g.stream.alloc(n) {
736            Ok(tmp_slice) => tmp_slice,
737            Err(err) => return Err(Error::Cuda(err)),
738        };
739        let cuda_array = CudaBackendArray { slice: Arc::new(Mutex::new(slice)), len: n, };
740        Ok(BackendArray::Cuda(cuda_array))
741    }
742
743    fn alloc_and_store_zeros(&self, n: usize) -> Result<BackendArray>
744    {
745        let inner_g = mutex_lock(&self.inner)?;
746        let slice: CudaSlice<f32> = match inner_g.stream.alloc_zeros(n) {
747            Ok(tmp_slice) => tmp_slice,
748            Err(err) => return Err(Error::Cuda(err)),
749        };
750        let cuda_array = CudaBackendArray { slice: Arc::new(Mutex::new(slice)), len: n, };
751        Ok(BackendArray::Cuda(cuda_array))
752    }
753    
754    fn alloc_and_store(&self, elems: &[f32]) -> Result<BackendArray>
755    {
756        let inner_g = mutex_lock(&self.inner)?;
757        let slice: CudaSlice<f32> = match inner_g.stream.clone_htod(elems) {
758            Ok(tmp_slice) => tmp_slice,
759            Err(err) => return Err(Error::Cuda(err)),
760        };
761        match inner_g.context.synchronize() {
762            Ok(()) => (),
763            Err(err) => return Err(Error::Cuda(err)),
764        };
765        let cuda_array = CudaBackendArray { slice: Arc::new(Mutex::new(slice)), len: elems.len(), };
766        Ok(BackendArray::Cuda(cuda_array))
767    }
768    
769    fn load(&self, a: &BackendArray, elems: &mut [f32]) -> Result<()>
770    {
771        #[allow(unreachable_patterns)]
772        match a {
773            BackendArray::Cuda(a2) => {
774                if a2.len != elems.len() {
775                    return Err(Error::BackendArrayElemCount(a2.len, elems.len()));
776                }
777                let inner_g = mutex_lock(&self.inner)?;
778                let a_slice_g = mutex_lock(&a2.slice)?;
779                match inner_g.stream.memcpy_dtoh(&(*a_slice_g), elems) {
780                    Ok(()) => (),
781                    Err(err) => return Err(Error::Cuda(err)),
782                };
783                match inner_g.context.synchronize() {
784                    Ok(()) => (),
785                    Err(err) => return Err(Error::Cuda(err)),
786                }
787            },
788            _ => return Err(Error::InvalidBackendArray),
789        }
790        Ok(())
791    }
792
793    fn store(&self, a: &BackendArray, elems: &[f32]) -> Result<()>
794    {
795        #[allow(unreachable_patterns)]
796        match a {
797            BackendArray::Cuda(a2) => {
798                if a2.len != elems.len() {
799                    return Err(Error::BackendArrayElemCount(a2.len, elems.len()));
800                }
801                let inner_g = mutex_lock(&self.inner)?;
802                let mut a_slice_g = mutex_lock(&a2.slice)?;
803                match inner_g.stream.memcpy_htod(elems, &mut (*a_slice_g)) {
804                    Ok(()) => (),
805                    Err(err) => return Err(Error::Cuda(err)),
806                };
807                match inner_g.context.synchronize() {
808                    Ok(()) => (),
809                    Err(err) => return Err(Error::Cuda(err)),
810                }
811            },
812            _ => return Err(Error::InvalidBackendArray),
813        }
814        Ok(())
815    }
816    
817    fn copy(&self, a: &BackendArray, b: &BackendArray) -> Result<()>
818    {
819        #[allow(unreachable_patterns)]
820        match (a, b) {
821            (BackendArray::Cuda(a2), BackendArray::Cuda(b2)) => {
822                if Arc::ptr_eq(&a2.slice, &b2.slice) {
823                    return Ok(());
824                }
825                if a2.len != b2.len {
826                    return Err(Error::TwoBackendArrayElemCounts(a2.len, b2.len));
827                }
828                let inner_g = mutex_lock(&self.inner)?;
829                let a_slice_g = mutex_lock(&a2.slice)?;
830                let mut b_slice_g = mutex_lock(&b2.slice)?;
831                match inner_g.stream.memcpy_dtod(&(*a_slice_g), &mut (*b_slice_g)) {
832                    Ok(()) => (),
833                    Err(err) => return Err(Error::Cuda(err)),
834                }
835                match inner_g.context.synchronize() {
836                    Ok(()) => (),
837                    Err(err) => return Err(Error::Cuda(err)),
838                }
839            },
840            _ => return Err(Error::InvalidBackendArray),
841        }
842        Ok(())
843    }
844
845    fn transpose_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
846    { self.check_and_launch_for_fun("transpose_a", a, b, n, m, 2, 2, true) }
847
848    fn add_a_b(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
849    { self.check_and_launch_for_op("add_a_b", a, b, c, n, m, 2, 2, true) }
850
851    fn add_at_b(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
852    { self.check_and_launch_for_op("add_at_b", a, b, c, n, m, 2, 2, true) }
853    
854    fn add_a_bt(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
855    { self.check_and_launch_for_op("add_a_bt", a, b, c, n, m, 2, 2, true) }
856
857    fn add_at_bt(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
858    { self.check_and_launch_for_op("add_at_bt", a, b, c, n, m, 2, 2, true) }
859
860    fn sub_a_b(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
861    { self.check_and_launch_for_op("sub_a_b", a, b, c, n, m, 2, 2, true) }
862
863    fn sub_at_b(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
864    { self.check_and_launch_for_op("sub_at_b", a, b, c, n, m, 2, 2, true) }
865    
866    fn sub_a_bt(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
867    { self.check_and_launch_for_op("sub_a_bt", a, b, c, n, m, 2, 2, true) }
868
869    fn sub_at_bt(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>    
870    { self.check_and_launch_for_op("sub_at_bt", a, b, c, n, m, 2, 2, true) }
871    
872    fn mul_a_b(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize, l: usize) -> Result<()>
873    {
874        if self.has_cublas {
875            self.check_and_launch_cublas_for_mul(a, b, c, n, m, l, false, false)
876        } else {
877            if self.has_ptx {
878                self.check_and_launch_for_ptx_mul("ptx_mul_a_b", a, b, c, n, m, l, 4, 4, true)
879            } else {
880                self.check_and_launch_for_mul("mul_a_b", a, b, c, n, m, l, 8, 4, true)
881            }
882        }
883    }
884
885    fn mul_at_b(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize, l: usize) -> Result<()>
886    {
887        if self.has_cublas {
888            self.check_and_launch_cublas_for_mul(a, b, c, n, m, l, true, false)
889        } else {
890            if self.has_ptx {
891                self.check_and_launch_for_ptx_mul("ptx_mul_at_b", a, b, c, n, m, l, 4, 4, false)
892            } else {
893                self.check_and_launch_for_mul("mul_at_b", a, b, c, n, m, l, 8, 4, false)
894            }
895        }
896    }
897
898    fn mul_a_bt(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize, l: usize) -> Result<()>
899    {
900        if self.has_cublas {
901            self.check_and_launch_cublas_for_mul(a, b, c, n, m, l, false, true)
902        } else {
903            if self.has_ptx {
904                self.check_and_launch_for_ptx_mul("ptx_mul_a_bt", a, b, c, n, m, l, 4, 4, true) 
905            } else {
906                self.check_and_launch_for_mul("mul_a_bt", a, b, c, n, m, l, 8, 4, true) 
907            }
908        }
909    }
910
911    fn mul_at_bt(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize, l: usize) -> Result<()>
912    {
913        if self.has_cublas {
914            self.check_and_launch_cublas_for_mul(a, b, c, n, m, l, true, true)
915        } else {
916            if self.has_ptx {
917                self.check_and_launch_for_ptx_mul("ptx_mul_at_bt", a, b, c, n, m, l, 4, 4, false)
918            } else {
919                self.check_and_launch_for_mul("mul_at_bt", a, b, c, n, m, l, 8, 4, false)
920            }
921        }
922    }
923
924    fn mul_a_b_for_elems(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
925    { self.check_and_launch_for_op("mul_a_b_for_elems", a, b, c, n, m, 2, 2, true) }
926
927    fn mul_at_b_for_elems(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
928    { self.check_and_launch_for_op("mul_at_b_for_elems", a, b, c, n, m, 2, 2, true) }
929    
930    fn mul_a_bt_for_elems(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
931    { self.check_and_launch_for_op("mul_a_bt_for_elems", a, b, c, n, m, 2, 2, true) }
932    
933    fn mul_at_bt_for_elems(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
934    { self.check_and_launch_for_op("mul_at_bt_for_elems", a, b, c, n, m, 2, 2, true) }
935
936    fn div_a_b_for_elems(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
937    { self.check_and_launch_for_op("div_a_b_for_elems", a, b, c, n, m, 2, 2, true) }
938
939    fn div_at_b_for_elems(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
940    { self.check_and_launch_for_op("div_at_b_for_elems", a, b, c, n, m, 2, 2, true) }
941    
942    fn div_a_bt_for_elems(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
943    { self.check_and_launch_for_op("div_a_bt_for_elems", a, b, c, n, m, 2, 2, true) }
944    
945    fn div_at_bt_for_elems(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
946    { self.check_and_launch_for_op("div_at_bt_for_elems", a, b, c, n, m, 2, 2, true) }
947
948    fn add_a_b_for_scalar(&self, a: &BackendArray, b: f32, c: &BackendArray, n: usize, m: usize) -> Result<()>
949    { self.check_and_launch_for_scalar("add_a_b_for_scalar", a, b, c, n, m, 2, 2, true) }
950
951    fn add_at_b_for_scalar(&self, a: &BackendArray, b: f32, c: &BackendArray, n: usize, m: usize) -> Result<()>
952    { self.check_and_launch_for_scalar("add_at_b_for_scalar", a, b, c, n, m, 2, 2, true) }
953
954    fn sub_a_b_for_scalar(&self, a: &BackendArray, b: f32, c: &BackendArray, n: usize, m: usize) -> Result<()>
955    { self.check_and_launch_for_scalar("sub_a_b_for_scalar", a, b, c, n, m, 2, 2, true) }
956
957    fn sub_at_b_for_scalar(&self, a: &BackendArray, b: f32, c: &BackendArray, n: usize, m: usize) -> Result<()>
958    { self.check_and_launch_for_scalar("sub_at_b_for_scalar", a, b, c, n, m, 2, 2, true) }
959
960    fn rsub_a_b_for_scalar(&self, a: &BackendArray, b: f32, c: &BackendArray, n: usize, m: usize) -> Result<()>
961    { self.check_and_launch_for_scalar("rsub_a_b_for_scalar", a, b, c, n, m, 2, 2, true) }
962
963    fn rsub_at_b_for_scalar(&self, a: &BackendArray, b: f32, c: &BackendArray, n: usize, m: usize) -> Result<()>
964    { self.check_and_launch_for_scalar("rsub_at_b_for_scalar", a, b, c, n, m, 2, 2, true) }
965    
966    fn mul_a_b_for_scalar(&self, a: &BackendArray, b: f32, c: &BackendArray, n: usize, m: usize) -> Result<()>
967    { self.check_and_launch_for_scalar("mul_a_b_for_scalar", a, b, c, n, m, 2, 2, true) }
968
969    fn mul_at_b_for_scalar(&self, a: &BackendArray, b: f32, c: &BackendArray, n: usize, m: usize) -> Result<()>
970    { self.check_and_launch_for_scalar("mul_at_b_for_scalar", a, b, c, n, m, 2, 2, true) }
971
972    fn div_a_b_for_scalar(&self, a: &BackendArray, b: f32, c: &BackendArray, n: usize, m: usize) -> Result<()>
973    { self.check_and_launch_for_scalar("div_a_b_for_scalar", a, b, c, n, m, 2, 2, true) }
974
975    fn div_at_b_for_scalar(&self, a: &BackendArray, b: f32, c: &BackendArray, n: usize, m: usize) -> Result<()>
976    { self.check_and_launch_for_scalar("div_at_b_for_scalar", a, b, c, n, m, 2, 2, true) }
977
978    fn rdiv_a_b_for_scalar(&self, a: &BackendArray, b: f32, c: &BackendArray, n: usize, m: usize) -> Result<()>
979    { self.check_and_launch_for_scalar("rdiv_a_b_for_scalar", a, b, c, n, m, 2, 2, true) }
980
981    fn rdiv_at_b_for_scalar(&self, a: &BackendArray, b: f32, c: &BackendArray, n: usize, m: usize) -> Result<()>
982    { self.check_and_launch_for_scalar("rdiv_at_b_for_scalar", a, b, c, n, m, 2, 2, true) }
983
984    fn sigmoid_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
985    { self.check_and_launch_for_fun("sigmoid_a", a, b, n, m, 2, 2, true) }
986
987    fn sigmoid_at(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
988    { self.check_and_launch_for_fun("sigmoid_at", a, b, n, m, 2, 2, true) }
989
990    fn tanh_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
991    { self.check_and_launch_for_fun("tanh_a", a, b, n, m, 2, 2, true) }
992
993    fn tanh_at(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
994    { self.check_and_launch_for_fun("tanh_at", a, b, n, m, 2, 2, true) }
995
996    fn swish_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
997    { self.check_and_launch_for_fun("swish_a", a, b, n, m, 2, 2, true) }
998
999    fn swish_at(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1000    { self.check_and_launch_for_fun("swish_at", a, b, n, m, 2, 2, true) }
1001
1002    fn softmax_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1003    { self.check_and_launch_for_fun_and_tiles("softmax_a", a, b, n, m, 2, 2, true) }
1004
1005    fn softmax_at(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1006    { self.check_and_launch_for_fun_and_tiles("softmax_at", a, b, n, m, 2, 2, false) }
1007
1008    fn sqrt_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1009    { self.check_and_launch_for_fun("sqrt_a", a, b, n, m, 2, 2, true) }
1010
1011    fn sqrt_at(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1012    { self.check_and_launch_for_fun("sqrt_at", a, b, n, m, 2, 2, true) }
1013
1014    fn repeat_col_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1015    { self.check_and_launch_for_repeat_col("repeat_col_a", a, b, n, m, 2, 2, true) }
1016
1017    fn repeat_row_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1018    { self.check_and_launch_for_repeat_row("repeat_row_a", a, b, n, m, 2, 2, true) }
1019
1020    fn abs_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1021    { self.check_and_launch_for_fun("abs_a", a, b, n, m, 2, 2, true) }
1022
1023    fn abs_at(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1024    { self.check_and_launch_for_fun("abs_at", a, b, n, m, 2, 2, true) }
1025
1026    fn pow_a_b(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
1027    { self.check_and_launch_for_op("pow_a_b", a, b, c, n, m, 2, 2, true) }
1028
1029    fn pow_at_b(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
1030    { self.check_and_launch_for_op("pow_at_b", a, b, c, n, m, 2, 2, true) }
1031    
1032    fn pow_a_bt(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
1033    { self.check_and_launch_for_op("pow_a_bt", a, b, c, n, m, 2, 2, true) }
1034    
1035    fn pow_at_bt(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
1036    { self.check_and_launch_for_op("pow_at_bt", a, b, c, n, m, 2, 2, true) }
1037
1038    fn pow_a_b_for_scalar(&self, a: &BackendArray, b: f32, c: &BackendArray, n: usize, m: usize) -> Result<()>
1039    { self.check_and_launch_for_scalar("pow_a_b_for_scalar", a, b, c, n, m, 2, 2, true) }
1040
1041    fn pow_at_b_for_scalar(&self, a: &BackendArray, b: f32, c: &BackendArray, n: usize, m: usize) -> Result<()>
1042    { self.check_and_launch_for_scalar("pow_at_b_for_scalar", a, b, c, n, m, 2, 2, true) }
1043
1044    fn rpow_a_b_for_scalar(&self, a: &BackendArray, b: f32, c: &BackendArray, n: usize, m: usize) -> Result<()>
1045    { self.check_and_launch_for_scalar("rpow_a_b_for_scalar", a, b, c, n, m, 2, 2, true) }
1046
1047    fn rpow_at_b_for_scalar(&self, a: &BackendArray, b: f32, c: &BackendArray, n: usize, m: usize) -> Result<()>
1048    { self.check_and_launch_for_scalar("rpow_at_b_for_scalar", a, b, c, n, m, 2, 2, true) }
1049
1050    fn exp_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1051    { self.check_and_launch_for_fun("exp_a", a, b, n, m, 2, 2, true) }
1052
1053    fn exp_at(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1054    { self.check_and_launch_for_fun("exp_at", a, b, n, m, 2, 2, true) }
1055
1056    fn ln_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1057    { self.check_and_launch_for_fun("ln_a", a, b, n, m, 2, 2, true) }
1058
1059    fn ln_at(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1060    { self.check_and_launch_for_fun("ln_at", a, b, n, m, 2, 2, true) }
1061
1062    fn log2_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1063    { self.check_and_launch_for_fun("log2_a", a, b, n, m, 2, 2, true) }
1064
1065    fn log2_at(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1066    { self.check_and_launch_for_fun("log2_at", a, b, n, m, 2, 2, true) }
1067
1068    fn log10_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1069    { self.check_and_launch_for_fun("log10_a", a, b, n, m, 2, 2, true) }
1070
1071    fn log10_at(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1072    { self.check_and_launch_for_fun("log10_at", a, b, n, m, 2, 2, true) }
1073
1074    fn sin_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1075    { self.check_and_launch_for_fun("sin_a", a, b, n, m, 2, 2, true) }
1076
1077    fn sin_at(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1078    { self.check_and_launch_for_fun("sin_at", a, b, n, m, 2, 2, true) }
1079
1080    fn cos_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1081    { self.check_and_launch_for_fun("cos_a", a, b, n, m, 2, 2, true) }
1082
1083    fn cos_at(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1084    { self.check_and_launch_for_fun("cos_at", a, b, n, m, 2, 2, true) }
1085
1086    fn tan_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1087    { self.check_and_launch_for_fun("tan_a", a, b, n, m, 2, 2, true) }
1088
1089    fn tan_at(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1090    { self.check_and_launch_for_fun("tan_at", a, b, n, m, 2, 2, true) }
1091
1092    fn asin_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1093    { self.check_and_launch_for_fun("asin_a", a, b, n, m, 2, 2, true) }
1094
1095    fn asin_at(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1096    { self.check_and_launch_for_fun("asin_at", a, b, n, m, 2, 2, true) }
1097
1098    fn acos_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1099    { self.check_and_launch_for_fun("acos_a", a, b, n, m, 2, 2, true) }
1100
1101    fn acos_at(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1102    { self.check_and_launch_for_fun("acos_at", a, b, n, m, 2, 2, true) }
1103
1104    fn atan_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1105    { self.check_and_launch_for_fun("atan_a", a, b, n, m, 2, 2, true) }
1106
1107    fn atan_at(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1108    { self.check_and_launch_for_fun("atan_at", a, b, n, m, 2, 2, true) }
1109
1110    fn atan2_a_b(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
1111    { self.check_and_launch_for_op("atan2_a_b", a, b, c, n, m, 2, 2, true) }
1112
1113    fn atan2_at_b(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
1114    { self.check_and_launch_for_op("atan2_at_b", a, b, c, n, m, 2, 2, true) }
1115    
1116    fn atan2_a_bt(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
1117    { self.check_and_launch_for_op("atan2_a_bt", a, b, c, n, m, 2, 2, true) }
1118    
1119    fn atan2_at_bt(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
1120    { self.check_and_launch_for_op("atan2_at_bt", a, b, c, n, m, 2, 2, true) }
1121
1122    fn atan2_a_b_for_scalar(&self, a: &BackendArray, b: f32, c: &BackendArray, n: usize, m: usize) -> Result<()>
1123    { self.check_and_launch_for_scalar("atan2_a_b_for_scalar", a, b, c, n, m, 2, 2, true) }
1124
1125    fn atan2_at_b_for_scalar(&self, a: &BackendArray, b: f32, c: &BackendArray, n: usize, m: usize) -> Result<()>
1126    { self.check_and_launch_for_scalar("atan2_at_b_for_scalar", a, b, c, n, m, 2, 2, true) }
1127
1128    fn ratan2_a_b_for_scalar(&self, a: &BackendArray, b: f32, c: &BackendArray, n: usize, m: usize) -> Result<()>
1129    { self.check_and_launch_for_scalar("ratan2_a_b_for_scalar", a, b, c, n, m, 2, 2, true) }
1130
1131    fn ratan2_at_b_for_scalar(&self, a: &BackendArray, b: f32, c: &BackendArray, n: usize, m: usize) -> Result<()>
1132    { self.check_and_launch_for_scalar("ratan2_at_b_for_scalar", a, b, c, n, m, 2, 2, true) }
1133
1134    fn sinh_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1135    { self.check_and_launch_for_fun("sinh_a", a, b, n, m, 2, 2, true) }
1136
1137    fn sinh_at(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1138    { self.check_and_launch_for_fun("sinh_at", a, b, n, m, 2, 2, true) }
1139
1140    fn cosh_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1141    { self.check_and_launch_for_fun("cosh_a", a, b, n, m, 2, 2, true) }
1142
1143    fn cosh_at(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1144    { self.check_and_launch_for_fun("cosh_at", a, b, n, m, 2, 2, true) }
1145
1146    fn asinh_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1147    { self.check_and_launch_for_fun("asinh_a", a, b, n, m, 2, 2, true) }
1148
1149    fn asinh_at(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1150    { self.check_and_launch_for_fun("asinh_at", a, b, n, m, 2, 2, true) }
1151
1152    fn acosh_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1153    { self.check_and_launch_for_fun("acosh_a", a, b, n, m, 2, 2, true) }
1154
1155    fn acosh_at(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1156    { self.check_and_launch_for_fun("acosh_at", a, b, n, m, 2, 2, true) }
1157
1158    fn atanh_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1159    { self.check_and_launch_for_fun("atanh_a", a, b, n, m, 2, 2, true) }
1160
1161    fn atanh_at(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1162    { self.check_and_launch_for_fun("atanh_at", a, b, n, m, 2, 2, true) }
1163
1164    fn signum_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1165    { self.check_and_launch_for_fun("signum_a", a, b, n, m, 2, 2, true) }
1166
1167    fn signum_at(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1168    { self.check_and_launch_for_fun("signum_at", a, b, n, m, 2, 2, true) }
1169
1170    fn ceil_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1171    { self.check_and_launch_for_fun("ceil_a", a, b, n, m, 2, 2, true) }
1172
1173    fn ceil_at(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1174    { self.check_and_launch_for_fun("ceil_at", a, b, n, m, 2, 2, true) }
1175
1176    fn floor_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1177    { self.check_and_launch_for_fun("floor_a", a, b, n, m, 2, 2, true) }
1178
1179    fn floor_at(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1180    { self.check_and_launch_for_fun("floor_at", a, b, n, m, 2, 2, true) }
1181
1182    fn round_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1183    { self.check_and_launch_for_fun("round_a", a, b, n, m, 2, 2, true) }
1184
1185    fn round_at(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1186    { self.check_and_launch_for_fun("round_at", a, b, n, m, 2, 2, true) }
1187
1188    fn trunc_a(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1189    { self.check_and_launch_for_fun("trunc_a", a, b, n, m, 2, 2, true) }
1190
1191    fn trunc_at(&self, a: &BackendArray, b: &BackendArray, n: usize, m: usize) -> Result<()>
1192    { self.check_and_launch_for_fun("trunc_at", a, b, n, m, 2, 2, true) }
1193
1194    fn max_a_b(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
1195    { self.check_and_launch_for_op("max_a_b", a, b, c, n, m, 2, 2, true) }
1196
1197    fn max_at_b(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
1198    { self.check_and_launch_for_op("max_at_b", a, b, c, n, m, 2, 2, true) }
1199    
1200    fn max_a_bt(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
1201    { self.check_and_launch_for_op("max_a_bt", a, b, c, n, m, 2, 2, true) }
1202    
1203    fn max_at_bt(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
1204    { self.check_and_launch_for_op("max_at_bt", a, b, c, n, m, 2, 2, true) }
1205
1206    fn max_a_b_for_scalar(&self, a: &BackendArray, b: f32, c: &BackendArray, n: usize, m: usize) -> Result<()>
1207    { self.check_and_launch_for_scalar("max_a_b_for_scalar", a, b, c, n, m, 2, 2, true) }
1208
1209    fn max_at_b_for_scalar(&self, a: &BackendArray, b: f32, c: &BackendArray, n: usize, m: usize) -> Result<()>
1210    { self.check_and_launch_for_scalar("max_at_b_for_scalar", a, b, c, n, m, 2, 2, true) }
1211
1212    fn min_a_b(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
1213    { self.check_and_launch_for_op("min_a_b", a, b, c, n, m, 2, 2, true) }
1214
1215    fn min_at_b(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
1216    { self.check_and_launch_for_op("min_at_b", a, b, c, n, m, 2, 2, true) }
1217    
1218    fn min_a_bt(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
1219    { self.check_and_launch_for_op("min_a_bt", a, b, c, n, m, 2, 2, true) }
1220    
1221    fn min_at_bt(&self, a: &BackendArray, b: &BackendArray, c: &BackendArray, n: usize, m: usize) -> Result<()>
1222    { self.check_and_launch_for_op("min_at_bt", a, b, c, n, m, 2, 2, true) }
1223
1224    fn min_a_b_for_scalar(&self, a: &BackendArray, b: f32, c: &BackendArray, n: usize, m: usize) -> Result<()>
1225    { self.check_and_launch_for_scalar("min_a_b_for_scalar", a, b, c, n, m, 2, 2, true) }
1226
1227    fn min_at_b_for_scalar(&self, a: &BackendArray, b: f32, c: &BackendArray, n: usize, m: usize) -> Result<()>
1228    { self.check_and_launch_for_scalar("min_at_b_for_scalar", a, b, c, n, m, 2, 2, true) }
1229}
1230
1231#[cfg(test)]
1232mod tests;