From 5f334da8134a9a9a80fc22a93042e014ed821b97 Mon Sep 17 00:00:00 2001 From: conrad Date: Thu, 17 Jul 2025 16:54:41 +1000 Subject: [PATCH] don't use divide-and-conquer when matrix size is too large for LAPACK with 32 bit ints --- include/armadillo_bits/auxlib_meat.hpp | 18 ++++------- include/armadillo_bits/fn_svd.hpp | 44 ++++++++++++++++++++++++-- 2 files changed, 49 insertions(+), 13 deletions(-) diff --git a/include/armadillo_bits/auxlib_meat.hpp b/include/armadillo_bits/auxlib_meat.hpp index 3b2c53db..7d8b37ef 100644 --- a/include/armadillo_bits/auxlib_meat.hpp +++ b/include/armadillo_bits/auxlib_meat.hpp @@ -3888,9 +3888,7 @@ auxlib::svd_dc(Mat& U, Col& S, Mat& V, Mat& A) blas_int lda = blas_int(A.n_rows); blas_int ldu = blas_int(U.n_rows); blas_int ldvt = blas_int(V.n_rows); - blas_int lwork1 = 3*min_mn*min_mn + (std::max)(max_mn, 4*min_mn*min_mn + 4*min_mn); // as per LAPACK 3.2 docs - blas_int lwork2 = 4*min_mn*min_mn + 6*min_mn + max_mn; // as per LAPACK 3.8 docs; consistent with LAPACK 3.4 docs - blas_int lwork_min = (std::max)(lwork1, lwork2); // due to differences between LAPACK 3.2 and 3.8 + blas_int lwork_min = 4*min_mn*min_mn + 6*min_mn + max_mn; // as per LAPACK 3.8 and 3.12 docs; consistent with LAPACK 3.4 docs blas_int info = 0; S.set_size( static_cast(min_mn) ); @@ -3968,8 +3966,8 @@ auxlib::svd_dc(Mat< std::complex >& U, Col& S, Mat< std::complex >& V, blas_int lda = blas_int(A.n_rows); blas_int ldu = blas_int(U.n_rows); blas_int ldvt = blas_int(V.n_rows); - blas_int lwork_min = min_mn*min_mn + 2*min_mn + max_mn; // as per LAPACK 3.2, 3.4, 3.8 docs - blas_int lrwork = min_mn * ((std::max)(5*min_mn+7, 2*max_mn + 2*min_mn+1)); // as per LAPACK 3.4 docs; LAPACK 3.8 uses 5*min_mn+5 instead of 5*min_mn+7 + blas_int lwork_min = min_mn*min_mn + 2*min_mn + max_mn; // as per LAPACK 3.2, 3.4, 3.8, 3.12 docs + blas_int lrwork = min_mn * ((std::max)(5*min_mn+5, 2*max_mn + 2*min_mn+1)); // as per LAPACK 3.8 and 3.12 docs blas_int info = 0; S.set_size( static_cast(min_mn) ); @@ -4037,13 +4035,11 @@ auxlib::svd_dc_econ(Mat& U, Col& S, Mat& V, Mat& A) blas_int m = blas_int(A.n_rows); blas_int n = blas_int(A.n_cols); blas_int min_mn = (std::min)(m,n); - blas_int max_mn = (std::max)(m,n); + // blas_int max_mn = (std::max)(m,n); blas_int lda = blas_int(A.n_rows); blas_int ldu = m; blas_int ldvt = min_mn; - blas_int lwork1 = 3*min_mn*min_mn + (std::max)( max_mn, 4*min_mn*min_mn + 4*min_mn ); // as per LAPACK 3.2 docs - blas_int lwork2 = 4*min_mn*min_mn + 6*min_mn + max_mn; // as per LAPACK 3.4 docs; LAPACK 3.8 requires 4*min_mn*min_mn + 7*min_mn - blas_int lwork_min = (std::max)(lwork1, lwork2); // due to differences between LAPACK 3.2 and 3.4 + blas_int lwork_min = 4*min_mn*min_mn + 7*min_mn; // as per LAPACK 3.8 and 3.12 docs blas_int info = 0; if(A.is_empty()) @@ -4128,8 +4124,8 @@ auxlib::svd_dc_econ(Mat< std::complex >& U, Col& S, Mat< std::complex > blas_int lda = blas_int(A.n_rows); blas_int ldu = m; blas_int ldvt = min_mn; - blas_int lwork_min = min_mn*min_mn + 2*min_mn + max_mn; // as per LAPACK 3.2 docs - blas_int lrwork = min_mn * ((std::max)(5*min_mn+7, 2*max_mn + 2*min_mn+1)); // LAPACK 3.8 uses 5*min_mn+5 instead of 5*min_mn+7 + blas_int lwork_min = min_mn*min_mn + 3*min_mn; // as per LAPACK 3.12 docs + blas_int lrwork = min_mn * ((std::max)(5*min_mn+5, 2*max_mn + 2*min_mn+1)); // as per LAPACK 3.8 and 3.12 docs blas_int info = 0; if(A.is_empty()) diff --git a/include/armadillo_bits/fn_svd.hpp b/include/armadillo_bits/fn_svd.hpp index 794faf64..1dfe76b1 100644 --- a/include/armadillo_bits/fn_svd.hpp +++ b/include/armadillo_bits/fn_svd.hpp @@ -114,7 +114,27 @@ svd Mat A(X.get_ref()); - const bool status = (sig == 'd') ? auxlib::svd_dc(U, S, V, A) : auxlib::svd(U, S, V, A); + bool status = false; + + if(sig == 'd') + { + const uword N = (std::min)(A.n_rows, A.n_cols); + + const uword N_limit = (is_cx::yes) ? uword(20000) : uword(23000); + + const bool allow_dc = (sizeof(blas_int) >= std::size_t(8)) ? true : (N <= N_limit); + + if(allow_dc == false) + { + arma_warn(3, "svd(): matrix size too large for divide-and-conquer algorithm; using standard algorithm instead"); + } + + status = (allow_dc) ? auxlib::svd_dc(U, S, V, A) : auxlib::svd(U, S, V, A); + } + else + { + status = auxlib::svd(U, S, V, A); + } if(status == false) { @@ -166,7 +186,27 @@ svd_econ Mat A(X.get_ref()); - const bool status = ((mode == 'b') && (sig == 'd')) ? auxlib::svd_dc_econ(U, S, V, A) : auxlib::svd_econ(U, S, V, A, mode); + bool status = false; + + if( (mode == 'b') && (sig == 'd') ) + { + const uword N = (std::min)(A.n_rows, A.n_cols); + + const uword N_limit = (is_cx::yes) ? uword(20000) : uword(23000); + + const bool allow_dc = (sizeof(blas_int) >= std::size_t(8)) ? true : (N <= N_limit); + + if(allow_dc == false) + { + arma_warn(3, "svd_econ(): matrix size too large for divide-and-conquer algorithm; using standard algorithm instead"); + } + + status = (allow_dc) ? auxlib::svd_dc_econ(U, S, V, A) : auxlib::svd_econ(U, S, V, A, mode); + } + else + { + status = auxlib::svd_econ(U, S, V, A, mode); + } if(status == false) {