use crate::{Mat, MatMut, MatRef};
use nalgebra::{
DMatrix, DMatrixView, DMatrixViewMut, Dyn, Matrix, U1, ViewStorage, ViewStorageMut,
};
use num_traits::Zero;
use oxiblas_core::scalar::Scalar;
pub fn dmatrix_to_mat<T: Scalar + Clone + Zero>(dm: &DMatrix<T>) -> Mat<T> {
let nrows = dm.nrows();
let ncols = dm.ncols();
let mut mat = Mat::filled(nrows, ncols, T::zero());
for j in 0..ncols {
for i in 0..nrows {
mat[(i, j)] = dm[(i, j)];
}
}
mat
}
pub fn mat_to_dmatrix<T: Scalar + Clone + Zero + nalgebra::Scalar>(mat: &Mat<T>) -> DMatrix<T> {
let nrows = mat.nrows();
let ncols = mat.ncols();
DMatrix::from_fn(nrows, ncols, |i, j| mat[(i, j)])
}
pub fn mat_ref_to_dmatrix<T: Scalar + Clone + Zero + nalgebra::Scalar>(
mat: MatRef<'_, T>,
) -> DMatrix<T> {
let nrows = mat.nrows();
let ncols = mat.ncols();
DMatrix::from_fn(nrows, ncols, |i, j| mat[(i, j)])
}
pub fn mat_ref_to_dmatrix_view<'a, T: Scalar + nalgebra::Scalar>(
mat: MatRef<'a, T>,
) -> DMatrixView<'a, T> {
let shape = (Dyn(mat.nrows()), Dyn(mat.ncols()));
let strides = (U1, Dyn(mat.row_stride()));
let storage: ViewStorage<'a, T, Dyn, Dyn, U1, Dyn> =
unsafe { ViewStorage::from_raw_parts(mat.as_ptr(), shape, strides) };
Matrix::from_data(storage)
}
pub fn mat_mut_to_dmatrix_view_mut<'a, T: Scalar + nalgebra::Scalar>(
mut mat: MatMut<'a, T>,
) -> DMatrixViewMut<'a, T> {
let shape = (Dyn(mat.nrows()), Dyn(mat.ncols()));
let strides = (U1, Dyn(mat.row_stride()));
let ptr = mat.as_mut_ptr();
let storage: ViewStorageMut<'a, T, Dyn, Dyn, U1, Dyn> =
unsafe { ViewStorageMut::from_raw_parts(ptr, shape, strides) };
Matrix::from_data(storage)
}
pub fn dmatrix_to_mat_ref<T: Scalar + nalgebra::Scalar>(dm: &DMatrix<T>) -> MatRef<'_, T> {
unsafe { MatRef::new(dm.as_ptr(), dm.nrows(), dm.ncols(), dm.nrows()) }
}
pub fn dmatrix_to_mat_mut<T: Scalar + nalgebra::Scalar>(dm: &mut DMatrix<T>) -> MatMut<'_, T> {
let nrows = dm.nrows();
let ncols = dm.ncols();
unsafe { MatMut::new(dm.as_mut_ptr(), nrows, ncols, nrows) }
}
pub fn dmatrix_view_to_mat<T: Scalar + Clone + Zero>(view: DMatrixView<'_, T>) -> Mat<T> {
let nrows = view.nrows();
let ncols = view.ncols();
let mut mat = Mat::filled(nrows, ncols, T::zero());
for j in 0..ncols {
for i in 0..nrows {
mat[(i, j)] = view[(i, j)];
}
}
mat
}
pub trait MatNalgebraExt<T: Scalar> {
fn to_dmatrix(&self) -> DMatrix<T>
where
T: Clone + Zero + nalgebra::Scalar;
fn to_dmatrix_view(&self) -> DMatrixView<'_, T>
where
T: nalgebra::Scalar;
}
impl<T: Scalar> MatNalgebraExt<T> for Mat<T> {
fn to_dmatrix(&self) -> DMatrix<T>
where
T: Clone + Zero + nalgebra::Scalar,
{
mat_to_dmatrix(self)
}
fn to_dmatrix_view(&self) -> DMatrixView<'_, T>
where
T: nalgebra::Scalar,
{
mat_ref_to_dmatrix_view(self.as_ref())
}
}
impl<T: Scalar> MatNalgebraExt<T> for MatRef<'_, T> {
fn to_dmatrix(&self) -> DMatrix<T>
where
T: Clone + Zero + nalgebra::Scalar,
{
mat_ref_to_dmatrix(*self)
}
fn to_dmatrix_view(&self) -> DMatrixView<'_, T>
where
T: nalgebra::Scalar,
{
mat_ref_to_dmatrix_view(*self)
}
}
pub trait DMatrixOxiblasExt<T: Scalar> {
fn to_mat(&self) -> Mat<T>
where
T: Clone + Zero;
fn to_mat_ref(&self) -> MatRef<'_, T>
where
T: nalgebra::Scalar;
}
impl<T: Scalar + nalgebra::Scalar> DMatrixOxiblasExt<T> for DMatrix<T> {
fn to_mat(&self) -> Mat<T>
where
T: Clone + Zero,
{
dmatrix_to_mat(self)
}
fn to_mat_ref(&self) -> MatRef<'_, T>
where
T: nalgebra::Scalar,
{
dmatrix_to_mat_ref(self)
}
}
impl<T: Scalar + nalgebra::Scalar> DMatrixOxiblasExt<T> for DMatrixView<'_, T> {
fn to_mat(&self) -> Mat<T>
where
T: Clone + Zero,
{
dmatrix_view_to_mat(*self)
}
fn to_mat_ref(&self) -> MatRef<'_, T>
where
T: nalgebra::Scalar,
{
unsafe { MatRef::new(self.as_ptr(), self.nrows(), self.ncols(), self.strides().1) }
}
}
impl<T: Scalar + Clone + Zero + nalgebra::Scalar> From<DMatrix<T>> for Mat<T> {
fn from(dm: DMatrix<T>) -> Self {
dmatrix_to_mat(&dm)
}
}
impl<T: Scalar + Clone + Zero + nalgebra::Scalar> From<&DMatrix<T>> for Mat<T> {
fn from(dm: &DMatrix<T>) -> Self {
dmatrix_to_mat(dm)
}
}
impl<T: Scalar + Clone + Zero + nalgebra::Scalar> From<Mat<T>> for DMatrix<T> {
fn from(mat: Mat<T>) -> Self {
mat_to_dmatrix(&mat)
}
}
impl<T: Scalar + Clone + Zero + nalgebra::Scalar> From<&Mat<T>> for DMatrix<T> {
fn from(mat: &Mat<T>) -> Self {
mat_to_dmatrix(mat)
}
}
use nalgebra::DVector;
pub fn dvector_to_mat<T: Scalar + Clone + Zero>(dv: &DVector<T>) -> Mat<T> {
let n = dv.len();
let mut mat = Mat::filled(n, 1, T::zero());
for i in 0..n {
mat[(i, 0)] = dv[i];
}
mat
}
pub fn mat_to_dvector<T: Scalar + Clone + Zero + nalgebra::Scalar>(mat: &Mat<T>) -> DVector<T> {
assert_eq!(mat.ncols(), 1, "Matrix must be a column vector");
let n = mat.nrows();
DVector::from_fn(n, |i, _| mat[(i, 0)])
}
impl<T: Scalar + Clone + Zero + nalgebra::Scalar> From<DVector<T>> for Mat<T> {
fn from(dv: DVector<T>) -> Self {
dvector_to_mat(&dv)
}
}
#[cfg(test)]
mod tests {
use super::*;
use approx::assert_relative_eq;
#[test]
fn test_dmatrix_to_mat_f64() {
let dm = DMatrix::from_fn(3, 4, |i, j| (i * 4 + j) as f64);
let mat = dmatrix_to_mat(&dm);
assert_eq!(mat.nrows(), 3);
assert_eq!(mat.ncols(), 4);
for j in 0..4 {
for i in 0..3 {
assert_relative_eq!(mat[(i, j)], dm[(i, j)], epsilon = 1e-10);
}
}
}
#[test]
fn test_mat_to_dmatrix_f64() {
let mat: Mat<f64> = Mat::from_rows(&[&[1.0, 2.0, 3.0], &[4.0, 5.0, 6.0]]);
let dm = mat_to_dmatrix(&mat);
assert_eq!(dm.nrows(), 2);
assert_eq!(dm.ncols(), 3);
for j in 0..3 {
for i in 0..2 {
assert_relative_eq!(dm[(i, j)], mat[(i, j)], epsilon = 1e-10);
}
}
}
#[test]
fn test_roundtrip_f64() {
let original: Mat<f64> = Mat::from_rows(&[&[1.0, 2.0], &[3.0, 4.0], &[5.0, 6.0]]);
let dm = mat_to_dmatrix(&original);
let recovered = dmatrix_to_mat(&dm);
assert_eq!(original.nrows(), recovered.nrows());
assert_eq!(original.ncols(), recovered.ncols());
for j in 0..original.ncols() {
for i in 0..original.nrows() {
assert_relative_eq!(original[(i, j)], recovered[(i, j)], epsilon = 1e-10);
}
}
}
#[test]
fn test_from_trait_dmatrix_to_mat() {
let dm = DMatrix::from_fn(2, 2, |i, j| (i + j) as f64);
let mat: Mat<f64> = dm.clone().into();
assert_eq!(mat[(0, 0)], 0.0);
assert_eq!(mat[(0, 1)], 1.0);
assert_eq!(mat[(1, 0)], 1.0);
assert_eq!(mat[(1, 1)], 2.0);
}
#[test]
fn test_from_trait_mat_to_dmatrix() {
let mat: Mat<f64> = Mat::from_rows(&[&[1.0, 2.0], &[3.0, 4.0]]);
let dm: DMatrix<f64> = mat.clone().into();
assert_eq!(dm[(0, 0)], 1.0);
assert_eq!(dm[(0, 1)], 2.0);
assert_eq!(dm[(1, 0)], 3.0);
assert_eq!(dm[(1, 1)], 4.0);
}
#[test]
fn test_extension_traits() {
let mat: Mat<f64> = Mat::from_rows(&[&[1.0, 2.0], &[3.0, 4.0]]);
let dm = mat.to_dmatrix();
let recovered = dm.to_mat();
assert_eq!(mat[(0, 0)], recovered[(0, 0)]);
assert_eq!(mat[(1, 1)], recovered[(1, 1)]);
}
#[test]
fn test_vector_conversions() {
let dv = DVector::from_vec(vec![1.0f64, 2.0, 3.0, 4.0]);
let mat = dvector_to_mat(&dv);
assert_eq!(mat.nrows(), 4);
assert_eq!(mat.ncols(), 1);
assert_eq!(mat[(0, 0)], 1.0);
assert_eq!(mat[(3, 0)], 4.0);
let dv2 = mat_to_dvector(&mat);
assert_eq!(dv2[0], 1.0);
assert_eq!(dv2[3], 4.0);
}
#[test]
fn test_f32_conversions() {
let dm = DMatrix::from_fn(2, 3, |i, j| (i + j) as f32);
let mat = dmatrix_to_mat(&dm);
let dm2 = mat_to_dmatrix(&mat);
for j in 0..3 {
for i in 0..2 {
assert_relative_eq!(dm[(i, j)], dm2[(i, j)], epsilon = 1e-6);
}
}
}
#[test]
fn test_complex_conversions() {
use num_complex::Complex64;
let dm = DMatrix::from_fn(2, 2, |i, j| Complex64::new((i + j) as f64, (i * j) as f64));
let mat = dmatrix_to_mat(&dm);
let dm2 = mat_to_dmatrix(&mat);
for j in 0..2 {
for i in 0..2 {
assert_relative_eq!(dm[(i, j)].re, dm2[(i, j)].re, epsilon = 1e-10);
assert_relative_eq!(dm[(i, j)].im, dm2[(i, j)].im, epsilon = 1e-10);
}
}
}
#[test]
fn test_empty_matrix() {
let dm: DMatrix<f64> = DMatrix::zeros(0, 0);
let mat = dmatrix_to_mat(&dm);
assert_eq!(mat.nrows(), 0);
assert_eq!(mat.ncols(), 0);
}
#[test]
fn test_single_element() {
let dm = DMatrix::from_element(1, 1, 42.0f64);
let mat = dmatrix_to_mat(&dm);
assert_eq!(mat.nrows(), 1);
assert_eq!(mat.ncols(), 1);
assert_eq!(mat[(0, 0)], 42.0);
}
#[test]
fn test_large_matrix() {
let dm = DMatrix::from_fn(100, 100, |i, j| (i * 100 + j) as f64);
let mat = dmatrix_to_mat(&dm);
let dm2 = mat_to_dmatrix(&mat);
assert_relative_eq!(dm[(0, 0)], dm2[(0, 0)], epsilon = 1e-10);
assert_relative_eq!(dm[(99, 99)], dm2[(99, 99)], epsilon = 1e-10);
assert_relative_eq!(dm[(50, 50)], dm2[(50, 50)], epsilon = 1e-10);
}
#[test]
fn test_mat_ref_to_dmatrix_view_zero_copy() {
let mat: Mat<f64> = Mat::from_rows(&[
&[1.0, 2.0, 3.0, 4.0, 5.0],
&[6.0, 7.0, 8.0, 9.0, 10.0],
&[11.0, 12.0, 13.0, 14.0, 15.0],
]);
assert!(mat.row_stride() > mat.nrows());
let original_ptr = mat.as_ptr();
let view = mat_ref_to_dmatrix_view(mat.as_ref());
assert_eq!(view.as_ptr(), original_ptr);
assert_eq!(view.nrows(), 3);
assert_eq!(view.ncols(), 5);
for j in 0..5 {
for i in 0..3 {
assert_eq!(view[(i, j)], mat[(i, j)]);
}
}
}
#[test]
fn test_mat_mut_to_dmatrix_view_mut_zero_copy() {
let mut mat: Mat<f64> = Mat::zeros(3, 5);
assert!(mat.row_stride() > mat.nrows());
let original_ptr = mat.as_ptr();
{
let mut view = mat_mut_to_dmatrix_view_mut(mat.as_mut());
assert_eq!(view.as_ptr(), original_ptr);
for j in 0..5 {
for i in 0..3 {
view[(i, j)] = (i * 10 + j) as f64;
}
}
}
for j in 0..5 {
for i in 0..3 {
assert_eq!(mat[(i, j)], (i * 10 + j) as f64);
}
}
}
#[test]
fn test_dmatrix_to_mat_ref_zero_copy() {
let dm = DMatrix::from_fn(4, 3, |i, j| (i * 3 + j) as f64);
let original_ptr = dm.as_ptr();
let view = dmatrix_to_mat_ref(&dm);
assert_eq!(view.as_ptr(), original_ptr);
assert_eq!(view.shape(), (4, 3));
for j in 0..3 {
for i in 0..4 {
assert_eq!(view[(i, j)], dm[(i, j)]);
}
}
}
#[test]
fn test_dmatrix_to_mat_mut_zero_copy() {
let mut dm = DMatrix::from_fn(3, 3, |_, _| 0.0f64);
let original_ptr = dm.as_ptr();
{
let mut view = dmatrix_to_mat_mut(&mut dm);
assert_eq!(view.as_ptr(), original_ptr);
view[(1, 2)] = 42.0;
}
assert_eq!(dm[(1, 2)], 42.0);
}
#[test]
fn test_nalgebra_ext_zero_copy_methods() {
let mat: Mat<f64> = Mat::from_rows(&[&[1.0, 2.0], &[3.0, 4.0]]);
let view = mat.to_dmatrix_view();
assert_eq!(view.as_ptr(), mat.as_ptr());
assert_eq!(view[(1, 0)], 3.0);
let dm = DMatrix::from_fn(2, 2, |i, j| (i + j) as f64);
let mref = dm.to_mat_ref();
assert_eq!(mref.as_ptr(), dm.as_ptr());
assert_eq!(mref[(1, 1)], 2.0);
}
}