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> {
67 let c = cusolver()?;
68 let cu = c.cusolver_dn_create()?;
69 let mut h: cusolverDnHandle_t = core::ptr::null_mut();
70 check(unsafe { cu(&mut h) })?;
71 Ok(Self { handle: h })
72 }
73
74 pub fn set_stream(&self, stream: &Stream) -> Result<()> {
75 let c = cusolver()?;
76 let cu = c.cusolver_dn_set_stream()?;
77 check(unsafe { cu(self.handle, stream.as_raw() as _) })
78 }
79
80 pub fn stream(&self) -> Result<*mut c_void> {
86 let c = cusolver()?;
87 let cu = c.cusolver_dn_get_stream()?;
88 let mut s: *mut c_void = core::ptr::null_mut();
89 check(unsafe { cu(self.handle, &mut s as *mut *mut c_void as *mut _) })?;
90 Ok(s)
91 }
92
93 pub fn version() -> Result<i32> {
94 let c = cusolver()?;
95 let cu = c.cusolver_get_version()?;
96 let mut v: c_int = 0;
97 check(unsafe { cu(&mut v) })?;
98 Ok(v)
99 }
100
101 #[inline]
102 pub fn as_raw(&self) -> cusolverDnHandle_t {
103 self.handle
104 }
105}
106
107impl Drop for DnHandle {
108 fn drop(&mut self) {
109 if let Ok(c) = cusolver() {
110 if let Ok(cu) = c.cusolver_dn_destroy() {
111 let _ = unsafe { cu(self.handle) };
112 }
113 }
114 }
115}
116
117#[derive(Copy, Clone, Debug, Eq, PartialEq, Default)]
119pub enum Op {
120 #[default]
121 N,
122 T,
123 C,
124}
125
126impl Op {
127 fn raw(self) -> cublasOperation_t {
128 match self {
129 Op::N => cublasOperation_t::N,
130 Op::T => cublasOperation_t::T,
131 Op::C => cublasOperation_t::C,
132 }
133 }
134}
135
136pub trait SolverScalar: DeviceRepr + Copy + 'static + sealed::Sealed {
140 type Real: DeviceRepr + Copy + 'static;
143
144 #[doc(hidden)]
146 unsafe fn getrf_buf(
147 h: cusolverDnHandle_t,
148 m: c_int,
149 n: c_int,
150 a: *mut Self,
151 lda: c_int,
152 lwork: *mut c_int,
153 ) -> cusolverStatus_t;
154
155 #[doc(hidden)]
157 #[allow(clippy::too_many_arguments)]
158 unsafe fn getrf(
159 h: cusolverDnHandle_t,
160 m: c_int,
161 n: c_int,
162 a: *mut Self,
163 lda: c_int,
164 workspace: *mut Self,
165 ipiv: *mut c_int,
166 info: *mut c_int,
167 ) -> cusolverStatus_t;
168
169 #[doc(hidden)]
171 #[allow(clippy::too_many_arguments)]
172 unsafe fn getrs(
173 h: cusolverDnHandle_t,
174 trans: cublasOperation_t,
175 n: c_int,
176 nrhs: c_int,
177 a: *const Self,
178 lda: c_int,
179 ipiv: *const c_int,
180 b: *mut Self,
181 ldb: c_int,
182 info: *mut c_int,
183 ) -> cusolverStatus_t;
184
185 #[doc(hidden)]
187 unsafe fn geqrf_buf(
188 h: cusolverDnHandle_t,
189 m: c_int,
190 n: c_int,
191 a: *mut Self,
192 lda: c_int,
193 lwork: *mut c_int,
194 ) -> cusolverStatus_t;
195
196 #[doc(hidden)]
198 #[allow(clippy::too_many_arguments)]
199 unsafe fn geqrf(
200 h: cusolverDnHandle_t,
201 m: c_int,
202 n: c_int,
203 a: *mut Self,
204 lda: c_int,
205 tau: *mut Self,
206 workspace: *mut Self,
207 lwork: c_int,
208 info: *mut c_int,
209 ) -> cusolverStatus_t;
210
211 #[doc(hidden)]
213 unsafe fn potrf_buf(
214 h: cusolverDnHandle_t,
215 uplo: cublasFillMode_t,
216 n: c_int,
217 a: *mut Self,
218 lda: c_int,
219 lwork: *mut c_int,
220 ) -> cusolverStatus_t;
221
222 #[doc(hidden)]
224 #[allow(clippy::too_many_arguments)]
225 unsafe fn potrf(
226 h: cusolverDnHandle_t,
227 uplo: cublasFillMode_t,
228 n: c_int,
229 a: *mut Self,
230 lda: c_int,
231 workspace: *mut Self,
232 lwork: c_int,
233 info: *mut c_int,
234 ) -> cusolverStatus_t;
235
236 #[doc(hidden)]
238 #[allow(clippy::too_many_arguments)]
239 unsafe fn potrs(
240 h: cusolverDnHandle_t,
241 uplo: cublasFillMode_t,
242 n: c_int,
243 nrhs: c_int,
244 a: *const Self,
245 lda: c_int,
246 b: *mut Self,
247 ldb: c_int,
248 info: *mut c_int,
249 ) -> cusolverStatus_t;
250
251 #[doc(hidden)]
253 unsafe fn gesvd_buf(
254 h: cusolverDnHandle_t,
255 m: c_int,
256 n: c_int,
257 lwork: *mut c_int,
258 ) -> cusolverStatus_t;
259
260 #[doc(hidden)]
262 #[allow(clippy::too_many_arguments)]
263 unsafe fn gesvd(
264 h: cusolverDnHandle_t,
265 jobu: u8,
266 jobvt: u8,
267 m: c_int,
268 n: c_int,
269 a: *mut Self,
270 lda: c_int,
271 s: *mut Self::Real,
272 u: *mut Self,
273 ldu: c_int,
274 vt: *mut Self,
275 ldvt: c_int,
276 work: *mut Self,
277 lwork: c_int,
278 rwork: *mut Self::Real,
279 info: *mut c_int,
280 ) -> cusolverStatus_t;
281
282 #[doc(hidden)]
284 #[allow(clippy::too_many_arguments)]
285 unsafe fn syevd_buf(
286 h: cusolverDnHandle_t,
287 jobz: cusolverEigMode_t,
288 uplo: cublasFillMode_t,
289 n: c_int,
290 a: *const Self,
291 lda: c_int,
292 w: *const Self::Real,
293 lwork: *mut c_int,
294 ) -> cusolverStatus_t;
295
296 #[doc(hidden)]
298 #[allow(clippy::too_many_arguments)]
299 unsafe fn syevd(
300 h: cusolverDnHandle_t,
301 jobz: cusolverEigMode_t,
302 uplo: cublasFillMode_t,
303 n: c_int,
304 a: *mut Self,
305 lda: c_int,
306 w: *mut Self::Real,
307 work: *mut Self,
308 lwork: c_int,
309 info: *mut c_int,
310 ) -> cusolverStatus_t;
311}
312
313mod sealed {
314 use baracuda_types::{Complex32, Complex64};
315 pub trait Sealed {}
316 impl Sealed for f32 {}
317 impl Sealed for f64 {}
318 impl Sealed for Complex32 {}
319 impl Sealed for Complex64 {}
320}
321
322macro_rules! real_impl {
323 ($t:ty, $getrf_buf:ident, $getrf:ident, $getrs:ident,
324 $geqrf_buf:ident, $geqrf:ident,
325 $potrf_buf:ident, $potrf:ident, $potrs:ident,
326 $gesvd_buf:ident, $gesvd:ident,
327 $syevd_buf:ident, $syevd:ident) => {
328 impl SolverScalar for $t {
329 type Real = $t;
330
331 unsafe fn getrf_buf(
332 h: cusolverDnHandle_t,
333 m: c_int,
334 n: c_int,
335 a: *mut $t,
336 lda: c_int,
337 lwork: *mut c_int,
338 ) -> cusolverStatus_t { unsafe {
339 match cusolver().and_then(|c| c.$getrf_buf()) {
340 Ok(f) => f(h, m, n, a, lda, lwork),
341 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
342 }
343 }}
344 unsafe fn getrf(
345 h: cusolverDnHandle_t,
346 m: c_int,
347 n: c_int,
348 a: *mut $t,
349 lda: c_int,
350 work: *mut $t,
351 ipiv: *mut c_int,
352 info: *mut c_int,
353 ) -> cusolverStatus_t { unsafe {
354 match cusolver().and_then(|c| c.$getrf()) {
355 Ok(f) => f(h, m, n, a, lda, work, ipiv, info),
356 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
357 }
358 }}
359 unsafe fn getrs(
360 h: cusolverDnHandle_t,
361 trans: cublasOperation_t,
362 n: c_int,
363 nrhs: c_int,
364 a: *const $t,
365 lda: c_int,
366 ipiv: *const c_int,
367 b: *mut $t,
368 ldb: c_int,
369 info: *mut c_int,
370 ) -> cusolverStatus_t { unsafe {
371 match cusolver().and_then(|c| c.$getrs()) {
372 Ok(f) => f(h, trans, n, nrhs, a, lda, ipiv, b, ldb, info),
373 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
374 }
375 }}
376 unsafe fn geqrf_buf(
377 h: cusolverDnHandle_t,
378 m: c_int,
379 n: c_int,
380 a: *mut $t,
381 lda: c_int,
382 lwork: *mut c_int,
383 ) -> cusolverStatus_t { unsafe {
384 match cusolver().and_then(|c| c.$geqrf_buf()) {
385 Ok(f) => f(h, m, n, a, lda, lwork),
386 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
387 }
388 }}
389 unsafe fn geqrf(
390 h: cusolverDnHandle_t,
391 m: c_int,
392 n: c_int,
393 a: *mut $t,
394 lda: c_int,
395 tau: *mut $t,
396 work: *mut $t,
397 lwork: c_int,
398 info: *mut c_int,
399 ) -> cusolverStatus_t { unsafe {
400 match cusolver().and_then(|c| c.$geqrf()) {
401 Ok(f) => f(h, m, n, a, lda, tau, work, lwork, info),
402 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
403 }
404 }}
405 unsafe fn potrf_buf(
406 h: cusolverDnHandle_t,
407 uplo: cublasFillMode_t,
408 n: c_int,
409 a: *mut $t,
410 lda: c_int,
411 lwork: *mut c_int,
412 ) -> cusolverStatus_t { unsafe {
413 match cusolver().and_then(|c| c.$potrf_buf()) {
414 Ok(f) => f(h, uplo, n, a, lda, lwork),
415 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
416 }
417 }}
418 unsafe fn potrf(
419 h: cusolverDnHandle_t,
420 uplo: cublasFillMode_t,
421 n: c_int,
422 a: *mut $t,
423 lda: c_int,
424 work: *mut $t,
425 lwork: c_int,
426 info: *mut c_int,
427 ) -> cusolverStatus_t { unsafe {
428 match cusolver().and_then(|c| c.$potrf()) {
429 Ok(f) => f(h, uplo, n, a, lda, work, lwork, info),
430 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
431 }
432 }}
433 unsafe fn potrs(
434 h: cusolverDnHandle_t,
435 uplo: cublasFillMode_t,
436 n: c_int,
437 nrhs: c_int,
438 a: *const $t,
439 lda: c_int,
440 b: *mut $t,
441 ldb: c_int,
442 info: *mut c_int,
443 ) -> cusolverStatus_t { unsafe {
444 match cusolver().and_then(|c| c.$potrs()) {
445 Ok(f) => f(h, uplo, n, nrhs, a, lda, b, ldb, info),
446 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
447 }
448 }}
449 unsafe fn gesvd_buf(
450 h: cusolverDnHandle_t,
451 m: c_int,
452 n: c_int,
453 lwork: *mut c_int,
454 ) -> cusolverStatus_t { unsafe {
455 match cusolver().and_then(|c| c.$gesvd_buf()) {
456 Ok(f) => f(h, m, n, lwork),
457 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
458 }
459 }}
460 unsafe fn gesvd(
461 h: cusolverDnHandle_t,
462 jobu: u8,
463 jobvt: u8,
464 m: c_int,
465 n: c_int,
466 a: *mut $t,
467 lda: c_int,
468 s: *mut $t,
469 u: *mut $t,
470 ldu: c_int,
471 vt: *mut $t,
472 ldvt: c_int,
473 work: *mut $t,
474 lwork: c_int,
475 rwork: *mut $t,
476 info: *mut c_int,
477 ) -> cusolverStatus_t { unsafe {
478 match cusolver().and_then(|c| c.$gesvd()) {
479 Ok(f) => f(
480 h, jobu, jobvt, m, n, a, lda, s, u, ldu, vt, ldvt, work, lwork, rwork, info,
481 ),
482 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
483 }
484 }}
485 unsafe fn syevd_buf(
486 h: cusolverDnHandle_t,
487 jobz: cusolverEigMode_t,
488 uplo: cublasFillMode_t,
489 n: c_int,
490 a: *const $t,
491 lda: c_int,
492 w: *const $t,
493 lwork: *mut c_int,
494 ) -> cusolverStatus_t { unsafe {
495 match cusolver().and_then(|c| c.$syevd_buf()) {
496 Ok(f) => f(h, jobz, uplo, n, a, lda, w, lwork),
497 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
498 }
499 }}
500 unsafe fn syevd(
501 h: cusolverDnHandle_t,
502 jobz: cusolverEigMode_t,
503 uplo: cublasFillMode_t,
504 n: c_int,
505 a: *mut $t,
506 lda: c_int,
507 w: *mut $t,
508 work: *mut $t,
509 lwork: c_int,
510 info: *mut c_int,
511 ) -> cusolverStatus_t { unsafe {
512 match cusolver().and_then(|c| c.$syevd()) {
513 Ok(f) => f(h, jobz, uplo, n, a, lda, w, work, lwork, info),
514 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
515 }
516 }}
517 }
518 };
519}
520
521macro_rules! complex_impl {
522 ($t:ty, $real:ty, $raw:ty,
523 $getrf_buf:ident, $getrf:ident, $getrs:ident,
524 $geqrf_buf:ident, $geqrf:ident,
525 $potrf_buf:ident, $potrf:ident, $potrs:ident,
526 $gesvd_buf:ident, $gesvd:ident,
527 $heevd_buf:ident, $heevd:ident) => {
528 impl SolverScalar for $t {
529 type Real = $real;
530
531 unsafe fn getrf_buf(
532 h: cusolverDnHandle_t,
533 m: c_int,
534 n: c_int,
535 a: *mut $t,
536 lda: c_int,
537 lwork: *mut c_int,
538 ) -> cusolverStatus_t { unsafe {
539 match cusolver().and_then(|c| c.$getrf_buf()) {
540 Ok(f) => f(h, m, n, a as *mut $raw, lda, lwork),
541 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
542 }
543 }}
544 unsafe fn getrf(
545 h: cusolverDnHandle_t,
546 m: c_int,
547 n: c_int,
548 a: *mut $t,
549 lda: c_int,
550 work: *mut $t,
551 ipiv: *mut c_int,
552 info: *mut c_int,
553 ) -> cusolverStatus_t { unsafe {
554 match cusolver().and_then(|c| c.$getrf()) {
555 Ok(f) => f(
556 h,
557 m,
558 n,
559 a as *mut $raw,
560 lda,
561 work as *mut $raw,
562 ipiv,
563 info,
564 ),
565 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
566 }
567 }}
568 unsafe fn getrs(
569 h: cusolverDnHandle_t,
570 trans: cublasOperation_t,
571 n: c_int,
572 nrhs: c_int,
573 a: *const $t,
574 lda: c_int,
575 ipiv: *const c_int,
576 b: *mut $t,
577 ldb: c_int,
578 info: *mut c_int,
579 ) -> cusolverStatus_t { unsafe {
580 match cusolver().and_then(|c| c.$getrs()) {
581 Ok(f) => f(
582 h,
583 trans,
584 n,
585 nrhs,
586 a as *const $raw,
587 lda,
588 ipiv,
589 b as *mut $raw,
590 ldb,
591 info,
592 ),
593 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
594 }
595 }}
596 unsafe fn geqrf_buf(
597 h: cusolverDnHandle_t,
598 m: c_int,
599 n: c_int,
600 a: *mut $t,
601 lda: c_int,
602 lwork: *mut c_int,
603 ) -> cusolverStatus_t { unsafe {
604 match cusolver().and_then(|c| c.$geqrf_buf()) {
605 Ok(f) => f(h, m, n, a as *mut $raw, lda, lwork),
606 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
607 }
608 }}
609 unsafe fn geqrf(
610 h: cusolverDnHandle_t,
611 m: c_int,
612 n: c_int,
613 a: *mut $t,
614 lda: c_int,
615 tau: *mut $t,
616 work: *mut $t,
617 lwork: c_int,
618 info: *mut c_int,
619 ) -> cusolverStatus_t { unsafe {
620 match cusolver().and_then(|c| c.$geqrf()) {
621 Ok(f) => f(
622 h,
623 m,
624 n,
625 a as *mut $raw,
626 lda,
627 tau as *mut $raw,
628 work as *mut $raw,
629 lwork,
630 info,
631 ),
632 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
633 }
634 }}
635 unsafe fn potrf_buf(
636 h: cusolverDnHandle_t,
637 uplo: cublasFillMode_t,
638 n: c_int,
639 a: *mut $t,
640 lda: c_int,
641 lwork: *mut c_int,
642 ) -> cusolverStatus_t { unsafe {
643 match cusolver().and_then(|c| c.$potrf_buf()) {
644 Ok(f) => f(h, uplo, n, a as *mut $raw, lda, lwork),
645 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
646 }
647 }}
648 unsafe fn potrf(
649 h: cusolverDnHandle_t,
650 uplo: cublasFillMode_t,
651 n: c_int,
652 a: *mut $t,
653 lda: c_int,
654 work: *mut $t,
655 lwork: c_int,
656 info: *mut c_int,
657 ) -> cusolverStatus_t { unsafe {
658 match cusolver().and_then(|c| c.$potrf()) {
659 Ok(f) => f(
660 h,
661 uplo,
662 n,
663 a as *mut $raw,
664 lda,
665 work as *mut $raw,
666 lwork,
667 info,
668 ),
669 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
670 }
671 }}
672 unsafe fn potrs(
673 h: cusolverDnHandle_t,
674 uplo: cublasFillMode_t,
675 n: c_int,
676 nrhs: c_int,
677 a: *const $t,
678 lda: c_int,
679 b: *mut $t,
680 ldb: c_int,
681 info: *mut c_int,
682 ) -> cusolverStatus_t { unsafe {
683 match cusolver().and_then(|c| c.$potrs()) {
684 Ok(f) => f(
685 h,
686 uplo,
687 n,
688 nrhs,
689 a as *const $raw,
690 lda,
691 b as *mut $raw,
692 ldb,
693 info,
694 ),
695 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
696 }
697 }}
698 unsafe fn gesvd_buf(
699 h: cusolverDnHandle_t,
700 m: c_int,
701 n: c_int,
702 lwork: *mut c_int,
703 ) -> cusolverStatus_t { unsafe {
704 match cusolver().and_then(|c| c.$gesvd_buf()) {
705 Ok(f) => f(h, m, n, lwork),
706 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
707 }
708 }}
709 unsafe fn gesvd(
710 h: cusolverDnHandle_t,
711 jobu: u8,
712 jobvt: u8,
713 m: c_int,
714 n: c_int,
715 a: *mut $t,
716 lda: c_int,
717 s: *mut $real,
718 u: *mut $t,
719 ldu: c_int,
720 vt: *mut $t,
721 ldvt: c_int,
722 work: *mut $t,
723 lwork: c_int,
724 rwork: *mut $real,
725 info: *mut c_int,
726 ) -> cusolverStatus_t { unsafe {
727 match cusolver().and_then(|c| c.$gesvd()) {
728 Ok(f) => f(
729 h,
730 jobu,
731 jobvt,
732 m,
733 n,
734 a as *mut $raw,
735 lda,
736 s,
737 u as *mut $raw,
738 ldu,
739 vt as *mut $raw,
740 ldvt,
741 work as *mut $raw,
742 lwork,
743 rwork,
744 info,
745 ),
746 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
747 }
748 }}
749 unsafe fn syevd_buf(
750 h: cusolverDnHandle_t,
751 jobz: cusolverEigMode_t,
752 uplo: cublasFillMode_t,
753 n: c_int,
754 a: *const $t,
755 lda: c_int,
756 w: *const $real,
757 lwork: *mut c_int,
758 ) -> cusolverStatus_t { unsafe {
759 match cusolver().and_then(|c| c.$heevd_buf()) {
760 Ok(f) => f(h, jobz, uplo, n, a as *const $raw, lda, w, lwork),
761 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
762 }
763 }}
764 unsafe fn syevd(
765 h: cusolverDnHandle_t,
766 jobz: cusolverEigMode_t,
767 uplo: cublasFillMode_t,
768 n: c_int,
769 a: *mut $t,
770 lda: c_int,
771 w: *mut $real,
772 work: *mut $t,
773 lwork: c_int,
774 info: *mut c_int,
775 ) -> cusolverStatus_t { unsafe {
776 match cusolver().and_then(|c| c.$heevd()) {
777 Ok(f) => f(
778 h,
779 jobz,
780 uplo,
781 n,
782 a as *mut $raw,
783 lda,
784 w,
785 work as *mut $raw,
786 lwork,
787 info,
788 ),
789 Err(_) => cusolverStatus_t::NOT_INITIALIZED,
790 }
791 }}
792 }
793 };
794}
795
796real_impl!(
797 f32,
798 cusolver_dn_sgetrf_buffer_size,
799 cusolver_dn_sgetrf,
800 cusolver_dn_sgetrs,
801 cusolver_dn_sgeqrf_buffer_size,
802 cusolver_dn_sgeqrf,
803 cusolver_dn_spotrf_buffer_size,
804 cusolver_dn_spotrf,
805 cusolver_dn_spotrs,
806 cusolver_dn_sgesvd_buffer_size,
807 cusolver_dn_sgesvd,
808 cusolver_dn_ssyevd_buffer_size,
809 cusolver_dn_ssyevd
810);
811
812real_impl!(
813 f64,
814 cusolver_dn_dgetrf_buffer_size,
815 cusolver_dn_dgetrf,
816 cusolver_dn_dgetrs,
817 cusolver_dn_dgeqrf_buffer_size,
818 cusolver_dn_dgeqrf,
819 cusolver_dn_dpotrf_buffer_size,
820 cusolver_dn_dpotrf,
821 cusolver_dn_dpotrs,
822 cusolver_dn_dgesvd_buffer_size,
823 cusolver_dn_dgesvd,
824 cusolver_dn_dsyevd_buffer_size,
825 cusolver_dn_dsyevd
826);
827
828complex_impl!(
829 Complex32,
830 f32,
831 cuComplex,
832 cusolver_dn_cgetrf_buffer_size,
833 cusolver_dn_cgetrf,
834 cusolver_dn_cgetrs,
835 cusolver_dn_cgeqrf_buffer_size,
836 cusolver_dn_cgeqrf,
837 cusolver_dn_cpotrf_buffer_size,
838 cusolver_dn_cpotrf,
839 cusolver_dn_cpotrs,
840 cusolver_dn_cgesvd_buffer_size,
841 cusolver_dn_cgesvd,
842 cusolver_dn_cheevd_buffer_size,
843 cusolver_dn_cheevd
844);
845
846complex_impl!(
847 Complex64,
848 f64,
849 cuDoubleComplex,
850 cusolver_dn_zgetrf_buffer_size,
851 cusolver_dn_zgetrf,
852 cusolver_dn_zgetrs,
853 cusolver_dn_zgeqrf_buffer_size,
854 cusolver_dn_zgeqrf,
855 cusolver_dn_zpotrf_buffer_size,
856 cusolver_dn_zpotrf,
857 cusolver_dn_zpotrs,
858 cusolver_dn_zgesvd_buffer_size,
859 cusolver_dn_zgesvd,
860 cusolver_dn_zheevd_buffer_size,
861 cusolver_dn_zheevd
862);
863
864#[allow(clippy::too_many_arguments)]
868pub fn getrf<T: SolverScalar>(
869 handle: &DnHandle,
870 m: i32,
871 n: i32,
872 a: &mut DeviceBuffer<T>,
873 lda: i32,
874 ipiv: &mut DeviceBuffer<i32>,
875 info: &mut DeviceBuffer<i32>,
876) -> Result<()> {
877 let mut lwork: c_int = 0;
878 check(unsafe { T::getrf_buf(handle.handle, m, n, a.as_raw().0 as *mut T, lda, &mut lwork) })?;
879 let workspace =
880 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
881 check(unsafe {
882 T::getrf(
883 handle.handle,
884 m,
885 n,
886 a.as_raw().0 as *mut T,
887 lda,
888 workspace.as_raw().0 as *mut T,
889 ipiv.as_raw().0 as *mut c_int,
890 info.as_raw().0 as *mut c_int,
891 )
892 })
893}
894
895#[allow(clippy::too_many_arguments)]
897pub fn getrs<T: SolverScalar>(
898 handle: &DnHandle,
899 trans: Op,
900 n: i32,
901 nrhs: i32,
902 a: &DeviceBuffer<T>,
903 lda: i32,
904 ipiv: &DeviceBuffer<i32>,
905 b: &mut DeviceBuffer<T>,
906 ldb: i32,
907 info: &mut DeviceBuffer<i32>,
908) -> Result<()> {
909 check(unsafe {
910 T::getrs(
911 handle.handle,
912 trans.raw(),
913 n,
914 nrhs,
915 a.as_raw().0 as *const T,
916 lda,
917 ipiv.as_raw().0 as *const c_int,
918 b.as_raw().0 as *mut T,
919 ldb,
920 info.as_raw().0 as *mut c_int,
921 )
922 })
923}
924
925#[allow(clippy::too_many_arguments)]
928pub fn geqrf<T: SolverScalar>(
929 handle: &DnHandle,
930 m: i32,
931 n: i32,
932 a: &mut DeviceBuffer<T>,
933 lda: i32,
934 tau: &mut DeviceBuffer<T>,
935 info: &mut DeviceBuffer<i32>,
936) -> Result<()> {
937 let mut lwork: c_int = 0;
938 check(unsafe { T::geqrf_buf(handle.handle, m, n, a.as_raw().0 as *mut T, lda, &mut lwork) })?;
939 let workspace =
940 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
941 check(unsafe {
942 T::geqrf(
943 handle.handle,
944 m,
945 n,
946 a.as_raw().0 as *mut T,
947 lda,
948 tau.as_raw().0 as *mut T,
949 workspace.as_raw().0 as *mut T,
950 lwork,
951 info.as_raw().0 as *mut c_int,
952 )
953 })
954}
955
956pub fn potrf<T: SolverScalar>(
958 handle: &DnHandle,
959 uplo: Fill,
960 n: i32,
961 a: &mut DeviceBuffer<T>,
962 lda: i32,
963 info: &mut DeviceBuffer<i32>,
964) -> Result<()> {
965 let mut lwork: c_int = 0;
966 check(unsafe { T::potrf_buf(handle.handle, uplo, n, a.as_raw().0 as *mut T, lda, &mut lwork) })?;
967 let workspace =
968 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
969 check(unsafe {
970 T::potrf(
971 handle.handle,
972 uplo,
973 n,
974 a.as_raw().0 as *mut T,
975 lda,
976 workspace.as_raw().0 as *mut T,
977 lwork,
978 info.as_raw().0 as *mut c_int,
979 )
980 })
981}
982
983#[allow(clippy::too_many_arguments)]
985pub fn potrs<T: SolverScalar>(
986 handle: &DnHandle,
987 uplo: Fill,
988 n: i32,
989 nrhs: i32,
990 a: &DeviceBuffer<T>,
991 lda: i32,
992 b: &mut DeviceBuffer<T>,
993 ldb: i32,
994 info: &mut DeviceBuffer<i32>,
995) -> Result<()> {
996 check(unsafe {
997 T::potrs(
998 handle.handle,
999 uplo,
1000 n,
1001 nrhs,
1002 a.as_raw().0 as *const T,
1003 lda,
1004 b.as_raw().0 as *mut T,
1005 ldb,
1006 info.as_raw().0 as *mut c_int,
1007 )
1008 })
1009}
1010
1011#[allow(clippy::too_many_arguments)]
1017pub fn gesvd<T: SolverScalar>(
1018 handle: &DnHandle,
1019 jobu: u8,
1020 jobvt: u8,
1021 m: i32,
1022 n: i32,
1023 a: &mut DeviceBuffer<T>,
1024 lda: i32,
1025 s: &mut DeviceBuffer<T::Real>,
1026 u: &mut DeviceBuffer<T>,
1027 ldu: i32,
1028 vt: &mut DeviceBuffer<T>,
1029 ldvt: i32,
1030 rwork: &mut DeviceBuffer<T::Real>,
1031 info: &mut DeviceBuffer<i32>,
1032) -> Result<()> {
1033 let mut lwork: c_int = 0;
1034 check(unsafe { T::gesvd_buf(handle.handle, m, n, &mut lwork) })?;
1035 let workspace =
1036 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1037 check(unsafe {
1038 T::gesvd(
1039 handle.handle,
1040 jobu,
1041 jobvt,
1042 m,
1043 n,
1044 a.as_raw().0 as *mut T,
1045 lda,
1046 s.as_raw().0 as *mut T::Real,
1047 u.as_raw().0 as *mut T,
1048 ldu,
1049 vt.as_raw().0 as *mut T,
1050 ldvt,
1051 workspace.as_raw().0 as *mut T,
1052 lwork,
1053 rwork.as_raw().0 as *mut T::Real,
1054 info.as_raw().0 as *mut c_int,
1055 )
1056 })
1057}
1058
1059#[allow(clippy::too_many_arguments)]
1061pub fn syevd<T: SolverScalar>(
1062 handle: &DnHandle,
1063 jobz: EigMode,
1064 uplo: Fill,
1065 n: i32,
1066 a: &mut DeviceBuffer<T>,
1067 lda: i32,
1068 w: &mut DeviceBuffer<T::Real>,
1069 info: &mut DeviceBuffer<i32>,
1070) -> Result<()> {
1071 let mut lwork: c_int = 0;
1072 check(unsafe {
1073 T::syevd_buf(
1074 handle.handle,
1075 jobz,
1076 uplo,
1077 n,
1078 a.as_raw().0 as *const T,
1079 lda,
1080 w.as_raw().0 as *const T::Real,
1081 &mut lwork,
1082 )
1083 })?;
1084 let workspace =
1085 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1086 check(unsafe {
1087 T::syevd(
1088 handle.handle,
1089 jobz,
1090 uplo,
1091 n,
1092 a.as_raw().0 as *mut T,
1093 lda,
1094 w.as_raw().0 as *mut T::Real,
1095 workspace.as_raw().0 as *mut T,
1096 lwork,
1097 info.as_raw().0 as *mut c_int,
1098 )
1099 })
1100}
1101
1102pub use baracuda_cusolver_sys::{gesvdjInfo_t as GesvdjInfoRaw, syevjInfo_t as SyevjInfoRaw};
1105
1106#[derive(Debug)]
1108pub struct SyevjInfo {
1109 raw: SyevjInfoRaw,
1110}
1111
1112impl SyevjInfo {
1113 pub fn new() -> Result<Self> {
1114 let c = cusolver()?;
1115 let cu = c.cusolver_dn_create_syevj_info()?;
1116 let mut raw: SyevjInfoRaw = core::ptr::null_mut();
1117 check(unsafe { cu(&mut raw) })?;
1118 Ok(Self { raw })
1119 }
1120
1121 pub fn set_tolerance(&self, tol: f64) -> Result<()> {
1122 let c = cusolver()?;
1123 let cu = c.cusolver_dn_xsyevj_set_tolerance()?;
1124 check(unsafe { cu(self.raw, tol) })
1125 }
1126
1127 pub fn set_max_sweeps(&self, n: i32) -> Result<()> {
1128 let c = cusolver()?;
1129 let cu = c.cusolver_dn_xsyevj_set_max_sweeps()?;
1130 check(unsafe { cu(self.raw, n) })
1131 }
1132
1133 pub fn as_raw(&self) -> SyevjInfoRaw {
1134 self.raw
1135 }
1136}
1137
1138impl Drop for SyevjInfo {
1139 fn drop(&mut self) {
1140 if let Ok(c) = cusolver() {
1141 if let Ok(cu) = c.cusolver_dn_destroy_syevj_info() {
1142 let _ = unsafe { cu(self.raw) };
1143 }
1144 }
1145 }
1146}
1147
1148#[derive(Debug)]
1150pub struct GesvdjInfo {
1151 raw: GesvdjInfoRaw,
1152}
1153
1154impl GesvdjInfo {
1155 pub fn new() -> Result<Self> {
1156 let c = cusolver()?;
1157 let cu = c.cusolver_dn_create_gesvdj_info()?;
1158 let mut raw: GesvdjInfoRaw = core::ptr::null_mut();
1159 check(unsafe { cu(&mut raw) })?;
1160 Ok(Self { raw })
1161 }
1162
1163 pub fn as_raw(&self) -> GesvdjInfoRaw {
1164 self.raw
1165 }
1166}
1167
1168impl Drop for GesvdjInfo {
1169 fn drop(&mut self) {
1170 if let Ok(c) = cusolver() {
1171 if let Ok(cu) = c.cusolver_dn_destroy_gesvdj_info() {
1172 let _ = unsafe { cu(self.raw) };
1173 }
1174 }
1175 }
1176}
1177
1178#[allow(clippy::too_many_arguments)]
1181pub fn syevj<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 params: &SyevjInfo,
1191) -> Result<()> {
1192 use baracuda_cusolver_sys::{
1193 cuComplex, cuDoubleComplex,
1194 };
1195 use core::mem;
1196
1197 let mut lwork: c_int = 0;
1198
1199 macro_rules! dispatch_real {
1202 ($t:ty, $bufsize:ident, $solve:ident) => {{
1203 let c = cusolver()?;
1204 check(unsafe {
1205 (c.$bufsize()?)(
1206 handle.as_raw(),
1207 jobz,
1208 uplo,
1209 n,
1210 a.as_raw().0 as *const $t,
1211 lda,
1212 w.as_raw().0 as *const $t,
1213 &mut lwork,
1214 params.raw,
1215 )
1216 })?;
1217 let workspace =
1218 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1219 check(unsafe {
1220 (c.$solve()?)(
1221 handle.as_raw(),
1222 jobz,
1223 uplo,
1224 n,
1225 a.as_raw().0 as *mut $t,
1226 lda,
1227 w.as_raw().0 as *mut $t,
1228 workspace.as_raw().0 as *mut $t,
1229 lwork,
1230 info.as_raw().0 as *mut c_int,
1231 params.raw,
1232 )
1233 })
1234 }};
1235 }
1236 macro_rules! dispatch_complex {
1237 ($t:ty, $real:ty, $raw:ty, $bufsize:ident, $solve:ident) => {{
1238 let c = cusolver()?;
1239 check(unsafe {
1240 (c.$bufsize()?)(
1241 handle.as_raw(),
1242 jobz,
1243 uplo,
1244 n,
1245 a.as_raw().0 as *const $raw,
1246 lda,
1247 w.as_raw().0 as *const $real,
1248 &mut lwork,
1249 params.raw,
1250 )
1251 })?;
1252 let workspace =
1253 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1254 check(unsafe {
1255 (c.$solve()?)(
1256 handle.as_raw(),
1257 jobz,
1258 uplo,
1259 n,
1260 a.as_raw().0 as *mut $raw,
1261 lda,
1262 w.as_raw().0 as *mut $real,
1263 workspace.as_raw().0 as *mut $raw,
1264 lwork,
1265 info.as_raw().0 as *mut c_int,
1266 params.raw,
1267 )
1268 })
1269 }};
1270 }
1271
1272 if mem::size_of::<T>() == mem::size_of::<f32>() && mem::size_of::<T::Real>() == 4 {
1273 dispatch_real!(f32, cusolver_dn_ssyevj_buffer_size, cusolver_dn_ssyevj)
1274 } else if mem::size_of::<T>() == mem::size_of::<f64>() && mem::size_of::<T::Real>() == 8 {
1275 dispatch_real!(f64, cusolver_dn_dsyevj_buffer_size, cusolver_dn_dsyevj)
1276 } else if mem::size_of::<T>() == mem::size_of::<Complex32>() {
1277 dispatch_complex!(
1278 Complex32,
1279 f32,
1280 cuComplex,
1281 cusolver_dn_cheevj_buffer_size,
1282 cusolver_dn_cheevj
1283 )
1284 } else {
1285 dispatch_complex!(
1286 Complex64,
1287 f64,
1288 cuDoubleComplex,
1289 cusolver_dn_zheevj_buffer_size,
1290 cusolver_dn_zheevj
1291 )
1292 }
1293}
1294
1295#[allow(clippy::too_many_arguments)]
1297pub fn gesvdj<T: SolverScalar>(
1298 handle: &DnHandle,
1299 jobz: EigMode,
1300 econ: bool,
1301 m: i32,
1302 n: i32,
1303 a: &mut DeviceBuffer<T>,
1304 lda: i32,
1305 s: &mut DeviceBuffer<T::Real>,
1306 u: &mut DeviceBuffer<T>,
1307 ldu: i32,
1308 v: &mut DeviceBuffer<T>,
1309 ldv: i32,
1310 info: &mut DeviceBuffer<i32>,
1311 params: &GesvdjInfo,
1312) -> Result<()> {
1313 use baracuda_cusolver_sys::{cuComplex, cuDoubleComplex};
1314 use core::mem;
1315
1316 let mut lwork: c_int = 0;
1317 let econ_i = if econ { 1 } else { 0 };
1318
1319 macro_rules! dispatch_real {
1320 ($t:ty, $bufsize:ident, $solve:ident) => {{
1321 let c = cusolver()?;
1322 check(unsafe {
1323 (c.$bufsize()?)(
1324 handle.as_raw(),
1325 jobz,
1326 econ_i,
1327 m,
1328 n,
1329 a.as_raw().0 as *const $t,
1330 lda,
1331 s.as_raw().0 as *const $t,
1332 u.as_raw().0 as *const $t,
1333 ldu,
1334 v.as_raw().0 as *const $t,
1335 ldv,
1336 &mut lwork,
1337 params.raw,
1338 )
1339 })?;
1340 let workspace =
1341 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1342 check(unsafe {
1343 (c.$solve()?)(
1344 handle.as_raw(),
1345 jobz,
1346 econ_i,
1347 m,
1348 n,
1349 a.as_raw().0 as *mut $t,
1350 lda,
1351 s.as_raw().0 as *mut $t,
1352 u.as_raw().0 as *mut $t,
1353 ldu,
1354 v.as_raw().0 as *mut $t,
1355 ldv,
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 econ_i,
1372 m,
1373 n,
1374 a.as_raw().0 as *const $raw,
1375 lda,
1376 s.as_raw().0 as *const $real,
1377 u.as_raw().0 as *const $raw,
1378 ldu,
1379 v.as_raw().0 as *const $raw,
1380 ldv,
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 econ_i,
1392 m,
1393 n,
1394 a.as_raw().0 as *mut $raw,
1395 lda,
1396 s.as_raw().0 as *mut $real,
1397 u.as_raw().0 as *mut $raw,
1398 ldu,
1399 v.as_raw().0 as *mut $raw,
1400 ldv,
1401 workspace.as_raw().0 as *mut $raw,
1402 lwork,
1403 info.as_raw().0 as *mut c_int,
1404 params.raw,
1405 )
1406 })
1407 }};
1408 }
1409
1410 if mem::size_of::<T>() == mem::size_of::<f32>() && mem::size_of::<T::Real>() == 4 {
1411 dispatch_real!(f32, cusolver_dn_sgesvdj_buffer_size, cusolver_dn_sgesvdj)
1412 } else if mem::size_of::<T>() == mem::size_of::<f64>() && mem::size_of::<T::Real>() == 8 {
1413 dispatch_real!(f64, cusolver_dn_dgesvdj_buffer_size, cusolver_dn_dgesvdj)
1414 } else if mem::size_of::<T>() == mem::size_of::<Complex32>() {
1415 dispatch_complex!(
1416 Complex32,
1417 f32,
1418 cuComplex,
1419 cusolver_dn_cgesvdj_buffer_size,
1420 cusolver_dn_cgesvdj
1421 )
1422 } else {
1423 dispatch_complex!(
1424 Complex64,
1425 f64,
1426 cuDoubleComplex,
1427 cusolver_dn_zgesvdj_buffer_size,
1428 cusolver_dn_zgesvdj
1429 )
1430 }
1431}
1432
1433#[allow(clippy::too_many_arguments)]
1438pub fn orgqr<T: SolverScalar>(
1439 handle: &DnHandle,
1440 m: i32,
1441 n: i32,
1442 k: i32,
1443 a: &mut DeviceBuffer<T>,
1444 lda: i32,
1445 tau: &DeviceBuffer<T>,
1446 info: &mut DeviceBuffer<i32>,
1447) -> Result<()> {
1448 use baracuda_cusolver_sys::{cuComplex, cuDoubleComplex};
1449 use core::mem;
1450
1451 let mut lwork: c_int = 0;
1452 macro_rules! dispatch {
1453 ($t:ty, $raw:ty, $bufsize:ident, $solve:ident) => {{
1454 let c = cusolver()?;
1455 check(unsafe {
1456 (c.$bufsize()?)(
1457 handle.as_raw(),
1458 m,
1459 n,
1460 k,
1461 a.as_raw().0 as *const $raw,
1462 lda,
1463 tau.as_raw().0 as *const $raw,
1464 &mut lwork,
1465 )
1466 })?;
1467 let workspace =
1468 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1469 check(unsafe {
1470 (c.$solve()?)(
1471 handle.as_raw(),
1472 m,
1473 n,
1474 k,
1475 a.as_raw().0 as *mut $raw,
1476 lda,
1477 tau.as_raw().0 as *const $raw,
1478 workspace.as_raw().0 as *mut $raw,
1479 lwork,
1480 info.as_raw().0 as *mut c_int,
1481 )
1482 })
1483 }};
1484 }
1485
1486 if mem::size_of::<T>() == mem::size_of::<f32>() && mem::size_of::<T::Real>() == 4 {
1487 dispatch!(f32, f32, cusolver_dn_sorgqr_buffer_size, cusolver_dn_sorgqr)
1488 } else if mem::size_of::<T>() == mem::size_of::<f64>() && mem::size_of::<T::Real>() == 8 {
1489 dispatch!(f64, f64, cusolver_dn_dorgqr_buffer_size, cusolver_dn_dorgqr)
1490 } else if mem::size_of::<T>() == mem::size_of::<Complex32>() {
1491 dispatch!(
1492 Complex32,
1493 cuComplex,
1494 cusolver_dn_cungqr_buffer_size,
1495 cusolver_dn_cungqr
1496 )
1497 } else {
1498 dispatch!(
1499 Complex64,
1500 cuDoubleComplex,
1501 cusolver_dn_zungqr_buffer_size,
1502 cusolver_dn_zungqr
1503 )
1504 }
1505}
1506
1507#[derive(Copy, Clone, Debug, Eq, PartialEq)]
1509pub enum Side {
1510 Left,
1511 Right,
1512}
1513
1514impl Side {
1515 fn raw(self) -> core::ffi::c_int {
1516 match self {
1517 Side::Left => 0,
1518 Side::Right => 1,
1519 }
1520 }
1521}
1522
1523#[allow(clippy::too_many_arguments)]
1526pub fn ormqr<T: SolverScalar>(
1527 handle: &DnHandle,
1528 side: Side,
1529 trans: Op,
1530 m: i32,
1531 n: i32,
1532 k: i32,
1533 a: &DeviceBuffer<T>,
1534 lda: i32,
1535 tau: &DeviceBuffer<T>,
1536 c_mat: &mut DeviceBuffer<T>,
1537 ldc: i32,
1538 info: &mut DeviceBuffer<i32>,
1539) -> Result<()> {
1540 use baracuda_cusolver_sys::{cuComplex, cuDoubleComplex};
1541 use core::mem;
1542
1543 let mut lwork: c_int = 0;
1544 let side_i = side.raw();
1545 macro_rules! dispatch {
1546 ($t:ty, $raw:ty, $bufsize:ident, $solve:ident) => {{
1547 let ca = cusolver()?;
1548 check(unsafe {
1549 (ca.$bufsize()?)(
1550 handle.as_raw(),
1551 side_i,
1552 trans.raw(),
1553 m,
1554 n,
1555 k,
1556 a.as_raw().0 as *const $raw,
1557 lda,
1558 tau.as_raw().0 as *const $raw,
1559 c_mat.as_raw().0 as *const $raw,
1560 ldc,
1561 &mut lwork,
1562 )
1563 })?;
1564 let workspace =
1565 DeviceBuffer::<T>::new(c_mat.context(), lwork as usize).map_err(alloc_fail)?;
1566 check(unsafe {
1567 (ca.$solve()?)(
1568 handle.as_raw(),
1569 side_i,
1570 trans.raw(),
1571 m,
1572 n,
1573 k,
1574 a.as_raw().0 as *const $raw,
1575 lda,
1576 tau.as_raw().0 as *const $raw,
1577 c_mat.as_raw().0 as *mut $raw,
1578 ldc,
1579 workspace.as_raw().0 as *mut $raw,
1580 lwork,
1581 info.as_raw().0 as *mut c_int,
1582 )
1583 })
1584 }};
1585 }
1586
1587 if mem::size_of::<T>() == mem::size_of::<f32>() && mem::size_of::<T::Real>() == 4 {
1588 dispatch!(f32, f32, cusolver_dn_sormqr_buffer_size, cusolver_dn_sormqr)
1589 } else if mem::size_of::<T>() == mem::size_of::<f64>() && mem::size_of::<T::Real>() == 8 {
1590 dispatch!(f64, f64, cusolver_dn_dormqr_buffer_size, cusolver_dn_dormqr)
1591 } else if mem::size_of::<T>() == mem::size_of::<Complex32>() {
1592 dispatch!(
1593 Complex32,
1594 cuComplex,
1595 cusolver_dn_cunmqr_buffer_size,
1596 cusolver_dn_cunmqr
1597 )
1598 } else {
1599 dispatch!(
1600 Complex64,
1601 cuDoubleComplex,
1602 cusolver_dn_zunmqr_buffer_size,
1603 cusolver_dn_zunmqr
1604 )
1605 }
1606}
1607
1608#[allow(clippy::too_many_arguments)]
1615pub fn gels<T: SolverScalar>(
1616 handle: &DnHandle,
1617 m: i32,
1618 n: i32,
1619 nrhs: i32,
1620 a: &mut DeviceBuffer<T>,
1621 lda: i32,
1622 b: &mut DeviceBuffer<T>,
1623 ldb: i32,
1624 x: &mut DeviceBuffer<T>,
1625 ldx: i32,
1626 info: &mut DeviceBuffer<i32>,
1627) -> Result<i32> {
1628 use baracuda_cusolver_sys::{cuComplex, cuDoubleComplex};
1629 use core::mem;
1630
1631 let mut bytes: usize = 0;
1632
1633 macro_rules! dispatch {
1634 ($t:ty, $raw:ty, $bufsize:ident, $solve:ident) => {{
1635 let cs = cusolver()?;
1636 check(unsafe {
1637 (cs.$bufsize()?)(
1638 handle.as_raw(),
1639 m,
1640 n,
1641 nrhs,
1642 a.as_raw().0 as *mut $raw,
1643 lda,
1644 b.as_raw().0 as *mut $raw,
1645 ldb,
1646 x.as_raw().0 as *mut $raw,
1647 ldx,
1648 core::ptr::null_mut(),
1649 &mut bytes,
1650 )
1651 })?;
1652 let units = bytes.div_ceil(mem::size_of::<T>());
1654 let workspace =
1655 DeviceBuffer::<T>::new(a.context(), units).map_err(alloc_fail)?;
1656 let mut iter: c_int = 0;
1657 check(unsafe {
1658 (cs.$solve()?)(
1659 handle.as_raw(),
1660 m,
1661 n,
1662 nrhs,
1663 a.as_raw().0 as *mut $raw,
1664 lda,
1665 b.as_raw().0 as *mut $raw,
1666 ldb,
1667 x.as_raw().0 as *mut $raw,
1668 ldx,
1669 workspace.as_raw().0 as *mut c_void,
1670 bytes,
1671 &mut iter,
1672 info.as_raw().0 as *mut c_int,
1673 )
1674 })?;
1675 Ok(iter)
1676 }};
1677 }
1678
1679 if mem::size_of::<T>() == mem::size_of::<f32>() && mem::size_of::<T::Real>() == 4 {
1680 dispatch!(f32, f32, cusolver_dn_ssgels_buffer_size, cusolver_dn_ssgels)
1681 } else if mem::size_of::<T>() == mem::size_of::<f64>() && mem::size_of::<T::Real>() == 8 {
1682 dispatch!(f64, f64, cusolver_dn_ddgels_buffer_size, cusolver_dn_ddgels)
1683 } else if mem::size_of::<T>() == mem::size_of::<Complex32>() {
1684 dispatch!(
1685 Complex32,
1686 cuComplex,
1687 cusolver_dn_ccgels_buffer_size,
1688 cusolver_dn_ccgels
1689 )
1690 } else {
1691 dispatch!(
1692 Complex64,
1693 cuDoubleComplex,
1694 cusolver_dn_zzgels_buffer_size,
1695 cusolver_dn_zzgels
1696 )
1697 }
1698}
1699
1700pub fn potri<T: SolverScalar>(
1706 handle: &DnHandle,
1707 uplo: Fill,
1708 n: i32,
1709 a: &mut DeviceBuffer<T>,
1710 lda: i32,
1711 info: &mut DeviceBuffer<i32>,
1712) -> Result<()> {
1713 use baracuda_cusolver_sys::{cuComplex, cuDoubleComplex};
1714 use core::mem;
1715
1716 let mut lwork: c_int = 0;
1717 macro_rules! dispatch {
1718 ($t:ty, $raw:ty, $bufsize:ident, $solve:ident) => {{
1719 let cs = cusolver()?;
1720 check(unsafe {
1721 (cs.$bufsize()?)(
1722 handle.as_raw(),
1723 uplo,
1724 n,
1725 a.as_raw().0 as *mut $raw,
1726 lda,
1727 &mut lwork,
1728 )
1729 })?;
1730 let workspace =
1731 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1732 check(unsafe {
1733 (cs.$solve()?)(
1734 handle.as_raw(),
1735 uplo,
1736 n,
1737 a.as_raw().0 as *mut $raw,
1738 lda,
1739 workspace.as_raw().0 as *mut $raw,
1740 lwork,
1741 info.as_raw().0 as *mut c_int,
1742 )
1743 })
1744 }};
1745 }
1746
1747 if mem::size_of::<T>() == mem::size_of::<f32>() && mem::size_of::<T::Real>() == 4 {
1748 dispatch!(f32, f32, cusolver_dn_spotri_buffer_size, cusolver_dn_spotri)
1749 } else if mem::size_of::<T>() == mem::size_of::<f64>() && mem::size_of::<T::Real>() == 8 {
1750 dispatch!(f64, f64, cusolver_dn_dpotri_buffer_size, cusolver_dn_dpotri)
1751 } else if mem::size_of::<T>() == mem::size_of::<Complex32>() {
1752 dispatch!(
1753 Complex32,
1754 cuComplex,
1755 cusolver_dn_cpotri_buffer_size,
1756 cusolver_dn_cpotri
1757 )
1758 } else {
1759 dispatch!(
1760 Complex64,
1761 cuDoubleComplex,
1762 cusolver_dn_zpotri_buffer_size,
1763 cusolver_dn_zpotri
1764 )
1765 }
1766}
1767
1768#[allow(clippy::too_many_arguments)]
1774pub fn syevj_batched<T: SolverScalar>(
1775 handle: &DnHandle,
1776 jobz: EigMode,
1777 uplo: Fill,
1778 n: i32,
1779 a: &mut DeviceBuffer<T>,
1780 lda: i32,
1781 w: &mut DeviceBuffer<T::Real>,
1782 info: &mut DeviceBuffer<i32>,
1783 params: &SyevjInfo,
1784 batch_size: i32,
1785) -> Result<()> {
1786 use baracuda_cusolver_sys::{cuComplex, cuDoubleComplex};
1787 use core::mem;
1788
1789 let mut lwork: c_int = 0;
1790 macro_rules! dispatch_real {
1791 ($t:ty, $bufsize:ident, $solve:ident) => {{
1792 let c = cusolver()?;
1793 check(unsafe {
1794 (c.$bufsize()?)(
1795 handle.as_raw(),
1796 jobz,
1797 uplo,
1798 n,
1799 a.as_raw().0 as *const $t,
1800 lda,
1801 w.as_raw().0 as *const $t,
1802 &mut lwork,
1803 params.as_raw(),
1804 batch_size,
1805 )
1806 })?;
1807 let workspace =
1808 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1809 check(unsafe {
1810 (c.$solve()?)(
1811 handle.as_raw(),
1812 jobz,
1813 uplo,
1814 n,
1815 a.as_raw().0 as *mut $t,
1816 lda,
1817 w.as_raw().0 as *mut $t,
1818 workspace.as_raw().0 as *mut $t,
1819 lwork,
1820 info.as_raw().0 as *mut c_int,
1821 params.as_raw(),
1822 batch_size,
1823 )
1824 })
1825 }};
1826 }
1827 macro_rules! dispatch_complex {
1828 ($t:ty, $real:ty, $raw:ty, $bufsize:ident, $solve:ident) => {{
1829 let c = cusolver()?;
1830 check(unsafe {
1831 (c.$bufsize()?)(
1832 handle.as_raw(),
1833 jobz,
1834 uplo,
1835 n,
1836 a.as_raw().0 as *const $raw,
1837 lda,
1838 w.as_raw().0 as *const $real,
1839 &mut lwork,
1840 params.as_raw(),
1841 batch_size,
1842 )
1843 })?;
1844 let workspace =
1845 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1846 check(unsafe {
1847 (c.$solve()?)(
1848 handle.as_raw(),
1849 jobz,
1850 uplo,
1851 n,
1852 a.as_raw().0 as *mut $raw,
1853 lda,
1854 w.as_raw().0 as *mut $real,
1855 workspace.as_raw().0 as *mut $raw,
1856 lwork,
1857 info.as_raw().0 as *mut c_int,
1858 params.as_raw(),
1859 batch_size,
1860 )
1861 })
1862 }};
1863 }
1864
1865 if mem::size_of::<T>() == mem::size_of::<f32>() && mem::size_of::<T::Real>() == 4 {
1866 dispatch_real!(
1867 f32,
1868 cusolver_dn_ssyevj_batched_buffer_size,
1869 cusolver_dn_ssyevj_batched
1870 )
1871 } else if mem::size_of::<T>() == mem::size_of::<f64>() && mem::size_of::<T::Real>() == 8 {
1872 dispatch_real!(
1873 f64,
1874 cusolver_dn_dsyevj_batched_buffer_size,
1875 cusolver_dn_dsyevj_batched
1876 )
1877 } else if mem::size_of::<T>() == mem::size_of::<Complex32>() {
1878 dispatch_complex!(
1879 Complex32,
1880 f32,
1881 cuComplex,
1882 cusolver_dn_cheevj_batched_buffer_size,
1883 cusolver_dn_cheevj_batched
1884 )
1885 } else {
1886 dispatch_complex!(
1887 Complex64,
1888 f64,
1889 cuDoubleComplex,
1890 cusolver_dn_zheevj_batched_buffer_size,
1891 cusolver_dn_zheevj_batched
1892 )
1893 }
1894}
1895
1896#[allow(clippy::too_many_arguments)]
1898pub fn gesvdj_batched<T: SolverScalar>(
1899 handle: &DnHandle,
1900 jobz: EigMode,
1901 m: i32,
1902 n: i32,
1903 a: &mut DeviceBuffer<T>,
1904 lda: i32,
1905 s: &mut DeviceBuffer<T::Real>,
1906 u: &mut DeviceBuffer<T>,
1907 ldu: i32,
1908 v: &mut DeviceBuffer<T>,
1909 ldv: i32,
1910 info: &mut DeviceBuffer<i32>,
1911 params: &GesvdjInfo,
1912 batch_size: i32,
1913) -> Result<()> {
1914 use baracuda_cusolver_sys::{cuComplex, cuDoubleComplex};
1915 use core::mem;
1916
1917 let mut lwork: c_int = 0;
1918 macro_rules! dispatch_real {
1919 ($t:ty, $bufsize:ident, $solve:ident) => {{
1920 let c = cusolver()?;
1921 check(unsafe {
1922 (c.$bufsize()?)(
1923 handle.as_raw(),
1924 jobz,
1925 m,
1926 n,
1927 a.as_raw().0 as *const $t,
1928 lda,
1929 s.as_raw().0 as *const $t,
1930 u.as_raw().0 as *const $t,
1931 ldu,
1932 v.as_raw().0 as *const $t,
1933 ldv,
1934 &mut lwork,
1935 params.as_raw(),
1936 batch_size,
1937 )
1938 })?;
1939 let workspace =
1940 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1941 check(unsafe {
1942 (c.$solve()?)(
1943 handle.as_raw(),
1944 jobz,
1945 m,
1946 n,
1947 a.as_raw().0 as *mut $t,
1948 lda,
1949 s.as_raw().0 as *mut $t,
1950 u.as_raw().0 as *mut $t,
1951 ldu,
1952 v.as_raw().0 as *mut $t,
1953 ldv,
1954 workspace.as_raw().0 as *mut $t,
1955 lwork,
1956 info.as_raw().0 as *mut c_int,
1957 params.as_raw(),
1958 batch_size,
1959 )
1960 })
1961 }};
1962 }
1963 macro_rules! dispatch_complex {
1964 ($t:ty, $real:ty, $raw:ty, $bufsize:ident, $solve:ident) => {{
1965 let c = cusolver()?;
1966 check(unsafe {
1967 (c.$bufsize()?)(
1968 handle.as_raw(),
1969 jobz,
1970 m,
1971 n,
1972 a.as_raw().0 as *const $raw,
1973 lda,
1974 s.as_raw().0 as *const $real,
1975 u.as_raw().0 as *const $raw,
1976 ldu,
1977 v.as_raw().0 as *const $raw,
1978 ldv,
1979 &mut lwork,
1980 params.as_raw(),
1981 batch_size,
1982 )
1983 })?;
1984 let workspace =
1985 DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1986 check(unsafe {
1987 (c.$solve()?)(
1988 handle.as_raw(),
1989 jobz,
1990 m,
1991 n,
1992 a.as_raw().0 as *mut $raw,
1993 lda,
1994 s.as_raw().0 as *mut $real,
1995 u.as_raw().0 as *mut $raw,
1996 ldu,
1997 v.as_raw().0 as *mut $raw,
1998 ldv,
1999 workspace.as_raw().0 as *mut $raw,
2000 lwork,
2001 info.as_raw().0 as *mut c_int,
2002 params.as_raw(),
2003 batch_size,
2004 )
2005 })
2006 }};
2007 }
2008
2009 if mem::size_of::<T>() == mem::size_of::<f32>() && mem::size_of::<T::Real>() == 4 {
2010 dispatch_real!(
2011 f32,
2012 cusolver_dn_sgesvdj_batched_buffer_size,
2013 cusolver_dn_sgesvdj_batched
2014 )
2015 } else if mem::size_of::<T>() == mem::size_of::<f64>() && mem::size_of::<T::Real>() == 8 {
2016 dispatch_real!(
2017 f64,
2018 cusolver_dn_dgesvdj_batched_buffer_size,
2019 cusolver_dn_dgesvdj_batched
2020 )
2021 } else if mem::size_of::<T>() == mem::size_of::<Complex32>() {
2022 dispatch_complex!(
2023 Complex32,
2024 f32,
2025 cuComplex,
2026 cusolver_dn_cgesvdj_batched_buffer_size,
2027 cusolver_dn_cgesvdj_batched
2028 )
2029 } else {
2030 dispatch_complex!(
2031 Complex64,
2032 f64,
2033 cuDoubleComplex,
2034 cusolver_dn_zgesvdj_batched_buffer_size,
2035 cusolver_dn_zgesvdj_batched
2036 )
2037 }
2038}
2039
2040pub mod mg {
2043 use core::ffi::{c_int, c_void};
2048
2049 use baracuda_cusolver_sys::{
2050 cudaDataType, cudaLibMgGrid_t, cudaLibMgMatrixDesc_t, cusolver_mg, cusolverMgHandle_t,
2051 };
2052
2053 use super::{alloc_fail, check, EigMode, Fill, Result};
2054
2055 #[derive(Debug)]
2057 pub struct Handle {
2058 raw: cusolverMgHandle_t,
2059 }
2060
2061 impl Handle {
2062 pub fn new() -> Result<Self> {
2063 let mg = cusolver_mg()?;
2064 let cu = mg.cusolver_mg_create()?;
2065 let mut h: cusolverMgHandle_t = core::ptr::null_mut();
2066 check(unsafe { cu(&mut h) })?;
2067 Ok(Self { raw: h })
2068 }
2069
2070 pub fn device_select(&self, devices: &[i32]) -> Result<()> {
2073 let mg = cusolver_mg()?;
2074 let cu = mg.cusolver_mg_device_select()?;
2075 check(unsafe { cu(self.raw, devices.len() as c_int, devices.as_ptr()) })
2076 }
2077
2078 pub fn as_raw(&self) -> cusolverMgHandle_t {
2079 self.raw
2080 }
2081 }
2082
2083 impl Drop for Handle {
2084 fn drop(&mut self) {
2085 if let Ok(mg) = cusolver_mg() {
2086 if let Ok(cu) = mg.cusolver_mg_destroy() {
2087 let _ = unsafe { cu(self.raw) };
2088 }
2089 }
2090 }
2091 }
2092
2093 #[derive(Debug)]
2095 pub struct DeviceGrid {
2096 raw: cudaLibMgGrid_t,
2097 }
2098
2099 impl DeviceGrid {
2100 pub fn new(num_row_devices: i32, num_col_devices: i32, devices: &[i32], mapping: i32) -> Result<Self> {
2102 let mg = cusolver_mg()?;
2103 let cu = mg.cusolver_mg_create_device_grid()?;
2104 let mut raw: cudaLibMgGrid_t = core::ptr::null_mut();
2105 check(unsafe {
2106 cu(
2107 &mut raw,
2108 num_row_devices,
2109 num_col_devices,
2110 devices.as_ptr(),
2111 mapping,
2112 )
2113 })?;
2114 Ok(Self { raw })
2115 }
2116
2117 pub fn as_raw(&self) -> cudaLibMgGrid_t {
2118 self.raw
2119 }
2120 }
2121
2122 impl Drop for DeviceGrid {
2123 fn drop(&mut self) {
2124 if let Ok(mg) = cusolver_mg() {
2125 if let Ok(cu) = mg.cusolver_mg_destroy_grid() {
2126 let _ = unsafe { cu(self.raw) };
2127 }
2128 }
2129 }
2130 }
2131
2132 #[derive(Debug)]
2134 pub struct MatrixDesc {
2135 raw: cudaLibMgMatrixDesc_t,
2136 }
2137
2138 impl MatrixDesc {
2139 pub fn new(
2140 num_rows: i64,
2141 num_cols: i64,
2142 row_block_size: i64,
2143 col_block_size: i64,
2144 data_type: cudaDataType,
2145 grid: &DeviceGrid,
2146 ) -> Result<Self> {
2147 let mg = cusolver_mg()?;
2148 let cu = mg.cusolver_mg_create_matrix_desc()?;
2149 let mut raw: cudaLibMgMatrixDesc_t = core::ptr::null_mut();
2150 check(unsafe {
2151 cu(
2152 &mut raw,
2153 num_rows,
2154 num_cols,
2155 row_block_size,
2156 col_block_size,
2157 data_type,
2158 grid.as_raw(),
2159 )
2160 })?;
2161 Ok(Self { raw })
2162 }
2163
2164 pub fn as_raw(&self) -> cudaLibMgMatrixDesc_t {
2165 self.raw
2166 }
2167 }
2168
2169 impl Drop for MatrixDesc {
2170 fn drop(&mut self) {
2171 if let Ok(mg) = cusolver_mg() {
2172 if let Ok(cu) = mg.cusolver_mg_destroy_matrix_desc() {
2173 let _ = unsafe { cu(self.raw) };
2174 }
2175 }
2176 }
2177 }
2178
2179 #[allow(clippy::too_many_arguments)]
2185 pub unsafe fn getrf_buffer_size(
2186 handle: &Handle,
2187 m: i32,
2188 n: i32,
2189 array_d_a: *mut *mut c_void,
2190 ia: i32,
2191 ja: i32,
2192 desc_a: &MatrixDesc,
2193 array_d_ipiv: *mut *mut c_int,
2194 compute_type: cudaDataType,
2195 ) -> Result<i64> { unsafe {
2196 let mg = cusolver_mg()?;
2197 let cu = mg.cusolver_mg_getrf_buffer_size()?;
2198 let mut lwork: i64 = 0;
2199 check(cu(
2200 handle.as_raw(),
2201 m,
2202 n,
2203 array_d_a,
2204 ia,
2205 ja,
2206 desc_a.as_raw(),
2207 array_d_ipiv,
2208 compute_type,
2209 &mut lwork,
2210 ))?;
2211 Ok(lwork)
2212 }}
2213
2214 #[allow(clippy::too_many_arguments)]
2217 pub unsafe fn getrf(
2218 handle: &Handle,
2219 m: i32,
2220 n: i32,
2221 array_d_a: *mut *mut c_void,
2222 ia: i32,
2223 ja: i32,
2224 desc_a: &MatrixDesc,
2225 array_d_ipiv: *mut *mut c_int,
2226 compute_type: cudaDataType,
2227 array_d_work: *mut *mut c_void,
2228 lwork: i64,
2229 info: &mut [c_int],
2230 ) -> Result<()> { unsafe {
2231 let mg = cusolver_mg()?;
2232 let cu = mg.cusolver_mg_getrf()?;
2233 let _ = alloc_fail::<()>; check(cu(
2235 handle.as_raw(),
2236 m,
2237 n,
2238 array_d_a,
2239 ia,
2240 ja,
2241 desc_a.as_raw(),
2242 array_d_ipiv,
2243 compute_type,
2244 array_d_work,
2245 lwork,
2246 info.as_mut_ptr(),
2247 ))
2248 }}
2249
2250 #[allow(clippy::too_many_arguments)]
2255 pub unsafe fn potrf_buffer_size(
2256 handle: &Handle,
2257 uplo: Fill,
2258 n: i32,
2259 array_d_a: *mut *mut c_void,
2260 ia: i32,
2261 ja: i32,
2262 desc_a: &MatrixDesc,
2263 compute_type: cudaDataType,
2264 ) -> Result<i64> { unsafe {
2265 let mg = cusolver_mg()?;
2266 let cu = mg.cusolver_mg_potrf_buffer_size()?;
2267 let mut lwork: i64 = 0;
2268 check(cu(
2269 handle.as_raw(),
2270 uplo,
2271 n,
2272 array_d_a,
2273 ia,
2274 ja,
2275 desc_a.as_raw(),
2276 compute_type,
2277 &mut lwork,
2278 ))?;
2279 Ok(lwork)
2280 }}
2281
2282 #[allow(clippy::too_many_arguments)]
2285 pub unsafe fn potrf(
2286 handle: &Handle,
2287 uplo: Fill,
2288 n: i32,
2289 array_d_a: *mut *mut c_void,
2290 ia: i32,
2291 ja: i32,
2292 desc_a: &MatrixDesc,
2293 compute_type: cudaDataType,
2294 array_d_work: *mut *mut c_void,
2295 lwork: i64,
2296 info: &mut [c_int],
2297 ) -> Result<()> { unsafe {
2298 let mg = cusolver_mg()?;
2299 let cu = mg.cusolver_mg_potrf()?;
2300 check(cu(
2301 handle.as_raw(),
2302 uplo,
2303 n,
2304 array_d_a,
2305 ia,
2306 ja,
2307 desc_a.as_raw(),
2308 compute_type,
2309 array_d_work,
2310 lwork,
2311 info.as_mut_ptr(),
2312 ))
2313 }}
2314
2315 #[allow(clippy::too_many_arguments)]
2320 pub unsafe fn syevd_buffer_size(
2321 handle: &Handle,
2322 jobz: EigMode,
2323 uplo: Fill,
2324 n: i32,
2325 array_d_a: *mut *mut c_void,
2326 ia: i32,
2327 ja: i32,
2328 desc_a: &MatrixDesc,
2329 w: *mut c_void,
2330 data_type_w: cudaDataType,
2331 compute_type: cudaDataType,
2332 ) -> Result<i64> { unsafe {
2333 let mg = cusolver_mg()?;
2334 let cu = mg.cusolver_mg_syevd_buffer_size()?;
2335 let mut lwork: i64 = 0;
2336 check(cu(
2337 handle.as_raw(),
2338 jobz,
2339 uplo,
2340 n,
2341 array_d_a,
2342 ia,
2343 ja,
2344 desc_a.as_raw(),
2345 w,
2346 data_type_w,
2347 compute_type,
2348 &mut lwork,
2349 ))?;
2350 Ok(lwork)
2351 }}
2352
2353 #[allow(clippy::too_many_arguments)]
2356 pub unsafe fn syevd(
2357 handle: &Handle,
2358 jobz: EigMode,
2359 uplo: Fill,
2360 n: i32,
2361 array_d_a: *mut *mut c_void,
2362 ia: i32,
2363 ja: i32,
2364 desc_a: &MatrixDesc,
2365 w: *mut c_void,
2366 data_type_w: cudaDataType,
2367 compute_type: cudaDataType,
2368 array_d_work: *mut *mut c_void,
2369 lwork: i64,
2370 info: &mut [c_int],
2371 ) -> Result<()> { unsafe {
2372 let mg = cusolver_mg()?;
2373 let cu = mg.cusolver_mg_syevd()?;
2374 check(cu(
2375 handle.as_raw(),
2376 jobz,
2377 uplo,
2378 n,
2379 array_d_a,
2380 ia,
2381 ja,
2382 desc_a.as_raw(),
2383 w,
2384 data_type_w,
2385 compute_type,
2386 array_d_work,
2387 lwork,
2388 info.as_mut_ptr(),
2389 ))
2390 }}
2391}
2392
2393pub fn sgetrf(
2397 handle: &DnHandle,
2398 m: i32,
2399 n: i32,
2400 a: &mut DeviceBuffer<f32>,
2401 lda: i32,
2402 ipiv: &mut DeviceBuffer<i32>,
2403 info: &mut DeviceBuffer<i32>,
2404) -> Result<()> {
2405 getrf::<f32>(handle, m, n, a, lda, ipiv, info)
2406}
2407
2408#[allow(clippy::too_many_arguments)]
2410pub fn sgetrs(
2411 handle: &DnHandle,
2412 trans: Op,
2413 n: i32,
2414 nrhs: i32,
2415 a: &DeviceBuffer<f32>,
2416 lda: i32,
2417 ipiv: &DeviceBuffer<i32>,
2418 b: &mut DeviceBuffer<f32>,
2419 ldb: i32,
2420 info: &mut DeviceBuffer<i32>,
2421) -> Result<()> {
2422 getrs::<f32>(handle, trans, n, nrhs, a, lda, ipiv, b, ldb, info)
2423}
2424
2425pub mod xapi {
2428 use super::*;
2434 use baracuda_cusolver_sys::{cudaDataType, cusolverDnParams_t};
2435
2436 #[derive(Debug)]
2437 pub struct Params {
2438 raw: cusolverDnParams_t,
2439 }
2440
2441 impl Params {
2442 pub fn new() -> Result<Self> {
2443 let c = cusolver()?;
2444 let cu = c.cusolver_dn_create_params()?;
2445 let mut p: cusolverDnParams_t = core::ptr::null_mut();
2446 check(unsafe { cu(&mut p) })?;
2447 Ok(Self { raw: p })
2448 }
2449
2450 pub fn as_raw(&self) -> cusolverDnParams_t {
2451 self.raw
2452 }
2453 }
2454
2455 impl Drop for Params {
2456 fn drop(&mut self) {
2457 if let Ok(c) = cusolver() {
2458 if let Ok(cu) = c.cusolver_dn_destroy_params() {
2459 let _ = unsafe { cu(self.raw) };
2460 }
2461 }
2462 }
2463 }
2464
2465 #[allow(clippy::too_many_arguments)]
2477 pub unsafe fn xgetrf_buffer_size(
2478 handle: &DnHandle,
2479 params: &Params,
2480 m: i64,
2481 n: i64,
2482 data_type_a: cudaDataType,
2483 a: *const c_void,
2484 lda: i64,
2485 compute_type: cudaDataType,
2486 ) -> Result<(usize, usize)> {
2487 let c = cusolver()?;
2488 let cu = c.cusolver_dn_xgetrf_buffer_size()?;
2489 let (mut dev, mut host) = (0usize, 0usize);
2490 check(unsafe {
2491 cu(
2492 handle.as_raw(),
2493 params.raw,
2494 m,
2495 n,
2496 data_type_a,
2497 a,
2498 lda,
2499 compute_type,
2500 &mut dev,
2501 &mut host,
2502 )
2503 })?;
2504 Ok((dev, host))
2505 }
2506
2507 #[allow(clippy::too_many_arguments)]
2516 pub unsafe fn xgetrf(
2517 handle: &DnHandle,
2518 params: &Params,
2519 m: i64,
2520 n: i64,
2521 data_type_a: cudaDataType,
2522 a: *mut c_void,
2523 lda: i64,
2524 ipiv: *mut i64,
2525 compute_type: cudaDataType,
2526 device_buf: *mut c_void,
2527 device_bytes: usize,
2528 host_buf: *mut c_void,
2529 host_bytes: usize,
2530 info: *mut c_int,
2531 ) -> Result<()> { unsafe {
2532 let c = cusolver()?;
2533 let cu = c.cusolver_dn_xgetrf()?;
2534 check(cu(
2535 handle.as_raw(),
2536 params.raw,
2537 m,
2538 n,
2539 data_type_a,
2540 a,
2541 lda,
2542 ipiv,
2543 compute_type,
2544 device_buf,
2545 device_bytes,
2546 host_buf,
2547 host_bytes,
2548 info,
2549 ))
2550 }}
2551
2552 #[allow(clippy::too_many_arguments)]
2561 pub unsafe fn xgetrs(
2562 handle: &DnHandle,
2563 params: &Params,
2564 trans: Op,
2565 n: i64,
2566 nrhs: i64,
2567 data_type_a: cudaDataType,
2568 a: *const c_void,
2569 lda: i64,
2570 ipiv: *const i64,
2571 data_type_b: cudaDataType,
2572 b: *mut c_void,
2573 ldb: i64,
2574 info: *mut c_int,
2575 ) -> Result<()> { unsafe {
2576 let c = cusolver()?;
2577 let cu = c.cusolver_dn_xgetrs()?;
2578 check(cu(
2579 handle.as_raw(),
2580 params.raw,
2581 trans.raw(),
2582 n,
2583 nrhs,
2584 data_type_a,
2585 a,
2586 lda,
2587 ipiv,
2588 data_type_b,
2589 b,
2590 ldb,
2591 info,
2592 ))
2593 }}
2594
2595 #[allow(clippy::too_many_arguments)]
2603 pub unsafe fn xpotrs(
2604 handle: &DnHandle,
2605 params: &Params,
2606 uplo: Fill,
2607 n: i64,
2608 nrhs: i64,
2609 data_type_a: cudaDataType,
2610 a: *const c_void,
2611 lda: i64,
2612 data_type_b: cudaDataType,
2613 b: *mut c_void,
2614 ldb: i64,
2615 info: *mut c_int,
2616 ) -> Result<()> { unsafe {
2617 let c = cusolver()?;
2618 let cu = c.cusolver_dn_xpotrs()?;
2619 check(cu(
2620 handle.as_raw(),
2621 params.raw,
2622 uplo,
2623 n,
2624 nrhs,
2625 data_type_a,
2626 a,
2627 lda,
2628 data_type_b,
2629 b,
2630 ldb,
2631 info,
2632 ))
2633 }}
2634}
2635
2636pub mod sparse {
2639 use super::*;
2642 use baracuda_cusolver_sys::cusolverSpHandle_t;
2643 use core::ffi::c_int;
2644
2645 #[derive(Debug)]
2646 pub struct SpHandle {
2647 raw: cusolverSpHandle_t,
2648 _not_send: PhantomData<*mut ()>,
2649 }
2650
2651 impl SpHandle {
2652 pub fn new() -> Result<Self> {
2653 let c = cusolver()?;
2654 let cu = c.cusolver_sp_create()?;
2655 let mut h: cusolverSpHandle_t = core::ptr::null_mut();
2656 check(unsafe { cu(&mut h) })?;
2657 Ok(Self {
2658 raw: h,
2659 _not_send: PhantomData,
2660 })
2661 }
2662
2663 pub fn set_stream(&self, stream: &Stream) -> Result<()> {
2664 let c = cusolver()?;
2665 let cu = c.cusolver_sp_set_stream()?;
2666 check(unsafe { cu(self.raw, stream.as_raw() as _) })
2667 }
2668
2669 pub fn as_raw(&self) -> cusolverSpHandle_t {
2670 self.raw
2671 }
2672 }
2673
2674 impl Drop for SpHandle {
2675 fn drop(&mut self) {
2676 if let Ok(c) = cusolver() {
2677 if let Ok(cu) = c.cusolver_sp_destroy() {
2678 let _ = unsafe { cu(self.raw) };
2679 }
2680 }
2681 }
2682 }
2683
2684 #[allow(clippy::too_many_arguments)]
2691 pub unsafe fn scsrlsvchol(
2692 handle: &SpHandle,
2693 m: i32,
2694 nnz: i32,
2695 descr_a: *mut c_void,
2696 csr_val: *const f32,
2697 csr_row_ptr: *const c_int,
2698 csr_col_ind: *const c_int,
2699 b: *const f32,
2700 tol: f32,
2701 reorder: i32,
2702 x: *mut f32,
2703 singularity: *mut c_int,
2704 ) -> Result<()> { unsafe {
2705 let c = cusolver()?;
2706 let cu = c.cusolver_sp_scsrlsvchol()?;
2707 check(cu(
2708 handle.raw,
2709 m,
2710 nnz,
2711 descr_a,
2712 csr_val,
2713 csr_row_ptr,
2714 csr_col_ind,
2715 b,
2716 tol,
2717 reorder,
2718 x,
2719 singularity,
2720 ))
2721 }}
2722
2723 #[allow(clippy::too_many_arguments)]
2728 pub unsafe fn scsrlsvqr(
2729 handle: &SpHandle,
2730 m: i32,
2731 nnz: i32,
2732 descr_a: *mut c_void,
2733 csr_val: *const f32,
2734 csr_row_ptr: *const c_int,
2735 csr_col_ind: *const c_int,
2736 b: *const f32,
2737 tol: f32,
2738 reorder: i32,
2739 x: *mut f32,
2740 singularity: *mut c_int,
2741 ) -> Result<()> { unsafe {
2742 let c = cusolver()?;
2743 let cu = c.cusolver_sp_scsrlsvqr()?;
2744 check(cu(
2745 handle.raw,
2746 m,
2747 nnz,
2748 descr_a,
2749 csr_val,
2750 csr_row_ptr,
2751 csr_col_ind,
2752 b,
2753 tol,
2754 reorder,
2755 x,
2756 singularity,
2757 ))
2758 }}
2759}
2760
2761pub mod refactor {
2764 use super::*;
2768 use baracuda_cusolver_sys::cusolverRfHandle_t;
2769
2770 #[derive(Debug)]
2771 pub struct RfHandle {
2772 raw: cusolverRfHandle_t,
2773 _not_send: PhantomData<*mut ()>,
2774 }
2775
2776 impl RfHandle {
2777 pub fn new() -> Result<Self> {
2778 let c = cusolver()?;
2779 let cu = c.cusolver_rf_create()?;
2780 let mut h: cusolverRfHandle_t = core::ptr::null_mut();
2781 check(unsafe { cu(&mut h) })?;
2782 Ok(Self {
2783 raw: h,
2784 _not_send: PhantomData,
2785 })
2786 }
2787
2788 pub fn as_raw(&self) -> cusolverRfHandle_t {
2789 self.raw
2790 }
2791
2792 pub fn analyze(&self) -> Result<()> {
2793 let c = cusolver()?;
2794 let cu = c.cusolver_rf_analyze()?;
2795 check(unsafe { cu(self.raw) })
2796 }
2797
2798 pub fn refactor(&self) -> Result<()> {
2799 let c = cusolver()?;
2800 let cu = c.cusolver_rf_refactor()?;
2801 check(unsafe { cu(self.raw) })
2802 }
2803 }
2804
2805 impl Drop for RfHandle {
2806 fn drop(&mut self) {
2807 if let Ok(c) = cusolver() {
2808 if let Ok(cu) = c.cusolver_rf_destroy() {
2809 let _ = unsafe { cu(self.raw) };
2810 }
2811 }
2812 }
2813 }
2814}