From c1958d2aa743697ac1bbb3305c61129d60a7fade Mon Sep 17 00:00:00 2001 From: Lloyd Chan Date: Thu, 17 Aug 2017 15:44:44 +0900 Subject: [PATCH] Adds blas support for transposed row matrices (Closes #340) Adds a few unit tests and enables running unit tests for the main lib. --- Cargo.toml | 2 +- src/linalg/impl_linalg.rs | 47 +++++++++++++++++++++++++++++++++++++-- 2 files changed, 46 insertions(+), 3 deletions(-) diff --git a/Cargo.toml b/Cargo.toml index 02c5d723c..d78dfd4dd 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -22,7 +22,7 @@ exclude = ["docgen/images/*"] [lib] name = "ndarray" bench = false -test = false +test = true [dependencies.num-traits] version = "0.1.32" diff --git a/src/linalg/impl_linalg.rs b/src/linalg/impl_linalg.rs index 2b56f36c1..0128000ca 100644 --- a/src/linalg/impl_linalg.rs +++ b/src/linalg/impl_linalg.rs @@ -626,9 +626,10 @@ fn blas_row_major_2d(a: &ArrayBase) -> bool if !same_type::() { return false; } + let (m, n) = a.dim(); let s0 = a.strides()[0]; let s1 = a.strides()[1]; - if s1 != 1 { + if !(s1 == 1 || n == 1) { return false; } if s0 < 1 || s1 < 1 { @@ -639,7 +640,6 @@ fn blas_row_major_2d(a: &ArrayBase) -> bool { return false; } - let (m, n) = a.dim(); if m > blas_index::max_value() as usize || n > blas_index::max_value() as usize { @@ -647,3 +647,46 @@ fn blas_row_major_2d(a: &ArrayBase) -> bool } true } + +#[cfg(test)] +mod tests { + + use super::*; + + #[test] + #[cfg(feature="blas")] + fn blas_row_major_2d_normal_matrix() { + let m: Array2 = Array2::zeros((3, 5)); + assert!(blas_row_major_2d::(&m)); + } + + #[test] + #[cfg(feature="blas")] + fn blas_row_major_2d_row_matrix() { + let m: Array2 = Array2::zeros((1, 5)); + assert!(blas_row_major_2d::(&m)); + } + + #[test] + #[cfg(feature="blas")] + fn blas_row_major_2d_column_matrix() { + let m: Array2 = Array2::zeros((5, 1)); + assert!(blas_row_major_2d::(&m)); + } + + #[test] + #[cfg(feature="blas")] + fn blas_row_major_2d_transposed_row_matrix() { + let m: Array2 = Array2::zeros((1, 5)); + let m_t = m.t(); + assert!(blas_row_major_2d::(&m_t)); + } + + #[test] + #[cfg(feature="blas")] + fn blas_row_major_2d_transposed_column_matrix() { + let m: Array2 = Array2::zeros((5, 1)); + let m_t = m.t(); + assert!(blas_row_major_2d::(&m_t)); + } +} \ No newline at end of file