use std::uniform_real_distribution with promoted type

This commit is contained in:
conrad
2025-07-08 23:32:03 +10:00
parent 98df75f7d7
commit 46de13b363
4 changed files with 20 additions and 23 deletions
@@ -56,6 +56,7 @@ op_chi2rnd::apply_noalias(Mat<typename T1::elem_type>& out, const Proxy<T1>& P)
arma_debug_sigprint();
typedef typename T1::elem_type eT;
// we can only make a generator for float/double/long double types
typedef typename promote_type<eT, float>::result gT;
+8 -4
View File
@@ -24,8 +24,10 @@
template<typename eT>
struct norm2est_randu_filler
{
std::mt19937_64 local_engine;
std::uniform_real_distribution<eT> local_u_distr;
typedef typename promote_type<eT, float>::result eTp;
std::mt19937_64 local_engine;
std::uniform_real_distribution<eTp> local_u_distr;
inline norm2est_randu_filler();
@@ -36,8 +38,10 @@ struct norm2est_randu_filler
template<typename T>
struct norm2est_randu_filler< std::complex<T> >
{
std::mt19937_64 local_engine;
std::uniform_real_distribution<T> local_u_distr;
typedef typename promote_type<T, float>::result Tp;
std::mt19937_64 local_engine;
std::uniform_real_distribution<Tp> local_u_distr;
inline norm2est_randu_filler();
+8 -16
View File
@@ -27,11 +27,13 @@ norm2est_randu_filler<eT>::norm2est_randu_filler()
{
arma_debug_sigprint();
typedef typename promote_type<eT, float>::result eTp;
typedef typename std::mt19937_64::result_type local_seed_type;
local_engine.seed(local_seed_type(123));
typedef typename std::uniform_real_distribution<eT>::param_type local_param_type;
typedef typename std::uniform_real_distribution<eTp>::param_type local_param_type;
local_u_distr.param(local_param_type(-1.0, +1.0));
}
@@ -57,11 +59,13 @@ norm2est_randu_filler< std::complex<T> >::norm2est_randu_filler()
{
arma_debug_sigprint();
typedef typename promote_type<T, float>::result Tp;
typedef typename std::mt19937_64::result_type local_seed_type;
local_engine.seed(local_seed_type(123));
typedef typename std::uniform_real_distribution<T>::param_type local_param_type;
typedef typename std::uniform_real_distribution<Tp>::param_type local_param_type;
local_u_distr.param(local_param_type(-1.0, +1.0));
}
@@ -120,24 +124,12 @@ op_norm2est::norm2est
if((A.n_rows == 1) || (A.n_cols == 1)) { return op_norm::vec_norm_2( Proxy< Mat<eT> >(A) ); }
// low-precision types cannot be used for norm2est_randu_filler
// (std::uniform_real_distribution is undefined for types not float/double/long double)
norm2est_randu_filler< typename promote_type<eT, float>::result > randu_filler;
norm2est_randu_filler<eT> randu_filler;
Col<eT> x(A.n_rows, fill::none);
Col<eT> y(A.n_cols, fill::none);
if(is_fp16<eT>::yes)
{
// randu_filler can only fill floats, so do that and then convert
Col<float> tmp(y.n_elem);
randu_filler.fill(tmp.memptr(), tmp.n_elem);
arrayops::convert(y.memptr(), tmp.memptr(), tmp.n_elem);
}
else
{
randu_filler.fill(y.memptr(), y.n_elem);
}
randu_filler.fill(y.memptr(), y.n_elem);
T est_old = 0;
T est_cur = 0;
+3 -3
View File
@@ -956,9 +956,9 @@ typename get_pod_type<eT>::result
op_norm::mat_norm_2(const Mat<eT>& X, typename arma_fp16_only<eT>::result* junk)
{
arma_debug_sigprint();
arma_stop_logic_error("norm(): matrix 2-norm currently not supported for fp16 type");
arma_stop_logic_error("norm(): matrix 2-norm currently not supported for fp16; try norm2est() instead");
return typename get_pod_type<eT>::result(0);
}