diff --git a/include/armadillo_bits/arma_rng.hpp b/include/armadillo_bits/arma_rng.hpp index de084d16..584fc64b 100644 --- a/include/armadillo_bits/arma_rng.hpp +++ b/include/armadillo_bits/arma_rng.hpp @@ -231,7 +231,16 @@ struct arma_rng::randi } #else { - arma_rng_cxx98::randi_fill(mem, N, a, b); + if(N == uword(1)) { arma_rng_cxx98::randi_fill(mem, uword(1), a, b); return; } + + typedef std::mt19937_64::result_type seed_type; + + std::mt19937_64 local_engine; + std::uniform_int_distribution local_i_distr(a, b); + + local_engine.seed( seed_type(std::rand()) ); + + for(uword i=0; i() ); + for(uword i=0; i < N; ++i) { mem[i] = eT( arma_rng::randu() ); } } + #else + { + if(N == uword(1)) { mem[0] = eT( arma_rng_cxx98::randu_val() ); return; } + + typedef std::mt19937_64::result_type seed_type; + + std::mt19937_64 local_engine; + std::uniform_real_distribution local_u_distr; + + local_engine.seed( seed_type(std::rand()) ); + + for(uword i=0; i < N; ++i) { mem[i] = eT( local_u_distr(local_engine) ); } + } + #endif } }; @@ -293,13 +316,44 @@ struct arma_rng::randu< std::complex > void fill(std::complex* mem, const uword N) { - for(uword i=0; i < N; ++i) + #if defined(ARMA_RNG_ALT) || defined(ARMA_USE_EXTERN_RNG) { - const T a = T( arma_rng::randu() ); - const T b = T( arma_rng::randu() ); - - mem[i] = std::complex(a, b); + for(uword i=0; i < N; ++i) + { + const T a = T( arma_rng::randu() ); + const T b = T( arma_rng::randu() ); + + mem[i] = std::complex(a, b); + } } + #else + { + if(N == uword(1)) + { + const T a = T( arma_rng_cxx98::randu_val() ); + const T b = T( arma_rng_cxx98::randu_val() ); + + mem[0] = std::complex(a, b); + + return; + } + + typedef std::mt19937_64::result_type seed_type; + + std::mt19937_64 local_engine; + std::uniform_real_distribution local_u_distr; + + local_engine.seed( seed_type(std::rand()) ); + + 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) ); + + mem[i] = std::complex(a, b); + } + } + #endif } }; @@ -353,17 +407,34 @@ struct arma_rng::randn void fill_simple(eT* mem, const uword N) { - uword i, j; - - for(i=0, j=1; j < N; i+=2, j+=2) + #if defined(ARMA_RNG_ALT) || defined(ARMA_USE_EXTERN_RNG) { - arma_rng::randn::dual_val( mem[i], mem[j] ); + uword i, j; + + for(i=0, j=1; j < N; i+=2, j+=2) + { + arma_rng::randn::dual_val( mem[i], mem[j] ); + } + + if(i < N) + { + mem[i] = eT( arma_rng::randn() ); + } } - - if(i < N) + #else { - mem[i] = eT( arma_rng::randn() ); + if(N == uword(1)) { mem[0] = eT( arma_rng_cxx98::randn_val() ); return; } + + typedef std::mt19937_64::result_type seed_type; + + std::mt19937_64 local_engine; + std::normal_distribution local_n_distr; + + local_engine.seed( seed_type(std::rand()) ); + + for(uword i=0; i < N; ++i) { mem[i] = eT( local_n_distr(local_engine) ); } } + #endif } @@ -468,10 +539,40 @@ struct arma_rng::randn< std::complex > void fill_simple(std::complex* mem, const uword N) { - for(uword i=0; i < N; ++i) + #if defined(ARMA_RNG_ALT) || defined(ARMA_USE_EXTERN_RNG) { - mem[i] = std::complex( arma_rng::randn< std::complex >() ); + for(uword i=0; i < N; ++i) { mem[i] = std::complex( arma_rng::randn< std::complex >() ); } } + #else + { + if(N == uword(1)) + { + T a = T(0); + T b = T(0); + + arma_rng_cxx98::randn_dual_val(a,b); + + mem[0] = std::complex(a,b); + + return; + } + + typedef std::mt19937_64::result_type seed_type; + + std::mt19937_64 local_engine; + std::normal_distribution local_n_distr; + + local_engine.seed( seed_type(std::rand()) ); + + 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) ); + + mem[i] = std::complex(a,b); + } + } + #endif }