1#![warn(missing_debug_implementations)]
16
17use core::ffi::{c_int, c_void};
18use std::marker::PhantomData;
19
20use baracuda_cusolver_sys::{
21 cublasFillMode_t, cublasOperation_t, cuComplex, cuDoubleComplex, cusolver,
22 cusolverDnHandle_t, cusolverEigMode_t, cusolverStatus_t,
23};
24use baracuda_driver::{DeviceBuffer, Stream};
25use baracuda_types::{Complex32, Complex64, DeviceRepr};
26
27pub use baracuda_cusolver_sys::{
28 cublasFillMode_t as Fill, cusolverEigMode_t as EigMode,
29};
30
31pub type Error = baracuda_core::Error<cusolverStatus_t>;
33pub type Result<T, E = Error> = core::result::Result<T, E>;
35
36#[inline]
37fn check(status: cusolverStatus_t) -> Result<()> {
38 Error::check(status)
39}
40
41fn alloc_fail<E>(_e: E) -> Error {
43 Error::Status {
44 status: cusolverStatus_t::ALLOC_FAILED,
45 }
46}
47
48pub struct DnHandle {
52 handle: cusolverDnHandle_t,
53}
54
55unsafe impl Send for DnHandle {}
56
57impl core::fmt::Debug for DnHandle {
58 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
59 f.debug_struct("cusolver::DnHandle")
60 .field("handle", &self.handle)
61 .finish()
62 }
63}
64
65impl DnHandle {
66 pub fn new() -> Result<Self> {
68 let c = cusolver()?;
69 let cu = c.cusolver_dn_create()?;
70 let mut h: cusolverDnHandle_t = core::ptr::null_mut();
71 check(unsafe { cu(&mut h) })?;
72 Ok(Self { handle: h })
73 }
74
75 pub fn set_stream(&self, stream: &Stream) -> Result<()> {
78 let c = cusolver()?;
79 let cu = c.cusolver_dn_set_stream()?;
80 check(unsafe { cu(self.handle, stream.as_raw() as _) })
81 }
82
83 pub fn stream(&self) -> Result<*mut c_void> {
89 let c = cusolver()?;
90 let cu = c.cusolver_dn_get_stream()?;
91 let mut s: *mut c_void = core::ptr::null_mut();
92 check(unsafe { cu(self.handle, &mut s as *mut *mut c_void as *mut _) })?;
93 Ok(s)
94 }
95
96 pub fn version() -> Result<i32> {
99 let c = cusolver()?;
100 let cu = c.cusolver_get_version()?;
101 let mut v: c_int = 0;
102 check(unsafe { cu(&mut v) })?;
103 Ok(v)
104 }
105
106 #[inline]
108 pub fn as_raw(&self) -> cusolverDnHandle_t {
109 self.handle
110 }
111}
112
113impl Drop for DnHandle {
114 fn drop(&mut self) {
115 if let Ok(c) = cusolver() {
116 if let Ok(cu) = c.cusolver_dn_destroy() {
117 let _ = unsafe { cu(self.handle) };
118 }
119 }
120 }
121}
122
123#[derive(Copy, Clone, Debug, Eq, PartialEq, Default)]
125pub enum Op {
126 #[default]
128 N,
129 T,
131 C,
133}
134
135impl Op {
136 fn raw(self) -> cublasOperation_t {
137 match self {
138 Op::N => cublasOperation_t::N,
139 Op::T => cublasOperation_t::T,
140 Op::C => cublasOperation_t::C,
141 }
142 }
143}
144
145pub trait SolverScalar: DeviceRepr + Copy + 'static + sealed::Sealed {
149 type Real: DeviceRepr + Copy + 'static;
152
153 #[doc(hidden)]
155 unsafe fn getrf_buf(
156 h: cusolverDnHandle_t,
157 m: c_int,
158 n: c_int,
159 a: *mut Self,
160 lda: c_int,
161 lwork: *mut c_int,
162 ) -> cusolverStatus_t;
163
164 #[doc(hidden)]
166 #[allow(clippy::too_many_arguments)]
167 unsafe fn getrf(
168 h: cusolverDnHandle_t,
169 m: c_int,
170 n: c_int,
171 a: *mut Self,
172 lda: c_int,
173 workspace: *mut Self,
174 ipiv: *mut c_int,
175 info: *mut c_int,
176 ) -> cusolverStatus_t;
177
178 #[doc(hidden)]
180 #[allow(clippy::too_many_arguments)]
181 unsafe fn getrs(
182 h: cusolverDnHandle_t,
183 trans: cublasOperation_t,
184 n: c_int,
185 nrhs: c_int,
186 a: *const Self,
187 lda: c_int,
188 ipiv: *const c_int,
189 b: *mut Self,
190 ldb: c_int,
191 info: *mut c_int,
192 ) -> cusolverStatus_t;
193
194 #[doc(hidden)]
196 unsafe fn geqrf_buf(
197 h: cusolverDnHandle_t,
198 m: c_int,
199 n: c_int,
200 a: *mut Self,
201 lda: c_int,
202 lwork: *mut c_int,
203 ) -> cusolverStatus_t;
204
205 #[doc(hidden)]
207 #[allow(clippy::too_many_arguments)]
208 unsafe fn geqrf(
209 h: cusolverDnHandle_t,
210 m: c_int,
211 n: c_int,
212 a: *mut Self,
213 lda: c_int,
214 tau: *mut Self,
215 workspace: *mut Self,
216 lwork: c_int,
217 info: *mut c_int,
218 ) -> cusolverStatus_t;
219
220 #[doc(hidden)]
222 unsafe fn potrf_buf(
223 h: cusolverDnHandle_t,
224 uplo: cublasFillMode_t,
225 n: c_int,
226 a: *mut Self,
227 lda: c_int,
228 lwork: *mut c_int,
229 ) -> cusolverStatus_t;
230
231 #[doc(hidden)]
233 #[allow(clippy::too_many_arguments)]
234 unsafe fn potrf(
235 h: cusolverDnHandle_t,
236 uplo: cublasFillMode_t,
237 n: c_int,
238 a: *mut Self,
239 lda: c_int,
240 workspace: *mut Self,
241 lwork: c_int,
242 info: *mut c_int,
243 ) -> cusolverStatus_t;
244
245 #[doc(hidden)]
247 #[allow(clippy::too_many_arguments)]
248 unsafe fn potrs(
249 h: cusolverDnHandle_t,
250 uplo: cublasFillMode_t,
251 n: c_int,
252 nrhs: c_int,
253 a: *const Self,
254 lda: c_int,
255 b: *mut Self,
256 ldb: c_int,
257 info: *mut c_int,
258 ) -> cusolverStatus_t;
259
260 #[doc(hidden)]
262 unsafe fn gesvd_buf(
263 h: cusolverDnHandle_t,
264 m: c_int,
265 n: c_int,
266 lwork: *mut c_int,
267 ) -> cusolverStatus_t;
268
269 #[doc(hidden)]
271 #[allow(clippy::too_many_arguments)]
272 unsafe fn gesvd(
273 h: cusolverDnHandle_t,
274 jobu: u8,
275 jobvt: u8,
276 m: c_int,
277 n: c_int,
278 a: *mut Self,
279 lda: c_int,
280 s: *mut Self::Real,
281 u: *mut Self,
282 ldu: c_int,
283 vt: *mut Self,
284 ldvt: c_int,
285 work: *mut Self,
286 lwork: c_int,
287 rwork: *mut Self::Real,
288 info: *mut c_int,
289 ) -> cusolverStatus_t;
290
291 #[doc(hidden)]
293 #[allow(clippy::too_many_arguments)]
294 unsafe fn syevd_buf(
295 h: cusolverDnHandle_t,
296 jobz: cusolverEigMode_t,
297 uplo: cublasFillMode_t,
298 n: c_int,
299 a: *const Self,
300 lda: c_int,
301 w: *const Self::Real,
302 lwork: *mut c_int,
303 ) -> cusolverStatus_t;
304
305 #[doc(hidden)]
307 #[allow(clippy::too_many_arguments)]
308 unsafe fn syevd(
309 h: cusolverDnHandle_t,
310 jobz: cusolverEigMode_t,
311 uplo: cublasFillMode_t,
312 n: c_int,
313 a: *mut Self,
314 lda: c_int,
315 w: *mut Self::Real,
316 work: *mut Self,
317 lwork: c_int,
318 info: *mut c_int,
319 ) -> cusolverStatus_t;
320}
321
322mod sealed {
323 use baracuda_types::{Complex32, Complex64};
324 pub trait Sealed {}
325 impl Sealed for f32 {}
326 impl Sealed for f64 {}
327 impl Sealed for Complex32 {}
328 impl Sealed for Complex64 {}
329}
330
331macro_rules! real_impl {
332 ($t:ty, $getrf_buf:ident, $getrf:ident, $getrs:ident,
333 $geqrf_buf:ident, $geqrf:ident,
334 $potrf_buf:ident, $potrf:ident, $potrs:ident,
335 $gesvd_buf:ident, $gesvd:ident,
336 $syevd_buf:ident, $syevd:ident) => {
337 impl SolverScalar for $t {
338 type Real = $t;
339
340 unsafe fn getrf_buf(
341 h: cusolverDnHandle_t,
342 m: c_int,
343 n: c_int,
344 a: *mut $t,
345 lda: c_int,
346 lwork: *mut c_int,
347 ) -> cusolverStatus_t { unsafe {
348 match cusolver().and_then(|c| c.$getrf_buf()) {
349 Ok(f) => f(h, m, n, a, lda, lwork),
350 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
351 }
352 }}
353 unsafe fn getrf(
354 h: cusolverDnHandle_t,
355 m: c_int,
356 n: c_int,
357 a: *mut $t,
358 lda: c_int,
359 work: *mut $t,
360 ipiv: *mut c_int,
361 info: *mut c_int,
362 ) -> cusolverStatus_t { unsafe {
363 match cusolver().and_then(|c| c.$getrf()) {
364 Ok(f) => f(h, m, n, a, lda, work, ipiv, info),
365 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
366 }
367 }}
368 unsafe fn getrs(
369 h: cusolverDnHandle_t,
370 trans: cublasOperation_t,
371 n: c_int,
372 nrhs: c_int,
373 a: *const $t,
374 lda: c_int,
375 ipiv: *const c_int,
376 b: *mut $t,
377 ldb: c_int,
378 info: *mut c_int,
379 ) -> cusolverStatus_t { unsafe {
380 match cusolver().and_then(|c| c.$getrs()) {
381 Ok(f) => f(h, trans, n, nrhs, a, lda, ipiv, b, ldb, info),
382 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
383 }
384 }}
385 unsafe fn geqrf_buf(
386 h: cusolverDnHandle_t,
387 m: c_int,
388 n: c_int,
389 a: *mut $t,
390 lda: c_int,
391 lwork: *mut c_int,
392 ) -> cusolverStatus_t { unsafe {
393 match cusolver().and_then(|c| c.$geqrf_buf()) {
394 Ok(f) => f(h, m, n, a, lda, lwork),
395 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
396 }
397 }}
398 unsafe fn geqrf(
399 h: cusolverDnHandle_t,
400 m: c_int,
401 n: c_int,
402 a: *mut $t,
403 lda: c_int,
404 tau: *mut $t,
405 work: *mut $t,
406 lwork: c_int,
407 info: *mut c_int,
408 ) -> cusolverStatus_t { unsafe {
409 match cusolver().and_then(|c| c.$geqrf()) {
410 Ok(f) => f(h, m, n, a, lda, tau, work, lwork, info),
411 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
412 }
413 }}
414 unsafe fn potrf_buf(
415 h: cusolverDnHandle_t,
416 uplo: cublasFillMode_t,
417 n: c_int,
418 a: *mut $t,
419 lda: c_int,
420 lwork: *mut c_int,
421 ) -> cusolverStatus_t { unsafe {
422 match cusolver().and_then(|c| c.$potrf_buf()) {
423 Ok(f) => f(h, uplo, n, a, lda, lwork),
424 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
425 }
426 }}
427 unsafe fn potrf(
428 h: cusolverDnHandle_t,
429 uplo: cublasFillMode_t,
430 n: c_int,
431 a: *mut $t,
432 lda: c_int,
433 work: *mut $t,
434 lwork: c_int,
435 info: *mut c_int,
436 ) -> cusolverStatus_t { unsafe {
437 match cusolver().and_then(|c| c.$potrf()) {
438 Ok(f) => f(h, uplo, n, a, lda, work, lwork, info),
439 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
440 }
441 }}
442 unsafe fn potrs(
443 h: cusolverDnHandle_t,
444 uplo: cublasFillMode_t,
445 n: c_int,
446 nrhs: c_int,
447 a: *const $t,
448 lda: c_int,
449 b: *mut $t,
450 ldb: c_int,
451 info: *mut c_int,
452 ) -> cusolverStatus_t { unsafe {
453 match cusolver().and_then(|c| c.$potrs()) {
454 Ok(f) => f(h, uplo, n, nrhs, a, lda, b, ldb, info),
455 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
456 }
457 }}
458 unsafe fn gesvd_buf(
459 h: cusolverDnHandle_t,
460 m: c_int,
461 n: c_int,
462 lwork: *mut c_int,
463 ) -> cusolverStatus_t { unsafe {
464 match cusolver().and_then(|c| c.$gesvd_buf()) {
465 Ok(f) => f(h, m, n, lwork),
466 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
467 }
468 }}
469 unsafe fn gesvd(
470 h: cusolverDnHandle_t,
471 jobu: u8,
472 jobvt: u8,
473 m: c_int,
474 n: c_int,
475 a: *mut $t,
476 lda: c_int,
477 s: *mut $t,
478 u: *mut $t,
479 ldu: c_int,
480 vt: *mut $t,
481 ldvt: c_int,
482 work: *mut $t,
483 lwork: c_int,
484 rwork: *mut $t,
485 info: *mut c_int,
486 ) -> cusolverStatus_t { unsafe {
487 match cusolver().and_then(|c| c.$gesvd()) {
488 Ok(f) => f(
489 h, jobu, jobvt, m, n, a, lda, s, u, ldu, vt, ldvt, work, lwork, rwork, info,
490 ),
491 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
492 }
493 }}
494 unsafe fn syevd_buf(
495 h: cusolverDnHandle_t,
496 jobz: cusolverEigMode_t,
497 uplo: cublasFillMode_t,
498 n: c_int,
499 a: *const $t,
500 lda: c_int,
501 w: *const $t,
502 lwork: *mut c_int,
503 ) -> cusolverStatus_t { unsafe {
504 match cusolver().and_then(|c| c.$syevd_buf()) {
505 Ok(f) => f(h, jobz, uplo, n, a, lda, w, lwork),
506 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
507 }
508 }}
509 unsafe fn syevd(
510 h: cusolverDnHandle_t,
511 jobz: cusolverEigMode_t,
512 uplo: cublasFillMode_t,
513 n: c_int,
514 a: *mut $t,
515 lda: c_int,
516 w: *mut $t,
517 work: *mut $t,
518 lwork: c_int,
519 info: *mut c_int,
520 ) -> cusolverStatus_t { unsafe {
521 match cusolver().and_then(|c| c.$syevd()) {
522 Ok(f) => f(h, jobz, uplo, n, a, lda, w, work, lwork, info),
523 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
524 }
525 }}
526 }
527 };
528}
529
530macro_rules! complex_impl {
531 ($t:ty, $real:ty, $raw:ty,
532 $getrf_buf:ident, $getrf:ident, $getrs:ident,
533 $geqrf_buf:ident, $geqrf:ident,
534 $potrf_buf:ident, $potrf:ident, $potrs:ident,
535 $gesvd_buf:ident, $gesvd:ident,
536 $heevd_buf:ident, $heevd:ident) => {
537 impl SolverScalar for $t {
538 type Real = $real;
539
540 unsafe fn getrf_buf(
541 h: cusolverDnHandle_t,
542 m: c_int,
543 n: c_int,
544 a: *mut $t,
545 lda: c_int,
546 lwork: *mut c_int,
547 ) -> cusolverStatus_t { unsafe {
548 match cusolver().and_then(|c| c.$getrf_buf()) {
549 Ok(f) => f(h, m, n, a as *mut $raw, lda, lwork),
550 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
551 }
552 }}
553 unsafe fn getrf(
554 h: cusolverDnHandle_t,
555 m: c_int,
556 n: c_int,
557 a: *mut $t,
558 lda: c_int,
559 work: *mut $t,
560 ipiv: *mut c_int,
561 info: *mut c_int,
562 ) -> cusolverStatus_t { unsafe {
563 match cusolver().and_then(|c| c.$getrf()) {
564 Ok(f) => f(
565 h,
566 m,
567 n,
568 a as *mut $raw,
569 lda,
570 work as *mut $raw,
571 ipiv,
572 info,
573 ),
574 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
575 }
576 }}
577 unsafe fn getrs(
578 h: cusolverDnHandle_t,
579 trans: cublasOperation_t,
580 n: c_int,
581 nrhs: c_int,
582 a: *const $t,
583 lda: c_int,
584 ipiv: *const c_int,
585 b: *mut $t,
586 ldb: c_int,
587 info: *mut c_int,
588 ) -> cusolverStatus_t { unsafe {
589 match cusolver().and_then(|c| c.$getrs()) {
590 Ok(f) => f(
591 h,
592 trans,
593 n,
594 nrhs,
595 a as *const $raw,
596 lda,
597 ipiv,
598 b as *mut $raw,
599 ldb,
600 info,
601 ),
602 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
603 }
604 }}
605 unsafe fn geqrf_buf(
606 h: cusolverDnHandle_t,
607 m: c_int,
608 n: c_int,
609 a: *mut $t,
610 lda: c_int,
611 lwork: *mut c_int,
612 ) -> cusolverStatus_t { unsafe {
613 match cusolver().and_then(|c| c.$geqrf_buf()) {
614 Ok(f) => f(h, m, n, a as *mut $raw, lda, lwork),
615 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
616 }
617 }}
618 unsafe fn geqrf(
619 h: cusolverDnHandle_t,
620 m: c_int,
621 n: c_int,
622 a: *mut $t,
623 lda: c_int,
624 tau: *mut $t,
625 work: *mut $t,
626 lwork: c_int,
627 info: *mut c_int,
628 ) -> cusolverStatus_t { unsafe {
629 match cusolver().and_then(|c| c.$geqrf()) {
630 Ok(f) => f(
631 h,
632 m,
633 n,
634 a as *mut $raw,
635 lda,
636 tau as *mut $raw,
637 work as *mut $raw,
638 lwork,
639 info,
640 ),
641 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
642 }
643 }}
644 unsafe fn potrf_buf(
645 h: cusolverDnHandle_t,
646 uplo: cublasFillMode_t,
647 n: c_int,
648 a: *mut $t,
649 lda: c_int,
650 lwork: *mut c_int,
651 ) -> cusolverStatus_t { unsafe {
652 match cusolver().and_then(|c| c.$potrf_buf()) {
653 Ok(f) => f(h, uplo, n, a as *mut $raw, lda, lwork),
654 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
655 }
656 }}
657 unsafe fn potrf(
658 h: cusolverDnHandle_t,
659 uplo: cublasFillMode_t,
660 n: c_int,
661 a: *mut $t,
662 lda: c_int,
663 work: *mut $t,
664 lwork: c_int,
665 info: *mut c_int,
666 ) -> cusolverStatus_t { unsafe {
667 match cusolver().and_then(|c| c.$potrf()) {
668 Ok(f) => f(
669 h,
670 uplo,
671 n,
672 a as *mut $raw,
673 lda,
674 work as *mut $raw,
675 lwork,
676 info,
677 ),
678 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
679 }
680 }}
681 unsafe fn potrs(
682 h: cusolverDnHandle_t,
683 uplo: cublasFillMode_t,
684 n: c_int,
685 nrhs: c_int,
686 a: *const $t,
687 lda: c_int,
688 b: *mut $t,
689 ldb: c_int,
690 info: *mut c_int,
691 ) -> cusolverStatus_t { unsafe {
692 match cusolver().and_then(|c| c.$potrs()) {
693 Ok(f) => f(
694 h,
695 uplo,
696 n,
697 nrhs,
698 a as *const $raw,
699 lda,
700 b as *mut $raw,
701 ldb,
702 info,
703 ),
704 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
705 }
706 }}
707 unsafe fn gesvd_buf(
708 h: cusolverDnHandle_t,
709 m: c_int,
710 n: c_int,
711 lwork: *mut c_int,
712 ) -> cusolverStatus_t { unsafe {
713 match cusolver().and_then(|c| c.$gesvd_buf()) {
714 Ok(f) => f(h, m, n, lwork),
715 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
716 }
717 }}
718 unsafe fn gesvd(
719 h: cusolverDnHandle_t,
720 jobu: u8,
721 jobvt: u8,
722 m: c_int,
723 n: c_int,
724 a: *mut $t,
725 lda: c_int,
726 s: *mut $real,
727 u: *mut $t,
728 ldu: c_int,
729 vt: *mut $t,
730 ldvt: c_int,
731 work: *mut $t,
732 lwork: c_int,
733 rwork: *mut $real,
734 info: *mut c_int,
735 ) -> cusolverStatus_t { unsafe {
736 match cusolver().and_then(|c| c.$gesvd()) {
737 Ok(f) => f(
738 h,
739 jobu,
740 jobvt,
741 m,
742 n,
743 a as *mut $raw,
744 lda,
745 s,
746 u as *mut $raw,
747 ldu,
748 vt as *mut $raw,
749 ldvt,
750 work as *mut $raw,
751 lwork,
752 rwork,
753 info,
754 ),
755 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
756 }
757 }}
758 unsafe fn syevd_buf(
759 h: cusolverDnHandle_t,
760 jobz: cusolverEigMode_t,
761 uplo: cublasFillMode_t,
762 n: c_int,
763 a: *const $t,
764 lda: c_int,
765 w: *const $real,
766 lwork: *mut c_int,
767 ) -> cusolverStatus_t { unsafe {
768 match cusolver().and_then(|c| c.$heevd_buf()) {
769 Ok(f) => f(h, jobz, uplo, n, a as *const $raw, lda, w, lwork),
770 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
771 }
772 }}
773 unsafe fn syevd(
774 h: cusolverDnHandle_t,
775 jobz: cusolverEigMode_t,
776 uplo: cublasFillMode_t,
777 n: c_int,
778 a: *mut $t,
779 lda: c_int,
780 w: *mut $real,
781 work: *mut $t,
782 lwork: c_int,
783 info: *mut c_int,
784 ) -> cusolverStatus_t { unsafe {
785 match cusolver().and_then(|c| c.$heevd()) {
786 Ok(f) => f(
787 h,
788 jobz,
789 uplo,
790 n,
791 a as *mut $raw,
792 lda,
793 w,
794 work as *mut $raw,
795 lwork,
796 info,
797 ),
798 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
799 }
800 }}
801 }
802 };
803}
804
805real_impl!(
806 f32,
807 cusolver_dn_sgetrf_buffer_size,
808 cusolver_dn_sgetrf,
809 cusolver_dn_sgetrs,
810 cusolver_dn_sgeqrf_buffer_size,
811 cusolver_dn_sgeqrf,
812 cusolver_dn_spotrf_buffer_size,
813 cusolver_dn_spotrf,
814 cusolver_dn_spotrs,
815 cusolver_dn_sgesvd_buffer_size,
816 cusolver_dn_sgesvd,
817 cusolver_dn_ssyevd_buffer_size,
818 cusolver_dn_ssyevd
819);
820
821real_impl!(
822 f64,
823 cusolver_dn_dgetrf_buffer_size,
824 cusolver_dn_dgetrf,
825 cusolver_dn_dgetrs,
826 cusolver_dn_dgeqrf_buffer_size,
827 cusolver_dn_dgeqrf,
828 cusolver_dn_dpotrf_buffer_size,
829 cusolver_dn_dpotrf,
830 cusolver_dn_dpotrs,
831 cusolver_dn_dgesvd_buffer_size,
832 cusolver_dn_dgesvd,
833 cusolver_dn_dsyevd_buffer_size,
834 cusolver_dn_dsyevd
835);
836
837complex_impl!(
838 Complex32,
839 f32,
840 cuComplex,
841 cusolver_dn_cgetrf_buffer_size,
842 cusolver_dn_cgetrf,
843 cusolver_dn_cgetrs,
844 cusolver_dn_cgeqrf_buffer_size,
845 cusolver_dn_cgeqrf,
846 cusolver_dn_cpotrf_buffer_size,
847 cusolver_dn_cpotrf,
848 cusolver_dn_cpotrs,
849 cusolver_dn_cgesvd_buffer_size,
850 cusolver_dn_cgesvd,
851 cusolver_dn_cheevd_buffer_size,
852 cusolver_dn_cheevd
853);
854
855complex_impl!(
856 Complex64,
857 f64,
858 cuDoubleComplex,
859 cusolver_dn_zgetrf_buffer_size,
860 cusolver_dn_zgetrf,
861 cusolver_dn_zgetrs,
862 cusolver_dn_zgeqrf_buffer_size,
863 cusolver_dn_zgeqrf,
864 cusolver_dn_zpotrf_buffer_size,
865 cusolver_dn_zpotrf,
866 cusolver_dn_zpotrs,
867 cusolver_dn_zgesvd_buffer_size,
868 cusolver_dn_zgesvd,
869 cusolver_dn_zheevd_buffer_size,
870 cusolver_dn_zheevd
871);
872
873#[allow(clippy::too_many_arguments)]
905pub fn getrf<T: SolverScalar>(
906 handle: &DnHandle,
907 m: i32,
908 n: i32,
909 a: &mut DeviceBuffer<T>,
910 lda: i32,
911 ipiv: &mut DeviceBuffer<i32>,
912 info: &mut DeviceBuffer<i32>,
913) -> Result<()> {
914 let mut lwork: c_int = 0;
915 check(unsafe { T::getrf_buf(handle.handle, m, n, a.as_raw().0 as *mut T, lda, &mut lwork) })?;
916 let workspace =
917 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
918 check(unsafe {
919 T::getrf(
920 handle.handle,
921 m,
922 n,
923 a.as_raw().0 as *mut T,
924 lda,
925 workspace.as_raw().0 as *mut T,
926 ipiv.as_raw().0 as *mut c_int,
927 info.as_raw().0 as *mut c_int,
928 )
929 })
930}
931
932#[allow(clippy::too_many_arguments)]
934pub fn getrs<T: SolverScalar>(
935 handle: &DnHandle,
936 trans: Op,
937 n: i32,
938 nrhs: i32,
939 a: &DeviceBuffer<T>,
940 lda: i32,
941 ipiv: &DeviceBuffer<i32>,
942 b: &mut DeviceBuffer<T>,
943 ldb: i32,
944 info: &mut DeviceBuffer<i32>,
945) -> Result<()> {
946 check(unsafe {
947 T::getrs(
948 handle.handle,
949 trans.raw(),
950 n,
951 nrhs,
952 a.as_raw().0 as *const T,
953 lda,
954 ipiv.as_raw().0 as *const c_int,
955 b.as_raw().0 as *mut T,
956 ldb,
957 info.as_raw().0 as *mut c_int,
958 )
959 })
960}
961
962#[allow(clippy::too_many_arguments)]
990pub fn geqrf<T: SolverScalar>(
991 handle: &DnHandle,
992 m: i32,
993 n: i32,
994 a: &mut DeviceBuffer<T>,
995 lda: i32,
996 tau: &mut DeviceBuffer<T>,
997 info: &mut DeviceBuffer<i32>,
998) -> Result<()> {
999 let mut lwork: c_int = 0;
1000 check(unsafe { T::geqrf_buf(handle.handle, m, n, a.as_raw().0 as *mut T, lda, &mut lwork) })?;
1001 let workspace =
1002 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1003 check(unsafe {
1004 T::geqrf(
1005 handle.handle,
1006 m,
1007 n,
1008 a.as_raw().0 as *mut T,
1009 lda,
1010 tau.as_raw().0 as *mut T,
1011 workspace.as_raw().0 as *mut T,
1012 lwork,
1013 info.as_raw().0 as *mut c_int,
1014 )
1015 })
1016}
1017
1018pub fn potrf<T: SolverScalar>(
1045 handle: &DnHandle,
1046 uplo: Fill,
1047 n: i32,
1048 a: &mut DeviceBuffer<T>,
1049 lda: i32,
1050 info: &mut DeviceBuffer<i32>,
1051) -> Result<()> {
1052 let mut lwork: c_int = 0;
1053 check(unsafe { T::potrf_buf(handle.handle, uplo, n, a.as_raw().0 as *mut T, lda, &mut lwork) })?;
1054 let workspace =
1055 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1056 check(unsafe {
1057 T::potrf(
1058 handle.handle,
1059 uplo,
1060 n,
1061 a.as_raw().0 as *mut T,
1062 lda,
1063 workspace.as_raw().0 as *mut T,
1064 lwork,
1065 info.as_raw().0 as *mut c_int,
1066 )
1067 })
1068}
1069
1070#[allow(clippy::too_many_arguments)]
1072pub fn potrs<T: SolverScalar>(
1073 handle: &DnHandle,
1074 uplo: Fill,
1075 n: i32,
1076 nrhs: i32,
1077 a: &DeviceBuffer<T>,
1078 lda: i32,
1079 b: &mut DeviceBuffer<T>,
1080 ldb: i32,
1081 info: &mut DeviceBuffer<i32>,
1082) -> Result<()> {
1083 check(unsafe {
1084 T::potrs(
1085 handle.handle,
1086 uplo,
1087 n,
1088 nrhs,
1089 a.as_raw().0 as *const T,
1090 lda,
1091 b.as_raw().0 as *mut T,
1092 ldb,
1093 info.as_raw().0 as *mut c_int,
1094 )
1095 })
1096}
1097
1098#[allow(clippy::too_many_arguments)]
1137pub fn gesvd<T: SolverScalar>(
1138 handle: &DnHandle,
1139 jobu: u8,
1140 jobvt: u8,
1141 m: i32,
1142 n: i32,
1143 a: &mut DeviceBuffer<T>,
1144 lda: i32,
1145 s: &mut DeviceBuffer<T::Real>,
1146 u: &mut DeviceBuffer<T>,
1147 ldu: i32,
1148 vt: &mut DeviceBuffer<T>,
1149 ldvt: i32,
1150 rwork: &mut DeviceBuffer<T::Real>,
1151 info: &mut DeviceBuffer<i32>,
1152) -> Result<()> {
1153 let mut lwork: c_int = 0;
1154 check(unsafe { T::gesvd_buf(handle.handle, m, n, &mut lwork) })?;
1155 let workspace =
1156 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1157 check(unsafe {
1158 T::gesvd(
1159 handle.handle,
1160 jobu,
1161 jobvt,
1162 m,
1163 n,
1164 a.as_raw().0 as *mut T,
1165 lda,
1166 s.as_raw().0 as *mut T::Real,
1167 u.as_raw().0 as *mut T,
1168 ldu,
1169 vt.as_raw().0 as *mut T,
1170 ldvt,
1171 workspace.as_raw().0 as *mut T,
1172 lwork,
1173 rwork.as_raw().0 as *mut T::Real,
1174 info.as_raw().0 as *mut c_int,
1175 )
1176 })
1177}
1178
1179#[allow(clippy::too_many_arguments)]
1181pub fn syevd<T: SolverScalar>(
1182 handle: &DnHandle,
1183 jobz: EigMode,
1184 uplo: Fill,
1185 n: i32,
1186 a: &mut DeviceBuffer<T>,
1187 lda: i32,
1188 w: &mut DeviceBuffer<T::Real>,
1189 info: &mut DeviceBuffer<i32>,
1190) -> Result<()> {
1191 let mut lwork: c_int = 0;
1192 check(unsafe {
1193 T::syevd_buf(
1194 handle.handle,
1195 jobz,
1196 uplo,
1197 n,
1198 a.as_raw().0 as *const T,
1199 lda,
1200 w.as_raw().0 as *const T::Real,
1201 &mut lwork,
1202 )
1203 })?;
1204 let workspace =
1205 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1206 check(unsafe {
1207 T::syevd(
1208 handle.handle,
1209 jobz,
1210 uplo,
1211 n,
1212 a.as_raw().0 as *mut T,
1213 lda,
1214 w.as_raw().0 as *mut T::Real,
1215 workspace.as_raw().0 as *mut T,
1216 lwork,
1217 info.as_raw().0 as *mut c_int,
1218 )
1219 })
1220}
1221
1222pub use baracuda_cusolver_sys::{gesvdjInfo_t as GesvdjInfoRaw, syevjInfo_t as SyevjInfoRaw};
1225
1226#[derive(Debug)]
1228pub struct SyevjInfo {
1229 raw: SyevjInfoRaw,
1230}
1231
1232impl SyevjInfo {
1233 pub fn new() -> Result<Self> {
1235 let c = cusolver()?;
1236 let cu = c.cusolver_dn_create_syevj_info()?;
1237 let mut raw: SyevjInfoRaw = core::ptr::null_mut();
1238 check(unsafe { cu(&mut raw) })?;
1239 Ok(Self { raw })
1240 }
1241
1242 pub fn set_tolerance(&self, tol: f64) -> Result<()> {
1245 let c = cusolver()?;
1246 let cu = c.cusolver_dn_xsyevj_set_tolerance()?;
1247 check(unsafe { cu(self.raw, tol) })
1248 }
1249
1250 pub fn set_max_sweeps(&self, n: i32) -> Result<()> {
1253 let c = cusolver()?;
1254 let cu = c.cusolver_dn_xsyevj_set_max_sweeps()?;
1255 check(unsafe { cu(self.raw, n) })
1256 }
1257
1258 pub fn as_raw(&self) -> SyevjInfoRaw {
1260 self.raw
1261 }
1262}
1263
1264impl Drop for SyevjInfo {
1265 fn drop(&mut self) {
1266 if let Ok(c) = cusolver() {
1267 if let Ok(cu) = c.cusolver_dn_destroy_syevj_info() {
1268 let _ = unsafe { cu(self.raw) };
1269 }
1270 }
1271 }
1272}
1273
1274#[derive(Debug)]
1276pub struct GesvdjInfo {
1277 raw: GesvdjInfoRaw,
1278}
1279
1280impl GesvdjInfo {
1281 pub fn new() -> Result<Self> {
1283 let c = cusolver()?;
1284 let cu = c.cusolver_dn_create_gesvdj_info()?;
1285 let mut raw: GesvdjInfoRaw = core::ptr::null_mut();
1286 check(unsafe { cu(&mut raw) })?;
1287 Ok(Self { raw })
1288 }
1289
1290 pub fn as_raw(&self) -> GesvdjInfoRaw {
1292 self.raw
1293 }
1294}
1295
1296impl Drop for GesvdjInfo {
1297 fn drop(&mut self) {
1298 if let Ok(c) = cusolver() {
1299 if let Ok(cu) = c.cusolver_dn_destroy_gesvdj_info() {
1300 let _ = unsafe { cu(self.raw) };
1301 }
1302 }
1303 }
1304}
1305
1306#[allow(clippy::too_many_arguments)]
1309pub fn syevj<T: SolverScalar>(
1310 handle: &DnHandle,
1311 jobz: EigMode,
1312 uplo: Fill,
1313 n: i32,
1314 a: &mut DeviceBuffer<T>,
1315 lda: i32,
1316 w: &mut DeviceBuffer<T::Real>,
1317 info: &mut DeviceBuffer<i32>,
1318 params: &SyevjInfo,
1319) -> Result<()> {
1320 use baracuda_cusolver_sys::{
1321 cuComplex, cuDoubleComplex,
1322 };
1323 use core::mem;
1324
1325 let mut lwork: c_int = 0;
1326
1327 macro_rules! dispatch_real {
1330 ($t:ty, $bufsize:ident, $solve:ident) => {{
1331 let c = cusolver()?;
1332 check(unsafe {
1333 (c.$bufsize()?)(
1334 handle.as_raw(),
1335 jobz,
1336 uplo,
1337 n,
1338 a.as_raw().0 as *const $t,
1339 lda,
1340 w.as_raw().0 as *const $t,
1341 &mut lwork,
1342 params.raw,
1343 )
1344 })?;
1345 let workspace =
1346 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1347 check(unsafe {
1348 (c.$solve()?)(
1349 handle.as_raw(),
1350 jobz,
1351 uplo,
1352 n,
1353 a.as_raw().0 as *mut $t,
1354 lda,
1355 w.as_raw().0 as *mut $t,
1356 workspace.as_raw().0 as *mut $t,
1357 lwork,
1358 info.as_raw().0 as *mut c_int,
1359 params.raw,
1360 )
1361 })
1362 }};
1363 }
1364 macro_rules! dispatch_complex {
1365 ($t:ty, $real:ty, $raw:ty, $bufsize:ident, $solve:ident) => {{
1366 let c = cusolver()?;
1367 check(unsafe {
1368 (c.$bufsize()?)(
1369 handle.as_raw(),
1370 jobz,
1371 uplo,
1372 n,
1373 a.as_raw().0 as *const $raw,
1374 lda,
1375 w.as_raw().0 as *const $real,
1376 &mut lwork,
1377 params.raw,
1378 )
1379 })?;
1380 let workspace =
1381 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1382 check(unsafe {
1383 (c.$solve()?)(
1384 handle.as_raw(),
1385 jobz,
1386 uplo,
1387 n,
1388 a.as_raw().0 as *mut $raw,
1389 lda,
1390 w.as_raw().0 as *mut $real,
1391 workspace.as_raw().0 as *mut $raw,
1392 lwork,
1393 info.as_raw().0 as *mut c_int,
1394 params.raw,
1395 )
1396 })
1397 }};
1398 }
1399
1400 if mem::size_of::<T>() == mem::size_of::<f32>() && mem::size_of::<T::Real>() == 4 {
1401 dispatch_real!(f32, cusolver_dn_ssyevj_buffer_size, cusolver_dn_ssyevj)
1402 } else if mem::size_of::<T>() == mem::size_of::<f64>() && mem::size_of::<T::Real>() == 8 {
1403 dispatch_real!(f64, cusolver_dn_dsyevj_buffer_size, cusolver_dn_dsyevj)
1404 } else if mem::size_of::<T>() == mem::size_of::<Complex32>() {
1405 dispatch_complex!(
1406 Complex32,
1407 f32,
1408 cuComplex,
1409 cusolver_dn_cheevj_buffer_size,
1410 cusolver_dn_cheevj
1411 )
1412 } else {
1413 dispatch_complex!(
1414 Complex64,
1415 f64,
1416 cuDoubleComplex,
1417 cusolver_dn_zheevj_buffer_size,
1418 cusolver_dn_zheevj
1419 )
1420 }
1421}
1422
1423#[allow(clippy::too_many_arguments)]
1425pub fn gesvdj<T: SolverScalar>(
1426 handle: &DnHandle,
1427 jobz: EigMode,
1428 econ: bool,
1429 m: i32,
1430 n: i32,
1431 a: &mut DeviceBuffer<T>,
1432 lda: i32,
1433 s: &mut DeviceBuffer<T::Real>,
1434 u: &mut DeviceBuffer<T>,
1435 ldu: i32,
1436 v: &mut DeviceBuffer<T>,
1437 ldv: i32,
1438 info: &mut DeviceBuffer<i32>,
1439 params: &GesvdjInfo,
1440) -> Result<()> {
1441 use baracuda_cusolver_sys::{cuComplex, cuDoubleComplex};
1442 use core::mem;
1443
1444 let mut lwork: c_int = 0;
1445 let econ_i = if econ { 1 } else { 0 };
1446
1447 macro_rules! dispatch_real {
1448 ($t:ty, $bufsize:ident, $solve:ident) => {{
1449 let c = cusolver()?;
1450 check(unsafe {
1451 (c.$bufsize()?)(
1452 handle.as_raw(),
1453 jobz,
1454 econ_i,
1455 m,
1456 n,
1457 a.as_raw().0 as *const $t,
1458 lda,
1459 s.as_raw().0 as *const $t,
1460 u.as_raw().0 as *const $t,
1461 ldu,
1462 v.as_raw().0 as *const $t,
1463 ldv,
1464 &mut lwork,
1465 params.raw,
1466 )
1467 })?;
1468 let workspace =
1469 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1470 check(unsafe {
1471 (c.$solve()?)(
1472 handle.as_raw(),
1473 jobz,
1474 econ_i,
1475 m,
1476 n,
1477 a.as_raw().0 as *mut $t,
1478 lda,
1479 s.as_raw().0 as *mut $t,
1480 u.as_raw().0 as *mut $t,
1481 ldu,
1482 v.as_raw().0 as *mut $t,
1483 ldv,
1484 workspace.as_raw().0 as *mut $t,
1485 lwork,
1486 info.as_raw().0 as *mut c_int,
1487 params.raw,
1488 )
1489 })
1490 }};
1491 }
1492 macro_rules! dispatch_complex {
1493 ($t:ty, $real:ty, $raw:ty, $bufsize:ident, $solve:ident) => {{
1494 let c = cusolver()?;
1495 check(unsafe {
1496 (c.$bufsize()?)(
1497 handle.as_raw(),
1498 jobz,
1499 econ_i,
1500 m,
1501 n,
1502 a.as_raw().0 as *const $raw,
1503 lda,
1504 s.as_raw().0 as *const $real,
1505 u.as_raw().0 as *const $raw,
1506 ldu,
1507 v.as_raw().0 as *const $raw,
1508 ldv,
1509 &mut lwork,
1510 params.raw,
1511 )
1512 })?;
1513 let workspace =
1514 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1515 check(unsafe {
1516 (c.$solve()?)(
1517 handle.as_raw(),
1518 jobz,
1519 econ_i,
1520 m,
1521 n,
1522 a.as_raw().0 as *mut $raw,
1523 lda,
1524 s.as_raw().0 as *mut $real,
1525 u.as_raw().0 as *mut $raw,
1526 ldu,
1527 v.as_raw().0 as *mut $raw,
1528 ldv,
1529 workspace.as_raw().0 as *mut $raw,
1530 lwork,
1531 info.as_raw().0 as *mut c_int,
1532 params.raw,
1533 )
1534 })
1535 }};
1536 }
1537
1538 if mem::size_of::<T>() == mem::size_of::<f32>() && mem::size_of::<T::Real>() == 4 {
1539 dispatch_real!(f32, cusolver_dn_sgesvdj_buffer_size, cusolver_dn_sgesvdj)
1540 } else if mem::size_of::<T>() == mem::size_of::<f64>() && mem::size_of::<T::Real>() == 8 {
1541 dispatch_real!(f64, cusolver_dn_dgesvdj_buffer_size, cusolver_dn_dgesvdj)
1542 } else if mem::size_of::<T>() == mem::size_of::<Complex32>() {
1543 dispatch_complex!(
1544 Complex32,
1545 f32,
1546 cuComplex,
1547 cusolver_dn_cgesvdj_buffer_size,
1548 cusolver_dn_cgesvdj
1549 )
1550 } else {
1551 dispatch_complex!(
1552 Complex64,
1553 f64,
1554 cuDoubleComplex,
1555 cusolver_dn_zgesvdj_buffer_size,
1556 cusolver_dn_zgesvdj
1557 )
1558 }
1559}
1560
1561#[allow(clippy::too_many_arguments)]
1566pub fn orgqr<T: SolverScalar>(
1567 handle: &DnHandle,
1568 m: i32,
1569 n: i32,
1570 k: i32,
1571 a: &mut DeviceBuffer<T>,
1572 lda: i32,
1573 tau: &DeviceBuffer<T>,
1574 info: &mut DeviceBuffer<i32>,
1575) -> Result<()> {
1576 use baracuda_cusolver_sys::{cuComplex, cuDoubleComplex};
1577 use core::mem;
1578
1579 let mut lwork: c_int = 0;
1580 macro_rules! dispatch {
1581 ($t:ty, $raw:ty, $bufsize:ident, $solve:ident) => {{
1582 let c = cusolver()?;
1583 check(unsafe {
1584 (c.$bufsize()?)(
1585 handle.as_raw(),
1586 m,
1587 n,
1588 k,
1589 a.as_raw().0 as *const $raw,
1590 lda,
1591 tau.as_raw().0 as *const $raw,
1592 &mut lwork,
1593 )
1594 })?;
1595 let workspace =
1596 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1597 check(unsafe {
1598 (c.$solve()?)(
1599 handle.as_raw(),
1600 m,
1601 n,
1602 k,
1603 a.as_raw().0 as *mut $raw,
1604 lda,
1605 tau.as_raw().0 as *const $raw,
1606 workspace.as_raw().0 as *mut $raw,
1607 lwork,
1608 info.as_raw().0 as *mut c_int,
1609 )
1610 })
1611 }};
1612 }
1613
1614 if mem::size_of::<T>() == mem::size_of::<f32>() && mem::size_of::<T::Real>() == 4 {
1615 dispatch!(f32, f32, cusolver_dn_sorgqr_buffer_size, cusolver_dn_sorgqr)
1616 } else if mem::size_of::<T>() == mem::size_of::<f64>() && mem::size_of::<T::Real>() == 8 {
1617 dispatch!(f64, f64, cusolver_dn_dorgqr_buffer_size, cusolver_dn_dorgqr)
1618 } else if mem::size_of::<T>() == mem::size_of::<Complex32>() {
1619 dispatch!(
1620 Complex32,
1621 cuComplex,
1622 cusolver_dn_cungqr_buffer_size,
1623 cusolver_dn_cungqr
1624 )
1625 } else {
1626 dispatch!(
1627 Complex64,
1628 cuDoubleComplex,
1629 cusolver_dn_zungqr_buffer_size,
1630 cusolver_dn_zungqr
1631 )
1632 }
1633}
1634
1635#[derive(Copy, Clone, Debug, Eq, PartialEq)]
1637pub enum Side {
1638 Left,
1640 Right,
1642}
1643
1644impl Side {
1645 fn raw(self) -> core::ffi::c_int {
1646 match self {
1647 Side::Left => 0,
1648 Side::Right => 1,
1649 }
1650 }
1651}
1652
1653#[allow(clippy::too_many_arguments)]
1656pub fn ormqr<T: SolverScalar>(
1657 handle: &DnHandle,
1658 side: Side,
1659 trans: Op,
1660 m: i32,
1661 n: i32,
1662 k: i32,
1663 a: &DeviceBuffer<T>,
1664 lda: i32,
1665 tau: &DeviceBuffer<T>,
1666 c_mat: &mut DeviceBuffer<T>,
1667 ldc: i32,
1668 info: &mut DeviceBuffer<i32>,
1669) -> Result<()> {
1670 use baracuda_cusolver_sys::{cuComplex, cuDoubleComplex};
1671 use core::mem;
1672
1673 let mut lwork: c_int = 0;
1674 let side_i = side.raw();
1675 macro_rules! dispatch {
1676 ($t:ty, $raw:ty, $bufsize:ident, $solve:ident) => {{
1677 let ca = cusolver()?;
1678 check(unsafe {
1679 (ca.$bufsize()?)(
1680 handle.as_raw(),
1681 side_i,
1682 trans.raw(),
1683 m,
1684 n,
1685 k,
1686 a.as_raw().0 as *const $raw,
1687 lda,
1688 tau.as_raw().0 as *const $raw,
1689 c_mat.as_raw().0 as *const $raw,
1690 ldc,
1691 &mut lwork,
1692 )
1693 })?;
1694 let workspace =
1695 DeviceBuffer::<T>::new(c_mat.context(), lwork as usize).map_err(alloc_fail)?;
1696 check(unsafe {
1697 (ca.$solve()?)(
1698 handle.as_raw(),
1699 side_i,
1700 trans.raw(),
1701 m,
1702 n,
1703 k,
1704 a.as_raw().0 as *const $raw,
1705 lda,
1706 tau.as_raw().0 as *const $raw,
1707 c_mat.as_raw().0 as *mut $raw,
1708 ldc,
1709 workspace.as_raw().0 as *mut $raw,
1710 lwork,
1711 info.as_raw().0 as *mut c_int,
1712 )
1713 })
1714 }};
1715 }
1716
1717 if mem::size_of::<T>() == mem::size_of::<f32>() && mem::size_of::<T::Real>() == 4 {
1718 dispatch!(f32, f32, cusolver_dn_sormqr_buffer_size, cusolver_dn_sormqr)
1719 } else if mem::size_of::<T>() == mem::size_of::<f64>() && mem::size_of::<T::Real>() == 8 {
1720 dispatch!(f64, f64, cusolver_dn_dormqr_buffer_size, cusolver_dn_dormqr)
1721 } else if mem::size_of::<T>() == mem::size_of::<Complex32>() {
1722 dispatch!(
1723 Complex32,
1724 cuComplex,
1725 cusolver_dn_cunmqr_buffer_size,
1726 cusolver_dn_cunmqr
1727 )
1728 } else {
1729 dispatch!(
1730 Complex64,
1731 cuDoubleComplex,
1732 cusolver_dn_zunmqr_buffer_size,
1733 cusolver_dn_zunmqr
1734 )
1735 }
1736}
1737
1738#[allow(clippy::too_many_arguments)]
1745pub fn gels<T: SolverScalar>(
1746 handle: &DnHandle,
1747 m: i32,
1748 n: i32,
1749 nrhs: i32,
1750 a: &mut DeviceBuffer<T>,
1751 lda: i32,
1752 b: &mut DeviceBuffer<T>,
1753 ldb: i32,
1754 x: &mut DeviceBuffer<T>,
1755 ldx: i32,
1756 info: &mut DeviceBuffer<i32>,
1757) -> Result<i32> {
1758 use baracuda_cusolver_sys::{cuComplex, cuDoubleComplex};
1759 use core::mem;
1760
1761 let mut bytes: usize = 0;
1762
1763 macro_rules! dispatch {
1764 ($t:ty, $raw:ty, $bufsize:ident, $solve:ident) => {{
1765 let cs = cusolver()?;
1766 check(unsafe {
1767 (cs.$bufsize()?)(
1768 handle.as_raw(),
1769 m,
1770 n,
1771 nrhs,
1772 a.as_raw().0 as *mut $raw,
1773 lda,
1774 b.as_raw().0 as *mut $raw,
1775 ldb,
1776 x.as_raw().0 as *mut $raw,
1777 ldx,
1778 core::ptr::null_mut(),
1779 &mut bytes,
1780 )
1781 })?;
1782 let units = bytes.div_ceil(mem::size_of::<T>());
1784 let workspace =
1785 DeviceBuffer::<T>::new(a.context(), units).map_err(alloc_fail)?;
1786 let mut iter: c_int = 0;
1787 check(unsafe {
1788 (cs.$solve()?)(
1789 handle.as_raw(),
1790 m,
1791 n,
1792 nrhs,
1793 a.as_raw().0 as *mut $raw,
1794 lda,
1795 b.as_raw().0 as *mut $raw,
1796 ldb,
1797 x.as_raw().0 as *mut $raw,
1798 ldx,
1799 workspace.as_raw().0 as *mut c_void,
1800 bytes,
1801 &mut iter,
1802 info.as_raw().0 as *mut c_int,
1803 )
1804 })?;
1805 Ok(iter)
1806 }};
1807 }
1808
1809 if mem::size_of::<T>() == mem::size_of::<f32>() && mem::size_of::<T::Real>() == 4 {
1810 dispatch!(f32, f32, cusolver_dn_ssgels_buffer_size, cusolver_dn_ssgels)
1811 } else if mem::size_of::<T>() == mem::size_of::<f64>() && mem::size_of::<T::Real>() == 8 {
1812 dispatch!(f64, f64, cusolver_dn_ddgels_buffer_size, cusolver_dn_ddgels)
1813 } else if mem::size_of::<T>() == mem::size_of::<Complex32>() {
1814 dispatch!(
1815 Complex32,
1816 cuComplex,
1817 cusolver_dn_ccgels_buffer_size,
1818 cusolver_dn_ccgels
1819 )
1820 } else {
1821 dispatch!(
1822 Complex64,
1823 cuDoubleComplex,
1824 cusolver_dn_zzgels_buffer_size,
1825 cusolver_dn_zzgels
1826 )
1827 }
1828}
1829
1830pub fn potri<T: SolverScalar>(
1836 handle: &DnHandle,
1837 uplo: Fill,
1838 n: i32,
1839 a: &mut DeviceBuffer<T>,
1840 lda: i32,
1841 info: &mut DeviceBuffer<i32>,
1842) -> Result<()> {
1843 use baracuda_cusolver_sys::{cuComplex, cuDoubleComplex};
1844 use core::mem;
1845
1846 let mut lwork: c_int = 0;
1847 macro_rules! dispatch {
1848 ($t:ty, $raw:ty, $bufsize:ident, $solve:ident) => {{
1849 let cs = cusolver()?;
1850 check(unsafe {
1851 (cs.$bufsize()?)(
1852 handle.as_raw(),
1853 uplo,
1854 n,
1855 a.as_raw().0 as *mut $raw,
1856 lda,
1857 &mut lwork,
1858 )
1859 })?;
1860 let workspace =
1861 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1862 check(unsafe {
1863 (cs.$solve()?)(
1864 handle.as_raw(),
1865 uplo,
1866 n,
1867 a.as_raw().0 as *mut $raw,
1868 lda,
1869 workspace.as_raw().0 as *mut $raw,
1870 lwork,
1871 info.as_raw().0 as *mut c_int,
1872 )
1873 })
1874 }};
1875 }
1876
1877 if mem::size_of::<T>() == mem::size_of::<f32>() && mem::size_of::<T::Real>() == 4 {
1878 dispatch!(f32, f32, cusolver_dn_spotri_buffer_size, cusolver_dn_spotri)
1879 } else if mem::size_of::<T>() == mem::size_of::<f64>() && mem::size_of::<T::Real>() == 8 {
1880 dispatch!(f64, f64, cusolver_dn_dpotri_buffer_size, cusolver_dn_dpotri)
1881 } else if mem::size_of::<T>() == mem::size_of::<Complex32>() {
1882 dispatch!(
1883 Complex32,
1884 cuComplex,
1885 cusolver_dn_cpotri_buffer_size,
1886 cusolver_dn_cpotri
1887 )
1888 } else {
1889 dispatch!(
1890 Complex64,
1891 cuDoubleComplex,
1892 cusolver_dn_zpotri_buffer_size,
1893 cusolver_dn_zpotri
1894 )
1895 }
1896}
1897
1898#[allow(clippy::too_many_arguments)]
1904pub fn syevj_batched<T: SolverScalar>(
1905 handle: &DnHandle,
1906 jobz: EigMode,
1907 uplo: Fill,
1908 n: i32,
1909 a: &mut DeviceBuffer<T>,
1910 lda: i32,
1911 w: &mut DeviceBuffer<T::Real>,
1912 info: &mut DeviceBuffer<i32>,
1913 params: &SyevjInfo,
1914 batch_size: i32,
1915) -> Result<()> {
1916 use baracuda_cusolver_sys::{cuComplex, cuDoubleComplex};
1917 use core::mem;
1918
1919 let mut lwork: c_int = 0;
1920 macro_rules! dispatch_real {
1921 ($t:ty, $bufsize:ident, $solve:ident) => {{
1922 let c = cusolver()?;
1923 check(unsafe {
1924 (c.$bufsize()?)(
1925 handle.as_raw(),
1926 jobz,
1927 uplo,
1928 n,
1929 a.as_raw().0 as *const $t,
1930 lda,
1931 w.as_raw().0 as *const $t,
1932 &mut lwork,
1933 params.as_raw(),
1934 batch_size,
1935 )
1936 })?;
1937 let workspace =
1938 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1939 check(unsafe {
1940 (c.$solve()?)(
1941 handle.as_raw(),
1942 jobz,
1943 uplo,
1944 n,
1945 a.as_raw().0 as *mut $t,
1946 lda,
1947 w.as_raw().0 as *mut $t,
1948 workspace.as_raw().0 as *mut $t,
1949 lwork,
1950 info.as_raw().0 as *mut c_int,
1951 params.as_raw(),
1952 batch_size,
1953 )
1954 })
1955 }};
1956 }
1957 macro_rules! dispatch_complex {
1958 ($t:ty, $real:ty, $raw:ty, $bufsize:ident, $solve:ident) => {{
1959 let c = cusolver()?;
1960 check(unsafe {
1961 (c.$bufsize()?)(
1962 handle.as_raw(),
1963 jobz,
1964 uplo,
1965 n,
1966 a.as_raw().0 as *const $raw,
1967 lda,
1968 w.as_raw().0 as *const $real,
1969 &mut lwork,
1970 params.as_raw(),
1971 batch_size,
1972 )
1973 })?;
1974 let workspace =
1975 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1976 check(unsafe {
1977 (c.$solve()?)(
1978 handle.as_raw(),
1979 jobz,
1980 uplo,
1981 n,
1982 a.as_raw().0 as *mut $raw,
1983 lda,
1984 w.as_raw().0 as *mut $real,
1985 workspace.as_raw().0 as *mut $raw,
1986 lwork,
1987 info.as_raw().0 as *mut c_int,
1988 params.as_raw(),
1989 batch_size,
1990 )
1991 })
1992 }};
1993 }
1994
1995 if mem::size_of::<T>() == mem::size_of::<f32>() && mem::size_of::<T::Real>() == 4 {
1996 dispatch_real!(
1997 f32,
1998 cusolver_dn_ssyevj_batched_buffer_size,
1999 cusolver_dn_ssyevj_batched
2000 )
2001 } else if mem::size_of::<T>() == mem::size_of::<f64>() && mem::size_of::<T::Real>() == 8 {
2002 dispatch_real!(
2003 f64,
2004 cusolver_dn_dsyevj_batched_buffer_size,
2005 cusolver_dn_dsyevj_batched
2006 )
2007 } else if mem::size_of::<T>() == mem::size_of::<Complex32>() {
2008 dispatch_complex!(
2009 Complex32,
2010 f32,
2011 cuComplex,
2012 cusolver_dn_cheevj_batched_buffer_size,
2013 cusolver_dn_cheevj_batched
2014 )
2015 } else {
2016 dispatch_complex!(
2017 Complex64,
2018 f64,
2019 cuDoubleComplex,
2020 cusolver_dn_zheevj_batched_buffer_size,
2021 cusolver_dn_zheevj_batched
2022 )
2023 }
2024}
2025
2026#[allow(clippy::too_many_arguments)]
2028pub fn gesvdj_batched<T: SolverScalar>(
2029 handle: &DnHandle,
2030 jobz: EigMode,
2031 m: i32,
2032 n: i32,
2033 a: &mut DeviceBuffer<T>,
2034 lda: i32,
2035 s: &mut DeviceBuffer<T::Real>,
2036 u: &mut DeviceBuffer<T>,
2037 ldu: i32,
2038 v: &mut DeviceBuffer<T>,
2039 ldv: i32,
2040 info: &mut DeviceBuffer<i32>,
2041 params: &GesvdjInfo,
2042 batch_size: i32,
2043) -> Result<()> {
2044 use baracuda_cusolver_sys::{cuComplex, cuDoubleComplex};
2045 use core::mem;
2046
2047 let mut lwork: c_int = 0;
2048 macro_rules! dispatch_real {
2049 ($t:ty, $bufsize:ident, $solve:ident) => {{
2050 let c = cusolver()?;
2051 check(unsafe {
2052 (c.$bufsize()?)(
2053 handle.as_raw(),
2054 jobz,
2055 m,
2056 n,
2057 a.as_raw().0 as *const $t,
2058 lda,
2059 s.as_raw().0 as *const $t,
2060 u.as_raw().0 as *const $t,
2061 ldu,
2062 v.as_raw().0 as *const $t,
2063 ldv,
2064 &mut lwork,
2065 params.as_raw(),
2066 batch_size,
2067 )
2068 })?;
2069 let workspace =
2070 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
2071 check(unsafe {
2072 (c.$solve()?)(
2073 handle.as_raw(),
2074 jobz,
2075 m,
2076 n,
2077 a.as_raw().0 as *mut $t,
2078 lda,
2079 s.as_raw().0 as *mut $t,
2080 u.as_raw().0 as *mut $t,
2081 ldu,
2082 v.as_raw().0 as *mut $t,
2083 ldv,
2084 workspace.as_raw().0 as *mut $t,
2085 lwork,
2086 info.as_raw().0 as *mut c_int,
2087 params.as_raw(),
2088 batch_size,
2089 )
2090 })
2091 }};
2092 }
2093 macro_rules! dispatch_complex {
2094 ($t:ty, $real:ty, $raw:ty, $bufsize:ident, $solve:ident) => {{
2095 let c = cusolver()?;
2096 check(unsafe {
2097 (c.$bufsize()?)(
2098 handle.as_raw(),
2099 jobz,
2100 m,
2101 n,
2102 a.as_raw().0 as *const $raw,
2103 lda,
2104 s.as_raw().0 as *const $real,
2105 u.as_raw().0 as *const $raw,
2106 ldu,
2107 v.as_raw().0 as *const $raw,
2108 ldv,
2109 &mut lwork,
2110 params.as_raw(),
2111 batch_size,
2112 )
2113 })?;
2114 let workspace =
2115 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
2116 check(unsafe {
2117 (c.$solve()?)(
2118 handle.as_raw(),
2119 jobz,
2120 m,
2121 n,
2122 a.as_raw().0 as *mut $raw,
2123 lda,
2124 s.as_raw().0 as *mut $real,
2125 u.as_raw().0 as *mut $raw,
2126 ldu,
2127 v.as_raw().0 as *mut $raw,
2128 ldv,
2129 workspace.as_raw().0 as *mut $raw,
2130 lwork,
2131 info.as_raw().0 as *mut c_int,
2132 params.as_raw(),
2133 batch_size,
2134 )
2135 })
2136 }};
2137 }
2138
2139 if mem::size_of::<T>() == mem::size_of::<f32>() && mem::size_of::<T::Real>() == 4 {
2140 dispatch_real!(
2141 f32,
2142 cusolver_dn_sgesvdj_batched_buffer_size,
2143 cusolver_dn_sgesvdj_batched
2144 )
2145 } else if mem::size_of::<T>() == mem::size_of::<f64>() && mem::size_of::<T::Real>() == 8 {
2146 dispatch_real!(
2147 f64,
2148 cusolver_dn_dgesvdj_batched_buffer_size,
2149 cusolver_dn_dgesvdj_batched
2150 )
2151 } else if mem::size_of::<T>() == mem::size_of::<Complex32>() {
2152 dispatch_complex!(
2153 Complex32,
2154 f32,
2155 cuComplex,
2156 cusolver_dn_cgesvdj_batched_buffer_size,
2157 cusolver_dn_cgesvdj_batched
2158 )
2159 } else {
2160 dispatch_complex!(
2161 Complex64,
2162 f64,
2163 cuDoubleComplex,
2164 cusolver_dn_zgesvdj_batched_buffer_size,
2165 cusolver_dn_zgesvdj_batched
2166 )
2167 }
2168}
2169
2170pub mod mg {
2173 use core::ffi::{c_int, c_void};
2178
2179 use baracuda_cusolver_sys::{
2180 cudaDataType, cudaLibMgGrid_t, cudaLibMgMatrixDesc_t, cusolver_mg, cusolverMgHandle_t,
2181 };
2182
2183 use super::{alloc_fail, check, EigMode, Fill, Result};
2184
2185 #[derive(Debug)]
2187 pub struct Handle {
2188 raw: cusolverMgHandle_t,
2189 }
2190
2191 impl Handle {
2192 pub fn new() -> Result<Self> {
2195 let mg = cusolver_mg()?;
2196 let cu = mg.cusolver_mg_create()?;
2197 let mut h: cusolverMgHandle_t = core::ptr::null_mut();
2198 check(unsafe { cu(&mut h) })?;
2199 Ok(Self { raw: h })
2200 }
2201
2202 pub fn device_select(&self, devices: &[i32]) -> Result<()> {
2205 let mg = cusolver_mg()?;
2206 let cu = mg.cusolver_mg_device_select()?;
2207 check(unsafe { cu(self.raw, devices.len() as c_int, devices.as_ptr()) })
2208 }
2209
2210 pub fn as_raw(&self) -> cusolverMgHandle_t {
2212 self.raw
2213 }
2214 }
2215
2216 impl Drop for Handle {
2217 fn drop(&mut self) {
2218 if let Ok(mg) = cusolver_mg() {
2219 if let Ok(cu) = mg.cusolver_mg_destroy() {
2220 let _ = unsafe { cu(self.raw) };
2221 }
2222 }
2223 }
2224 }
2225
2226 #[derive(Debug)]
2228 pub struct DeviceGrid {
2229 raw: cudaLibMgGrid_t,
2230 }
2231
2232 impl DeviceGrid {
2233 pub fn new(num_row_devices: i32, num_col_devices: i32, devices: &[i32], mapping: i32) -> Result<Self> {
2237 let mg = cusolver_mg()?;
2238 let cu = mg.cusolver_mg_create_device_grid()?;
2239 let mut raw: cudaLibMgGrid_t = core::ptr::null_mut();
2240 check(unsafe {
2241 cu(
2242 &mut raw,
2243 num_row_devices,
2244 num_col_devices,
2245 devices.as_ptr(),
2246 mapping,
2247 )
2248 })?;
2249 Ok(Self { raw })
2250 }
2251
2252 pub fn as_raw(&self) -> cudaLibMgGrid_t {
2254 self.raw
2255 }
2256 }
2257
2258 impl Drop for DeviceGrid {
2259 fn drop(&mut self) {
2260 if let Ok(mg) = cusolver_mg() {
2261 if let Ok(cu) = mg.cusolver_mg_destroy_grid() {
2262 let _ = unsafe { cu(self.raw) };
2263 }
2264 }
2265 }
2266 }
2267
2268 #[derive(Debug)]
2270 pub struct MatrixDesc {
2271 raw: cudaLibMgMatrixDesc_t,
2272 }
2273
2274 impl MatrixDesc {
2275 pub fn new(
2278 num_rows: i64,
2279 num_cols: i64,
2280 row_block_size: i64,
2281 col_block_size: i64,
2282 data_type: cudaDataType,
2283 grid: &DeviceGrid,
2284 ) -> Result<Self> {
2285 let mg = cusolver_mg()?;
2286 let cu = mg.cusolver_mg_create_matrix_desc()?;
2287 let mut raw: cudaLibMgMatrixDesc_t = core::ptr::null_mut();
2288 check(unsafe {
2289 cu(
2290 &mut raw,
2291 num_rows,
2292 num_cols,
2293 row_block_size,
2294 col_block_size,
2295 data_type,
2296 grid.as_raw(),
2297 )
2298 })?;
2299 Ok(Self { raw })
2300 }
2301
2302 pub fn as_raw(&self) -> cudaLibMgMatrixDesc_t {
2304 self.raw
2305 }
2306 }
2307
2308 impl Drop for MatrixDesc {
2309 fn drop(&mut self) {
2310 if let Ok(mg) = cusolver_mg() {
2311 if let Ok(cu) = mg.cusolver_mg_destroy_matrix_desc() {
2312 let _ = unsafe { cu(self.raw) };
2313 }
2314 }
2315 }
2316 }
2317
2318 #[allow(clippy::too_many_arguments)]
2324 pub unsafe fn getrf_buffer_size(
2325 handle: &Handle,
2326 m: i32,
2327 n: i32,
2328 array_d_a: *mut *mut c_void,
2329 ia: i32,
2330 ja: i32,
2331 desc_a: &MatrixDesc,
2332 array_d_ipiv: *mut *mut c_int,
2333 compute_type: cudaDataType,
2334 ) -> Result<i64> { unsafe {
2335 let mg = cusolver_mg()?;
2336 let cu = mg.cusolver_mg_getrf_buffer_size()?;
2337 let mut lwork: i64 = 0;
2338 check(cu(
2339 handle.as_raw(),
2340 m,
2341 n,
2342 array_d_a,
2343 ia,
2344 ja,
2345 desc_a.as_raw(),
2346 array_d_ipiv,
2347 compute_type,
2348 &mut lwork,
2349 ))?;
2350 Ok(lwork)
2351 }}
2352
2353 #[allow(clippy::too_many_arguments)]
2356 pub unsafe fn getrf(
2357 handle: &Handle,
2358 m: i32,
2359 n: i32,
2360 array_d_a: *mut *mut c_void,
2361 ia: i32,
2362 ja: i32,
2363 desc_a: &MatrixDesc,
2364 array_d_ipiv: *mut *mut c_int,
2365 compute_type: cudaDataType,
2366 array_d_work: *mut *mut c_void,
2367 lwork: i64,
2368 info: &mut [c_int],
2369 ) -> Result<()> { unsafe {
2370 let mg = cusolver_mg()?;
2371 let cu = mg.cusolver_mg_getrf()?;
2372 let _ = alloc_fail::<()>; check(cu(
2374 handle.as_raw(),
2375 m,
2376 n,
2377 array_d_a,
2378 ia,
2379 ja,
2380 desc_a.as_raw(),
2381 array_d_ipiv,
2382 compute_type,
2383 array_d_work,
2384 lwork,
2385 info.as_mut_ptr(),
2386 ))
2387 }}
2388
2389 #[allow(clippy::too_many_arguments)]
2394 pub unsafe fn potrf_buffer_size(
2395 handle: &Handle,
2396 uplo: Fill,
2397 n: i32,
2398 array_d_a: *mut *mut c_void,
2399 ia: i32,
2400 ja: i32,
2401 desc_a: &MatrixDesc,
2402 compute_type: cudaDataType,
2403 ) -> Result<i64> { unsafe {
2404 let mg = cusolver_mg()?;
2405 let cu = mg.cusolver_mg_potrf_buffer_size()?;
2406 let mut lwork: i64 = 0;
2407 check(cu(
2408 handle.as_raw(),
2409 uplo,
2410 n,
2411 array_d_a,
2412 ia,
2413 ja,
2414 desc_a.as_raw(),
2415 compute_type,
2416 &mut lwork,
2417 ))?;
2418 Ok(lwork)
2419 }}
2420
2421 #[allow(clippy::too_many_arguments)]
2424 pub unsafe fn potrf(
2425 handle: &Handle,
2426 uplo: Fill,
2427 n: i32,
2428 array_d_a: *mut *mut c_void,
2429 ia: i32,
2430 ja: i32,
2431 desc_a: &MatrixDesc,
2432 compute_type: cudaDataType,
2433 array_d_work: *mut *mut c_void,
2434 lwork: i64,
2435 info: &mut [c_int],
2436 ) -> Result<()> { unsafe {
2437 let mg = cusolver_mg()?;
2438 let cu = mg.cusolver_mg_potrf()?;
2439 check(cu(
2440 handle.as_raw(),
2441 uplo,
2442 n,
2443 array_d_a,
2444 ia,
2445 ja,
2446 desc_a.as_raw(),
2447 compute_type,
2448 array_d_work,
2449 lwork,
2450 info.as_mut_ptr(),
2451 ))
2452 }}
2453
2454 #[allow(clippy::too_many_arguments)]
2459 pub unsafe fn syevd_buffer_size(
2460 handle: &Handle,
2461 jobz: EigMode,
2462 uplo: Fill,
2463 n: i32,
2464 array_d_a: *mut *mut c_void,
2465 ia: i32,
2466 ja: i32,
2467 desc_a: &MatrixDesc,
2468 w: *mut c_void,
2469 data_type_w: cudaDataType,
2470 compute_type: cudaDataType,
2471 ) -> Result<i64> { unsafe {
2472 let mg = cusolver_mg()?;
2473 let cu = mg.cusolver_mg_syevd_buffer_size()?;
2474 let mut lwork: i64 = 0;
2475 check(cu(
2476 handle.as_raw(),
2477 jobz,
2478 uplo,
2479 n,
2480 array_d_a,
2481 ia,
2482 ja,
2483 desc_a.as_raw(),
2484 w,
2485 data_type_w,
2486 compute_type,
2487 &mut lwork,
2488 ))?;
2489 Ok(lwork)
2490 }}
2491
2492 #[allow(clippy::too_many_arguments)]
2495 pub unsafe fn syevd(
2496 handle: &Handle,
2497 jobz: EigMode,
2498 uplo: Fill,
2499 n: i32,
2500 array_d_a: *mut *mut c_void,
2501 ia: i32,
2502 ja: i32,
2503 desc_a: &MatrixDesc,
2504 w: *mut c_void,
2505 data_type_w: cudaDataType,
2506 compute_type: cudaDataType,
2507 array_d_work: *mut *mut c_void,
2508 lwork: i64,
2509 info: &mut [c_int],
2510 ) -> Result<()> { unsafe {
2511 let mg = cusolver_mg()?;
2512 let cu = mg.cusolver_mg_syevd()?;
2513 check(cu(
2514 handle.as_raw(),
2515 jobz,
2516 uplo,
2517 n,
2518 array_d_a,
2519 ia,
2520 ja,
2521 desc_a.as_raw(),
2522 w,
2523 data_type_w,
2524 compute_type,
2525 array_d_work,
2526 lwork,
2527 info.as_mut_ptr(),
2528 ))
2529 }}
2530}
2531
2532pub fn sgetrf(
2536 handle: &DnHandle,
2537 m: i32,
2538 n: i32,
2539 a: &mut DeviceBuffer<f32>,
2540 lda: i32,
2541 ipiv: &mut DeviceBuffer<i32>,
2542 info: &mut DeviceBuffer<i32>,
2543) -> Result<()> {
2544 getrf::<f32>(handle, m, n, a, lda, ipiv, info)
2545}
2546
2547#[allow(clippy::too_many_arguments)]
2549pub fn sgetrs(
2550 handle: &DnHandle,
2551 trans: Op,
2552 n: i32,
2553 nrhs: i32,
2554 a: &DeviceBuffer<f32>,
2555 lda: i32,
2556 ipiv: &DeviceBuffer<i32>,
2557 b: &mut DeviceBuffer<f32>,
2558 ldb: i32,
2559 info: &mut DeviceBuffer<i32>,
2560) -> Result<()> {
2561 getrs::<f32>(handle, trans, n, nrhs, a, lda, ipiv, b, ldb, info)
2562}
2563
2564pub mod xapi {
2567 use super::*;
2573 use baracuda_cusolver_sys::{cudaDataType, cusolverDnParams_t};
2574
2575 #[derive(Debug)]
2578 pub struct Params {
2579 raw: cusolverDnParams_t,
2580 }
2581
2582 impl Params {
2583 pub fn new() -> Result<Self> {
2586 let c = cusolver()?;
2587 let cu = c.cusolver_dn_create_params()?;
2588 let mut p: cusolverDnParams_t = core::ptr::null_mut();
2589 check(unsafe { cu(&mut p) })?;
2590 Ok(Self { raw: p })
2591 }
2592
2593 pub fn as_raw(&self) -> cusolverDnParams_t {
2595 self.raw
2596 }
2597 }
2598
2599 impl Drop for Params {
2600 fn drop(&mut self) {
2601 if let Ok(c) = cusolver() {
2602 if let Ok(cu) = c.cusolver_dn_destroy_params() {
2603 let _ = unsafe { cu(self.raw) };
2604 }
2605 }
2606 }
2607 }
2608
2609 #[allow(clippy::too_many_arguments)]
2621 pub unsafe fn xgetrf_buffer_size(
2622 handle: &DnHandle,
2623 params: &Params,
2624 m: i64,
2625 n: i64,
2626 data_type_a: cudaDataType,
2627 a: *const c_void,
2628 lda: i64,
2629 compute_type: cudaDataType,
2630 ) -> Result<(usize, usize)> {
2631 let c = cusolver()?;
2632 let cu = c.cusolver_dn_xgetrf_buffer_size()?;
2633 let (mut dev, mut host) = (0usize, 0usize);
2634 check(unsafe {
2635 cu(
2636 handle.as_raw(),
2637 params.raw,
2638 m,
2639 n,
2640 data_type_a,
2641 a,
2642 lda,
2643 compute_type,
2644 &mut dev,
2645 &mut host,
2646 )
2647 })?;
2648 Ok((dev, host))
2649 }
2650
2651 #[allow(clippy::too_many_arguments)]
2660 pub unsafe fn xgetrf(
2661 handle: &DnHandle,
2662 params: &Params,
2663 m: i64,
2664 n: i64,
2665 data_type_a: cudaDataType,
2666 a: *mut c_void,
2667 lda: i64,
2668 ipiv: *mut i64,
2669 compute_type: cudaDataType,
2670 device_buf: *mut c_void,
2671 device_bytes: usize,
2672 host_buf: *mut c_void,
2673 host_bytes: usize,
2674 info: *mut c_int,
2675 ) -> Result<()> { unsafe {
2676 let c = cusolver()?;
2677 let cu = c.cusolver_dn_xgetrf()?;
2678 check(cu(
2679 handle.as_raw(),
2680 params.raw,
2681 m,
2682 n,
2683 data_type_a,
2684 a,
2685 lda,
2686 ipiv,
2687 compute_type,
2688 device_buf,
2689 device_bytes,
2690 host_buf,
2691 host_bytes,
2692 info,
2693 ))
2694 }}
2695
2696 #[allow(clippy::too_many_arguments)]
2705 pub unsafe fn xgetrs(
2706 handle: &DnHandle,
2707 params: &Params,
2708 trans: Op,
2709 n: i64,
2710 nrhs: i64,
2711 data_type_a: cudaDataType,
2712 a: *const c_void,
2713 lda: i64,
2714 ipiv: *const i64,
2715 data_type_b: cudaDataType,
2716 b: *mut c_void,
2717 ldb: i64,
2718 info: *mut c_int,
2719 ) -> Result<()> { unsafe {
2720 let c = cusolver()?;
2721 let cu = c.cusolver_dn_xgetrs()?;
2722 check(cu(
2723 handle.as_raw(),
2724 params.raw,
2725 trans.raw(),
2726 n,
2727 nrhs,
2728 data_type_a,
2729 a,
2730 lda,
2731 ipiv,
2732 data_type_b,
2733 b,
2734 ldb,
2735 info,
2736 ))
2737 }}
2738
2739 #[allow(clippy::too_many_arguments)]
2747 pub unsafe fn xpotrs(
2748 handle: &DnHandle,
2749 params: &Params,
2750 uplo: Fill,
2751 n: i64,
2752 nrhs: i64,
2753 data_type_a: cudaDataType,
2754 a: *const c_void,
2755 lda: i64,
2756 data_type_b: cudaDataType,
2757 b: *mut c_void,
2758 ldb: i64,
2759 info: *mut c_int,
2760 ) -> Result<()> { unsafe {
2761 let c = cusolver()?;
2762 let cu = c.cusolver_dn_xpotrs()?;
2763 check(cu(
2764 handle.as_raw(),
2765 params.raw,
2766 uplo,
2767 n,
2768 nrhs,
2769 data_type_a,
2770 a,
2771 lda,
2772 data_type_b,
2773 b,
2774 ldb,
2775 info,
2776 ))
2777 }}
2778}
2779
2780pub mod sparse {
2783 use super::*;
2786 use baracuda_cusolver_sys::cusolverSpHandle_t;
2787 use core::ffi::c_int;
2788
2789 #[derive(Debug)]
2791 pub struct SpHandle {
2792 raw: cusolverSpHandle_t,
2793 _not_send: PhantomData<*mut ()>,
2794 }
2795
2796 impl SpHandle {
2797 pub fn new() -> Result<Self> {
2800 let c = cusolver()?;
2801 let cu = c.cusolver_sp_create()?;
2802 let mut h: cusolverSpHandle_t = core::ptr::null_mut();
2803 check(unsafe { cu(&mut h) })?;
2804 Ok(Self {
2805 raw: h,
2806 _not_send: PhantomData,
2807 })
2808 }
2809
2810 pub fn set_stream(&self, stream: &Stream) -> Result<()> {
2813 let c = cusolver()?;
2814 let cu = c.cusolver_sp_set_stream()?;
2815 check(unsafe { cu(self.raw, stream.as_raw() as _) })
2816 }
2817
2818 pub fn as_raw(&self) -> cusolverSpHandle_t {
2820 self.raw
2821 }
2822 }
2823
2824 impl Drop for SpHandle {
2825 fn drop(&mut self) {
2826 if let Ok(c) = cusolver() {
2827 if let Ok(cu) = c.cusolver_sp_destroy() {
2828 let _ = unsafe { cu(self.raw) };
2829 }
2830 }
2831 }
2832 }
2833
2834 #[allow(clippy::too_many_arguments)]
2841 pub unsafe fn scsrlsvchol(
2842 handle: &SpHandle,
2843 m: i32,
2844 nnz: i32,
2845 descr_a: *mut c_void,
2846 csr_val: *const f32,
2847 csr_row_ptr: *const c_int,
2848 csr_col_ind: *const c_int,
2849 b: *const f32,
2850 tol: f32,
2851 reorder: i32,
2852 x: *mut f32,
2853 singularity: *mut c_int,
2854 ) -> Result<()> { unsafe {
2855 let c = cusolver()?;
2856 let cu = c.cusolver_sp_scsrlsvchol()?;
2857 check(cu(
2858 handle.raw,
2859 m,
2860 nnz,
2861 descr_a,
2862 csr_val,
2863 csr_row_ptr,
2864 csr_col_ind,
2865 b,
2866 tol,
2867 reorder,
2868 x,
2869 singularity,
2870 ))
2871 }}
2872
2873 #[allow(clippy::too_many_arguments)]
2878 pub unsafe fn scsrlsvqr(
2879 handle: &SpHandle,
2880 m: i32,
2881 nnz: i32,
2882 descr_a: *mut c_void,
2883 csr_val: *const f32,
2884 csr_row_ptr: *const c_int,
2885 csr_col_ind: *const c_int,
2886 b: *const f32,
2887 tol: f32,
2888 reorder: i32,
2889 x: *mut f32,
2890 singularity: *mut c_int,
2891 ) -> Result<()> { unsafe {
2892 let c = cusolver()?;
2893 let cu = c.cusolver_sp_scsrlsvqr()?;
2894 check(cu(
2895 handle.raw,
2896 m,
2897 nnz,
2898 descr_a,
2899 csr_val,
2900 csr_row_ptr,
2901 csr_col_ind,
2902 b,
2903 tol,
2904 reorder,
2905 x,
2906 singularity,
2907 ))
2908 }}
2909}
2910
2911pub mod refactor {
2914 use super::*;
2918 use baracuda_cusolver_sys::cusolverRfHandle_t;
2919
2920 #[derive(Debug)]
2922 pub struct RfHandle {
2923 raw: cusolverRfHandle_t,
2924 _not_send: PhantomData<*mut ()>,
2925 }
2926
2927 impl RfHandle {
2928 pub fn new() -> Result<Self> {
2931 let c = cusolver()?;
2932 let cu = c.cusolver_rf_create()?;
2933 let mut h: cusolverRfHandle_t = core::ptr::null_mut();
2934 check(unsafe { cu(&mut h) })?;
2935 Ok(Self {
2936 raw: h,
2937 _not_send: PhantomData,
2938 })
2939 }
2940
2941 pub fn as_raw(&self) -> cusolverRfHandle_t {
2943 self.raw
2944 }
2945
2946 pub fn analyze(&self) -> Result<()> {
2949 let c = cusolver()?;
2950 let cu = c.cusolver_rf_analyze()?;
2951 check(unsafe { cu(self.raw) })
2952 }
2953
2954 pub fn refactor(&self) -> Result<()> {
2957 let c = cusolver()?;
2958 let cu = c.cusolver_rf_refactor()?;
2959 check(unsafe { cu(self.raw) })
2960 }
2961 }
2962
2963 impl Drop for RfHandle {
2964 fn drop(&mut self) {
2965 if let Ok(c) = cusolver() {
2966 if let Ok(cu) = c.cusolver_rf_destroy() {
2967 let _ = unsafe { cu(self.raw) };
2968 }
2969 }
2970 }
2971 }
2972}