1use 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#[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
62pub 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 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 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 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;