use oxiblas_core::scalar::{Field, Scalar};
use oxiblas_matrix::{Mat, MatRef};
pub fn kron<T: Scalar + Clone + Field + bytemuck::Zeroable>(
a: MatRef<'_, T>,
b: MatRef<'_, T>,
) -> Mat<T> {
let m = a.nrows();
let n = a.ncols();
let p = b.nrows();
let q = b.ncols();
let mut result = Mat::zeros(m * p, n * q);
for i in 0..m {
for j in 0..n {
let aij = a[(i, j)];
for k in 0..p {
for l in 0..q {
result[(i * p + k, j * q + l)] = aij * b[(k, l)];
}
}
}
}
result
}
pub fn khatri_rao<T: Scalar + Clone + Field + bytemuck::Zeroable>(
a: MatRef<'_, T>,
b: MatRef<'_, T>,
) -> Mat<T> {
assert_eq!(
a.ncols(),
b.ncols(),
"Matrices must have the same number of columns for Khatri-Rao product"
);
let m = a.nrows();
let p = b.nrows();
let n = a.ncols();
let mut result = Mat::zeros(m * p, n);
for j in 0..n {
for i in 0..m {
let aij = a[(i, j)];
for k in 0..p {
result[(i * p + k, j)] = aij * b[(k, j)];
}
}
}
result
}
pub fn kron_sum<T: Scalar + Clone + Field + bytemuck::Zeroable>(
a: MatRef<'_, T>,
b: MatRef<'_, T>,
) -> Mat<T> {
assert_eq!(a.nrows(), a.ncols(), "First matrix must be square");
assert_eq!(b.nrows(), b.ncols(), "Second matrix must be square");
let m = a.nrows();
let p = b.nrows();
let n = m * p;
let eye_m = Mat::<T>::eye(m);
let eye_p = Mat::<T>::eye(p);
let a_kron_ip = kron(a, eye_p.as_ref());
let im_kron_b = kron(eye_m.as_ref(), b);
let mut result = Mat::zeros(n, n);
for i in 0..n {
for j in 0..n {
result[(i, j)] = a_kron_ip[(i, j)] + im_kron_b[(i, j)];
}
}
result
}
pub fn vec_mat<T: Scalar + Clone + bytemuck::Zeroable>(a: MatRef<'_, T>) -> Mat<T> {
let m = a.nrows();
let n = a.ncols();
let mut result = Mat::zeros(m * n, 1);
for j in 0..n {
for i in 0..m {
result[(j * m + i, 0)] = a[(i, j)];
}
}
result
}
pub fn unvec<T: Scalar + Clone + bytemuck::Zeroable>(
v: MatRef<'_, T>,
m: usize,
n: usize,
) -> Mat<T> {
assert_eq!(v.nrows(), m * n, "Vector length must equal m*n");
assert_eq!(v.ncols(), 1, "Input must be a column vector");
let mut result = Mat::zeros(m, n);
for j in 0..n {
for i in 0..m {
result[(i, j)] = v[(j * m + i, 0)];
}
}
result
}
pub fn commutation_matrix<T: Scalar + Clone + Field + bytemuck::Zeroable>(
m: usize,
n: usize,
) -> Mat<T> {
let mn = m * n;
let mut result = Mat::zeros(mn, mn);
for i in 0..m {
for j in 0..n {
let idx1 = j * m + i; let idx2 = i * n + j; result[(idx2, idx1)] = T::one();
}
}
result
}
pub fn duplication_matrix<T: Scalar + Clone + Field + bytemuck::Zeroable>(n: usize) -> Mat<T> {
let n_sq = n * n;
let n_half = n * (n + 1) / 2;
let mut result = Mat::zeros(n_sq, n_half);
let mut col = 0;
for j in 0..n {
for i in j..n {
let row1 = j * n + i;
result[(row1, col)] = T::one();
if i != j {
let row2 = i * n + j;
result[(row2, col)] = T::one();
}
col += 1;
}
}
result
}
pub fn elimination_matrix<T: Scalar + Clone + Field + bytemuck::Zeroable>(n: usize) -> Mat<T> {
let n_sq = n * n;
let n_half = n * (n + 1) / 2;
let mut result = Mat::zeros(n_half, n_sq);
let mut row = 0;
for j in 0..n {
for i in j..n {
let col = j * n + i;
result[(row, col)] = T::one();
row += 1;
}
}
result
}
pub fn kron_vec<T: Scalar + Clone + Field + bytemuck::Zeroable>(
a: MatRef<'_, T>,
b: MatRef<'_, T>,
x: &[T],
) -> Vec<T> {
let m = a.nrows();
let n = a.ncols();
let p = b.nrows();
let q = b.ncols();
assert_eq!(x.len(), n * q, "Vector length must match A.ncols * B.ncols");
let mut y = vec![T::zero(); q * m];
for i in 0..q {
for j in 0..m {
let mut sum = T::zero();
for k in 0..n {
sum = sum + x[k * q + i] * a[(j, k)]; }
y[j * q + i] = sum;
}
}
let mut result = vec![T::zero(); p * m];
for i in 0..p {
for j in 0..m {
let mut sum = T::zero();
for k in 0..q {
sum = sum + b[(i, k)] * y[j * q + k];
}
result[j * p + i] = sum;
}
}
result
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_kron_identity() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0]]);
let b = Mat::from_rows(&[&[1.0f64, 0.0], &[0.0, 1.0]]);
let c = kron(a.as_ref(), b.as_ref());
assert_eq!(c.nrows(), 4);
assert_eq!(c.ncols(), 4);
assert!((c[(0, 0)] - 1.0).abs() < 1e-10);
assert!((c[(0, 1)] - 0.0).abs() < 1e-10);
assert!((c[(0, 2)] - 2.0).abs() < 1e-10);
assert!((c[(0, 3)] - 0.0).abs() < 1e-10);
}
#[test]
fn test_kron_2x2() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0]]);
let b = Mat::from_rows(&[&[0.0f64, 5.0], &[6.0, 7.0]]);
let c = kron(a.as_ref(), b.as_ref());
assert_eq!(c.nrows(), 4);
assert_eq!(c.ncols(), 4);
assert!((c[(0, 0)] - 0.0).abs() < 1e-10);
assert!((c[(0, 1)] - 5.0).abs() < 1e-10);
assert!((c[(1, 0)] - 6.0).abs() < 1e-10);
assert!((c[(1, 1)] - 7.0).abs() < 1e-10);
assert!((c[(0, 2)] - 0.0).abs() < 1e-10);
assert!((c[(0, 3)] - 10.0).abs() < 1e-10);
assert!((c[(1, 2)] - 12.0).abs() < 1e-10);
assert!((c[(1, 3)] - 14.0).abs() < 1e-10);
}
#[test]
fn test_kron_rectangular() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0]]);
let b = Mat::from_rows(&[&[1.0f64], &[2.0]]);
let c = kron(a.as_ref(), b.as_ref());
assert_eq!(c.nrows(), 2);
assert_eq!(c.ncols(), 3);
assert!((c[(0, 0)] - 1.0).abs() < 1e-10);
assert!((c[(1, 0)] - 2.0).abs() < 1e-10);
assert!((c[(0, 1)] - 2.0).abs() < 1e-10);
assert!((c[(1, 1)] - 4.0).abs() < 1e-10);
assert!((c[(0, 2)] - 3.0).abs() < 1e-10);
assert!((c[(1, 2)] - 6.0).abs() < 1e-10);
}
#[test]
fn test_khatri_rao() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0]]);
let b = Mat::from_rows(&[&[5.0f64, 6.0], &[7.0, 8.0]]);
let c = khatri_rao(a.as_ref(), b.as_ref());
assert_eq!(c.nrows(), 4);
assert_eq!(c.ncols(), 2);
assert!((c[(0, 0)] - 5.0).abs() < 1e-10); assert!((c[(1, 0)] - 7.0).abs() < 1e-10); assert!((c[(2, 0)] - 15.0).abs() < 1e-10); assert!((c[(3, 0)] - 21.0).abs() < 1e-10);
assert!((c[(0, 1)] - 12.0).abs() < 1e-10); assert!((c[(1, 1)] - 16.0).abs() < 1e-10); assert!((c[(2, 1)] - 24.0).abs() < 1e-10); assert!((c[(3, 1)] - 32.0).abs() < 1e-10); }
#[test]
fn test_kron_sum() {
let a = Mat::from_rows(&[&[1.0f64, 0.0], &[0.0, 2.0]]);
let b = Mat::from_rows(&[&[3.0f64, 0.0], &[0.0, 4.0]]);
let c = kron_sum(a.as_ref(), b.as_ref());
assert_eq!(c.nrows(), 4);
assert_eq!(c.ncols(), 4);
assert!((c[(0, 0)] - 4.0).abs() < 1e-10);
assert!((c[(1, 1)] - 5.0).abs() < 1e-10);
assert!((c[(2, 2)] - 5.0).abs() < 1e-10);
assert!((c[(3, 3)] - 6.0).abs() < 1e-10);
}
#[test]
fn test_vec_unvec() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0]]);
let v = vec_mat(a.as_ref());
assert!((v[(0, 0)] - 1.0).abs() < 1e-10);
assert!((v[(1, 0)] - 3.0).abs() < 1e-10);
assert!((v[(2, 0)] - 2.0).abs() < 1e-10);
assert!((v[(3, 0)] - 4.0).abs() < 1e-10);
let b = unvec(v.as_ref(), 2, 2);
for i in 0..2 {
for j in 0..2 {
assert!((a[(i, j)] - b[(i, j)]).abs() < 1e-10);
}
}
}
#[test]
fn test_commutation_matrix() {
let k = commutation_matrix::<f64>(2, 3);
assert_eq!(k.nrows(), 6);
assert_eq!(k.ncols(), 6);
for i in 0..6 {
let mut row_sum = 0.0;
for j in 0..6 {
row_sum += k[(i, j)];
}
assert!((row_sum - 1.0).abs() < 1e-10);
}
}
#[test]
fn test_duplication_matrix() {
let d = duplication_matrix::<f64>(2);
assert_eq!(d.nrows(), 4);
assert_eq!(d.ncols(), 3);
}
#[test]
fn test_elimination_matrix() {
let l = elimination_matrix::<f64>(2);
assert_eq!(l.nrows(), 3);
assert_eq!(l.ncols(), 4);
}
#[test]
fn test_kron_vec() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0]]);
let b = Mat::from_rows(&[&[1.0f64, 0.0], &[0.0, 1.0]]);
let x = vec![1.0f64, 0.0, 0.0, 1.0];
let y = kron_vec(a.as_ref(), b.as_ref(), &x);
let k = kron(a.as_ref(), b.as_ref());
for i in 0..4 {
let mut expected = 0.0;
for j in 0..4 {
expected += k[(i, j)] * x[j];
}
assert!((y[i] - expected).abs() < 1e-10);
}
}
}