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