From 1eda07f74ff98af9450dbd8a33f0e04b61ad2720 Mon Sep 17 00:00:00 2001 From: conrad Date: Tue, 8 Mar 2022 14:34:28 +1000 Subject: [PATCH] don't query optimum size for small matrices --- .../newarp_TridiagEigen_meat.hpp | 40 +++++++++++-------- 1 file changed, 23 insertions(+), 17 deletions(-) diff --git a/include/armadillo_bits/newarp_TridiagEigen_meat.hpp b/include/armadillo_bits/newarp_TridiagEigen_meat.hpp index 3ae74194..1dde1b15 100644 --- a/include/armadillo_bits/newarp_TridiagEigen_meat.hpp +++ b/include/armadillo_bits/newarp_TridiagEigen_meat.hpp @@ -60,26 +60,32 @@ TridiagEigen::compute(const Mat& mat_obj) evecs.set_size(n, n); - char compz = 'I'; - blas_int lwork = 1 + 4*n + n*n; - blas_int liwork = 3 + 5*n; - blas_int info = blas_int(0); + char compz = 'I'; + blas_int lwork_min = 1 + 4*n + n*n; + blas_int liwork_min = 3 + 5*n; + blas_int info = blas_int(0); - eT work_query[2] = {}; - blas_int lwork_query = blas_int(-1); + blas_int lwork_proposed = 0; + blas_int liwork_proposed = 0; - blas_int iwork_query[2] = {}; - blas_int liwork_query = blas_int(-1); + if(n >= 32) + { + eT work_query[2] = {}; + blas_int lwork_query = blas_int(-1); + + blas_int iwork_query[2] = {}; + blas_int liwork_query = blas_int(-1); + + lapack::stedc(&compz, &n, main_diag.memptr(), sub_diag.memptr(), evecs.memptr(), &n, &work_query[0], &lwork_query, &iwork_query[0], &liwork_query, &info); + + if(info != 0) { arma_stop_runtime_error("lapack::stedc(): couldn't get size of work arrays"); return; } + + lwork_proposed = static_cast( work_query[0] ); + liwork_proposed = iwork_query[0]; + } - lapack::stedc(&compz, &n, main_diag.memptr(), sub_diag.memptr(), evecs.memptr(), &n, &work_query[0], &lwork_query, &iwork_query[0], &liwork_query, &info); - - if(info != 0) { arma_stop_runtime_error("lapack::stedc(): couldn't get size of work arrays"); return; } - - blas_int lwork_proposed = static_cast( work_query[0] ); - blas_int liwork_proposed = iwork_query[0]; - - lwork = (std::max)( lwork, lwork_proposed); - liwork = (std::max)(liwork, liwork_proposed); + blas_int lwork = (std::max)( lwork_min, lwork_proposed); + blas_int liwork = (std::max)(liwork_min, liwork_proposed); podarray work( static_cast( lwork) ); podarray iwork( static_cast(liwork) );