use std::uniform_real_distribution with promoted type
This commit is contained in:
@@ -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;
|
||||
|
||||
|
||||
@@ -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();
|
||||
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user