From 7737e1be065e6f2afc0cc411babe1d4eb5ced5a2 Mon Sep 17 00:00:00 2001 From: conrad Date: Mon, 22 Mar 2021 16:14:41 +1000 Subject: [PATCH] use extern rng seeed instead of extern mt19937_64 --- include/armadillo_bits/arma_rng.hpp | 146 +++++++++++++++------------- src/wrapper1.cpp | 3 +- 2 files changed, 81 insertions(+), 68 deletions(-) diff --git a/include/armadillo_bits/arma_rng.hpp b/include/armadillo_bits/arma_rng.hpp index da4bdce0..a9d089f3 100644 --- a/include/armadillo_bits/arma_rng.hpp +++ b/include/armadillo_bits/arma_rng.hpp @@ -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 struct randi; template struct randu; template 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::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::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 local_i_distr(a, b); - - mem[0] = eT(local_i_distr(mt19937_64_instance)); - - return; - } - - std::mt19937_64 local_engine; std::uniform_int_distribution 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() ); return; } - - typedef typename std::mt19937_64::result_type local_seed_type; - - std::mt19937_64 local_engine; std::uniform_real_distribution 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 > arma_inline operator std::complex () { - const T a = T( arma_rng::randu() ); - const T b = T( arma_rng::randu() ); - - return std::complex(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(a, b); + } + #elif defined(ARMA_USE_EXTERN_RNG) + { + std::uniform_real_distribution 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(a, b); + } + #else + { + const T a = T( arma_rng_cxx98::randu_val() ); + const T b = T( arma_rng_cxx98::randu_val() ); + + return std::complex(a, b); + } + #endif } @@ -369,27 +408,14 @@ struct arma_rng::randu< std::complex > } #elif defined(ARMA_USE_EXTERN_RNG) { - if(N == uword(1)) - { - const T a = T( arma_rng::randu() ); - const T b = T( arma_rng::randu() ); - - mem[0] = std::complex(a, b); - - return; - } - - typedef typename std::mt19937_64::result_type local_seed_type; - - std::mt19937_64 local_engine; std::uniform_real_distribution 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(a, b); } @@ -445,6 +471,8 @@ struct arma_rng::randn { std::normal_distribution 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 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 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 > } #elif defined(ARMA_USE_EXTERN_RNG) { - typedef typename std::mt19937_64::result_type local_seed_type; - - std::mt19937_64 local_engine; std::normal_distribution 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(a,b); } @@ -749,23 +773,11 @@ struct arma_rng::randg { #if defined(ARMA_USE_EXTERN_RNG) { - if(N == uword(1)) - { - std::gamma_distribution 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 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