avoid unnecessary alias checks

This commit is contained in:
conrad
2025-10-09 19:54:41 +10:00
parent db0f9bb7eb
commit c33258d10c
3 changed files with 60 additions and 36 deletions
+1 -1
View File
@@ -5869,7 +5869,7 @@ Mat<eT>::Mat(const Glue<T1, T2, glue_type>& X)
arma_type_check(( is_same_type< eT, typename T1::elem_type >::no ));
arma_type_check(( is_same_type< eT, typename T2::elem_type >::no ));
glue_type::apply(*this, X);
glue_type::apply(static_cast< Mat_noalias<eT>& >(*this), X);
}
+16 -13
View File
@@ -40,7 +40,7 @@ struct depth_lhs< glue_type, Glue<T1,T2,glue_type> >
template<bool do_inv_detect>
template<bool do_inv_detect, bool check_alias>
struct glue_times_redirect2_helper
{
template<typename T1, typename T2>
@@ -48,8 +48,8 @@ struct glue_times_redirect2_helper
};
template<>
struct glue_times_redirect2_helper<true>
template<bool check_alias>
struct glue_times_redirect2_helper<true, check_alias>
{
template<typename T1, typename T2>
arma_hot inline static void apply(Mat<typename T1::elem_type>& out, const Glue<T1,T2,glue_times>& X);
@@ -57,7 +57,7 @@ struct glue_times_redirect2_helper<true>
template<bool do_inv_detect>
template<bool do_inv_detect, bool check_alias>
struct glue_times_redirect3_helper
{
template<typename T1, typename T2, typename T3>
@@ -65,8 +65,8 @@ struct glue_times_redirect3_helper
};
template<>
struct glue_times_redirect3_helper<true>
template<bool check_alias>
struct glue_times_redirect3_helper<true, check_alias>
{
template<typename T1, typename T2, typename T3>
arma_hot inline static void apply(Mat<typename T1::elem_type>& out, const Glue< Glue<T1,T2,glue_times>,T3,glue_times>& X);
@@ -74,7 +74,7 @@ struct glue_times_redirect3_helper<true>
template<uword N>
template<uword N, bool check_alias>
struct glue_times_redirect
{
template<typename T1, typename T2>
@@ -82,24 +82,24 @@ struct glue_times_redirect
};
template<>
struct glue_times_redirect<2>
template<bool check_alias>
struct glue_times_redirect<2, check_alias>
{
template<typename T1, typename T2>
arma_hot inline static void apply(Mat<typename T1::elem_type>& out, const Glue<T1,T2,glue_times>& X);
};
template<>
struct glue_times_redirect<3>
template<bool check_alias>
struct glue_times_redirect<3, check_alias>
{
template<typename T1, typename T2, typename T3>
arma_hot inline static void apply(Mat<typename T1::elem_type>& out, const Glue< Glue<T1,T2,glue_times>,T3,glue_times>& X);
};
template<>
struct glue_times_redirect<4>
template<bool check_alias>
struct glue_times_redirect<4, check_alias>
{
template<typename T1, typename T2, typename T3, typename T4>
arma_hot inline static void apply(Mat<typename T1::elem_type>& out, const Glue< Glue< Glue<T1,T2,glue_times>, T3, glue_times>, T4, glue_times>& X);
@@ -121,6 +121,9 @@ struct glue_times
template<typename T1, typename T2>
arma_hot inline static void apply(Mat<typename T1::elem_type>& out, const Glue<T1,T2,glue_times>& X);
template<typename T1, typename T2>
arma_hot inline static void apply(Mat_noalias<typename T1::elem_type>& out, const Glue<T1,T2,glue_times>& X);
template<typename T1>
arma_hot inline static void apply_inplace(Mat<typename T1::elem_type>& out, const T1& X);
+43 -22
View File
@@ -21,11 +21,11 @@
template<bool do_inv_detect>
template<bool do_inv_detect, bool check_alias>
template<typename T1, typename T2>
inline
void
glue_times_redirect2_helper<do_inv_detect>::apply(Mat<typename T1::elem_type>& out, const Glue<T1,T2,glue_times>& X)
glue_times_redirect2_helper<do_inv_detect, check_alias>::apply(Mat<typename T1::elem_type>& out, const Glue<T1,T2,glue_times>& X)
{
arma_debug_sigprint();
@@ -55,7 +55,7 @@ glue_times_redirect2_helper<do_inv_detect>::apply(Mat<typename T1::elem_type>& o
return;
}
const bool alias = U1.is_alias(out) || U2.is_alias(out);
const bool alias = (check_alias) && (U1.is_alias(out) || U2.is_alias(out));
if(alias == false)
{
@@ -87,10 +87,11 @@ glue_times_redirect2_helper<do_inv_detect>::apply(Mat<typename T1::elem_type>& o
template<bool check_alias>
template<typename T1, typename T2>
inline
void
glue_times_redirect2_helper<true>::apply(Mat<typename T1::elem_type>& out, const Glue<T1,T2,glue_times>& X)
glue_times_redirect2_helper<true, check_alias>::apply(Mat<typename T1::elem_type>& out, const Glue<T1,T2,glue_times>& X)
{
arma_debug_sigprint();
@@ -148,7 +149,7 @@ glue_times_redirect2_helper<true>::apply(Mat<typename T1::elem_type>& out, const
if(is_cx<eT>::yes) { arma_warn(1, "inv_sympd(): given matrix is not hermitian"); }
}
const unwrap_check<T2> B_tmp(X.B, out);
const unwrap_check<T2> B_tmp(X.B, out); // TODO: refactor to use quasi_unwrap
const Mat<eT>& B = B_tmp.M;
arma_conform_assert_mul_size(A, B, "matrix multiplication");
@@ -202,16 +203,16 @@ glue_times_redirect2_helper<true>::apply(Mat<typename T1::elem_type>& out, const
return;
}
glue_times_redirect2_helper<false>::apply(out, X);
glue_times_redirect2_helper<false, check_alias>::apply(out, X);
}
template<bool do_inv_detect>
template<bool do_inv_detect, bool check_alias>
template<typename T1, typename T2, typename T3>
inline
void
glue_times_redirect3_helper<do_inv_detect>::apply(Mat<typename T1::elem_type>& out, const Glue< Glue<T1,T2,glue_times>, T3, glue_times>& X)
glue_times_redirect3_helper<do_inv_detect, check_alias>::apply(Mat<typename T1::elem_type>& out, const Glue< Glue<T1,T2,glue_times>, T3, glue_times>& X)
{
arma_debug_sigprint();
@@ -231,7 +232,7 @@ glue_times_redirect3_helper<do_inv_detect>::apply(Mat<typename T1::elem_type>& o
constexpr bool use_alpha = partial_unwrap<T1>::do_times || partial_unwrap<T2>::do_times || partial_unwrap<T3>::do_times;
const eT alpha = use_alpha ? (U1.get_val() * U2.get_val() * U3.get_val()) : eT(0);
const bool alias = U1.is_alias(out) || U2.is_alias(out) || U3.is_alias(out);
const bool alias = (check_alias) && (U1.is_alias(out) || U2.is_alias(out) || U3.is_alias(out));
if(alias == false)
{
@@ -265,10 +266,11 @@ glue_times_redirect3_helper<do_inv_detect>::apply(Mat<typename T1::elem_type>& o
template<bool check_alias>
template<typename T1, typename T2, typename T3>
inline
void
glue_times_redirect3_helper<true>::apply(Mat<typename T1::elem_type>& out, const Glue< Glue<T1,T2,glue_times>, T3, glue_times>& X)
glue_times_redirect3_helper<true, check_alias>::apply(Mat<typename T1::elem_type>& out, const Glue< Glue<T1,T2,glue_times>, T3, glue_times>& X)
{
arma_debug_sigprint();
@@ -371,7 +373,7 @@ glue_times_redirect3_helper<true>::apply(Mat<typename T1::elem_type>& out, const
constexpr bool use_alpha = partial_unwrap<T1>::do_times;
const eT alpha = use_alpha ? U1.get_val() : eT(0);
if(U1.is_alias(out))
if( (check_alias) && U1.is_alias(out) )
{
Mat<eT> tmp;
@@ -388,16 +390,16 @@ glue_times_redirect3_helper<true>::apply(Mat<typename T1::elem_type>& out, const
}
glue_times_redirect3_helper<false>::apply(out, X);
glue_times_redirect3_helper<false, check_alias>::apply(out, X);
}
template<uword N>
template<uword N, bool check_alias>
template<typename T1, typename T2>
inline
void
glue_times_redirect<N>::apply(Mat<typename T1::elem_type>& out, const Glue<T1,T2,glue_times>& X)
glue_times_redirect<N, check_alias>::apply(Mat<typename T1::elem_type>& out, const Glue<T1,T2,glue_times>& X)
{
arma_debug_sigprint();
@@ -412,7 +414,7 @@ glue_times_redirect<N>::apply(Mat<typename T1::elem_type>& out, const Glue<T1,T2
constexpr bool use_alpha = partial_unwrap<T1>::do_times || partial_unwrap<T2>::do_times;
const eT alpha = use_alpha ? (U1.get_val() * U2.get_val()) : eT(0);
const bool alias = U1.is_alias(out) || U2.is_alias(out);
const bool alias = (check_alias) && (U1.is_alias(out) || U2.is_alias(out));
if(alias == false)
{
@@ -444,38 +446,41 @@ glue_times_redirect<N>::apply(Mat<typename T1::elem_type>& out, const Glue<T1,T2
template<bool check_alias>
template<typename T1, typename T2>
inline
void
glue_times_redirect<2>::apply(Mat<typename T1::elem_type>& out, const Glue<T1,T2,glue_times>& X)
glue_times_redirect<2, check_alias>::apply(Mat<typename T1::elem_type>& out, const Glue<T1,T2,glue_times>& X)
{
arma_debug_sigprint();
typedef typename T1::elem_type eT;
glue_times_redirect2_helper< is_blas_type<eT>::value >::apply(out, X);
glue_times_redirect2_helper< is_blas_type<eT>::value, check_alias >::apply(out, X);
}
template<bool check_alias>
template<typename T1, typename T2, typename T3>
inline
void
glue_times_redirect<3>::apply(Mat<typename T1::elem_type>& out, const Glue< Glue<T1,T2,glue_times>, T3, glue_times>& X)
glue_times_redirect<3, check_alias>::apply(Mat<typename T1::elem_type>& out, const Glue< Glue<T1,T2,glue_times>, T3, glue_times>& X)
{
arma_debug_sigprint();
typedef typename T1::elem_type eT;
glue_times_redirect3_helper< is_blas_type<eT>::value >::apply(out, X);
glue_times_redirect3_helper< is_blas_type<eT>::value, check_alias >::apply(out, X);
}
template<bool check_alias>
template<typename T1, typename T2, typename T3, typename T4>
inline
void
glue_times_redirect<4>::apply(Mat<typename T1::elem_type>& out, const Glue< Glue< Glue<T1,T2,glue_times>, T3, glue_times>, T4, glue_times>& X)
glue_times_redirect<4, check_alias>::apply(Mat<typename T1::elem_type>& out, const Glue< Glue< Glue<T1,T2,glue_times>, T3, glue_times>, T4, glue_times>& X)
{
arma_debug_sigprint();
@@ -497,7 +502,7 @@ glue_times_redirect<4>::apply(Mat<typename T1::elem_type>& out, const Glue< Glue
constexpr bool use_alpha = partial_unwrap<T1>::do_times || partial_unwrap<T2>::do_times || partial_unwrap<T3>::do_times || partial_unwrap<T4>::do_times;
const eT alpha = use_alpha ? (U1.get_val() * U2.get_val() * U3.get_val() * U4.get_val()) : eT(0);
const bool alias = U1.is_alias(out) || U2.is_alias(out) || U3.is_alias(out) || U4.is_alias(out);
const bool alias = (check_alias) && (U1.is_alias(out) || U2.is_alias(out) || U3.is_alias(out) || U4.is_alias(out));
if(alias == false)
{
@@ -544,7 +549,23 @@ glue_times::apply(Mat<typename T1::elem_type>& out, const Glue<T1,T2,glue_times>
arma_debug_print(arma_str::format("glue_times::apply(): N_mat: %u") % N_mat);
glue_times_redirect<N_mat>::apply(out, X);
glue_times_redirect<N_mat, true>::apply(out, X);
}
template<typename T1, typename T2>
inline
void
glue_times::apply(Mat_noalias<typename T1::elem_type>& out, const Glue<T1,T2,glue_times>& X)
{
arma_debug_sigprint();
constexpr uword N_mat = 1 + depth_lhs< glue_times, Glue<T1,T2,glue_times> >::num;
arma_debug_print(arma_str::format("glue_times::apply(): N_mat: %u") % N_mat);
glue_times_redirect<N_mat, false>::apply(out, X);
}