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    /// Create a new cuSOLVER dense handle (`cusolverDnCreate`).
67    pub fn new() -> Result<Self> {
68        let c = cusolver()?;
69        let cu = c.cusolver_dn_create()?;
70        let mut h: cusolverDnHandle_t = core::ptr::null_mut();
71        check(unsafe { cu(&mut h) })?;
72        Ok(Self { handle: h })
73    }
74
75    /// Bind this handle to `stream`; subsequent cuSOLVER calls dispatch on
76    /// it (`cusolverDnSetStream`).
77    pub fn set_stream(&self, stream: &Stream) -> Result<()> {
78        let c = cusolver()?;
79        let cu = c.cusolver_dn_set_stream()?;
80        check(unsafe { cu(self.handle, stream.as_raw() as _) })
81    }
82
83    /// Read back the raw stream pointer this handle is currently bound
84    /// to. Returns a raw `*mut c_void` rather than an owned [`Stream`]
85    /// because the cuSOLVER handle didn't create the stream and isn't
86    /// entitled to destroy it. The pointer aliases `CUstream` /
87    /// `cudaStream_t`.
88    pub fn stream(&self) -> Result<*mut c_void> {
89        let c = cusolver()?;
90        let cu = c.cusolver_dn_get_stream()?;
91        let mut s: *mut c_void = core::ptr::null_mut();
92        check(unsafe { cu(self.handle, &mut s as *mut *mut c_void as *mut _) })?;
93        Ok(s)
94    }
95
96    /// cuSOLVER library version (`cusolverGetVersion`),
97    /// e.g. `12604` for cuSOLVER 12.6.4.
98    pub fn version() -> Result<i32> {
99        let c = cusolver()?;
100        let cu = c.cusolver_get_version()?;
101        let mut v: c_int = 0;
102        check(unsafe { cu(&mut v) })?;
103        Ok(v)
104    }
105
106    /// Raw `cusolverDnHandle_t`. Use with care.
107    #[inline]
108    pub fn as_raw(&self) -> cusolverDnHandle_t {
109        self.handle
110    }
111}
112
113impl Drop for DnHandle {
114    fn drop(&mut self) {
115        if let Ok(c) = cusolver() {
116            if let Ok(cu) = c.cusolver_dn_destroy() {
117                let _ = unsafe { cu(self.handle) };
118            }
119        }
120    }
121}
122
123/// Transposition selector for solve-step calls.
124#[derive(Copy, Clone, Debug, Eq, PartialEq, Default)]
125pub enum Op {
126    /// No transpose.
127    #[default]
128    N,
129    /// Transpose.
130    T,
131    /// Conjugate transpose (real types: same as `T`).
132    C,
133}
134
135impl Op {
136    fn raw(self) -> cublasOperation_t {
137        match self {
138            Op::N => cublasOperation_t::N,
139            Op::T => cublasOperation_t::T,
140            Op::C => cublasOperation_t::C,
141        }
142    }
143}
144
145// ---- Trait framework ----------------------------------------------------
146
147/// Scalars supported by cuSOLVER's Dn S/D/C/Z API.
148pub trait SolverScalar: DeviceRepr + Copy + 'static + sealed::Sealed {
149    /// Real-valued associate type for ops that mix scalar and norm
150    /// (f32 → f32, f64 → f64, Complex32 → f32, Complex64 → f64).
151    type Real: DeviceRepr + Copy + 'static;
152
153    /// LU buffer size.
154    #[doc(hidden)]
155    unsafe fn getrf_buf(
156        h: cusolverDnHandle_t,
157        m: c_int,
158        n: c_int,
159        a: *mut Self,
160        lda: c_int,
161        lwork: *mut c_int,
162    ) -> cusolverStatus_t;
163
164    /// LU factorization.
165    #[doc(hidden)]
166    #[allow(clippy::too_many_arguments)]
167    unsafe fn getrf(
168        h: cusolverDnHandle_t,
169        m: c_int,
170        n: c_int,
171        a: *mut Self,
172        lda: c_int,
173        workspace: *mut Self,
174        ipiv: *mut c_int,
175        info: *mut c_int,
176    ) -> cusolverStatus_t;
177
178    /// LU solve.
179    #[doc(hidden)]
180    #[allow(clippy::too_many_arguments)]
181    unsafe fn getrs(
182        h: cusolverDnHandle_t,
183        trans: cublasOperation_t,
184        n: c_int,
185        nrhs: c_int,
186        a: *const Self,
187        lda: c_int,
188        ipiv: *const c_int,
189        b: *mut Self,
190        ldb: c_int,
191        info: *mut c_int,
192    ) -> cusolverStatus_t;
193
194    /// QR factorization buffer size.
195    #[doc(hidden)]
196    unsafe fn geqrf_buf(
197        h: cusolverDnHandle_t,
198        m: c_int,
199        n: c_int,
200        a: *mut Self,
201        lda: c_int,
202        lwork: *mut c_int,
203    ) -> cusolverStatus_t;
204
205    /// QR factorization.
206    #[doc(hidden)]
207    #[allow(clippy::too_many_arguments)]
208    unsafe fn geqrf(
209        h: cusolverDnHandle_t,
210        m: c_int,
211        n: c_int,
212        a: *mut Self,
213        lda: c_int,
214        tau: *mut Self,
215        workspace: *mut Self,
216        lwork: c_int,
217        info: *mut c_int,
218    ) -> cusolverStatus_t;
219
220    /// Cholesky buffer size.
221    #[doc(hidden)]
222    unsafe fn potrf_buf(
223        h: cusolverDnHandle_t,
224        uplo: cublasFillMode_t,
225        n: c_int,
226        a: *mut Self,
227        lda: c_int,
228        lwork: *mut c_int,
229    ) -> cusolverStatus_t;
230
231    /// Cholesky factorization.
232    #[doc(hidden)]
233    #[allow(clippy::too_many_arguments)]
234    unsafe fn potrf(
235        h: cusolverDnHandle_t,
236        uplo: cublasFillMode_t,
237        n: c_int,
238        a: *mut Self,
239        lda: c_int,
240        workspace: *mut Self,
241        lwork: c_int,
242        info: *mut c_int,
243    ) -> cusolverStatus_t;
244
245    /// Cholesky solve.
246    #[doc(hidden)]
247    #[allow(clippy::too_many_arguments)]
248    unsafe fn potrs(
249        h: cusolverDnHandle_t,
250        uplo: cublasFillMode_t,
251        n: c_int,
252        nrhs: c_int,
253        a: *const Self,
254        lda: c_int,
255        b: *mut Self,
256        ldb: c_int,
257        info: *mut c_int,
258    ) -> cusolverStatus_t;
259
260    /// SVD buffer size.
261    #[doc(hidden)]
262    unsafe fn gesvd_buf(
263        h: cusolverDnHandle_t,
264        m: c_int,
265        n: c_int,
266        lwork: *mut c_int,
267    ) -> cusolverStatus_t;
268
269    /// SVD (generic, taking real-valued S + rwork for complex variants).
270    #[doc(hidden)]
271    #[allow(clippy::too_many_arguments)]
272    unsafe fn gesvd(
273        h: cusolverDnHandle_t,
274        jobu: u8,
275        jobvt: u8,
276        m: c_int,
277        n: c_int,
278        a: *mut Self,
279        lda: c_int,
280        s: *mut Self::Real,
281        u: *mut Self,
282        ldu: c_int,
283        vt: *mut Self,
284        ldvt: c_int,
285        work: *mut Self,
286        lwork: c_int,
287        rwork: *mut Self::Real,
288        info: *mut c_int,
289    ) -> cusolverStatus_t;
290
291    /// syevd / heevd buffer size.
292    #[doc(hidden)]
293    #[allow(clippy::too_many_arguments)]
294    unsafe fn syevd_buf(
295        h: cusolverDnHandle_t,
296        jobz: cusolverEigMode_t,
297        uplo: cublasFillMode_t,
298        n: c_int,
299        a: *const Self,
300        lda: c_int,
301        w: *const Self::Real,
302        lwork: *mut c_int,
303    ) -> cusolverStatus_t;
304
305    /// syevd / heevd.
306    #[doc(hidden)]
307    #[allow(clippy::too_many_arguments)]
308    unsafe fn syevd(
309        h: cusolverDnHandle_t,
310        jobz: cusolverEigMode_t,
311        uplo: cublasFillMode_t,
312        n: c_int,
313        a: *mut Self,
314        lda: c_int,
315        w: *mut Self::Real,
316        work: *mut Self,
317        lwork: c_int,
318        info: *mut c_int,
319    ) -> cusolverStatus_t;
320}
321
322mod sealed {
323    use baracuda_types::{Complex32, Complex64};
324    pub trait Sealed {}
325    impl Sealed for f32 {}
326    impl Sealed for f64 {}
327    impl Sealed for Complex32 {}
328    impl Sealed for Complex64 {}
329}
330
331macro_rules! real_impl {
332    ($t:ty, $getrf_buf:ident, $getrf:ident, $getrs:ident,
333           $geqrf_buf:ident, $geqrf:ident,
334           $potrf_buf:ident, $potrf:ident, $potrs:ident,
335           $gesvd_buf:ident, $gesvd:ident,
336           $syevd_buf:ident, $syevd:ident) => {
337        impl SolverScalar for $t {
338            type Real = $t;
339
340            unsafe fn getrf_buf(
341                h: cusolverDnHandle_t,
342                m: c_int,
343                n: c_int,
344                a: *mut $t,
345                lda: c_int,
346                lwork: *mut c_int,
347            ) -> cusolverStatus_t { unsafe {
348                match cusolver().and_then(|c| c.$getrf_buf()) {
349                    Ok(f) => f(h, m, n, a, lda, lwork),
350                    Err(_) => cusolverStatus_t::NOT_INITIALIZED,
351                }
352            }}
353            unsafe fn getrf(
354                h: cusolverDnHandle_t,
355                m: c_int,
356                n: c_int,
357                a: *mut $t,
358                lda: c_int,
359                work: *mut $t,
360                ipiv: *mut c_int,
361                info: *mut c_int,
362            ) -> cusolverStatus_t { unsafe {
363                match cusolver().and_then(|c| c.$getrf()) {
364                    Ok(f) => f(h, m, n, a, lda, work, ipiv, info),
365                    Err(_) => cusolverStatus_t::NOT_INITIALIZED,
366                }
367            }}
368            unsafe fn getrs(
369                h: cusolverDnHandle_t,
370                trans: cublasOperation_t,
371                n: c_int,
372                nrhs: c_int,
373                a: *const $t,
374                lda: c_int,
375                ipiv: *const c_int,
376                b: *mut $t,
377                ldb: c_int,
378                info: *mut c_int,
379            ) -> cusolverStatus_t { unsafe {
380                match cusolver().and_then(|c| c.$getrs()) {
381                    Ok(f) => f(h, trans, n, nrhs, a, lda, ipiv, b, ldb, info),
382                    Err(_) => cusolverStatus_t::NOT_INITIALIZED,
383                }
384            }}
385            unsafe fn geqrf_buf(
386                h: cusolverDnHandle_t,
387                m: c_int,
388                n: c_int,
389                a: *mut $t,
390                lda: c_int,
391                lwork: *mut c_int,
392            ) -> cusolverStatus_t { unsafe {
393                match cusolver().and_then(|c| c.$geqrf_buf()) {
394                    Ok(f) => f(h, m, n, a, lda, lwork),
395                    Err(_) => cusolverStatus_t::NOT_INITIALIZED,
396                }
397            }}
398            unsafe fn geqrf(
399                h: cusolverDnHandle_t,
400                m: c_int,
401                n: c_int,
402                a: *mut $t,
403                lda: c_int,
404                tau: *mut $t,
405                work: *mut $t,
406                lwork: c_int,
407                info: *mut c_int,
408            ) -> cusolverStatus_t { unsafe {
409                match cusolver().and_then(|c| c.$geqrf()) {
410                    Ok(f) => f(h, m, n, a, lda, tau, work, lwork, info),
411                    Err(_) => cusolverStatus_t::NOT_INITIALIZED,
412                }
413            }}
414            unsafe fn potrf_buf(
415                h: cusolverDnHandle_t,
416                uplo: cublasFillMode_t,
417                n: c_int,
418                a: *mut $t,
419                lda: c_int,
420                lwork: *mut c_int,
421            ) -> cusolverStatus_t { unsafe {
422                match cusolver().and_then(|c| c.$potrf_buf()) {
423                    Ok(f) => f(h, uplo, n, a, lda, lwork),
424                    Err(_) => cusolverStatus_t::NOT_INITIALIZED,
425                }
426            }}
427            unsafe fn potrf(
428                h: cusolverDnHandle_t,
429                uplo: cublasFillMode_t,
430                n: c_int,
431                a: *mut $t,
432                lda: c_int,
433                work: *mut $t,
434                lwork: c_int,
435                info: *mut c_int,
436            ) -> cusolverStatus_t { unsafe {
437                match cusolver().and_then(|c| c.$potrf()) {
438                    Ok(f) => f(h, uplo, n, a, lda, work, lwork, info),
439                    Err(_) => cusolverStatus_t::NOT_INITIALIZED,
440                }
441            }}
442            unsafe fn potrs(
443                h: cusolverDnHandle_t,
444                uplo: cublasFillMode_t,
445                n: c_int,
446                nrhs: c_int,
447                a: *const $t,
448                lda: c_int,
449                b: *mut $t,
450                ldb: c_int,
451                info: *mut c_int,
452            ) -> cusolverStatus_t { unsafe {
453                match cusolver().and_then(|c| c.$potrs()) {
454                    Ok(f) => f(h, uplo, n, nrhs, a, lda, b, ldb, info),
455                    Err(_) => cusolverStatus_t::NOT_INITIALIZED,
456                }
457            }}
458            unsafe fn gesvd_buf(
459                h: cusolverDnHandle_t,
460                m: c_int,
461                n: c_int,
462                lwork: *mut c_int,
463            ) -> cusolverStatus_t { unsafe {
464                match cusolver().and_then(|c| c.$gesvd_buf()) {
465                    Ok(f) => f(h, m, n, lwork),
466                    Err(_) => cusolverStatus_t::NOT_INITIALIZED,
467                }
468            }}
469            unsafe fn gesvd(
470                h: cusolverDnHandle_t,
471                jobu: u8,
472                jobvt: u8,
473                m: c_int,
474                n: c_int,
475                a: *mut $t,
476                lda: c_int,
477                s: *mut $t,
478                u: *mut $t,
479                ldu: c_int,
480                vt: *mut $t,
481                ldvt: c_int,
482                work: *mut $t,
483                lwork: c_int,
484                rwork: *mut $t,
485                info: *mut c_int,
486            ) -> cusolverStatus_t { unsafe {
487                match cusolver().and_then(|c| c.$gesvd()) {
488                    Ok(f) => f(
489                        h, jobu, jobvt, m, n, a, lda, s, u, ldu, vt, ldvt, work, lwork, rwork, info,
490                    ),
491                    Err(_) => cusolverStatus_t::NOT_INITIALIZED,
492                }
493            }}
494            unsafe fn syevd_buf(
495                h: cusolverDnHandle_t,
496                jobz: cusolverEigMode_t,
497                uplo: cublasFillMode_t,
498                n: c_int,
499                a: *const $t,
500                lda: c_int,
501                w: *const $t,
502                lwork: *mut c_int,
503            ) -> cusolverStatus_t { unsafe {
504                match cusolver().and_then(|c| c.$syevd_buf()) {
505                    Ok(f) => f(h, jobz, uplo, n, a, lda, w, lwork),
506                    Err(_) => cusolverStatus_t::NOT_INITIALIZED,
507                }
508            }}
509            unsafe fn syevd(
510                h: cusolverDnHandle_t,
511                jobz: cusolverEigMode_t,
512                uplo: cublasFillMode_t,
513                n: c_int,
514                a: *mut $t,
515                lda: c_int,
516                w: *mut $t,
517                work: *mut $t,
518                lwork: c_int,
519                info: *mut c_int,
520            ) -> cusolverStatus_t { unsafe {
521                match cusolver().and_then(|c| c.$syevd()) {
522                    Ok(f) => f(h, jobz, uplo, n, a, lda, w, work, lwork, info),
523                    Err(_) => cusolverStatus_t::NOT_INITIALIZED,
524                }
525            }}
526        }
527    };
528}
529
530macro_rules! complex_impl {
531    ($t:ty, $real:ty, $raw:ty,
532     $getrf_buf:ident, $getrf:ident, $getrs:ident,
533     $geqrf_buf:ident, $geqrf:ident,
534     $potrf_buf:ident, $potrf:ident, $potrs:ident,
535     $gesvd_buf:ident, $gesvd:ident,
536     $heevd_buf:ident, $heevd:ident) => {
537        impl SolverScalar for $t {
538            type Real = $real;
539
540            unsafe fn getrf_buf(
541                h: cusolverDnHandle_t,
542                m: c_int,
543                n: c_int,
544                a: *mut $t,
545                lda: c_int,
546                lwork: *mut c_int,
547            ) -> cusolverStatus_t { unsafe {
548                match cusolver().and_then(|c| c.$getrf_buf()) {
549                    Ok(f) => f(h, m, n, a as *mut $raw, lda, lwork),
550                    Err(_) => cusolverStatus_t::NOT_INITIALIZED,
551                }
552            }}
553            unsafe fn getrf(
554                h: cusolverDnHandle_t,
555                m: c_int,
556                n: c_int,
557                a: *mut $t,
558                lda: c_int,
559                work: *mut $t,
560                ipiv: *mut c_int,
561                info: *mut c_int,
562            ) -> cusolverStatus_t { unsafe {
563                match cusolver().and_then(|c| c.$getrf()) {
564                    Ok(f) => f(
565                        h,
566                        m,
567                        n,
568                        a as *mut $raw,
569                        lda,
570                        work as *mut $raw,
571                        ipiv,
572                        info,
573                    ),
574                    Err(_) => cusolverStatus_t::NOT_INITIALIZED,
575                }
576            }}
577            unsafe fn getrs(
578                h: cusolverDnHandle_t,
579                trans: cublasOperation_t,
580                n: c_int,
581                nrhs: c_int,
582                a: *const $t,
583                lda: c_int,
584                ipiv: *const c_int,
585                b: *mut $t,
586                ldb: c_int,
587                info: *mut c_int,
588            ) -> cusolverStatus_t { unsafe {
589                match cusolver().and_then(|c| c.$getrs()) {
590                    Ok(f) => f(
591                        h,
592                        trans,
593                        n,
594                        nrhs,
595                        a as *const $raw,
596                        lda,
597                        ipiv,
598                        b as *mut $raw,
599                        ldb,
600                        info,
601                    ),
602                    Err(_) => cusolverStatus_t::NOT_INITIALIZED,
603                }
604            }}
605            unsafe fn geqrf_buf(
606                h: cusolverDnHandle_t,
607                m: c_int,
608                n: c_int,
609                a: *mut $t,
610                lda: c_int,
611                lwork: *mut c_int,
612            ) -> cusolverStatus_t { unsafe {
613                match cusolver().and_then(|c| c.$geqrf_buf()) {
614                    Ok(f) => f(h, m, n, a as *mut $raw, lda, lwork),
615                    Err(_) => cusolverStatus_t::NOT_INITIALIZED,
616                }
617            }}
618            unsafe fn geqrf(
619                h: cusolverDnHandle_t,
620                m: c_int,
621                n: c_int,
622                a: *mut $t,
623                lda: c_int,
624                tau: *mut $t,
625                work: *mut $t,
626                lwork: c_int,
627                info: *mut c_int,
628            ) -> cusolverStatus_t { unsafe {
629                match cusolver().and_then(|c| c.$geqrf()) {
630                    Ok(f) => f(
631                        h,
632                        m,
633                        n,
634                        a as *mut $raw,
635                        lda,
636                        tau as *mut $raw,
637                        work as *mut $raw,
638                        lwork,
639                        info,
640                    ),
641                    Err(_) => cusolverStatus_t::NOT_INITIALIZED,
642                }
643            }}
644            unsafe fn potrf_buf(
645                h: cusolverDnHandle_t,
646                uplo: cublasFillMode_t,
647                n: c_int,
648                a: *mut $t,
649                lda: c_int,
650                lwork: *mut c_int,
651            ) -> cusolverStatus_t { unsafe {
652                match cusolver().and_then(|c| c.$potrf_buf()) {
653                    Ok(f) => f(h, uplo, n, a as *mut $raw, lda, lwork),
654                    Err(_) => cusolverStatus_t::NOT_INITIALIZED,
655                }
656            }}
657            unsafe fn potrf(
658                h: cusolverDnHandle_t,
659                uplo: cublasFillMode_t,
660                n: c_int,
661                a: *mut $t,
662                lda: c_int,
663                work: *mut $t,
664                lwork: c_int,
665                info: *mut c_int,
666            ) -> cusolverStatus_t { unsafe {
667                match cusolver().and_then(|c| c.$potrf()) {
668                    Ok(f) => f(
669                        h,
670                        uplo,
671                        n,
672                        a as *mut $raw,
673                        lda,
674                        work as *mut $raw,
675                        lwork,
676                        info,
677                    ),
678                    Err(_) => cusolverStatus_t::NOT_INITIALIZED,
679                }
680            }}
681            unsafe fn potrs(
682                h: cusolverDnHandle_t,
683                uplo: cublasFillMode_t,
684                n: c_int,
685                nrhs: c_int,
686                a: *const $t,
687                lda: c_int,
688                b: *mut $t,
689                ldb: c_int,
690                info: *mut c_int,
691            ) -> cusolverStatus_t { unsafe {
692                match cusolver().and_then(|c| c.$potrs()) {
693                    Ok(f) => f(
694                        h,
695                        uplo,
696                        n,
697                        nrhs,
698                        a as *const $raw,
699                        lda,
700                        b as *mut $raw,
701                        ldb,
702                        info,
703                    ),
704                    Err(_) => cusolverStatus_t::NOT_INITIALIZED,
705                }
706            }}
707            unsafe fn gesvd_buf(
708                h: cusolverDnHandle_t,
709                m: c_int,
710                n: c_int,
711                lwork: *mut c_int,
712            ) -> cusolverStatus_t { unsafe {
713                match cusolver().and_then(|c| c.$gesvd_buf()) {
714                    Ok(f) => f(h, m, n, lwork),
715                    Err(_) => cusolverStatus_t::NOT_INITIALIZED,
716                }
717            }}
718            unsafe fn gesvd(
719                h: cusolverDnHandle_t,
720                jobu: u8,
721                jobvt: u8,
722                m: c_int,
723                n: c_int,
724                a: *mut $t,
725                lda: c_int,
726                s: *mut $real,
727                u: *mut $t,
728                ldu: c_int,
729                vt: *mut $t,
730                ldvt: c_int,
731                work: *mut $t,
732                lwork: c_int,
733                rwork: *mut $real,
734                info: *mut c_int,
735            ) -> cusolverStatus_t { unsafe {
736                match cusolver().and_then(|c| c.$gesvd()) {
737                    Ok(f) => f(
738                        h,
739                        jobu,
740                        jobvt,
741                        m,
742                        n,
743                        a as *mut $raw,
744                        lda,
745                        s,
746                        u as *mut $raw,
747                        ldu,
748                        vt as *mut $raw,
749                        ldvt,
750                        work as *mut $raw,
751                        lwork,
752                        rwork,
753                        info,
754                    ),
755                    Err(_) => cusolverStatus_t::NOT_INITIALIZED,
756                }
757            }}
758            unsafe fn syevd_buf(
759                h: cusolverDnHandle_t,
760                jobz: cusolverEigMode_t,
761                uplo: cublasFillMode_t,
762                n: c_int,
763                a: *const $t,
764                lda: c_int,
765                w: *const $real,
766                lwork: *mut c_int,
767            ) -> cusolverStatus_t { unsafe {
768                match cusolver().and_then(|c| c.$heevd_buf()) {
769                    Ok(f) => f(h, jobz, uplo, n, a as *const $raw, lda, w, lwork),
770                    Err(_) => cusolverStatus_t::NOT_INITIALIZED,
771                }
772            }}
773            unsafe fn syevd(
774                h: cusolverDnHandle_t,
775                jobz: cusolverEigMode_t,
776                uplo: cublasFillMode_t,
777                n: c_int,
778                a: *mut $t,
779                lda: c_int,
780                w: *mut $real,
781                work: *mut $t,
782                lwork: c_int,
783                info: *mut c_int,
784            ) -> cusolverStatus_t { unsafe {
785                match cusolver().and_then(|c| c.$heevd()) {
786                    Ok(f) => f(
787                        h,
788                        jobz,
789                        uplo,
790                        n,
791                        a as *mut $raw,
792                        lda,
793                        w,
794                        work as *mut $raw,
795                        lwork,
796                        info,
797                    ),
798                    Err(_) => cusolverStatus_t::NOT_INITIALIZED,
799                }
800            }}
801        }
802    };
803}
804
805real_impl!(
806    f32,
807    cusolver_dn_sgetrf_buffer_size,
808    cusolver_dn_sgetrf,
809    cusolver_dn_sgetrs,
810    cusolver_dn_sgeqrf_buffer_size,
811    cusolver_dn_sgeqrf,
812    cusolver_dn_spotrf_buffer_size,
813    cusolver_dn_spotrf,
814    cusolver_dn_spotrs,
815    cusolver_dn_sgesvd_buffer_size,
816    cusolver_dn_sgesvd,
817    cusolver_dn_ssyevd_buffer_size,
818    cusolver_dn_ssyevd
819);
820
821real_impl!(
822    f64,
823    cusolver_dn_dgetrf_buffer_size,
824    cusolver_dn_dgetrf,
825    cusolver_dn_dgetrs,
826    cusolver_dn_dgeqrf_buffer_size,
827    cusolver_dn_dgeqrf,
828    cusolver_dn_dpotrf_buffer_size,
829    cusolver_dn_dpotrf,
830    cusolver_dn_dpotrs,
831    cusolver_dn_dgesvd_buffer_size,
832    cusolver_dn_dgesvd,
833    cusolver_dn_dsyevd_buffer_size,
834    cusolver_dn_dsyevd
835);
836
837complex_impl!(
838    Complex32,
839    f32,
840    cuComplex,
841    cusolver_dn_cgetrf_buffer_size,
842    cusolver_dn_cgetrf,
843    cusolver_dn_cgetrs,
844    cusolver_dn_cgeqrf_buffer_size,
845    cusolver_dn_cgeqrf,
846    cusolver_dn_cpotrf_buffer_size,
847    cusolver_dn_cpotrf,
848    cusolver_dn_cpotrs,
849    cusolver_dn_cgesvd_buffer_size,
850    cusolver_dn_cgesvd,
851    cusolver_dn_cheevd_buffer_size,
852    cusolver_dn_cheevd
853);
854
855complex_impl!(
856    Complex64,
857    f64,
858    cuDoubleComplex,
859    cusolver_dn_zgetrf_buffer_size,
860    cusolver_dn_zgetrf,
861    cusolver_dn_zgetrs,
862    cusolver_dn_zgeqrf_buffer_size,
863    cusolver_dn_zgeqrf,
864    cusolver_dn_zpotrf_buffer_size,
865    cusolver_dn_zpotrf,
866    cusolver_dn_zpotrs,
867    cusolver_dn_zgesvd_buffer_size,
868    cusolver_dn_zgesvd,
869    cusolver_dn_zheevd_buffer_size,
870    cusolver_dn_zheevd
871);
872
873// ---- Public API ---------------------------------------------------------
874
875/// In-place LU factorization of a column-major matrix. Overwrites `a`.
876///
877/// On return `a` holds `L` (unit-diagonal lower) and `U` (upper) packed
878/// together; `ipiv` holds the row-pivot permutation (length = `min(m, n)`);
879/// `info` is a single-element device int (0 on success, k > 0 means U[k,k]
880/// is exactly zero).
881///
882/// # Example
883///
884/// LU-factorize a 4×4 matrix on the device. Pair with [`getrs`] to solve
885/// `A x = b`.
886///
887/// ```no_run
888/// use baracuda_driver::{Context, Device, DeviceBuffer};
889/// use baracuda_cusolver::{getrf, DnHandle};
890///
891/// # fn demo() -> Result<(), Box<dyn std::error::Error>> {
892/// let ctx = Context::new(&Device::get(0)?)?;
893/// let solver = DnHandle::new()?;
894///
895/// let n = 4i32;
896/// // Column-major 4×4 (caller initialises to a real, non-singular matrix).
897/// let mut a:    DeviceBuffer<f32> = DeviceBuffer::zeros(&ctx, (n * n) as usize)?;
898/// let mut ipiv: DeviceBuffer<i32> = DeviceBuffer::zeros(&ctx, n as usize)?;
899/// let mut info: DeviceBuffer<i32> = DeviceBuffer::zeros(&ctx, 1)?;
900///
901/// getrf::<f32>(&solver, n, n, &mut a, n, &mut ipiv, &mut info)?;
902/// # Ok(()) }
903/// ```
904#[allow(clippy::too_many_arguments)]
905pub fn getrf<T: SolverScalar>(
906    handle: &DnHandle,
907    m: i32,
908    n: i32,
909    a: &mut DeviceBuffer<T>,
910    lda: i32,
911    ipiv: &mut DeviceBuffer<i32>,
912    info: &mut DeviceBuffer<i32>,
913) -> Result<()> {
914    let mut lwork: c_int = 0;
915    check(unsafe { T::getrf_buf(handle.handle, m, n, a.as_raw().0 as *mut T, lda, &mut lwork) })?;
916    let workspace =
917        DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
918    check(unsafe {
919        T::getrf(
920            handle.handle,
921            m,
922            n,
923            a.as_raw().0 as *mut T,
924            lda,
925            workspace.as_raw().0 as *mut T,
926            ipiv.as_raw().0 as *mut c_int,
927            info.as_raw().0 as *mut c_int,
928        )
929    })
930}
931
932/// Solve `op(A) * X = B` using the LU factorization from [`getrf`].
933#[allow(clippy::too_many_arguments)]
934pub fn getrs<T: SolverScalar>(
935    handle: &DnHandle,
936    trans: Op,
937    n: i32,
938    nrhs: i32,
939    a: &DeviceBuffer<T>,
940    lda: i32,
941    ipiv: &DeviceBuffer<i32>,
942    b: &mut DeviceBuffer<T>,
943    ldb: i32,
944    info: &mut DeviceBuffer<i32>,
945) -> Result<()> {
946    check(unsafe {
947        T::getrs(
948            handle.handle,
949            trans.raw(),
950            n,
951            nrhs,
952            a.as_raw().0 as *const T,
953            lda,
954            ipiv.as_raw().0 as *const c_int,
955            b.as_raw().0 as *mut T,
956            ldb,
957            info.as_raw().0 as *mut c_int,
958        )
959    })
960}
961
962/// QR factorization: `A = Q * R`. Overwrites `a` (upper triangle = R,
963/// lower = Householder reflectors); `tau` receives reflector scalars.
964///
965/// To materialise `Q` explicitly, pair this with [`orgqr`]; to apply
966/// `Q` (or `Qᵀ`) to another matrix without forming it, use [`ormqr`].
967///
968/// # Example
969///
970/// QR-factorize a 6×4 tall-skinny matrix.
971///
972/// ```no_run
973/// use baracuda_driver::{Context, Device, DeviceBuffer};
974/// use baracuda_cusolver::{geqrf, DnHandle};
975///
976/// # fn demo() -> Result<(), Box<dyn std::error::Error>> {
977/// let ctx = Context::new(&Device::get(0)?)?;
978/// let solver = DnHandle::new()?;
979///
980/// let (m, n) = (6i32, 4i32);
981/// let mut a:    DeviceBuffer<f32> = DeviceBuffer::zeros(&ctx, (m * n) as usize)?;
982/// // `tau` is min(m, n) reflector scalars.
983/// let mut tau:  DeviceBuffer<f32> = DeviceBuffer::zeros(&ctx, n as usize)?;
984/// let mut info: DeviceBuffer<i32> = DeviceBuffer::zeros(&ctx, 1)?;
985///
986/// geqrf::<f32>(&solver, m, n, &mut a, m, &mut tau, &mut info)?;
987/// # Ok(()) }
988/// ```
989#[allow(clippy::too_many_arguments)]
990pub fn geqrf<T: SolverScalar>(
991    handle: &DnHandle,
992    m: i32,
993    n: i32,
994    a: &mut DeviceBuffer<T>,
995    lda: i32,
996    tau: &mut DeviceBuffer<T>,
997    info: &mut DeviceBuffer<i32>,
998) -> Result<()> {
999    let mut lwork: c_int = 0;
1000    check(unsafe { T::geqrf_buf(handle.handle, m, n, a.as_raw().0 as *mut T, lda, &mut lwork) })?;
1001    let workspace =
1002        DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1003    check(unsafe {
1004        T::geqrf(
1005            handle.handle,
1006            m,
1007            n,
1008            a.as_raw().0 as *mut T,
1009            lda,
1010            tau.as_raw().0 as *mut T,
1011            workspace.as_raw().0 as *mut T,
1012            lwork,
1013            info.as_raw().0 as *mut c_int,
1014        )
1015    })
1016}
1017
1018/// Cholesky factorization: `A = L * Lᵀ` (or `Uᵀ * U`). Overwrites `a`.
1019///
1020/// Requires `A` to be symmetric / Hermitian and positive-definite. `uplo`
1021/// selects which triangle of `a` holds the input (and receives the factor);
1022/// the other triangle is read- and write- ignored.
1023///
1024/// # Example
1025///
1026/// Cholesky of a 5×5 SPD matrix, storing `L` in the lower triangle.
1027///
1028/// ```no_run
1029/// use baracuda_driver::{Context, Device, DeviceBuffer};
1030/// use baracuda_cusolver::{potrf, DnHandle, Fill};
1031///
1032/// # fn demo() -> Result<(), Box<dyn std::error::Error>> {
1033/// let ctx = Context::new(&Device::get(0)?)?;
1034/// let solver = DnHandle::new()?;
1035///
1036/// let n = 5i32;
1037/// let mut a:    DeviceBuffer<f32> = DeviceBuffer::zeros(&ctx, (n * n) as usize)?;
1038/// let mut info: DeviceBuffer<i32> = DeviceBuffer::zeros(&ctx, 1)?;
1039///
1040/// potrf::<f32>(&solver, Fill::Lower, n, &mut a, n, &mut info)?;
1041/// // After return, lower triangle of `a` holds L such that A = L Lᵀ.
1042/// # Ok(()) }
1043/// ```
1044pub fn potrf<T: SolverScalar>(
1045    handle: &DnHandle,
1046    uplo: Fill,
1047    n: i32,
1048    a: &mut DeviceBuffer<T>,
1049    lda: i32,
1050    info: &mut DeviceBuffer<i32>,
1051) -> Result<()> {
1052    let mut lwork: c_int = 0;
1053    check(unsafe { T::potrf_buf(handle.handle, uplo, n, a.as_raw().0 as *mut T, lda, &mut lwork) })?;
1054    let workspace =
1055        DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1056    check(unsafe {
1057        T::potrf(
1058            handle.handle,
1059            uplo,
1060            n,
1061            a.as_raw().0 as *mut T,
1062            lda,
1063            workspace.as_raw().0 as *mut T,
1064            lwork,
1065            info.as_raw().0 as *mut c_int,
1066        )
1067    })
1068}
1069
1070/// Solve `A * X = B` using the Cholesky factorization from [`potrf`].
1071#[allow(clippy::too_many_arguments)]
1072pub fn potrs<T: SolverScalar>(
1073    handle: &DnHandle,
1074    uplo: Fill,
1075    n: i32,
1076    nrhs: i32,
1077    a: &DeviceBuffer<T>,
1078    lda: i32,
1079    b: &mut DeviceBuffer<T>,
1080    ldb: i32,
1081    info: &mut DeviceBuffer<i32>,
1082) -> Result<()> {
1083    check(unsafe {
1084        T::potrs(
1085            handle.handle,
1086            uplo,
1087            n,
1088            nrhs,
1089            a.as_raw().0 as *const T,
1090            lda,
1091            b.as_raw().0 as *mut T,
1092            ldb,
1093            info.as_raw().0 as *mut c_int,
1094        )
1095    })
1096}
1097
1098/// Full SVD: `A = U * Σ * Vᵀ`. `jobu`/`jobvt` are LAPACK-style single-byte
1099/// selectors (b'A' = all, b'S' = economy, b'N' = none, b'O' = overwrite A).
1100///
1101/// `rwork` must be provided for complex element types; pass an empty buffer
1102/// for real types (pointer is still non-null; cuSOLVER ignores it).
1103///
1104/// # Example
1105///
1106/// Full SVD of a 6×4 real matrix (`U` 6×6, `Σ` length-4, `Vᵀ` 4×4).
1107///
1108/// ```no_run
1109/// use baracuda_driver::{Context, Device, DeviceBuffer};
1110/// use baracuda_cusolver::{gesvd, DnHandle};
1111///
1112/// # fn demo() -> Result<(), Box<dyn std::error::Error>> {
1113/// let ctx = Context::new(&Device::get(0)?)?;
1114/// let solver = DnHandle::new()?;
1115///
1116/// let (m, n) = (6i32, 4i32);
1117/// let mut a:    DeviceBuffer<f32> = DeviceBuffer::zeros(&ctx, (m * n) as usize)?;
1118/// let mut s:    DeviceBuffer<f32> = DeviceBuffer::zeros(&ctx, n as usize)?;
1119/// let mut u:    DeviceBuffer<f32> = DeviceBuffer::zeros(&ctx, (m * m) as usize)?;
1120/// let mut vt:   DeviceBuffer<f32> = DeviceBuffer::zeros(&ctx, (n * n) as usize)?;
1121/// // Real-type SVD ignores `rwork`; allocate an empty buffer.
1122/// let mut rwork: DeviceBuffer<f32> = DeviceBuffer::zeros(&ctx, 1)?;
1123/// let mut info:  DeviceBuffer<i32> = DeviceBuffer::zeros(&ctx, 1)?;
1124///
1125/// gesvd::<f32>(
1126///     &solver, b'A', b'A', m, n,
1127///     &mut a, m,
1128///     &mut s,
1129///     &mut u, m,
1130///     &mut vt, n,
1131///     &mut rwork,
1132///     &mut info,
1133/// )?;
1134/// # Ok(()) }
1135/// ```
1136#[allow(clippy::too_many_arguments)]
1137pub fn gesvd<T: SolverScalar>(
1138    handle: &DnHandle,
1139    jobu: u8,
1140    jobvt: u8,
1141    m: i32,
1142    n: i32,
1143    a: &mut DeviceBuffer<T>,
1144    lda: i32,
1145    s: &mut DeviceBuffer<T::Real>,
1146    u: &mut DeviceBuffer<T>,
1147    ldu: i32,
1148    vt: &mut DeviceBuffer<T>,
1149    ldvt: i32,
1150    rwork: &mut DeviceBuffer<T::Real>,
1151    info: &mut DeviceBuffer<i32>,
1152) -> Result<()> {
1153    let mut lwork: c_int = 0;
1154    check(unsafe { T::gesvd_buf(handle.handle, m, n, &mut lwork) })?;
1155    let workspace =
1156        DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1157    check(unsafe {
1158        T::gesvd(
1159            handle.handle,
1160            jobu,
1161            jobvt,
1162            m,
1163            n,
1164            a.as_raw().0 as *mut T,
1165            lda,
1166            s.as_raw().0 as *mut T::Real,
1167            u.as_raw().0 as *mut T,
1168            ldu,
1169            vt.as_raw().0 as *mut T,
1170            ldvt,
1171            workspace.as_raw().0 as *mut T,
1172            lwork,
1173            rwork.as_raw().0 as *mut T::Real,
1174            info.as_raw().0 as *mut c_int,
1175        )
1176    })
1177}
1178
1179/// Symmetric / Hermitian eigenvalue decomposition: `A = Q * diag(w) * Qᵀ`.
1180#[allow(clippy::too_many_arguments)]
1181pub fn syevd<T: SolverScalar>(
1182    handle: &DnHandle,
1183    jobz: EigMode,
1184    uplo: Fill,
1185    n: i32,
1186    a: &mut DeviceBuffer<T>,
1187    lda: i32,
1188    w: &mut DeviceBuffer<T::Real>,
1189    info: &mut DeviceBuffer<i32>,
1190) -> Result<()> {
1191    let mut lwork: c_int = 0;
1192    check(unsafe {
1193        T::syevd_buf(
1194            handle.handle,
1195            jobz,
1196            uplo,
1197            n,
1198            a.as_raw().0 as *const T,
1199            lda,
1200            w.as_raw().0 as *const T::Real,
1201            &mut lwork,
1202        )
1203    })?;
1204    let workspace =
1205        DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1206    check(unsafe {
1207        T::syevd(
1208            handle.handle,
1209            jobz,
1210            uplo,
1211            n,
1212            a.as_raw().0 as *mut T,
1213            lda,
1214            w.as_raw().0 as *mut T::Real,
1215            workspace.as_raw().0 as *mut T,
1216            lwork,
1217            info.as_raw().0 as *mut c_int,
1218        )
1219    })
1220}
1221
1222// ---- Jacobi-based solvers (syevj / gesvdj) -----------------------------
1223
1224pub use baracuda_cusolver_sys::{gesvdjInfo_t as GesvdjInfoRaw, syevjInfo_t as SyevjInfoRaw};
1225
1226/// Jacobi-eigen tuning handle (tolerance + max sweeps).
1227#[derive(Debug)]
1228pub struct SyevjInfo {
1229    raw: SyevjInfoRaw,
1230}
1231
1232impl SyevjInfo {
1233    /// Create a tuning-info handle (`cusolverDnCreateSyevjInfo`).
1234    pub fn new() -> Result<Self> {
1235        let c = cusolver()?;
1236        let cu = c.cusolver_dn_create_syevj_info()?;
1237        let mut raw: SyevjInfoRaw = core::ptr::null_mut();
1238        check(unsafe { cu(&mut raw) })?;
1239        Ok(Self { raw })
1240    }
1241
1242    /// Set the per-sweep convergence tolerance
1243    /// (`cusolverDnXsyevjSetTolerance`).
1244    pub fn set_tolerance(&self, tol: f64) -> Result<()> {
1245        let c = cusolver()?;
1246        let cu = c.cusolver_dn_xsyevj_set_tolerance()?;
1247        check(unsafe { cu(self.raw, tol) })
1248    }
1249
1250    /// Cap the number of Jacobi sweeps
1251    /// (`cusolverDnXsyevjSetMaxSweeps`).
1252    pub fn set_max_sweeps(&self, n: i32) -> Result<()> {
1253        let c = cusolver()?;
1254        let cu = c.cusolver_dn_xsyevj_set_max_sweeps()?;
1255        check(unsafe { cu(self.raw, n) })
1256    }
1257
1258    /// Raw `syevjInfo_t`. Use with care.
1259    pub fn as_raw(&self) -> SyevjInfoRaw {
1260        self.raw
1261    }
1262}
1263
1264impl Drop for SyevjInfo {
1265    fn drop(&mut self) {
1266        if let Ok(c) = cusolver() {
1267            if let Ok(cu) = c.cusolver_dn_destroy_syevj_info() {
1268                let _ = unsafe { cu(self.raw) };
1269            }
1270        }
1271    }
1272}
1273
1274/// Jacobi-SVD tuning handle.
1275#[derive(Debug)]
1276pub struct GesvdjInfo {
1277    raw: GesvdjInfoRaw,
1278}
1279
1280impl GesvdjInfo {
1281    /// Create a tuning-info handle (`cusolverDnCreateGesvdjInfo`).
1282    pub fn new() -> Result<Self> {
1283        let c = cusolver()?;
1284        let cu = c.cusolver_dn_create_gesvdj_info()?;
1285        let mut raw: GesvdjInfoRaw = core::ptr::null_mut();
1286        check(unsafe { cu(&mut raw) })?;
1287        Ok(Self { raw })
1288    }
1289
1290    /// Raw `gesvdjInfo_t`. Use with care.
1291    pub fn as_raw(&self) -> GesvdjInfoRaw {
1292        self.raw
1293    }
1294}
1295
1296impl Drop for GesvdjInfo {
1297    fn drop(&mut self) {
1298        if let Ok(c) = cusolver() {
1299            if let Ok(cu) = c.cusolver_dn_destroy_gesvdj_info() {
1300                let _ = unsafe { cu(self.raw) };
1301            }
1302        }
1303    }
1304}
1305
1306/// Jacobi symmetric/Hermitian eigendecomposition (smaller matrices than
1307/// [`syevd`], faster convergence on well-conditioned problems).
1308#[allow(clippy::too_many_arguments)]
1309pub fn syevj<T: SolverScalar>(
1310    handle: &DnHandle,
1311    jobz: EigMode,
1312    uplo: Fill,
1313    n: i32,
1314    a: &mut DeviceBuffer<T>,
1315    lda: i32,
1316    w: &mut DeviceBuffer<T::Real>,
1317    info: &mut DeviceBuffer<i32>,
1318    params: &SyevjInfo,
1319) -> Result<()> {
1320    use baracuda_cusolver_sys::{
1321        cuComplex, cuDoubleComplex,
1322    };
1323    use core::mem;
1324
1325    let mut lwork: c_int = 0;
1326
1327    // Dispatch is simpler done via a type check, since syevj doesn't share
1328    // the generic trait shape (extra params: SyevjInfo).
1329    macro_rules! dispatch_real {
1330        ($t:ty, $bufsize:ident, $solve:ident) => {{
1331            let c = cusolver()?;
1332            check(unsafe {
1333                (c.$bufsize()?)(
1334                    handle.as_raw(),
1335                    jobz,
1336                    uplo,
1337                    n,
1338                    a.as_raw().0 as *const $t,
1339                    lda,
1340                    w.as_raw().0 as *const $t,
1341                    &mut lwork,
1342                    params.raw,
1343                )
1344            })?;
1345            let workspace =
1346                DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1347            check(unsafe {
1348                (c.$solve()?)(
1349                    handle.as_raw(),
1350                    jobz,
1351                    uplo,
1352                    n,
1353                    a.as_raw().0 as *mut $t,
1354                    lda,
1355                    w.as_raw().0 as *mut $t,
1356                    workspace.as_raw().0 as *mut $t,
1357                    lwork,
1358                    info.as_raw().0 as *mut c_int,
1359                    params.raw,
1360                )
1361            })
1362        }};
1363    }
1364    macro_rules! dispatch_complex {
1365        ($t:ty, $real:ty, $raw:ty, $bufsize:ident, $solve:ident) => {{
1366            let c = cusolver()?;
1367            check(unsafe {
1368                (c.$bufsize()?)(
1369                    handle.as_raw(),
1370                    jobz,
1371                    uplo,
1372                    n,
1373                    a.as_raw().0 as *const $raw,
1374                    lda,
1375                    w.as_raw().0 as *const $real,
1376                    &mut lwork,
1377                    params.raw,
1378                )
1379            })?;
1380            let workspace =
1381                DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1382            check(unsafe {
1383                (c.$solve()?)(
1384                    handle.as_raw(),
1385                    jobz,
1386                    uplo,
1387                    n,
1388                    a.as_raw().0 as *mut $raw,
1389                    lda,
1390                    w.as_raw().0 as *mut $real,
1391                    workspace.as_raw().0 as *mut $raw,
1392                    lwork,
1393                    info.as_raw().0 as *mut c_int,
1394                    params.raw,
1395                )
1396            })
1397        }};
1398    }
1399
1400    if mem::size_of::<T>() == mem::size_of::<f32>() && mem::size_of::<T::Real>() == 4 {
1401        dispatch_real!(f32, cusolver_dn_ssyevj_buffer_size, cusolver_dn_ssyevj)
1402    } else if mem::size_of::<T>() == mem::size_of::<f64>() && mem::size_of::<T::Real>() == 8 {
1403        dispatch_real!(f64, cusolver_dn_dsyevj_buffer_size, cusolver_dn_dsyevj)
1404    } else if mem::size_of::<T>() == mem::size_of::<Complex32>() {
1405        dispatch_complex!(
1406            Complex32,
1407            f32,
1408            cuComplex,
1409            cusolver_dn_cheevj_buffer_size,
1410            cusolver_dn_cheevj
1411        )
1412    } else {
1413        dispatch_complex!(
1414            Complex64,
1415            f64,
1416            cuDoubleComplex,
1417            cusolver_dn_zheevj_buffer_size,
1418            cusolver_dn_zheevj
1419        )
1420    }
1421}
1422
1423/// Jacobi SVD: `A = U * diag(s) * Vᴴ`. `econ` selects thin-SVD when set.
1424#[allow(clippy::too_many_arguments)]
1425pub fn gesvdj<T: SolverScalar>(
1426    handle: &DnHandle,
1427    jobz: EigMode,
1428    econ: bool,
1429    m: i32,
1430    n: i32,
1431    a: &mut DeviceBuffer<T>,
1432    lda: i32,
1433    s: &mut DeviceBuffer<T::Real>,
1434    u: &mut DeviceBuffer<T>,
1435    ldu: i32,
1436    v: &mut DeviceBuffer<T>,
1437    ldv: i32,
1438    info: &mut DeviceBuffer<i32>,
1439    params: &GesvdjInfo,
1440) -> Result<()> {
1441    use baracuda_cusolver_sys::{cuComplex, cuDoubleComplex};
1442    use core::mem;
1443
1444    let mut lwork: c_int = 0;
1445    let econ_i = if econ { 1 } else { 0 };
1446
1447    macro_rules! dispatch_real {
1448        ($t:ty, $bufsize:ident, $solve:ident) => {{
1449            let c = cusolver()?;
1450            check(unsafe {
1451                (c.$bufsize()?)(
1452                    handle.as_raw(),
1453                    jobz,
1454                    econ_i,
1455                    m,
1456                    n,
1457                    a.as_raw().0 as *const $t,
1458                    lda,
1459                    s.as_raw().0 as *const $t,
1460                    u.as_raw().0 as *const $t,
1461                    ldu,
1462                    v.as_raw().0 as *const $t,
1463                    ldv,
1464                    &mut lwork,
1465                    params.raw,
1466                )
1467            })?;
1468            let workspace =
1469                DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1470            check(unsafe {
1471                (c.$solve()?)(
1472                    handle.as_raw(),
1473                    jobz,
1474                    econ_i,
1475                    m,
1476                    n,
1477                    a.as_raw().0 as *mut $t,
1478                    lda,
1479                    s.as_raw().0 as *mut $t,
1480                    u.as_raw().0 as *mut $t,
1481                    ldu,
1482                    v.as_raw().0 as *mut $t,
1483                    ldv,
1484                    workspace.as_raw().0 as *mut $t,
1485                    lwork,
1486                    info.as_raw().0 as *mut c_int,
1487                    params.raw,
1488                )
1489            })
1490        }};
1491    }
1492    macro_rules! dispatch_complex {
1493        ($t:ty, $real:ty, $raw:ty, $bufsize:ident, $solve:ident) => {{
1494            let c = cusolver()?;
1495            check(unsafe {
1496                (c.$bufsize()?)(
1497                    handle.as_raw(),
1498                    jobz,
1499                    econ_i,
1500                    m,
1501                    n,
1502                    a.as_raw().0 as *const $raw,
1503                    lda,
1504                    s.as_raw().0 as *const $real,
1505                    u.as_raw().0 as *const $raw,
1506                    ldu,
1507                    v.as_raw().0 as *const $raw,
1508                    ldv,
1509                    &mut lwork,
1510                    params.raw,
1511                )
1512            })?;
1513            let workspace =
1514                DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1515            check(unsafe {
1516                (c.$solve()?)(
1517                    handle.as_raw(),
1518                    jobz,
1519                    econ_i,
1520                    m,
1521                    n,
1522                    a.as_raw().0 as *mut $raw,
1523                    lda,
1524                    s.as_raw().0 as *mut $real,
1525                    u.as_raw().0 as *mut $raw,
1526                    ldu,
1527                    v.as_raw().0 as *mut $raw,
1528                    ldv,
1529                    workspace.as_raw().0 as *mut $raw,
1530                    lwork,
1531                    info.as_raw().0 as *mut c_int,
1532                    params.raw,
1533                )
1534            })
1535        }};
1536    }
1537
1538    if mem::size_of::<T>() == mem::size_of::<f32>() && mem::size_of::<T::Real>() == 4 {
1539        dispatch_real!(f32, cusolver_dn_sgesvdj_buffer_size, cusolver_dn_sgesvdj)
1540    } else if mem::size_of::<T>() == mem::size_of::<f64>() && mem::size_of::<T::Real>() == 8 {
1541        dispatch_real!(f64, cusolver_dn_dgesvdj_buffer_size, cusolver_dn_dgesvdj)
1542    } else if mem::size_of::<T>() == mem::size_of::<Complex32>() {
1543        dispatch_complex!(
1544            Complex32,
1545            f32,
1546            cuComplex,
1547            cusolver_dn_cgesvdj_buffer_size,
1548            cusolver_dn_cgesvdj
1549        )
1550    } else {
1551        dispatch_complex!(
1552            Complex64,
1553            f64,
1554            cuDoubleComplex,
1555            cusolver_dn_zgesvdj_buffer_size,
1556            cusolver_dn_zgesvdj
1557        )
1558    }
1559}
1560
1561// ---- Generate / apply Q from QR (orgqr / ormqr) -------------------------
1562
1563/// Generate the orthogonal matrix `Q` from the factorization produced by
1564/// [`geqrf`]. After this, `a` holds the first `n` columns of `Q`.
1565#[allow(clippy::too_many_arguments)]
1566pub fn orgqr<T: SolverScalar>(
1567    handle: &DnHandle,
1568    m: i32,
1569    n: i32,
1570    k: i32,
1571    a: &mut DeviceBuffer<T>,
1572    lda: i32,
1573    tau: &DeviceBuffer<T>,
1574    info: &mut DeviceBuffer<i32>,
1575) -> Result<()> {
1576    use baracuda_cusolver_sys::{cuComplex, cuDoubleComplex};
1577    use core::mem;
1578
1579    let mut lwork: c_int = 0;
1580    macro_rules! dispatch {
1581        ($t:ty, $raw:ty, $bufsize:ident, $solve:ident) => {{
1582            let c = cusolver()?;
1583            check(unsafe {
1584                (c.$bufsize()?)(
1585                    handle.as_raw(),
1586                    m,
1587                    n,
1588                    k,
1589                    a.as_raw().0 as *const $raw,
1590                    lda,
1591                    tau.as_raw().0 as *const $raw,
1592                    &mut lwork,
1593                )
1594            })?;
1595            let workspace =
1596                DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1597            check(unsafe {
1598                (c.$solve()?)(
1599                    handle.as_raw(),
1600                    m,
1601                    n,
1602                    k,
1603                    a.as_raw().0 as *mut $raw,
1604                    lda,
1605                    tau.as_raw().0 as *const $raw,
1606                    workspace.as_raw().0 as *mut $raw,
1607                    lwork,
1608                    info.as_raw().0 as *mut c_int,
1609                )
1610            })
1611        }};
1612    }
1613
1614    if mem::size_of::<T>() == mem::size_of::<f32>() && mem::size_of::<T::Real>() == 4 {
1615        dispatch!(f32, f32, cusolver_dn_sorgqr_buffer_size, cusolver_dn_sorgqr)
1616    } else if mem::size_of::<T>() == mem::size_of::<f64>() && mem::size_of::<T::Real>() == 8 {
1617        dispatch!(f64, f64, cusolver_dn_dorgqr_buffer_size, cusolver_dn_dorgqr)
1618    } else if mem::size_of::<T>() == mem::size_of::<Complex32>() {
1619        dispatch!(
1620            Complex32,
1621            cuComplex,
1622            cusolver_dn_cungqr_buffer_size,
1623            cusolver_dn_cungqr
1624        )
1625    } else {
1626        dispatch!(
1627            Complex64,
1628            cuDoubleComplex,
1629            cusolver_dn_zungqr_buffer_size,
1630            cusolver_dn_zungqr
1631        )
1632    }
1633}
1634
1635/// Side argument for [`ormqr`].
1636#[derive(Copy, Clone, Debug, Eq, PartialEq)]
1637pub enum Side {
1638    /// Apply `Q` from the left: `C = op(Q) * C`.
1639    Left,
1640    /// Apply `Q` from the right: `C = C * op(Q)`.
1641    Right,
1642}
1643
1644impl Side {
1645    fn raw(self) -> core::ffi::c_int {
1646        match self {
1647            Side::Left => 0,
1648            Side::Right => 1,
1649        }
1650    }
1651}
1652
1653/// Apply `op(Q)` to `C`: `C = op(Q) * C` (Left) or `C = C * op(Q)` (Right),
1654/// where `Q` is packed in `a`+`tau` from [`geqrf`].
1655#[allow(clippy::too_many_arguments)]
1656pub fn ormqr<T: SolverScalar>(
1657    handle: &DnHandle,
1658    side: Side,
1659    trans: Op,
1660    m: i32,
1661    n: i32,
1662    k: i32,
1663    a: &DeviceBuffer<T>,
1664    lda: i32,
1665    tau: &DeviceBuffer<T>,
1666    c_mat: &mut DeviceBuffer<T>,
1667    ldc: i32,
1668    info: &mut DeviceBuffer<i32>,
1669) -> Result<()> {
1670    use baracuda_cusolver_sys::{cuComplex, cuDoubleComplex};
1671    use core::mem;
1672
1673    let mut lwork: c_int = 0;
1674    let side_i = side.raw();
1675    macro_rules! dispatch {
1676        ($t:ty, $raw:ty, $bufsize:ident, $solve:ident) => {{
1677            let ca = cusolver()?;
1678            check(unsafe {
1679                (ca.$bufsize()?)(
1680                    handle.as_raw(),
1681                    side_i,
1682                    trans.raw(),
1683                    m,
1684                    n,
1685                    k,
1686                    a.as_raw().0 as *const $raw,
1687                    lda,
1688                    tau.as_raw().0 as *const $raw,
1689                    c_mat.as_raw().0 as *const $raw,
1690                    ldc,
1691                    &mut lwork,
1692                )
1693            })?;
1694            let workspace =
1695                DeviceBuffer::<T>::new(c_mat.context(), lwork as usize).map_err(alloc_fail)?;
1696            check(unsafe {
1697                (ca.$solve()?)(
1698                    handle.as_raw(),
1699                    side_i,
1700                    trans.raw(),
1701                    m,
1702                    n,
1703                    k,
1704                    a.as_raw().0 as *const $raw,
1705                    lda,
1706                    tau.as_raw().0 as *const $raw,
1707                    c_mat.as_raw().0 as *mut $raw,
1708                    ldc,
1709                    workspace.as_raw().0 as *mut $raw,
1710                    lwork,
1711                    info.as_raw().0 as *mut c_int,
1712                )
1713            })
1714        }};
1715    }
1716
1717    if mem::size_of::<T>() == mem::size_of::<f32>() && mem::size_of::<T::Real>() == 4 {
1718        dispatch!(f32, f32, cusolver_dn_sormqr_buffer_size, cusolver_dn_sormqr)
1719    } else if mem::size_of::<T>() == mem::size_of::<f64>() && mem::size_of::<T::Real>() == 8 {
1720        dispatch!(f64, f64, cusolver_dn_dormqr_buffer_size, cusolver_dn_dormqr)
1721    } else if mem::size_of::<T>() == mem::size_of::<Complex32>() {
1722        dispatch!(
1723            Complex32,
1724            cuComplex,
1725            cusolver_dn_cunmqr_buffer_size,
1726            cusolver_dn_cunmqr
1727        )
1728    } else {
1729        dispatch!(
1730            Complex64,
1731            cuDoubleComplex,
1732            cusolver_dn_zunmqr_buffer_size,
1733            cusolver_dn_zunmqr
1734        )
1735    }
1736}
1737
1738// ---- gels: iterative-refinement least-squares solve ---------------------
1739
1740/// Solve `A * X = B` in the least-squares sense (iterative-refinement).
1741/// `A` is `m × n`, `B` is `m × nrhs`, `X` is `n × nrhs`. `A` and `B` may be
1742/// overwritten. Returns `iter`: number of refinement iterations used (-1 =
1743/// fallback to full precision).
1744#[allow(clippy::too_many_arguments)]
1745pub fn gels<T: SolverScalar>(
1746    handle: &DnHandle,
1747    m: i32,
1748    n: i32,
1749    nrhs: i32,
1750    a: &mut DeviceBuffer<T>,
1751    lda: i32,
1752    b: &mut DeviceBuffer<T>,
1753    ldb: i32,
1754    x: &mut DeviceBuffer<T>,
1755    ldx: i32,
1756    info: &mut DeviceBuffer<i32>,
1757) -> Result<i32> {
1758    use baracuda_cusolver_sys::{cuComplex, cuDoubleComplex};
1759    use core::mem;
1760
1761    let mut bytes: usize = 0;
1762
1763    macro_rules! dispatch {
1764        ($t:ty, $raw:ty, $bufsize:ident, $solve:ident) => {{
1765            let cs = cusolver()?;
1766            check(unsafe {
1767                (cs.$bufsize()?)(
1768                    handle.as_raw(),
1769                    m,
1770                    n,
1771                    nrhs,
1772                    a.as_raw().0 as *mut $raw,
1773                    lda,
1774                    b.as_raw().0 as *mut $raw,
1775                    ldb,
1776                    x.as_raw().0 as *mut $raw,
1777                    ldx,
1778                    core::ptr::null_mut(),
1779                    &mut bytes,
1780                )
1781            })?;
1782            // Allocate `bytes` worth of u8 workspace (rounding up to T units).
1783            let units = bytes.div_ceil(mem::size_of::<T>());
1784            let workspace =
1785                DeviceBuffer::<T>::new(a.context(), units).map_err(alloc_fail)?;
1786            let mut iter: c_int = 0;
1787            check(unsafe {
1788                (cs.$solve()?)(
1789                    handle.as_raw(),
1790                    m,
1791                    n,
1792                    nrhs,
1793                    a.as_raw().0 as *mut $raw,
1794                    lda,
1795                    b.as_raw().0 as *mut $raw,
1796                    ldb,
1797                    x.as_raw().0 as *mut $raw,
1798                    ldx,
1799                    workspace.as_raw().0 as *mut c_void,
1800                    bytes,
1801                    &mut iter,
1802                    info.as_raw().0 as *mut c_int,
1803                )
1804            })?;
1805            Ok(iter)
1806        }};
1807    }
1808
1809    if mem::size_of::<T>() == mem::size_of::<f32>() && mem::size_of::<T::Real>() == 4 {
1810        dispatch!(f32, f32, cusolver_dn_ssgels_buffer_size, cusolver_dn_ssgels)
1811    } else if mem::size_of::<T>() == mem::size_of::<f64>() && mem::size_of::<T::Real>() == 8 {
1812        dispatch!(f64, f64, cusolver_dn_ddgels_buffer_size, cusolver_dn_ddgels)
1813    } else if mem::size_of::<T>() == mem::size_of::<Complex32>() {
1814        dispatch!(
1815            Complex32,
1816            cuComplex,
1817            cusolver_dn_ccgels_buffer_size,
1818            cusolver_dn_ccgels
1819        )
1820    } else {
1821        dispatch!(
1822            Complex64,
1823            cuDoubleComplex,
1824            cusolver_dn_zzgels_buffer_size,
1825            cusolver_dn_zzgels
1826        )
1827    }
1828}
1829
1830// ---- potri: inverse from Cholesky factor --------------------------------
1831
1832/// Compute `A = (Lᵀ * L)⁻¹` or `A = (U * Uᵀ)⁻¹` given the Cholesky factor
1833/// already stored in the triangle selected by `uplo`. `a` must hold the
1834/// output of [`potrf`] in-place.
1835pub fn potri<T: SolverScalar>(
1836    handle: &DnHandle,
1837    uplo: Fill,
1838    n: i32,
1839    a: &mut DeviceBuffer<T>,
1840    lda: i32,
1841    info: &mut DeviceBuffer<i32>,
1842) -> Result<()> {
1843    use baracuda_cusolver_sys::{cuComplex, cuDoubleComplex};
1844    use core::mem;
1845
1846    let mut lwork: c_int = 0;
1847    macro_rules! dispatch {
1848        ($t:ty, $raw:ty, $bufsize:ident, $solve:ident) => {{
1849            let cs = cusolver()?;
1850            check(unsafe {
1851                (cs.$bufsize()?)(
1852                    handle.as_raw(),
1853                    uplo,
1854                    n,
1855                    a.as_raw().0 as *mut $raw,
1856                    lda,
1857                    &mut lwork,
1858                )
1859            })?;
1860            let workspace =
1861                DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1862            check(unsafe {
1863                (cs.$solve()?)(
1864                    handle.as_raw(),
1865                    uplo,
1866                    n,
1867                    a.as_raw().0 as *mut $raw,
1868                    lda,
1869                    workspace.as_raw().0 as *mut $raw,
1870                    lwork,
1871                    info.as_raw().0 as *mut c_int,
1872                )
1873            })
1874        }};
1875    }
1876
1877    if mem::size_of::<T>() == mem::size_of::<f32>() && mem::size_of::<T::Real>() == 4 {
1878        dispatch!(f32, f32, cusolver_dn_spotri_buffer_size, cusolver_dn_spotri)
1879    } else if mem::size_of::<T>() == mem::size_of::<f64>() && mem::size_of::<T::Real>() == 8 {
1880        dispatch!(f64, f64, cusolver_dn_dpotri_buffer_size, cusolver_dn_dpotri)
1881    } else if mem::size_of::<T>() == mem::size_of::<Complex32>() {
1882        dispatch!(
1883            Complex32,
1884            cuComplex,
1885            cusolver_dn_cpotri_buffer_size,
1886            cusolver_dn_cpotri
1887        )
1888    } else {
1889        dispatch!(
1890            Complex64,
1891            cuDoubleComplex,
1892            cusolver_dn_zpotri_buffer_size,
1893            cusolver_dn_zpotri
1894        )
1895    }
1896}
1897
1898// ---- Batched Jacobi eigen / SVD -----------------------------------------
1899
1900/// Batched Jacobi symmetric/Hermitian eigendecomposition. Every matrix in
1901/// the batch is `n × n` and stride `n × n`. `w` holds `n * batch_size`
1902/// eigenvalues, strided by `n`.
1903#[allow(clippy::too_many_arguments)]
1904pub fn syevj_batched<T: SolverScalar>(
1905    handle: &DnHandle,
1906    jobz: EigMode,
1907    uplo: Fill,
1908    n: i32,
1909    a: &mut DeviceBuffer<T>,
1910    lda: i32,
1911    w: &mut DeviceBuffer<T::Real>,
1912    info: &mut DeviceBuffer<i32>,
1913    params: &SyevjInfo,
1914    batch_size: i32,
1915) -> Result<()> {
1916    use baracuda_cusolver_sys::{cuComplex, cuDoubleComplex};
1917    use core::mem;
1918
1919    let mut lwork: c_int = 0;
1920    macro_rules! dispatch_real {
1921        ($t:ty, $bufsize:ident, $solve:ident) => {{
1922            let c = cusolver()?;
1923            check(unsafe {
1924                (c.$bufsize()?)(
1925                    handle.as_raw(),
1926                    jobz,
1927                    uplo,
1928                    n,
1929                    a.as_raw().0 as *const $t,
1930                    lda,
1931                    w.as_raw().0 as *const $t,
1932                    &mut lwork,
1933                    params.as_raw(),
1934                    batch_size,
1935                )
1936            })?;
1937            let workspace =
1938                DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1939            check(unsafe {
1940                (c.$solve()?)(
1941                    handle.as_raw(),
1942                    jobz,
1943                    uplo,
1944                    n,
1945                    a.as_raw().0 as *mut $t,
1946                    lda,
1947                    w.as_raw().0 as *mut $t,
1948                    workspace.as_raw().0 as *mut $t,
1949                    lwork,
1950                    info.as_raw().0 as *mut c_int,
1951                    params.as_raw(),
1952                    batch_size,
1953                )
1954            })
1955        }};
1956    }
1957    macro_rules! dispatch_complex {
1958        ($t:ty, $real:ty, $raw:ty, $bufsize:ident, $solve:ident) => {{
1959            let c = cusolver()?;
1960            check(unsafe {
1961                (c.$bufsize()?)(
1962                    handle.as_raw(),
1963                    jobz,
1964                    uplo,
1965                    n,
1966                    a.as_raw().0 as *const $raw,
1967                    lda,
1968                    w.as_raw().0 as *const $real,
1969                    &mut lwork,
1970                    params.as_raw(),
1971                    batch_size,
1972                )
1973            })?;
1974            let workspace =
1975                DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
1976            check(unsafe {
1977                (c.$solve()?)(
1978                    handle.as_raw(),
1979                    jobz,
1980                    uplo,
1981                    n,
1982                    a.as_raw().0 as *mut $raw,
1983                    lda,
1984                    w.as_raw().0 as *mut $real,
1985                    workspace.as_raw().0 as *mut $raw,
1986                    lwork,
1987                    info.as_raw().0 as *mut c_int,
1988                    params.as_raw(),
1989                    batch_size,
1990                )
1991            })
1992        }};
1993    }
1994
1995    if mem::size_of::<T>() == mem::size_of::<f32>() && mem::size_of::<T::Real>() == 4 {
1996        dispatch_real!(
1997            f32,
1998            cusolver_dn_ssyevj_batched_buffer_size,
1999            cusolver_dn_ssyevj_batched
2000        )
2001    } else if mem::size_of::<T>() == mem::size_of::<f64>() && mem::size_of::<T::Real>() == 8 {
2002        dispatch_real!(
2003            f64,
2004            cusolver_dn_dsyevj_batched_buffer_size,
2005            cusolver_dn_dsyevj_batched
2006        )
2007    } else if mem::size_of::<T>() == mem::size_of::<Complex32>() {
2008        dispatch_complex!(
2009            Complex32,
2010            f32,
2011            cuComplex,
2012            cusolver_dn_cheevj_batched_buffer_size,
2013            cusolver_dn_cheevj_batched
2014        )
2015    } else {
2016        dispatch_complex!(
2017            Complex64,
2018            f64,
2019            cuDoubleComplex,
2020            cusolver_dn_zheevj_batched_buffer_size,
2021            cusolver_dn_zheevj_batched
2022        )
2023    }
2024}
2025
2026/// Batched Jacobi SVD: batch of `m × n` matrices with stride `m×n`.
2027#[allow(clippy::too_many_arguments)]
2028pub fn gesvdj_batched<T: SolverScalar>(
2029    handle: &DnHandle,
2030    jobz: EigMode,
2031    m: i32,
2032    n: i32,
2033    a: &mut DeviceBuffer<T>,
2034    lda: i32,
2035    s: &mut DeviceBuffer<T::Real>,
2036    u: &mut DeviceBuffer<T>,
2037    ldu: i32,
2038    v: &mut DeviceBuffer<T>,
2039    ldv: i32,
2040    info: &mut DeviceBuffer<i32>,
2041    params: &GesvdjInfo,
2042    batch_size: i32,
2043) -> Result<()> {
2044    use baracuda_cusolver_sys::{cuComplex, cuDoubleComplex};
2045    use core::mem;
2046
2047    let mut lwork: c_int = 0;
2048    macro_rules! dispatch_real {
2049        ($t:ty, $bufsize:ident, $solve:ident) => {{
2050            let c = cusolver()?;
2051            check(unsafe {
2052                (c.$bufsize()?)(
2053                    handle.as_raw(),
2054                    jobz,
2055                    m,
2056                    n,
2057                    a.as_raw().0 as *const $t,
2058                    lda,
2059                    s.as_raw().0 as *const $t,
2060                    u.as_raw().0 as *const $t,
2061                    ldu,
2062                    v.as_raw().0 as *const $t,
2063                    ldv,
2064                    &mut lwork,
2065                    params.as_raw(),
2066                    batch_size,
2067                )
2068            })?;
2069            let workspace =
2070                DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
2071            check(unsafe {
2072                (c.$solve()?)(
2073                    handle.as_raw(),
2074                    jobz,
2075                    m,
2076                    n,
2077                    a.as_raw().0 as *mut $t,
2078                    lda,
2079                    s.as_raw().0 as *mut $t,
2080                    u.as_raw().0 as *mut $t,
2081                    ldu,
2082                    v.as_raw().0 as *mut $t,
2083                    ldv,
2084                    workspace.as_raw().0 as *mut $t,
2085                    lwork,
2086                    info.as_raw().0 as *mut c_int,
2087                    params.as_raw(),
2088                    batch_size,
2089                )
2090            })
2091        }};
2092    }
2093    macro_rules! dispatch_complex {
2094        ($t:ty, $real:ty, $raw:ty, $bufsize:ident, $solve:ident) => {{
2095            let c = cusolver()?;
2096            check(unsafe {
2097                (c.$bufsize()?)(
2098                    handle.as_raw(),
2099                    jobz,
2100                    m,
2101                    n,
2102                    a.as_raw().0 as *const $raw,
2103                    lda,
2104                    s.as_raw().0 as *const $real,
2105                    u.as_raw().0 as *const $raw,
2106                    ldu,
2107                    v.as_raw().0 as *const $raw,
2108                    ldv,
2109                    &mut lwork,
2110                    params.as_raw(),
2111                    batch_size,
2112                )
2113            })?;
2114            let workspace =
2115                DeviceBuffer::<T>::new(a.context(), lwork as usize).map_err(alloc_fail)?;
2116            check(unsafe {
2117                (c.$solve()?)(
2118                    handle.as_raw(),
2119                    jobz,
2120                    m,
2121                    n,
2122                    a.as_raw().0 as *mut $raw,
2123                    lda,
2124                    s.as_raw().0 as *mut $real,
2125                    u.as_raw().0 as *mut $raw,
2126                    ldu,
2127                    v.as_raw().0 as *mut $raw,
2128                    ldv,
2129                    workspace.as_raw().0 as *mut $raw,
2130                    lwork,
2131                    info.as_raw().0 as *mut c_int,
2132                    params.as_raw(),
2133                    batch_size,
2134                )
2135            })
2136        }};
2137    }
2138
2139    if mem::size_of::<T>() == mem::size_of::<f32>() && mem::size_of::<T::Real>() == 4 {
2140        dispatch_real!(
2141            f32,
2142            cusolver_dn_sgesvdj_batched_buffer_size,
2143            cusolver_dn_sgesvdj_batched
2144        )
2145    } else if mem::size_of::<T>() == mem::size_of::<f64>() && mem::size_of::<T::Real>() == 8 {
2146        dispatch_real!(
2147            f64,
2148            cusolver_dn_dgesvdj_batched_buffer_size,
2149            cusolver_dn_dgesvdj_batched
2150        )
2151    } else if mem::size_of::<T>() == mem::size_of::<Complex32>() {
2152        dispatch_complex!(
2153            Complex32,
2154            f32,
2155            cuComplex,
2156            cusolver_dn_cgesvdj_batched_buffer_size,
2157            cusolver_dn_cgesvdj_batched
2158        )
2159    } else {
2160        dispatch_complex!(
2161            Complex64,
2162            f64,
2163            cuDoubleComplex,
2164            cusolver_dn_zgesvdj_batched_buffer_size,
2165            cusolver_dn_zgesvdj_batched
2166        )
2167    }
2168}
2169
2170// ---- cuSOLVERMg: multi-GPU dense solvers --------------------------------
2171
2172pub mod mg {
2173    //! Multi-GPU dense solvers via `libcusolverMg`. Shares dimensions with
2174    //! the single-GPU API but takes arrays of device pointers (one per
2175    //! physical GPU after [`Handle::device_select`]).
2176
2177    use core::ffi::{c_int, c_void};
2178
2179    use baracuda_cusolver_sys::{
2180        cudaDataType, cudaLibMgGrid_t, cudaLibMgMatrixDesc_t, cusolver_mg, cusolverMgHandle_t,
2181    };
2182
2183    use super::{alloc_fail, check, EigMode, Fill, Result};
2184
2185    /// Multi-GPU cuSOLVER handle.
2186    #[derive(Debug)]
2187    pub struct Handle {
2188        raw: cusolverMgHandle_t,
2189    }
2190
2191    impl Handle {
2192        /// Create a new multi-GPU cuSOLVER handle
2193        /// (`cusolverMgCreate`).
2194        pub fn new() -> Result<Self> {
2195            let mg = cusolver_mg()?;
2196            let cu = mg.cusolver_mg_create()?;
2197            let mut h: cusolverMgHandle_t = core::ptr::null_mut();
2198            check(unsafe { cu(&mut h) })?;
2199            Ok(Self { raw: h })
2200        }
2201
2202        /// Assign a set of physical CUDA devices to this handle. Future
2203        /// factorizations will stripe across them.
2204        pub fn device_select(&self, devices: &[i32]) -> Result<()> {
2205            let mg = cusolver_mg()?;
2206            let cu = mg.cusolver_mg_device_select()?;
2207            check(unsafe { cu(self.raw, devices.len() as c_int, devices.as_ptr()) })
2208        }
2209
2210        /// Raw `cusolverMgHandle_t`. Use with care.
2211        pub fn as_raw(&self) -> cusolverMgHandle_t {
2212            self.raw
2213        }
2214    }
2215
2216    impl Drop for Handle {
2217        fn drop(&mut self) {
2218            if let Ok(mg) = cusolver_mg() {
2219                if let Ok(cu) = mg.cusolver_mg_destroy() {
2220                    let _ = unsafe { cu(self.raw) };
2221                }
2222            }
2223        }
2224    }
2225
2226    /// A device grid — assigns distribution roles to physical devices.
2227    #[derive(Debug)]
2228    pub struct DeviceGrid {
2229        raw: cudaLibMgGrid_t,
2230    }
2231
2232    impl DeviceGrid {
2233        /// Build a device grid (`cusolverMgCreateDeviceGrid`).
2234        ///
2235        /// `mapping` is typically `CUDALIBMG_GRID_MAPPING_COL_MAJOR (1)`.
2236        pub fn new(num_row_devices: i32, num_col_devices: i32, devices: &[i32], mapping: i32) -> Result<Self> {
2237            let mg = cusolver_mg()?;
2238            let cu = mg.cusolver_mg_create_device_grid()?;
2239            let mut raw: cudaLibMgGrid_t = core::ptr::null_mut();
2240            check(unsafe {
2241                cu(
2242                    &mut raw,
2243                    num_row_devices,
2244                    num_col_devices,
2245                    devices.as_ptr(),
2246                    mapping,
2247                )
2248            })?;
2249            Ok(Self { raw })
2250        }
2251
2252        /// Raw `cudaLibMgGrid_t`. Use with care.
2253        pub fn as_raw(&self) -> cudaLibMgGrid_t {
2254            self.raw
2255        }
2256    }
2257
2258    impl Drop for DeviceGrid {
2259        fn drop(&mut self) {
2260            if let Ok(mg) = cusolver_mg() {
2261                if let Ok(cu) = mg.cusolver_mg_destroy_grid() {
2262                    let _ = unsafe { cu(self.raw) };
2263                }
2264            }
2265        }
2266    }
2267
2268    /// Matrix-distribution descriptor.
2269    #[derive(Debug)]
2270    pub struct MatrixDesc {
2271        raw: cudaLibMgMatrixDesc_t,
2272    }
2273
2274    impl MatrixDesc {
2275        /// Create a distributed-matrix descriptor
2276        /// (`cusolverMgCreateMatrixDesc`).
2277        pub fn new(
2278            num_rows: i64,
2279            num_cols: i64,
2280            row_block_size: i64,
2281            col_block_size: i64,
2282            data_type: cudaDataType,
2283            grid: &DeviceGrid,
2284        ) -> Result<Self> {
2285            let mg = cusolver_mg()?;
2286            let cu = mg.cusolver_mg_create_matrix_desc()?;
2287            let mut raw: cudaLibMgMatrixDesc_t = core::ptr::null_mut();
2288            check(unsafe {
2289                cu(
2290                    &mut raw,
2291                    num_rows,
2292                    num_cols,
2293                    row_block_size,
2294                    col_block_size,
2295                    data_type,
2296                    grid.as_raw(),
2297                )
2298            })?;
2299            Ok(Self { raw })
2300        }
2301
2302        /// Raw `cudaLibMgMatrixDesc_t`. Use with care.
2303        pub fn as_raw(&self) -> cudaLibMgMatrixDesc_t {
2304            self.raw
2305        }
2306    }
2307
2308    impl Drop for MatrixDesc {
2309        fn drop(&mut self) {
2310            if let Ok(mg) = cusolver_mg() {
2311                if let Ok(cu) = mg.cusolver_mg_destroy_matrix_desc() {
2312                    let _ = unsafe { cu(self.raw) };
2313                }
2314            }
2315        }
2316    }
2317
2318    /// Multi-GPU LU buffer-size query.
2319    ///
2320    /// # Safety
2321    /// `array_d_a`, `array_d_ipiv` must be host arrays of device pointers
2322    /// matching the selected devices.
2323    #[allow(clippy::too_many_arguments)]
2324    pub unsafe fn getrf_buffer_size(
2325        handle: &Handle,
2326        m: i32,
2327        n: i32,
2328        array_d_a: *mut *mut c_void,
2329        ia: i32,
2330        ja: i32,
2331        desc_a: &MatrixDesc,
2332        array_d_ipiv: *mut *mut c_int,
2333        compute_type: cudaDataType,
2334    ) -> Result<i64> { unsafe {
2335        let mg = cusolver_mg()?;
2336        let cu = mg.cusolver_mg_getrf_buffer_size()?;
2337        let mut lwork: i64 = 0;
2338        check(cu(
2339            handle.as_raw(),
2340            m,
2341            n,
2342            array_d_a,
2343            ia,
2344            ja,
2345            desc_a.as_raw(),
2346            array_d_ipiv,
2347            compute_type,
2348            &mut lwork,
2349        ))?;
2350        Ok(lwork)
2351    }}
2352
2353    /// # Safety
2354    /// Same pointer-array requirements as [`getrf_buffer_size`].
2355    #[allow(clippy::too_many_arguments)]
2356    pub unsafe fn getrf(
2357        handle: &Handle,
2358        m: i32,
2359        n: i32,
2360        array_d_a: *mut *mut c_void,
2361        ia: i32,
2362        ja: i32,
2363        desc_a: &MatrixDesc,
2364        array_d_ipiv: *mut *mut c_int,
2365        compute_type: cudaDataType,
2366        array_d_work: *mut *mut c_void,
2367        lwork: i64,
2368        info: &mut [c_int],
2369    ) -> Result<()> { unsafe {
2370        let mg = cusolver_mg()?;
2371        let cu = mg.cusolver_mg_getrf()?;
2372        let _ = alloc_fail::<()>; // silence unused-import in release builds
2373        check(cu(
2374            handle.as_raw(),
2375            m,
2376            n,
2377            array_d_a,
2378            ia,
2379            ja,
2380            desc_a.as_raw(),
2381            array_d_ipiv,
2382            compute_type,
2383            array_d_work,
2384            lwork,
2385            info.as_mut_ptr(),
2386        ))
2387    }}
2388
2389    /// Multi-GPU Cholesky buffer-size.
2390    ///
2391    /// # Safety
2392    /// Same as [`getrf_buffer_size`].
2393    #[allow(clippy::too_many_arguments)]
2394    pub unsafe fn potrf_buffer_size(
2395        handle: &Handle,
2396        uplo: Fill,
2397        n: i32,
2398        array_d_a: *mut *mut c_void,
2399        ia: i32,
2400        ja: i32,
2401        desc_a: &MatrixDesc,
2402        compute_type: cudaDataType,
2403    ) -> Result<i64> { unsafe {
2404        let mg = cusolver_mg()?;
2405        let cu = mg.cusolver_mg_potrf_buffer_size()?;
2406        let mut lwork: i64 = 0;
2407        check(cu(
2408            handle.as_raw(),
2409            uplo,
2410            n,
2411            array_d_a,
2412            ia,
2413            ja,
2414            desc_a.as_raw(),
2415            compute_type,
2416            &mut lwork,
2417        ))?;
2418        Ok(lwork)
2419    }}
2420
2421    /// # Safety
2422    /// Same as [`getrf_buffer_size`].
2423    #[allow(clippy::too_many_arguments)]
2424    pub unsafe fn potrf(
2425        handle: &Handle,
2426        uplo: Fill,
2427        n: i32,
2428        array_d_a: *mut *mut c_void,
2429        ia: i32,
2430        ja: i32,
2431        desc_a: &MatrixDesc,
2432        compute_type: cudaDataType,
2433        array_d_work: *mut *mut c_void,
2434        lwork: i64,
2435        info: &mut [c_int],
2436    ) -> Result<()> { unsafe {
2437        let mg = cusolver_mg()?;
2438        let cu = mg.cusolver_mg_potrf()?;
2439        check(cu(
2440            handle.as_raw(),
2441            uplo,
2442            n,
2443            array_d_a,
2444            ia,
2445            ja,
2446            desc_a.as_raw(),
2447            compute_type,
2448            array_d_work,
2449            lwork,
2450            info.as_mut_ptr(),
2451        ))
2452    }}
2453
2454    /// Multi-GPU symmetric eigendecomposition buffer-size.
2455    ///
2456    /// # Safety
2457    /// Same as [`getrf_buffer_size`].
2458    #[allow(clippy::too_many_arguments)]
2459    pub unsafe fn syevd_buffer_size(
2460        handle: &Handle,
2461        jobz: EigMode,
2462        uplo: Fill,
2463        n: i32,
2464        array_d_a: *mut *mut c_void,
2465        ia: i32,
2466        ja: i32,
2467        desc_a: &MatrixDesc,
2468        w: *mut c_void,
2469        data_type_w: cudaDataType,
2470        compute_type: cudaDataType,
2471    ) -> Result<i64> { unsafe {
2472        let mg = cusolver_mg()?;
2473        let cu = mg.cusolver_mg_syevd_buffer_size()?;
2474        let mut lwork: i64 = 0;
2475        check(cu(
2476            handle.as_raw(),
2477            jobz,
2478            uplo,
2479            n,
2480            array_d_a,
2481            ia,
2482            ja,
2483            desc_a.as_raw(),
2484            w,
2485            data_type_w,
2486            compute_type,
2487            &mut lwork,
2488        ))?;
2489        Ok(lwork)
2490    }}
2491
2492    /// # Safety
2493    /// Same as [`getrf_buffer_size`].
2494    #[allow(clippy::too_many_arguments)]
2495    pub unsafe fn syevd(
2496        handle: &Handle,
2497        jobz: EigMode,
2498        uplo: Fill,
2499        n: i32,
2500        array_d_a: *mut *mut c_void,
2501        ia: i32,
2502        ja: i32,
2503        desc_a: &MatrixDesc,
2504        w: *mut c_void,
2505        data_type_w: cudaDataType,
2506        compute_type: cudaDataType,
2507        array_d_work: *mut *mut c_void,
2508        lwork: i64,
2509        info: &mut [c_int],
2510    ) -> Result<()> { unsafe {
2511        let mg = cusolver_mg()?;
2512        let cu = mg.cusolver_mg_syevd()?;
2513        check(cu(
2514            handle.as_raw(),
2515            jobz,
2516            uplo,
2517            n,
2518            array_d_a,
2519            ia,
2520            ja,
2521            desc_a.as_raw(),
2522            w,
2523            data_type_w,
2524            compute_type,
2525            array_d_work,
2526            lwork,
2527            info.as_mut_ptr(),
2528        ))
2529    }}
2530}
2531
2532// ---- Back-compat: single-precision shortcuts -----------------------------
2533
2534/// Shortcut for [`getrf`] on `f32`.
2535pub fn sgetrf(
2536    handle: &DnHandle,
2537    m: i32,
2538    n: i32,
2539    a: &mut DeviceBuffer<f32>,
2540    lda: i32,
2541    ipiv: &mut DeviceBuffer<i32>,
2542    info: &mut DeviceBuffer<i32>,
2543) -> Result<()> {
2544    getrf::<f32>(handle, m, n, a, lda, ipiv, info)
2545}
2546
2547/// Shortcut for [`getrs`] on `f32`.
2548#[allow(clippy::too_many_arguments)]
2549pub fn sgetrs(
2550    handle: &DnHandle,
2551    trans: Op,
2552    n: i32,
2553    nrhs: i32,
2554    a: &DeviceBuffer<f32>,
2555    lda: i32,
2556    ipiv: &DeviceBuffer<i32>,
2557    b: &mut DeviceBuffer<f32>,
2558    ldb: i32,
2559    info: &mut DeviceBuffer<i32>,
2560) -> Result<()> {
2561    getrs::<f32>(handle, trans, n, nrhs, a, lda, ipiv, b, ldb, info)
2562}
2563
2564// ---- Generic X... (64-bit-size, type-erased) ----------------------------
2565
2566pub mod xapi {
2567    //! The generic 64-bit cuSOLVER API (`cusolverDnX*`). Matrix dimensions
2568    //! are `i64`; element types are passed at call-time as
2569    //! [`cudaDataType`]. Workspace sizes are split between on-device and
2570    //! on-host buffers.
2571
2572    use super::*;
2573    use baracuda_cusolver_sys::{cudaDataType, cusolverDnParams_t};
2574
2575    /// Opaque parameter handle for the generic `cusolverDnX*` API
2576    /// (`cusolverDnParams_t`).
2577    #[derive(Debug)]
2578    pub struct Params {
2579        raw: cusolverDnParams_t,
2580    }
2581
2582    impl Params {
2583        /// Create a new parameters handle
2584        /// (`cusolverDnCreateParams`).
2585        pub fn new() -> Result<Self> {
2586            let c = cusolver()?;
2587            let cu = c.cusolver_dn_create_params()?;
2588            let mut p: cusolverDnParams_t = core::ptr::null_mut();
2589            check(unsafe { cu(&mut p) })?;
2590            Ok(Self { raw: p })
2591        }
2592
2593        /// Raw `cusolverDnParams_t`. Use with care.
2594        pub fn as_raw(&self) -> cusolverDnParams_t {
2595            self.raw
2596        }
2597    }
2598
2599    impl Drop for Params {
2600        fn drop(&mut self) {
2601            if let Ok(c) = cusolver() {
2602                if let Ok(cu) = c.cusolver_dn_destroy_params() {
2603                    let _ = unsafe { cu(self.raw) };
2604                }
2605            }
2606        }
2607    }
2608
2609    /// Buffer-size query for generic LU factorization. Returns
2610    /// `(workspace_bytes_on_device, workspace_bytes_on_host)`.
2611    ///
2612    /// # Safety
2613    ///
2614    /// `a` is a raw device pointer passed through to cuSOLVER's
2615    /// `cusolverDnXgetrf_bufferSize`. The pointer is not dereferenced by
2616    /// this Rust function but the underlying cuSOLVER call may sample
2617    /// alignment or layout from it; the caller must ensure it is either
2618    /// null or a valid device address for an `m x n` matrix with leading
2619    /// dimension `lda` of `data_type_a`.
2620    #[allow(clippy::too_many_arguments)]
2621    pub unsafe fn xgetrf_buffer_size(
2622        handle: &DnHandle,
2623        params: &Params,
2624        m: i64,
2625        n: i64,
2626        data_type_a: cudaDataType,
2627        a: *const c_void,
2628        lda: i64,
2629        compute_type: cudaDataType,
2630    ) -> Result<(usize, usize)> {
2631        let c = cusolver()?;
2632        let cu = c.cusolver_dn_xgetrf_buffer_size()?;
2633        let (mut dev, mut host) = (0usize, 0usize);
2634        check(unsafe {
2635            cu(
2636                handle.as_raw(),
2637                params.raw,
2638                m,
2639                n,
2640                data_type_a,
2641                a,
2642                lda,
2643                compute_type,
2644                &mut dev,
2645                &mut host,
2646            )
2647        })?;
2648        Ok((dev, host))
2649    }
2650
2651    /// Generic LU factorization (`Xgetrf`).
2652    ///
2653    /// # Safety
2654    ///
2655    /// All raw device pointers (`a`, `ipiv`, `info`, `dev_workspace`,
2656    /// `host_workspace`) must point to allocations of the appropriate type
2657    /// and size for the requested `(m, n, data_type_a, compute_type)`,
2658    /// and remain live until the call completes on `handle`'s stream.
2659    #[allow(clippy::too_many_arguments)]
2660    pub unsafe fn xgetrf(
2661        handle: &DnHandle,
2662        params: &Params,
2663        m: i64,
2664        n: i64,
2665        data_type_a: cudaDataType,
2666        a: *mut c_void,
2667        lda: i64,
2668        ipiv: *mut i64,
2669        compute_type: cudaDataType,
2670        device_buf: *mut c_void,
2671        device_bytes: usize,
2672        host_buf: *mut c_void,
2673        host_bytes: usize,
2674        info: *mut c_int,
2675    ) -> Result<()> { unsafe {
2676        let c = cusolver()?;
2677        let cu = c.cusolver_dn_xgetrf()?;
2678        check(cu(
2679            handle.as_raw(),
2680            params.raw,
2681            m,
2682            n,
2683            data_type_a,
2684            a,
2685            lda,
2686            ipiv,
2687            compute_type,
2688            device_buf,
2689            device_bytes,
2690            host_buf,
2691            host_bytes,
2692            info,
2693        ))
2694    }}
2695
2696    /// Solve `op(A) X = B` after [`xgetrf`] has factored A. `b`/`ldb`
2697    /// describe the right-hand-side matrix; `b` is overwritten with X.
2698    ///
2699    /// # Safety
2700    ///
2701    /// `a`, `ipiv`, `b`, and `info` must point to live device memory of
2702    /// the sizes implied by `n`, `nrhs`, `lda`, and `ldb`, with element
2703    /// types matching `data_type_a` / `data_type_b`.
2704    #[allow(clippy::too_many_arguments)]
2705    pub unsafe fn xgetrs(
2706        handle: &DnHandle,
2707        params: &Params,
2708        trans: Op,
2709        n: i64,
2710        nrhs: i64,
2711        data_type_a: cudaDataType,
2712        a: *const c_void,
2713        lda: i64,
2714        ipiv: *const i64,
2715        data_type_b: cudaDataType,
2716        b: *mut c_void,
2717        ldb: i64,
2718        info: *mut c_int,
2719    ) -> Result<()> { unsafe {
2720        let c = cusolver()?;
2721        let cu = c.cusolver_dn_xgetrs()?;
2722        check(cu(
2723            handle.as_raw(),
2724            params.raw,
2725            trans.raw(),
2726            n,
2727            nrhs,
2728            data_type_a,
2729            a,
2730            lda,
2731            ipiv,
2732            data_type_b,
2733            b,
2734            ldb,
2735            info,
2736        ))
2737    }}
2738
2739    /// Solve `A X = B` after [`xpotrf`] has Cholesky-factored A.
2740    /// `b`/`ldb` describe the right-hand-side matrix; `b` is overwritten
2741    /// with X.
2742    ///
2743    /// # Safety
2744    ///
2745    /// See [`xgetrs`].
2746    #[allow(clippy::too_many_arguments)]
2747    pub unsafe fn xpotrs(
2748        handle: &DnHandle,
2749        params: &Params,
2750        uplo: Fill,
2751        n: i64,
2752        nrhs: i64,
2753        data_type_a: cudaDataType,
2754        a: *const c_void,
2755        lda: i64,
2756        data_type_b: cudaDataType,
2757        b: *mut c_void,
2758        ldb: i64,
2759        info: *mut c_int,
2760    ) -> Result<()> { unsafe {
2761        let c = cusolver()?;
2762        let cu = c.cusolver_dn_xpotrs()?;
2763        check(cu(
2764            handle.as_raw(),
2765            params.raw,
2766            uplo,
2767            n,
2768            nrhs,
2769            data_type_a,
2770            a,
2771            lda,
2772            data_type_b,
2773            b,
2774            ldb,
2775            info,
2776        ))
2777    }}
2778}
2779
2780// ---- Sparse --------------------------------------------------------------
2781
2782pub mod sparse {
2783    //! `cusolverSp*` — solve sparse linear systems via Cholesky or QR.
2784
2785    use super::*;
2786    use baracuda_cusolver_sys::cusolverSpHandle_t;
2787    use core::ffi::c_int;
2788
2789    /// Sparse cuSOLVER handle (`cusolverSpHandle_t`).
2790    #[derive(Debug)]
2791    pub struct SpHandle {
2792        raw: cusolverSpHandle_t,
2793        _not_send: PhantomData<*mut ()>,
2794    }
2795
2796    impl SpHandle {
2797        /// Create a new sparse cuSOLVER handle
2798        /// (`cusolverSpCreate`).
2799        pub fn new() -> Result<Self> {
2800            let c = cusolver()?;
2801            let cu = c.cusolver_sp_create()?;
2802            let mut h: cusolverSpHandle_t = core::ptr::null_mut();
2803            check(unsafe { cu(&mut h) })?;
2804            Ok(Self {
2805                raw: h,
2806                _not_send: PhantomData,
2807            })
2808        }
2809
2810        /// Bind this handle to `stream`
2811        /// (`cusolverSpSetStream`).
2812        pub fn set_stream(&self, stream: &Stream) -> Result<()> {
2813            let c = cusolver()?;
2814            let cu = c.cusolver_sp_set_stream()?;
2815            check(unsafe { cu(self.raw, stream.as_raw() as _) })
2816        }
2817
2818        /// Raw `cusolverSpHandle_t`. Use with care.
2819        pub fn as_raw(&self) -> cusolverSpHandle_t {
2820            self.raw
2821        }
2822    }
2823
2824    impl Drop for SpHandle {
2825        fn drop(&mut self) {
2826            if let Ok(c) = cusolver() {
2827                if let Ok(cu) = c.cusolver_sp_destroy() {
2828                    let _ = unsafe { cu(self.raw) };
2829                }
2830            }
2831        }
2832    }
2833
2834    /// Sparse Cholesky solve: `A * x = b` for SPD `A`.
2835    ///
2836    /// # Safety
2837    /// `descr_a`, CSR arrays, b and x must live on-device (b + x on-device,
2838    /// CSR arrays + descriptor on-device) and satisfy cuSOLVER sparse
2839    /// format requirements.
2840    #[allow(clippy::too_many_arguments)]
2841    pub unsafe fn scsrlsvchol(
2842        handle: &SpHandle,
2843        m: i32,
2844        nnz: i32,
2845        descr_a: *mut c_void,
2846        csr_val: *const f32,
2847        csr_row_ptr: *const c_int,
2848        csr_col_ind: *const c_int,
2849        b: *const f32,
2850        tol: f32,
2851        reorder: i32,
2852        x: *mut f32,
2853        singularity: *mut c_int,
2854    ) -> Result<()> { unsafe {
2855        let c = cusolver()?;
2856        let cu = c.cusolver_sp_scsrlsvchol()?;
2857        check(cu(
2858            handle.raw,
2859            m,
2860            nnz,
2861            descr_a,
2862            csr_val,
2863            csr_row_ptr,
2864            csr_col_ind,
2865            b,
2866            tol,
2867            reorder,
2868            x,
2869            singularity,
2870        ))
2871    }}
2872
2873    /// Sparse QR solve (least-squares, handles non-SPD systems).
2874    ///
2875    /// # Safety
2876    /// Same as [`scsrlsvchol`].
2877    #[allow(clippy::too_many_arguments)]
2878    pub unsafe fn scsrlsvqr(
2879        handle: &SpHandle,
2880        m: i32,
2881        nnz: i32,
2882        descr_a: *mut c_void,
2883        csr_val: *const f32,
2884        csr_row_ptr: *const c_int,
2885        csr_col_ind: *const c_int,
2886        b: *const f32,
2887        tol: f32,
2888        reorder: i32,
2889        x: *mut f32,
2890        singularity: *mut c_int,
2891    ) -> Result<()> { unsafe {
2892        let c = cusolver()?;
2893        let cu = c.cusolver_sp_scsrlsvqr()?;
2894        check(cu(
2895            handle.raw,
2896            m,
2897            nnz,
2898            descr_a,
2899            csr_val,
2900            csr_row_ptr,
2901            csr_col_ind,
2902            b,
2903            tol,
2904            reorder,
2905            x,
2906            singularity,
2907        ))
2908    }}
2909}
2910
2911// ---- Refactor ------------------------------------------------------------
2912
2913pub mod refactor {
2914    //! `cusolverRf*` — fast re-factorization given a sparsity pattern, for
2915    //! solving many systems that differ only in numeric values.
2916
2917    use super::*;
2918    use baracuda_cusolver_sys::cusolverRfHandle_t;
2919
2920    /// Refactor cuSOLVER handle (`cusolverRfHandle_t`).
2921    #[derive(Debug)]
2922    pub struct RfHandle {
2923        raw: cusolverRfHandle_t,
2924        _not_send: PhantomData<*mut ()>,
2925    }
2926
2927    impl RfHandle {
2928        /// Create a new refactor handle
2929        /// (`cusolverRfCreate`).
2930        pub fn new() -> Result<Self> {
2931            let c = cusolver()?;
2932            let cu = c.cusolver_rf_create()?;
2933            let mut h: cusolverRfHandle_t = core::ptr::null_mut();
2934            check(unsafe { cu(&mut h) })?;
2935            Ok(Self {
2936                raw: h,
2937                _not_send: PhantomData,
2938            })
2939        }
2940
2941        /// Raw `cusolverRfHandle_t`. Use with care.
2942        pub fn as_raw(&self) -> cusolverRfHandle_t {
2943            self.raw
2944        }
2945
2946        /// Run the symbolic analysis phase
2947        /// (`cusolverRfAnalyze`).
2948        pub fn analyze(&self) -> Result<()> {
2949            let c = cusolver()?;
2950            let cu = c.cusolver_rf_analyze()?;
2951            check(unsafe { cu(self.raw) })
2952        }
2953
2954        /// Run the numeric refactorization step
2955        /// (`cusolverRfRefactor`).
2956        pub fn refactor(&self) -> Result<()> {
2957            let c = cusolver()?;
2958            let cu = c.cusolver_rf_refactor()?;
2959            check(unsafe { cu(self.raw) })
2960        }
2961    }
2962
2963    impl Drop for RfHandle {
2964        fn drop(&mut self) {
2965            if let Ok(c) = cusolver() {
2966                if let Ok(cu) = c.cusolver_rf_destroy() {
2967                    let _ = unsafe { cu(self.raw) };
2968                }
2969            }
2970        }
2971    }
2972}