diff --git a/include/armadillo_bits/op_dot_meat.hpp b/include/armadillo_bits/op_dot_meat.hpp index 84356732..2590637c 100644 --- a/include/armadillo_bits/op_dot_meat.hpp +++ b/include/armadillo_bits/op_dot_meat.hpp @@ -29,33 +29,35 @@ op_dot::direct_dot_generic(const uword n_elem, const eT* const A, const eT* cons { arma_debug_sigprint(); + typedef typename conditional_promote_type::value, eT, float>::result acc_eT; + #if defined(__FAST_MATH__) { - eT val = eT(0); + acc_eT val = acc_eT(0); - for(uword i=0; i < n_elem; ++i) { val += A[i] * B[i]; } + for(uword i=0; i < n_elem; ++i) { val += acc_eT( A[i] * B[i] ); } - return val; + return eT(val); } #else { - eT val1 = eT(0); - eT val2 = eT(0); + acc_eT val1 = acc_eT(0); + acc_eT val2 = acc_eT(0); uword i, j; for(i=0, j=1; j < n_elem; i+=2, j+=2) { - val1 += A[i] * B[i]; - val2 += A[j] * B[j]; + val1 += acc_eT( A[i] * B[i] ); + val2 += acc_eT( A[j] * B[j] ); } if(i < n_elem) { - val1 += A[i] * B[i]; + val1 += acc_eT( A[i] * B[i] ); } - return val1 + val2; + return eT( val1 + val2 ); } #endif } @@ -72,8 +74,10 @@ op_dot::direct_dot_generic(const uword n_elem, const eT* const A, const eT* cons typedef typename get_pod_type::result T; - T val_real = T(0); - T val_imag = T(0); + typedef typename conditional_promote_type::value, T, float>::result acc_T; + + acc_T val_real = acc_T(0); + acc_T val_imag = acc_T(0); for(uword i=0; i(val_real, val_imag); + return std::complex( T(val_real), T(val_imag) ); } @@ -165,25 +169,7 @@ op_dot::direct_dot(const uword n_elem, const eT* const A, const eT* const B) { arma_debug_sigprint(); - typedef typename promote_type::result acc_eT; - - acc_eT val1 = acc_eT(0); - acc_eT val2 = acc_eT(0); - - uword i, j; - - for(i=0, j=1; j < n_elem; i+=2, j+=2) - { - val1 += acc_eT(A[i] * B[i]); - val2 += acc_eT(A[j] * B[j]); - } - - if(i < n_elem) - { - val1 += acc_eT(A[i] * B[i]); - } - - return eT(val1 + val2); + return op_dot::direct_dot_generic(n_elem, A, B); } @@ -196,29 +182,7 @@ op_dot::direct_dot(const uword n_elem, const eT* const A, const eT* const B) { arma_debug_sigprint(); - typedef typename get_pod_type::result T; - - typedef typename promote_type::result acc_T; - - acc_T val_real = acc_T(0); - acc_T val_imag = acc_T(0); - - for(uword i=0; i& X = A[i]; - const std::complex& Y = B[i]; - - const T a = X.real(); - const T b = X.imag(); - - const T c = Y.real(); - const T d = Y.imag(); - - val_real += acc_T(a*c) - acc_T(b*d); - val_imag += acc_T(a*d) + acc_T(b*c); - } - - return std::complex( T(val_real), T(val_imag) ); + return op_dot::direct_dot_generic(n_elem, A, B); } @@ -369,7 +333,7 @@ op_dot::apply_proxy_linear(const Proxy& PA, const Proxy& PB) typedef typename T1::elem_type eT; - typedef typename promote_type::result acc_eT; + typedef typename conditional_promote_type::value, eT, float>::result acc_eT; typedef typename Proxy::ea_type ea_type1; typedef typename Proxy::ea_type ea_type2; @@ -386,16 +350,16 @@ op_dot::apply_proxy_linear(const Proxy& PA, const Proxy& PB) for(i=0, j=1; j& PA, const Proxy& PB) typedef typename T1::elem_type eT; typedef typename get_pod_type::result T; - typedef typename promote_type::result acc_T; + typedef typename conditional_promote_type::value, T, float>::result acc_T; typedef typename Proxy::ea_type ea_type1; typedef typename Proxy::ea_type ea_type2; diff --git a/include/armadillo_bits/promote_type.hpp b/include/armadillo_bits/promote_type.hpp index a4f76049..60adf49f 100644 --- a/include/armadillo_bits/promote_type.hpp +++ b/include/armadillo_bits/promote_type.hpp @@ -340,4 +340,11 @@ struct eT_promoter +template struct conditional_promote_type { }; + +template struct conditional_promote_type { typedef eT1 result; }; +template struct conditional_promote_type { typedef typename promote_type::result result; }; + + + //! @}