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