#[allow(unused_imports)]
use crate::prelude::*;
const BLOCK_SIZE: usize = 64;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Transpose {
NoTrans,
Trans,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Side {
Left,
Right,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum UpLo {
Upper,
Lower,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Diag {
NonUnit,
Unit,
}
#[inline]
pub fn ddot(x: &[f64], y: &[f64]) -> f64 {
assert_eq!(
x.len(),
y.len(),
"Vector lengths must match for dot product"
);
let n = x.len();
let mut sum = 0.0;
let chunks = n / 4;
let remainder = n % 4;
for i in 0..chunks {
let idx = i * 4;
sum += x[idx] * y[idx];
sum += x[idx + 1] * y[idx + 1];
sum += x[idx + 2] * y[idx + 2];
sum += x[idx + 3] * y[idx + 3];
}
for i in (chunks * 4)..n {
sum += x[i] * y[i];
}
let _ = remainder;
sum
}
#[inline]
pub fn dnrm2(x: &[f64]) -> f64 {
if x.is_empty() {
return 0.0;
}
let n = x.len();
let mut scale = 0.0f64;
for &xi in x {
let abs_xi = xi.abs();
if abs_xi > scale {
scale = abs_xi;
}
}
if scale == 0.0 {
return 0.0;
}
let mut sum = 0.0;
let inv_scale = 1.0 / scale;
let chunks = n / 4;
for i in 0..chunks {
let idx = i * 4;
let s0 = x[idx] * inv_scale;
let s1 = x[idx + 1] * inv_scale;
let s2 = x[idx + 2] * inv_scale;
let s3 = x[idx + 3] * inv_scale;
sum += s0 * s0 + s1 * s1 + s2 * s2 + s3 * s3;
}
for s in x.iter().skip(chunks * 4).take(n - chunks * 4) {
let s = s * inv_scale;
sum += s * s;
}
scale * sum.sqrt()
}
#[inline]
pub fn dscal(alpha: f64, x: &mut [f64]) {
if alpha == 1.0 {
return;
}
if alpha == 0.0 {
x.fill(0.0);
return;
}
let n = x.len();
let chunks = n / 4;
for i in 0..chunks {
let idx = i * 4;
x[idx] *= alpha;
x[idx + 1] *= alpha;
x[idx + 2] *= alpha;
x[idx + 3] *= alpha;
}
for x_val in x.iter_mut().skip(chunks * 4).take(n - chunks * 4) {
*x_val *= alpha;
}
}
#[inline]
pub fn daxpy(alpha: f64, x: &[f64], y: &mut [f64]) {
assert_eq!(x.len(), y.len(), "Vector lengths must match for DAXPY");
if alpha == 0.0 {
return;
}
let n = x.len();
let chunks = n / 4;
for i in 0..chunks {
let idx = i * 4;
y[idx] += alpha * x[idx];
y[idx + 1] += alpha * x[idx + 1];
y[idx + 2] += alpha * x[idx + 2];
y[idx + 3] += alpha * x[idx + 3];
}
for i in (chunks * 4)..n {
y[i] += alpha * x[i];
}
}
#[inline]
pub fn dcopy(x: &[f64], y: &mut [f64]) {
assert_eq!(x.len(), y.len(), "Vector lengths must match for DCOPY");
y.copy_from_slice(x);
}
#[inline]
pub fn dswap(x: &mut [f64], y: &mut [f64]) {
assert_eq!(x.len(), y.len(), "Vector lengths must match for DSWAP");
x.swap_with_slice(y);
}
#[inline]
pub fn idamax(x: &[f64]) -> usize {
if x.is_empty() {
return 0;
}
let mut max_idx = 0;
let mut max_val = x[0].abs();
for (i, &xi) in x.iter().enumerate().skip(1) {
let abs_xi = xi.abs();
if abs_xi > max_val {
max_val = abs_xi;
max_idx = i;
}
}
max_idx
}
#[inline]
pub fn dasum(x: &[f64]) -> f64 {
let n = x.len();
let mut sum = 0.0;
let chunks = n / 4;
for i in 0..chunks {
let idx = i * 4;
sum += x[idx].abs();
sum += x[idx + 1].abs();
sum += x[idx + 2].abs();
sum += x[idx + 3].abs();
}
for x_val in x.iter().skip(chunks * 4).take(n - chunks * 4) {
sum += x_val.abs();
}
sum
}
#[allow(clippy::too_many_arguments)]
pub fn dgemv(
trans: Transpose,
m: usize,
n: usize,
alpha: f64,
a: &[f64],
x: &[f64],
beta: f64,
y: &mut [f64],
) {
assert_eq!(a.len(), m * n, "Matrix A size must be m * n");
match trans {
Transpose::NoTrans => {
assert_eq!(x.len(), n, "Vector x length must be n for NoTrans");
assert_eq!(y.len(), m, "Vector y length must be m for NoTrans");
if beta == 0.0 {
y.fill(0.0);
} else if beta != 1.0 {
dscal(beta, y);
}
if alpha == 0.0 {
return;
}
for (i, y_val) in y.iter_mut().enumerate().take(m) {
let row_start = i * n;
let mut sum = 0.0;
let chunks = n / 4;
for j in 0..chunks {
let idx = j * 4;
sum += a[row_start + idx] * x[idx];
sum += a[row_start + idx + 1] * x[idx + 1];
sum += a[row_start + idx + 2] * x[idx + 2];
sum += a[row_start + idx + 3] * x[idx + 3];
}
for j in (chunks * 4)..n {
sum += a[row_start + j] * x[j];
}
*y_val += alpha * sum;
}
}
Transpose::Trans => {
assert_eq!(x.len(), m, "Vector x length must be m for Trans");
assert_eq!(y.len(), n, "Vector y length must be n for Trans");
if beta == 0.0 {
y.fill(0.0);
} else if beta != 1.0 {
dscal(beta, y);
}
if alpha == 0.0 {
return;
}
for (i, x_val) in x.iter().enumerate().take(m) {
let row_start = i * n;
let alpha_xi = alpha * x_val;
let chunks = n / 4;
for j in 0..chunks {
let idx = j * 4;
y[idx] += alpha_xi * a[row_start + idx];
y[idx + 1] += alpha_xi * a[row_start + idx + 1];
y[idx + 2] += alpha_xi * a[row_start + idx + 2];
y[idx + 3] += alpha_xi * a[row_start + idx + 3];
}
for j in (chunks * 4)..n {
y[j] += alpha_xi * a[row_start + j];
}
}
}
}
}
#[allow(clippy::too_many_arguments)]
pub fn dtrsv(uplo: UpLo, trans: Transpose, diag: Diag, n: usize, a: &[f64], b: &mut [f64]) {
assert_eq!(a.len(), n * n, "Matrix A size must be n * n");
assert_eq!(b.len(), n, "Vector b length must be n");
if n == 0 {
return;
}
match (uplo, trans) {
(UpLo::Lower, Transpose::NoTrans) | (UpLo::Upper, Transpose::Trans) => {
for i in 0..n {
let mut sum = b[i];
for j in 0..i {
let a_ij = if trans == Transpose::Trans {
a[j * n + i]
} else {
a[i * n + j]
};
sum -= a_ij * b[j];
}
if diag == Diag::NonUnit {
let a_ii = a[i * n + i];
b[i] = sum / a_ii;
} else {
b[i] = sum;
}
}
}
(UpLo::Upper, Transpose::NoTrans) | (UpLo::Lower, Transpose::Trans) => {
for i in (0..n).rev() {
let mut sum = b[i];
for j in (i + 1)..n {
let a_ij = if trans == Transpose::Trans {
a[j * n + i]
} else {
a[i * n + j]
};
sum -= a_ij * b[j];
}
if diag == Diag::NonUnit {
let a_ii = a[i * n + i];
b[i] = sum / a_ii;
} else {
b[i] = sum;
}
}
}
}
}
#[allow(clippy::too_many_arguments)]
pub fn dgemm(
trans_a: Transpose,
trans_b: Transpose,
m: usize,
n: usize,
k: usize,
alpha: f64,
a: &[f64],
b: &[f64],
beta: f64,
c: &mut [f64],
) {
let (a_rows, a_cols) = match trans_a {
Transpose::NoTrans => (m, k),
Transpose::Trans => (k, m),
};
let (b_rows, b_cols) = match trans_b {
Transpose::NoTrans => (k, n),
Transpose::Trans => (n, k),
};
assert_eq!(a.len(), a_rows * a_cols, "Matrix A size mismatch");
assert_eq!(b.len(), b_rows * b_cols, "Matrix B size mismatch");
assert_eq!(c.len(), m * n, "Matrix C size must be m * n");
if beta == 0.0 {
c.fill(0.0);
} else if beta != 1.0 {
for ci in c.iter_mut() {
*ci *= beta;
}
}
if alpha == 0.0 {
return;
}
if m * n * k > BLOCK_SIZE * BLOCK_SIZE * BLOCK_SIZE {
dgemm_blocked(trans_a, trans_b, m, n, k, alpha, a, b, c);
} else {
dgemm_simple(trans_a, trans_b, m, n, k, alpha, a, b, c);
}
}
#[allow(clippy::too_many_arguments)]
fn dgemm_simple(
trans_a: Transpose,
trans_b: Transpose,
m: usize,
n: usize,
k: usize,
alpha: f64,
a: &[f64],
b: &[f64],
c: &mut [f64],
) {
let (a_cols, b_cols) = match (trans_a, trans_b) {
(Transpose::NoTrans, Transpose::NoTrans) => (k, n),
(Transpose::NoTrans, Transpose::Trans) => (k, k),
(Transpose::Trans, Transpose::NoTrans) => (m, n),
(Transpose::Trans, Transpose::Trans) => (m, k),
};
for i in 0..m {
for j in 0..n {
let mut sum = 0.0;
for l in 0..k {
let a_il = match trans_a {
Transpose::NoTrans => a[i * a_cols + l],
Transpose::Trans => a[l * a_cols + i],
};
let b_lj = match trans_b {
Transpose::NoTrans => b[l * b_cols + j],
Transpose::Trans => b[j * b_cols + l],
};
sum += a_il * b_lj;
}
c[i * n + j] += alpha * sum;
}
}
}
#[allow(clippy::too_many_arguments)]
fn dgemm_blocked(
trans_a: Transpose,
trans_b: Transpose,
m: usize,
n: usize,
k: usize,
alpha: f64,
a: &[f64],
b: &[f64],
c: &mut [f64],
) {
let (a_cols, b_cols) = match (trans_a, trans_b) {
(Transpose::NoTrans, Transpose::NoTrans) => (k, n),
(Transpose::NoTrans, Transpose::Trans) => (k, k),
(Transpose::Trans, Transpose::NoTrans) => (m, n),
(Transpose::Trans, Transpose::Trans) => (m, k),
};
for i0 in (0..m).step_by(BLOCK_SIZE) {
let i1 = (i0 + BLOCK_SIZE).min(m);
for j0 in (0..n).step_by(BLOCK_SIZE) {
let j1 = (j0 + BLOCK_SIZE).min(n);
for l0 in (0..k).step_by(BLOCK_SIZE) {
let l1 = (l0 + BLOCK_SIZE).min(k);
for i in i0..i1 {
for j in j0..j1 {
let mut sum = 0.0;
for l in l0..l1 {
let a_il = match trans_a {
Transpose::NoTrans => a[i * a_cols + l],
Transpose::Trans => a[l * a_cols + i],
};
let b_lj = match trans_b {
Transpose::NoTrans => b[l * b_cols + j],
Transpose::Trans => b[j * b_cols + l],
};
sum += a_il * b_lj;
}
c[i * n + j] += alpha * sum;
}
}
}
}
}
}
#[allow(clippy::too_many_arguments)]
pub fn dtrsm(
side: Side,
uplo: UpLo,
trans: Transpose,
diag: Diag,
m: usize,
n: usize,
alpha: f64,
a: &[f64],
b: &mut [f64],
) {
let a_size = match side {
Side::Left => m,
Side::Right => n,
};
assert_eq!(a.len(), a_size * a_size, "Matrix A must be square");
assert_eq!(b.len(), m * n, "Matrix B size must be m * n");
if alpha != 1.0 {
for bi in b.iter_mut() {
*bi *= alpha;
}
}
match side {
Side::Left => dtrsm_left(uplo, trans, diag, m, n, a, b),
Side::Right => dtrsm_right(uplo, trans, diag, m, n, a, b),
}
}
fn dtrsm_left(
uplo: UpLo,
trans: Transpose,
diag: Diag,
m: usize,
n: usize,
a: &[f64],
b: &mut [f64],
) {
for col in 0..n {
match (uplo, trans) {
(UpLo::Lower, Transpose::NoTrans) | (UpLo::Upper, Transpose::Trans) => {
for i in 0..m {
let mut sum = b[i * n + col];
for j in 0..i {
let a_ij = if trans == Transpose::Trans {
a[j * m + i]
} else {
a[i * m + j]
};
sum -= a_ij * b[j * n + col];
}
if diag == Diag::NonUnit {
let a_ii = a[i * m + i];
b[i * n + col] = sum / a_ii;
} else {
b[i * n + col] = sum;
}
}
}
(UpLo::Upper, Transpose::NoTrans) | (UpLo::Lower, Transpose::Trans) => {
for i in (0..m).rev() {
let mut sum = b[i * n + col];
for j in (i + 1)..m {
let a_ij = if trans == Transpose::Trans {
a[j * m + i]
} else {
a[i * m + j]
};
sum -= a_ij * b[j * n + col];
}
if diag == Diag::NonUnit {
let a_ii = a[i * m + i];
b[i * n + col] = sum / a_ii;
} else {
b[i * n + col] = sum;
}
}
}
}
}
}
fn dtrsm_right(
uplo: UpLo,
trans: Transpose,
diag: Diag,
m: usize,
n: usize,
a: &[f64],
b: &mut [f64],
) {
for row in 0..m {
match (uplo, trans) {
(UpLo::Upper, Transpose::NoTrans) | (UpLo::Lower, Transpose::Trans) => {
for j in 0..n {
let mut sum = b[row * n + j];
for k in 0..j {
let a_kj = if trans == Transpose::Trans {
a[j * n + k]
} else {
a[k * n + j]
};
sum -= b[row * n + k] * a_kj;
}
if diag == Diag::NonUnit {
let a_jj = a[j * n + j];
b[row * n + j] = sum / a_jj;
} else {
b[row * n + j] = sum;
}
}
}
(UpLo::Lower, Transpose::NoTrans) | (UpLo::Upper, Transpose::Trans) => {
for j in (0..n).rev() {
let mut sum = b[row * n + j];
for k in (j + 1)..n {
let a_kj = if trans == Transpose::Trans {
a[j * n + k]
} else {
a[k * n + j]
};
sum -= b[row * n + k] * a_kj;
}
if diag == Diag::NonUnit {
let a_jj = a[j * n + j];
b[row * n + j] = sum / a_jj;
} else {
b[row * n + j] = sum;
}
}
}
}
}
}
#[derive(Debug, Clone)]
pub struct BlasLPConfig {
pub block_size: usize,
pub zero_tolerance: f64,
pub use_pivoting: bool,
}
impl Default for BlasLPConfig {
fn default() -> Self {
Self {
block_size: BLOCK_SIZE,
zero_tolerance: 1e-12,
use_pivoting: true,
}
}
}
#[allow(dead_code)]
pub fn compute_reduced_cost(
c: &[f64],
c_b: &[f64],
b_inv_a: &[f64],
m: usize,
n: usize,
) -> Vec<f64> {
let mut reduced = c.to_vec();
for j in 0..n {
for i in 0..m {
reduced[j] -= c_b[i] * b_inv_a[i * n + j];
}
}
reduced
}
#[allow(dead_code)]
pub fn solve_basis(b: &[f64], n: usize, rhs: &mut [f64], config: &BlasLPConfig) -> bool {
let mut lu = b.to_vec();
let mut perm: Vec<usize> = (0..n).collect();
for k in 0..n - 1 {
if config.use_pivoting {
let mut max_idx = k;
let mut max_val = lu[k * n + k].abs();
for i in (k + 1)..n {
let val = lu[i * n + k].abs();
if val > max_val {
max_val = val;
max_idx = i;
}
}
if max_val < config.zero_tolerance {
return false; }
if max_idx != k {
for j in 0..n {
lu.swap(k * n + j, max_idx * n + j);
}
perm.swap(k, max_idx);
}
}
let pivot = lu[k * n + k];
if pivot.abs() < config.zero_tolerance {
return false; }
for i in (k + 1)..n {
let factor = lu[i * n + k] / pivot;
lu[i * n + k] = factor;
for j in (k + 1)..n {
lu[i * n + j] -= factor * lu[k * n + j];
}
}
}
let mut tmp = vec![0.0; n];
for i in 0..n {
tmp[i] = rhs[perm[i]];
}
rhs.copy_from_slice(&tmp);
for i in 1..n {
for j in 0..i {
rhs[i] -= lu[i * n + j] * rhs[j];
}
}
for i in (0..n).rev() {
for j in (i + 1)..n {
rhs[i] -= lu[i * n + j] * rhs[j];
}
rhs[i] /= lu[i * n + i];
}
true
}
#[cfg(test)]
mod tests {
use super::*;
const EPSILON: f64 = 1e-10;
fn approx_eq(a: f64, b: f64) -> bool {
(a - b).abs() < EPSILON
}
fn approx_eq_vec(a: &[f64], b: &[f64]) -> bool {
a.len() == b.len() && a.iter().zip(b.iter()).all(|(&ai, &bi)| approx_eq(ai, bi))
}
#[test]
fn test_ddot() {
let x = vec![1.0, 2.0, 3.0, 4.0];
let y = vec![5.0, 6.0, 7.0, 8.0];
let result = ddot(&x, &y);
assert!(approx_eq(result, 70.0)); }
#[test]
fn test_ddot_empty() {
let x: Vec<f64> = vec![];
let y: Vec<f64> = vec![];
assert!(approx_eq(ddot(&x, &y), 0.0));
}
#[test]
fn test_dnrm2() {
let x = vec![3.0, 4.0];
assert!(approx_eq(dnrm2(&x), 5.0));
}
#[test]
fn test_dnrm2_large_values() {
let scale = 1e150;
let x = vec![3.0 * scale, 4.0 * scale];
assert!(approx_eq(dnrm2(&x), 5.0 * scale));
}
#[test]
fn test_dnrm2_small_values() {
let scale = 1e-150;
let x = vec![3.0 * scale, 4.0 * scale];
assert!(approx_eq(dnrm2(&x), 5.0 * scale));
}
#[test]
fn test_dscal() {
let mut x = vec![1.0, 2.0, 3.0, 4.0, 5.0];
dscal(2.0, &mut x);
assert!(approx_eq_vec(&x, &[2.0, 4.0, 6.0, 8.0, 10.0]));
}
#[test]
fn test_dscal_zero() {
let mut x = vec![1.0, 2.0, 3.0];
dscal(0.0, &mut x);
assert!(approx_eq_vec(&x, &[0.0, 0.0, 0.0]));
}
#[test]
fn test_dscal_one() {
let mut x = vec![1.0, 2.0, 3.0];
let original = x.clone();
dscal(1.0, &mut x);
assert!(approx_eq_vec(&x, &original));
}
#[test]
fn test_daxpy() {
let x = vec![1.0, 2.0, 3.0, 4.0, 5.0];
let mut y = vec![10.0, 20.0, 30.0, 40.0, 50.0];
daxpy(2.0, &x, &mut y);
assert!(approx_eq_vec(&y, &[12.0, 24.0, 36.0, 48.0, 60.0]));
}
#[test]
fn test_daxpy_zero_alpha() {
let x = vec![1.0, 2.0, 3.0];
let mut y = vec![10.0, 20.0, 30.0];
let original_y = y.clone();
daxpy(0.0, &x, &mut y);
assert!(approx_eq_vec(&y, &original_y));
}
#[test]
fn test_dcopy() {
let x = vec![1.0, 2.0, 3.0];
let mut y = vec![0.0, 0.0, 0.0];
dcopy(&x, &mut y);
assert!(approx_eq_vec(&y, &x));
}
#[test]
fn test_dswap() {
let mut x = vec![1.0, 2.0, 3.0];
let mut y = vec![4.0, 5.0, 6.0];
dswap(&mut x, &mut y);
assert!(approx_eq_vec(&x, &[4.0, 5.0, 6.0]));
assert!(approx_eq_vec(&y, &[1.0, 2.0, 3.0]));
}
#[test]
fn test_idamax() {
let x = vec![1.0, -5.0, 3.0, -2.0];
assert_eq!(idamax(&x), 1);
}
#[test]
fn test_idamax_empty() {
let x: Vec<f64> = vec![];
assert_eq!(idamax(&x), 0);
}
#[test]
fn test_dasum() {
let x = vec![1.0, -2.0, 3.0, -4.0, 5.0];
assert!(approx_eq(dasum(&x), 15.0));
}
#[test]
fn test_dgemv_notrans() {
let a = vec![1.0, 2.0, 3.0, 4.0];
let x = vec![5.0, 6.0];
let mut y = vec![0.0, 0.0];
dgemv(Transpose::NoTrans, 2, 2, 1.0, &a, &x, 0.0, &mut y);
assert!(approx_eq_vec(&y, &[17.0, 39.0]));
}
#[test]
fn test_dgemv_trans() {
let a = vec![1.0, 2.0, 3.0, 4.0];
let x = vec![5.0, 6.0];
let mut y = vec![0.0, 0.0];
dgemv(Transpose::Trans, 2, 2, 1.0, &a, &x, 0.0, &mut y);
assert!(approx_eq_vec(&y, &[23.0, 34.0]));
}
#[test]
fn test_dgemv_with_beta() {
let a = vec![1.0, 2.0, 3.0, 4.0];
let x = vec![1.0, 1.0];
let mut y = vec![10.0, 10.0];
dgemv(Transpose::NoTrans, 2, 2, 2.0, &a, &x, 3.0, &mut y);
assert!(approx_eq_vec(&y, &[36.0, 44.0]));
}
#[test]
fn test_dtrsv_lower() {
let l = vec![2.0, 0.0, 1.0, 3.0];
let mut b = vec![4.0, 5.0];
dtrsv(
UpLo::Lower,
Transpose::NoTrans,
Diag::NonUnit,
2,
&l,
&mut b,
);
assert!(approx_eq_vec(&b, &[2.0, 1.0]));
}
#[test]
fn test_dtrsv_upper() {
let u = vec![2.0, 1.0, 0.0, 3.0];
let mut b = vec![5.0, 6.0];
dtrsv(
UpLo::Upper,
Transpose::NoTrans,
Diag::NonUnit,
2,
&u,
&mut b,
);
assert!(approx_eq_vec(&b, &[1.5, 2.0]));
}
#[test]
fn test_dtrsv_unit_diagonal() {
let l = vec![1.0, 0.0, 2.0, 1.0];
let mut b = vec![3.0, 8.0];
dtrsv(UpLo::Lower, Transpose::NoTrans, Diag::Unit, 2, &l, &mut b);
assert!(approx_eq_vec(&b, &[3.0, 2.0]));
}
#[test]
fn test_dgemm_basic() {
let a = vec![1.0, 2.0, 3.0, 4.0];
let b = vec![5.0, 6.0, 7.0, 8.0];
let mut c = vec![0.0; 4];
dgemm(
Transpose::NoTrans,
Transpose::NoTrans,
2,
2,
2,
1.0,
&a,
&b,
0.0,
&mut c,
);
assert!(approx_eq_vec(&c, &[19.0, 22.0, 43.0, 50.0]));
}
#[test]
fn test_dgemm_with_transpose_a() {
let a = vec![1.0, 3.0, 2.0, 4.0];
let b = vec![5.0, 6.0, 7.0, 8.0];
let mut c = vec![0.0; 4];
dgemm(
Transpose::Trans,
Transpose::NoTrans,
2,
2,
2,
1.0,
&a,
&b,
0.0,
&mut c,
);
assert!(approx_eq_vec(&c, &[19.0, 22.0, 43.0, 50.0]));
}
#[test]
fn test_dgemm_with_transpose_b() {
let a = vec![1.0, 2.0, 3.0, 4.0];
let b = vec![5.0, 7.0, 6.0, 8.0];
let mut c = vec![0.0; 4];
dgemm(
Transpose::NoTrans,
Transpose::Trans,
2,
2,
2,
1.0,
&a,
&b,
0.0,
&mut c,
);
assert!(approx_eq_vec(&c, &[19.0, 22.0, 43.0, 50.0]));
}
#[test]
fn test_dgemm_with_alpha_beta() {
let a = vec![1.0, 0.0, 0.0, 1.0]; let b = vec![1.0, 2.0, 3.0, 4.0];
let mut c = vec![10.0, 20.0, 30.0, 40.0];
dgemm(
Transpose::NoTrans,
Transpose::NoTrans,
2,
2,
2,
2.0,
&a,
&b,
3.0,
&mut c,
);
assert!(approx_eq_vec(&c, &[32.0, 64.0, 96.0, 128.0]));
}
#[test]
fn test_dgemm_non_square() {
let a = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let b = vec![7.0, 8.0, 9.0, 10.0, 11.0, 12.0];
let mut c = vec![0.0; 4];
dgemm(
Transpose::NoTrans,
Transpose::NoTrans,
2,
2,
3,
1.0,
&a,
&b,
0.0,
&mut c,
);
assert!(approx_eq_vec(&c, &[58.0, 64.0, 139.0, 154.0]));
}
#[test]
fn test_dtrsm_left_lower() {
let l = vec![2.0, 0.0, 1.0, 3.0];
let mut b = vec![4.0, 6.0, 5.0, 9.0];
dtrsm(
Side::Left,
UpLo::Lower,
Transpose::NoTrans,
Diag::NonUnit,
2,
2,
1.0,
&l,
&mut b,
);
assert!(approx_eq_vec(&b, &[2.0, 3.0, 1.0, 2.0]));
}
#[test]
fn test_dtrsm_with_alpha() {
let l = vec![2.0, 0.0, 1.0, 3.0];
let mut b = vec![4.0, 6.0, 5.0, 9.0];
dtrsm(
Side::Left,
UpLo::Lower,
Transpose::NoTrans,
Diag::NonUnit,
2,
2,
2.0,
&l,
&mut b,
);
assert!(approx_eq_vec(&b, &[4.0, 6.0, 2.0, 4.0]));
}
#[test]
fn test_blas_lp_config() {
let config = BlasLPConfig::default();
assert_eq!(config.block_size, BLOCK_SIZE);
assert!(config.zero_tolerance > 0.0);
assert!(config.use_pivoting);
}
#[test]
fn test_solve_basis_simple() {
let b = vec![2.0, 1.0, 1.0, 3.0];
let mut rhs = vec![5.0, 7.0];
let config = BlasLPConfig::default();
let success = solve_basis(&b, 2, &mut rhs, &config);
assert!(success);
assert!(approx_eq(rhs[0], 1.6));
assert!(approx_eq(rhs[1], 1.8));
}
#[test]
fn test_identity_operations() {
let i = vec![1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0];
let x = vec![1.0, 2.0, 3.0];
let mut y = vec![0.0, 0.0, 0.0];
dgemv(Transpose::NoTrans, 3, 3, 1.0, &i, &x, 0.0, &mut y);
assert!(approx_eq_vec(&y, &x));
}
#[test]
fn test_large_matrix_blocked() {
let n = 100;
let a: Vec<f64> = (0..n * n).map(|i| (i % 7) as f64).collect();
let b: Vec<f64> = (0..n * n).map(|i| ((i + 3) % 5) as f64).collect();
let mut c = vec![0.0; n * n];
dgemm(
Transpose::NoTrans,
Transpose::NoTrans,
n,
n,
n,
1.0,
&a,
&b,
0.0,
&mut c,
);
let mut expected = 0.0;
for k in 0..n {
expected += a[k] * b[k * n];
}
assert!(approx_eq(c[0], expected));
}
}