use extern rng seeed instead of extern mt19937_64
This commit is contained in:
@@ -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
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user