Skip to main content

baracuda_cusolver/
lib.rs

1//! Safe Rust wrappers for NVIDIA cuSOLVER.
2//!
3//! Covers the dense API (`Dn`) for all four BLAS scalar types:
4//! - LU factorization: `getrf` + `getrs`
5//! - QR factorization: `geqrf`
6//! - Cholesky: `potrf` + `potrs`
7//! - SVD: `gesvd`
8//! - Symmetric / Hermitian eigendecomposition: `syevd` / `heevd`
9//!
10//! The generic 64-bit X… API (`xgetrf`, `xgeqrf`, `xpotrf`) gives
11//! type-erased data pointers and is exposed under [`xapi`]. The sparse API
12//! (`cusolverSp*`) is under [`sparse`]. The refactor API (`cusolverRf*`) is
13//! under [`refactor`].
14
15#![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
31/// Error type for cuSOLVER operations.
32pub type Error = baracuda_core::Error<cusolverStatus_t>;
33/// Result alias.
34pub 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
41/// Convert a driver allocation failure into a cuSOLVER ALLOC_FAILED.
42fn alloc_fail<E>(_e: E) -> Error {
43    Error::Status {
44        status: cusolverStatus_t::ALLOC_FAILED,
45    }
46}
47
48// ---- Handle -------------------------------------------------------------
49
50/// Dense cuSOLVER handle.
51pub 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    /// Read back the raw stream pointer this handle is currently bound
81    /// to. Returns a raw `*mut c_void` rather than an owned [`Stream`]
82    /// because the cuSOLVER handle didn't create the stream and isn't
83    /// entitled to destroy it. The pointer aliases `CUstream` /
84    /// `cudaStream_t`.
85    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/// Transposition selector for solve-step calls.
118#[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
136// ---- Trait framework ----------------------------------------------------
137
138/// Scalars supported by cuSOLVER's Dn S/D/C/Z API.
139pub trait SolverScalar: DeviceRepr + Copy + 'static + sealed::Sealed {
140    /// Real-valued associate type for ops that mix scalar and norm
141    /// (f32 → f32, f64 → f64, Complex32 → f32, Complex64 → f64).
142    type Real: DeviceRepr + Copy + 'static;
143
144    /// LU buffer size.
145    #[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    /// LU factorization.
156    #[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    /// LU solve.
170    #[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    /// QR factorization buffer size.
186    #[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    /// QR factorization.
197    #[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    /// Cholesky buffer size.
212    #[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    /// Cholesky factorization.
223    #[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    /// Cholesky solve.
237    #[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    /// SVD buffer size.
252    #[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    /// SVD (generic, taking real-valued S + rwork for complex variants).
261    #[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    /// syevd / heevd buffer size.
283    #[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    /// syevd / heevd.
297    #[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// ---- Public API ---------------------------------------------------------
865
866/// In-place LU factorization of a column-major matrix. Overwrites `a`.
867#[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/// Solve `op(A) * X = B` using the LU factorization from [`getrf`].
896#[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/// QR factorization: `A = Q * R`. Overwrites `a` (upper triangle = R,
926/// lower = Householder reflectors); `tau` receives reflector scalars.
927#[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
956/// Cholesky factorization: `A = L * Lᵀ` (or `Uᵀ * U`). Overwrites `a`.
957pub 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/// Solve `A * X = B` using the Cholesky factorization from [`potrf`].
984#[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/// Full SVD: `A = U * Σ * Vᵀ`. `jobu`/`jobvt` are LAPACK-style single-byte
1012/// selectors (b'A' = all, b'S' = economy, b'N' = none, b'O' = overwrite A).
1013///
1014/// `rwork` must be provided for complex element types; pass an empty buffer
1015/// for real types (pointer is still non-null; cuSOLVER ignores it).
1016#[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/// Symmetric / Hermitian eigenvalue decomposition: `A = Q * diag(w) * Qᵀ`.
1060#[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
1102// ---- Jacobi-based solvers (syevj / gesvdj) -----------------------------
1103
1104pub use baracuda_cusolver_sys::{gesvdjInfo_t as GesvdjInfoRaw, syevjInfo_t as SyevjInfoRaw};
1105
1106/// Jacobi-eigen tuning handle (tolerance + max sweeps).
1107#[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/// Jacobi-SVD tuning handle.
1149#[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/// Jacobi symmetric/Hermitian eigendecomposition (smaller matrices than
1179/// [`syevd`], faster convergence on well-conditioned problems).
1180#[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    // Dispatch is simpler done via a type check, since syevj doesn't share
1200    // the generic trait shape (extra params: SyevjInfo).
1201    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/// Jacobi SVD: `A = U * diag(s) * Vᴴ`. `econ` selects thin-SVD when set.
1296#[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// ---- Generate / apply Q from QR (orgqr / ormqr) -------------------------
1434
1435/// Generate the orthogonal matrix `Q` from the factorization produced by
1436/// [`geqrf`]. After this, `a` holds the first `n` columns of `Q`.
1437#[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/// Side argument for [`ormqr`].
1508#[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/// Apply `op(Q)` to `C`: `C = op(Q) * C` (Left) or `C = C * op(Q)` (Right),
1524/// where `Q` is packed in `a`+`tau` from [`geqrf`].
1525#[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// ---- gels: iterative-refinement least-squares solve ---------------------
1609
1610/// Solve `A * X = B` in the least-squares sense (iterative-refinement).
1611/// `A` is `m × n`, `B` is `m × nrhs`, `X` is `n × nrhs`. `A` and `B` may be
1612/// overwritten. Returns `iter`: number of refinement iterations used (-1 =
1613/// fallback to full precision).
1614#[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            // Allocate `bytes` worth of u8 workspace (rounding up to T units).
1653            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
1700// ---- potri: inverse from Cholesky factor --------------------------------
1701
1702/// Compute `A = (Lᵀ * L)⁻¹` or `A = (U * Uᵀ)⁻¹` given the Cholesky factor
1703/// already stored in the triangle selected by `uplo`. `a` must hold the
1704/// output of [`potrf`] in-place.
1705pub 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// ---- Batched Jacobi eigen / SVD -----------------------------------------
1769
1770/// Batched Jacobi symmetric/Hermitian eigendecomposition. Every matrix in
1771/// the batch is `n × n` and stride `n × n`. `w` holds `n * batch_size`
1772/// eigenvalues, strided by `n`.
1773#[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/// Batched Jacobi SVD: batch of `m × n` matrices with stride `m×n`.
1897#[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
2040// ---- cuSOLVERMg: multi-GPU dense solvers --------------------------------
2041
2042pub mod mg {
2043    //! Multi-GPU dense solvers via `libcusolverMg`. Shares dimensions with
2044    //! the single-GPU API but takes arrays of device pointers (one per
2045    //! physical GPU after [`Handle::device_select`]).
2046
2047    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    /// Multi-GPU cuSOLVER handle.
2056    #[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        /// Assign a set of physical CUDA devices to this handle. Future
2071        /// factorizations will stripe across them.
2072        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    /// A device grid — assigns distribution roles to physical devices.
2094    #[derive(Debug)]
2095    pub struct DeviceGrid {
2096        raw: cudaLibMgGrid_t,
2097    }
2098
2099    impl DeviceGrid {
2100        /// `mapping` is typically `CUDALIBMG_GRID_MAPPING_COL_MAJOR (1)`.
2101        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    /// Matrix-distribution descriptor.
2133    #[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    /// Multi-GPU LU buffer-size query.
2180    ///
2181    /// # Safety
2182    /// `array_d_a`, `array_d_ipiv` must be host arrays of device pointers
2183    /// matching the selected devices.
2184    #[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    /// # Safety
2215    /// Same pointer-array requirements as [`getrf_buffer_size`].
2216    #[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::<()>; // silence unused-import in release builds
2234        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    /// Multi-GPU Cholesky buffer-size.
2251    ///
2252    /// # Safety
2253    /// Same as [`getrf_buffer_size`].
2254    #[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    /// # Safety
2283    /// Same as [`getrf_buffer_size`].
2284    #[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    /// Multi-GPU symmetric eigendecomposition buffer-size.
2316    ///
2317    /// # Safety
2318    /// Same as [`getrf_buffer_size`].
2319    #[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    /// # Safety
2354    /// Same as [`getrf_buffer_size`].
2355    #[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
2393// ---- Back-compat: single-precision shortcuts -----------------------------
2394
2395/// Shortcut for [`getrf`] on `f32`.
2396pub 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/// Shortcut for [`getrs`] on `f32`.
2409#[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
2425// ---- Generic X... (64-bit-size, type-erased) ----------------------------
2426
2427pub mod xapi {
2428    //! The generic 64-bit cuSOLVER API (`cusolverDnX*`). Matrix dimensions
2429    //! are `i64`; element types are passed at call-time as
2430    //! [`cudaDataType`]. Workspace sizes are split between on-device and
2431    //! on-host buffers.
2432
2433    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    /// Buffer-size query for generic LU factorization. Returns
2466    /// `(workspace_bytes_on_device, workspace_bytes_on_host)`.
2467    ///
2468    /// # Safety
2469    ///
2470    /// `a` is a raw device pointer passed through to cuSOLVER's
2471    /// `cusolverDnXgetrf_bufferSize`. The pointer is not dereferenced by
2472    /// this Rust function but the underlying cuSOLVER call may sample
2473    /// alignment or layout from it; the caller must ensure it is either
2474    /// null or a valid device address for an `m x n` matrix with leading
2475    /// dimension `lda` of `data_type_a`.
2476    #[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    /// Generic LU factorization (`Xgetrf`).
2508    ///
2509    /// # Safety
2510    ///
2511    /// All raw device pointers (`a`, `ipiv`, `info`, `dev_workspace`,
2512    /// `host_workspace`) must point to allocations of the appropriate type
2513    /// and size for the requested `(m, n, data_type_a, compute_type)`,
2514    /// and remain live until the call completes on `handle`'s stream.
2515    #[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    /// Solve `op(A) X = B` after [`xgetrf`] has factored A. `b`/`ldb`
2553    /// describe the right-hand-side matrix; `b` is overwritten with X.
2554    ///
2555    /// # Safety
2556    ///
2557    /// `a`, `ipiv`, `b`, and `info` must point to live device memory of
2558    /// the sizes implied by `n`, `nrhs`, `lda`, and `ldb`, with element
2559    /// types matching `data_type_a` / `data_type_b`.
2560    #[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    /// Solve `A X = B` after [`xpotrf`] has Cholesky-factored A.
2596    /// `b`/`ldb` describe the right-hand-side matrix; `b` is overwritten
2597    /// with X.
2598    ///
2599    /// # Safety
2600    ///
2601    /// See [`xgetrs`].
2602    #[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
2636// ---- Sparse --------------------------------------------------------------
2637
2638pub mod sparse {
2639    //! `cusolverSp*` — solve sparse linear systems via Cholesky or QR.
2640
2641    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    /// Sparse Cholesky solve: `A * x = b` for SPD `A`.
2685    ///
2686    /// # Safety
2687    /// `descr_a`, CSR arrays, b and x must live on-device (b + x on-device,
2688    /// CSR arrays + descriptor on-device) and satisfy cuSOLVER sparse
2689    /// format requirements.
2690    #[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    /// Sparse QR solve (least-squares, handles non-SPD systems).
2724    ///
2725    /// # Safety
2726    /// Same as [`scsrlsvchol`].
2727    #[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
2761// ---- Refactor ------------------------------------------------------------
2762
2763pub mod refactor {
2764    //! `cusolverRf*` — fast re-factorization given a sparsity pattern, for
2765    //! solving many systems that differ only in numeric values.
2766
2767    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}