use extern rng seeed instead of extern mt19937_64

This commit is contained in:
conrad
2021-03-22 16:14:41 +10:00
parent b305164abb
commit 7737e1be06
2 changed files with 81 additions and 68 deletions
+79 -67
View File
@@ -36,10 +36,15 @@
#if defined(ARMA_USE_EXTERN_RNG)
extern thread_local std::mt19937_64 mt19937_64_instance;
//extern thread_local std::mt19937_64 mt19937_64_instance;
extern thread_local unsigned long long extern_rng_seed;
namespace { thread_local std::mt19937_64 mt19937_64_instance; }
#endif
class arma_rng
{
public:
@@ -63,6 +68,8 @@ class arma_rng
inline static void set_seed(const seed_type val);
inline static void set_seed_random();
inline static seed_type get_extern_rng_seed();
template<typename eT> struct randi;
template<typename eT> struct randu;
template<typename eT> struct randn;
@@ -81,6 +88,8 @@ arma_rng::set_seed(const arma_rng::seed_type val)
}
#elif defined(ARMA_USE_EXTERN_RNG)
{
extern_rng_seed = val;
mt19937_64_instance.seed(val);
}
#else
@@ -173,6 +182,27 @@ arma_rng::set_seed_random()
inline
arma_rng::seed_type
arma_rng::get_extern_rng_seed()
{
#if defined(ARMA_USE_EXTERN_RNG)
{
typedef unsigned long long extern_rng_seed_type;
const extern_rng_seed_type seed_val = extern_rng_seed;
extern_rng_seed = (seed_val < std::numeric_limits<extern_rng_seed_type>::max()) ? (seed_val + extern_rng_seed_type(1)) : extern_rng_seed_type(1);
return seed_type(seed_val);
}
#endif
return seed_type(0);
}
//
@@ -191,6 +221,8 @@ struct arma_rng::randi
{
constexpr double scale = double(std::numeric_limits<int>::max()) / double(std::mt19937_64::max());
mt19937_64_instance.seed(arma_rng::get_extern_rng_seed());
return eT( double(mt19937_64_instance()) * scale );
}
#else
@@ -233,23 +265,11 @@ struct arma_rng::randi
}
#elif defined(ARMA_USE_EXTERN_RNG)
{
if(N == uword(1))
{
std::uniform_int_distribution<int> local_i_distr(a, b);
mem[0] = eT(local_i_distr(mt19937_64_instance));
return;
}
std::mt19937_64 local_engine;
std::uniform_int_distribution<int> local_i_distr(a, b);
typedef typename std::mt19937_64::result_type local_seed_type;
mt19937_64_instance.seed(arma_rng::get_extern_rng_seed());
local_engine.seed( local_seed_type(mt19937_64_instance()) );
for(uword i=0; i<N; ++i) { mem[i] = eT(local_i_distr(local_engine)); }
for(uword i=0; i<N; ++i) { mem[i] = eT(local_i_distr(mt19937_64_instance)); }
}
#else
{
@@ -288,6 +308,8 @@ struct arma_rng::randu
{
constexpr double scale = double(1.0) / double(std::mt19937_64::max());
mt19937_64_instance.seed(arma_rng::get_extern_rng_seed());
return eT( double(mt19937_64_instance()) * scale );
}
#else
@@ -309,16 +331,11 @@ struct arma_rng::randu
}
#elif defined(ARMA_USE_EXTERN_RNG)
{
if(N == uword(1)) { mem[0] = eT( arma_rng::randu<eT>() ); return; }
typedef typename std::mt19937_64::result_type local_seed_type;
std::mt19937_64 local_engine;
std::uniform_real_distribution<double> local_u_distr;
local_engine.seed( local_seed_type(mt19937_64_instance()) );
mt19937_64_instance.seed(arma_rng::get_extern_rng_seed());
for(uword i=0; i < N; ++i) { mem[i] = eT( local_u_distr(local_engine) ); }
for(uword i=0; i < N; ++i) { mem[i] = eT( local_u_distr(mt19937_64_instance) ); }
}
#else
{
@@ -345,10 +362,32 @@ struct arma_rng::randu< std::complex<T> >
arma_inline
operator std::complex<T> ()
{
const T a = T( arma_rng::randu<T>() );
const T b = T( arma_rng::randu<T>() );
return std::complex<T>(a, b);
#if defined(ARMA_RNG_ALT)
{
const T a = T( arma_rng_alt::randu_val() );
const T b = T( arma_rng_alt::randu_val() );
return std::complex<T>(a, b);
}
#elif defined(ARMA_USE_EXTERN_RNG)
{
std::uniform_real_distribution<double> local_u_distr;
mt19937_64_instance.seed(arma_rng::get_extern_rng_seed());
const T a = T( local_u_distr(mt19937_64_instance) );
const T b = T( local_u_distr(mt19937_64_instance) );
return std::complex<T>(a, b);
}
#else
{
const T a = T( arma_rng_cxx98::randu_val() );
const T b = T( arma_rng_cxx98::randu_val() );
return std::complex<T>(a, b);
}
#endif
}
@@ -369,27 +408,14 @@ struct arma_rng::randu< std::complex<T> >
}
#elif defined(ARMA_USE_EXTERN_RNG)
{
if(N == uword(1))
{
const T a = T( arma_rng::randu<T>() );
const T b = T( arma_rng::randu<T>() );
mem[0] = std::complex<T>(a, b);
return;
}
typedef typename std::mt19937_64::result_type local_seed_type;
std::mt19937_64 local_engine;
std::uniform_real_distribution<double> local_u_distr;
local_engine.seed( local_seed_type(mt19937_64_instance()) );
mt19937_64_instance.seed(arma_rng::get_extern_rng_seed());
for(uword i=0; i < N; ++i)
{
const T a = T( local_u_distr(local_engine) );
const T b = T( local_u_distr(local_engine) );
const T a = T( local_u_distr(mt19937_64_instance) );
const T b = T( local_u_distr(mt19937_64_instance) );
mem[i] = std::complex<T>(a, b);
}
@@ -445,6 +471,8 @@ struct arma_rng::randn
{
std::normal_distribution<double> local_n_distr;
mt19937_64_instance.seed(arma_rng::get_extern_rng_seed());
return eT( local_n_distr(mt19937_64_instance) );
}
#else
@@ -468,6 +496,8 @@ struct arma_rng::randn
{
std::normal_distribution<double> local_n_distr;
mt19937_64_instance.seed(arma_rng::get_extern_rng_seed());
out1 = eT( local_n_distr(mt19937_64_instance) );
out2 = eT( local_n_distr(mt19937_64_instance) );
}
@@ -490,14 +520,11 @@ struct arma_rng::randn
}
#elif defined(ARMA_USE_EXTERN_RNG)
{
typedef typename std::mt19937_64::result_type local_seed_type;
std::mt19937_64 local_engine;
std::normal_distribution<double> local_n_distr;
local_engine.seed( local_seed_type(mt19937_64_instance()) );
mt19937_64_instance.seed(arma_rng::get_extern_rng_seed());
for(uword i=0; i < N; ++i) { mem[i] = eT( local_n_distr(local_engine) ); }
for(uword i=0; i < N; ++i) { mem[i] = eT( local_n_distr(mt19937_64_instance) ); }
}
#else
{
@@ -623,17 +650,14 @@ struct arma_rng::randn< std::complex<T> >
}
#elif defined(ARMA_USE_EXTERN_RNG)
{
typedef typename std::mt19937_64::result_type local_seed_type;
std::mt19937_64 local_engine;
std::normal_distribution<double> local_n_distr;
local_engine.seed( local_seed_type(mt19937_64_instance()) );
mt19937_64_instance.seed(arma_rng::get_extern_rng_seed());
for(uword i=0; i < N; ++i)
{
const T a = T( local_n_distr(local_engine) );
const T b = T( local_n_distr(local_engine) );
const T a = T( local_n_distr(mt19937_64_instance) );
const T b = T( local_n_distr(mt19937_64_instance) );
mem[i] = std::complex<T>(a,b);
}
@@ -749,23 +773,11 @@ struct arma_rng::randg
{
#if defined(ARMA_USE_EXTERN_RNG)
{
if(N == uword(1))
{
std::gamma_distribution<double> local_g_distr(a,b);
mem[0] = eT(local_g_distr(mt19937_64_instance));
return;
}
typedef typename std::mt19937_64::result_type local_seed_type;
std::mt19937_64 local_engine;
std::gamma_distribution<double> local_g_distr(a,b);
local_engine.seed( local_seed_type(mt19937_64_instance()) );
mt19937_64_instance.seed(arma_rng::get_extern_rng_seed());
for(uword i=0; i<N; ++i) { mem[i] = eT(local_g_distr(local_engine)); }
for(uword i=0; i<N; ++i) { mem[i] = eT(local_g_distr(mt19937_64_instance)); }
}
#else
{
+2 -1
View File
@@ -40,7 +40,8 @@
#include "armadillo_bits/arma_rng_cxx11.hpp"
thread_local arma_rng_cxx11 arma_rng_cxx11_instance;
thread_local std::mt19937_64 mt19937_64_instance;
// thread_local std::mt19937_64 mt19937_64_instance;
thread_local unsigned long long extern_rng_seed = 1u;
}
#endif