use mdarray::{Array, Dense, Dim, Layout, Shape, Slice};
use mdarray_linalg::solve::{Solve, SolveError};
use num_complex::ComplexFloat;
use super::{scalar::LapackScalar, simple::gesv};
use crate::Lapack;
impl<T, D: Dim> Solve<T, D> for Lapack
where
T: ComplexFloat + Default + LapackScalar,
T::Real: Into<T>,
{
fn solve_write<R: Dim, La: Layout, Lb: Layout>(
&self,
a: &mut Slice<T, (D, D), La>,
b: &mut Slice<T, (D, R), Lb>,
) -> Result<(), SolveError> {
gesv::<_, Lb, T, D, R>(a, b)
}
fn solve<R: Dim, La: Layout, Lb: Layout>(
&self,
a: &mut Slice<T, (D, D), La>,
b: &Slice<T, (D, R), Lb>,
) -> Result<Array<T, (D, R)>, SolveError> {
let ash = *a.shape();
let bsh = *b.shape();
let n = ash.dim(0);
let nrhs = bsh.dim(1);
let mut b_copy = Array::from_elem(<(D, R) as Shape>::from_dims(&[n, nrhs]), T::default());
for i in 0..n {
for j in 0..nrhs {
b_copy[[i, j]] = b[[i, j]];
}
}
gesv::<_, Dense, T, D, R>(a, &mut b_copy)?;
Ok(b_copy)
}
}