use super::miniblas::{ddot, daxpy};
#[derive(Debug)]
pub enum LinpackError {
NonPositiveDefinite(usize),
ZeroDiagonal(usize),
InvalidJob(i32)
}
pub fn dpofa(a: &mut [f64], lda: usize, n: usize) -> Result<usize, LinpackError> {
if n > lda {
return Err(LinpackError::NonPositiveDefinite(0));
}
for j in 0..n {
let mut s = 0.0;
for k in 0..j {
let mut t = a[k + j * lda];
if k > 0 {
t -= ddot(
k as i32,
&a[k * lda..k * lda + k],
1,
&a[j * lda..j * lda + k],
1
);
}
t /= a[k + k * lda];
a[k + j * lda] = t;
s += t * t;
}
s = a[j + j * lda] - s;
if s <= 0.0 {
return Ok(j + 1); }
a[j + j * lda] = s.sqrt();
}
Ok(0)
}
pub fn dtrsl(
t: &mut [f64],
ldt: usize,
n: usize,
b: &mut [f64],
job: i32
) -> Result<usize, LinpackError> {
for i in 0..n {
if t[i + i * ldt] == 0.0 {
return Err(LinpackError::ZeroDiagonal(i + 1));
}
}
let case = match (job % 10, job % 100 / 10) {
(0, 0) => 1, (1, 0) => 2, (0, 1) => 3, (1, 1) => 4, _ => return Err(LinpackError::InvalidJob(job)),
};
match case {
1 => { b[0] /= t[0];
for j in 1..n {
let temp = -b[j - 1];
daxpy(
(n - j) as i32,
temp,
&t[j + (j - 1) * ldt..],
1,
&mut b[j..],
1
);
b[j] /= t[j + j * ldt];
}
},
2 => { b[n - 1] /= t[(n - 1) + (n - 1) * ldt];
for jj in 1..n {
let j = n - jj - 1;
let temp = -b[j + 1];
daxpy(
(j + 1) as i32,
temp,
&t[(j + 1) * ldt..],
1,
&mut b[..j + 1],
1
);
b[j] /= t[j + j * ldt];
}
},
3 => { b[n - 1] /= t[(n - 1) + (n - 1) * ldt];
for jj in 1..n {
let j = n - jj - 1;
let mut sum = 0.0;
for k in (j + 1)..n {
sum += t[k + j * ldt] * b[k];
}
b[j] = (b[j] - sum) / t[j + j * ldt];
}
},
4 => { b[0] /= t[0];
for j in 1..n {
let mut sum = 0.0;
for k in 0..j {
sum += t[k + j * ldt] * b[k];
}
b[j] = (b[j] - sum) / t[j + j * ldt];
}
},
_ => unreachable!(),
}
Ok(0)
}