From 91dbb8c2ea26ffb1e7e4844882108ace495f2207 Mon Sep 17 00:00:00 2001 From: conrad Date: Mon, 7 Aug 2023 10:07:46 +1000 Subject: [PATCH] rework arna_rng to use thread-safe mersenne twister as default; ensure unique seeds for each thread --- CMakeLists.txt | 36 +-- include/armadillo | 2 +- include/armadillo_bits/arma_rng.hpp | 350 +++++++++++++++++++----- include/armadillo_bits/config.hpp | 17 -- include/armadillo_bits/config.hpp.cmake | 17 -- src/wrapper1.cpp | 15 +- 6 files changed, 293 insertions(+), 144 deletions(-) diff --git a/CMakeLists.txt b/CMakeLists.txt index 7d21cf2d..a67ef800 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -53,12 +53,11 @@ set(ARMA_USE_WRAPPER true) # the settings below will be automatically configured by the rest of this script -set(ARMA_USE_LAPACK false) -set(ARMA_USE_BLAS false) -set(ARMA_USE_ATLAS false) -set(ARMA_USE_ARPACK false) -set(ARMA_USE_EXTERN_RNG false) -set(ARMA_USE_SUPERLU false) # Caveat: only SuperLU version 5.x can be used! +set(ARMA_USE_LAPACK false) +set(ARMA_USE_BLAS false) +set(ARMA_USE_ATLAS false) +set(ARMA_USE_ARPACK false) +set(ARMA_USE_SUPERLU false) # Caveat: only SuperLU version 5.x can be used! ## extract version from sources @@ -84,7 +83,7 @@ if(NOT CXX_FLAGS_EMPTY) endif() -# NOTE: ARMA_USE_EXTERN_RNG requires compiler support for thread_local and C++11 +# NOTE: Armadillo requires compiler support for thread_local and C++11 # NOTE: for Linux, this is available with gcc 4.8.3 onwards # NOTE: for macOS, thread_local is supoported in Xcode 8 (mid 2016 onwards) in C++11 mode @@ -94,7 +93,6 @@ endif() if(DEFINED CMAKE_CXX_COMPILER_ID AND DEFINED CMAKE_CXX_COMPILER_VERSION) if(CMAKE_CXX_COMPILER_ID STREQUAL "GNU") if(NOT (${CMAKE_CXX_COMPILER_VERSION} VERSION_LESS 4.8.3)) - set(ARMA_USE_EXTERN_RNG true) message(STATUS "Detected gcc 4.8.3 or newer") if(${CMAKE_CXX_COMPILER_VERSION} VERSION_LESS 6.1.0) message(STATUS "*** WARNING: support for gcc versions older than 6.1 is deprecated") @@ -110,7 +108,6 @@ if(DEFINED CMAKE_CXX_COMPILER_ID AND DEFINED CMAKE_CXX_COMPILER_VERSION) if(NOT (${CMAKE_MAJOR_VERSION} LESS 3)) if(CMAKE_CXX_COMPILER_ID STREQUAL "Clang") if(NOT ${CMAKE_CXX_COMPILER_VERSION} VERSION_LESS 6.0) - set(ARMA_USE_EXTERN_RNG true) message(STATUS "Detected Clang 6.0 or newer") if(NOT DEFINED CMAKE_CXX_STANDARD) set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -std=c++14") @@ -121,7 +118,6 @@ if(DEFINED CMAKE_CXX_COMPILER_ID AND DEFINED CMAKE_CXX_COMPILER_VERSION) endif() elseif(CMAKE_CXX_COMPILER_ID STREQUAL "AppleClang") if(NOT ${CMAKE_CXX_COMPILER_VERSION} VERSION_LESS 8.0) - set(ARMA_USE_EXTERN_RNG true) message(STATUS "Detected AppleClang 8.0 or newer") if(NOT DEFINED CMAKE_CXX_STANDARD) set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -std=c++14") @@ -135,13 +131,6 @@ if(DEFINED CMAKE_CXX_COMPILER_ID AND DEFINED CMAKE_CXX_COMPILER_VERSION) endif() endif() -if(MINGW OR MSYS OR CYGWIN OR MSVC) - # MinGW doesn't correctly handle thread_local - set(ARMA_USE_EXTERN_RNG false) -endif() - -message(STATUS "ARMA_USE_EXTERN_RNG = ${ARMA_USE_EXTERN_RNG}") - # As Red Hat Enterprise Linux (and related systems such as Fedora) # does not search /usr/local/lib by default, we need to place the @@ -488,13 +477,12 @@ endif() message(STATUS "") message(STATUS "*** Result of configuration:") -message(STATUS "*** ARMA_USE_WRAPPER = ${ARMA_USE_WRAPPER}") -message(STATUS "*** ARMA_USE_LAPACK = ${ARMA_USE_LAPACK}") -message(STATUS "*** ARMA_USE_BLAS = ${ARMA_USE_BLAS}") -message(STATUS "*** ARMA_USE_ATLAS = ${ARMA_USE_ATLAS}") -message(STATUS "*** ARMA_USE_ARPACK = ${ARMA_USE_ARPACK}") -message(STATUS "*** ARMA_USE_EXTERN_RNG = ${ARMA_USE_EXTERN_RNG}") -message(STATUS "*** ARMA_USE_SUPERLU = ${ARMA_USE_SUPERLU}") +message(STATUS "*** ARMA_USE_WRAPPER = ${ARMA_USE_WRAPPER}") +message(STATUS "*** ARMA_USE_LAPACK = ${ARMA_USE_LAPACK}") +message(STATUS "*** ARMA_USE_BLAS = ${ARMA_USE_BLAS}") +message(STATUS "*** ARMA_USE_ATLAS = ${ARMA_USE_ATLAS}") +message(STATUS "*** ARMA_USE_ARPACK = ${ARMA_USE_ARPACK}") +message(STATUS "*** ARMA_USE_SUPERLU = ${ARMA_USE_SUPERLU}") message(STATUS "") message(STATUS "*** Armadillo wrapper library will use the following libraries:") message(STATUS "*** ARMA_LIBS = ${ARMA_LIBS}") diff --git a/include/armadillo b/include/armadillo index 31ca6dbf..59694d1c 100644 --- a/include/armadillo +++ b/include/armadillo @@ -50,10 +50,10 @@ #include #include #include +#include #if !defined(ARMA_DONT_USE_STD_MUTEX) #include - #include #endif // #if defined(ARMA_HAVE_CXX17) diff --git a/include/armadillo_bits/arma_rng.hpp b/include/armadillo_bits/arma_rng.hpp index 6c5eb058..da1b4f7a 100644 --- a/include/armadillo_bits/arma_rng.hpp +++ b/include/armadillo_bits/arma_rng.hpp @@ -20,55 +20,63 @@ //! @{ -#if defined(ARMA_RNG_ALT) - #undef ARMA_USE_EXTERN_RNG +#undef ARMA_USE_CXX11_RNG +#define ARMA_USE_CXX11_RNG + +#undef ARMA_USE_THREAD_LOCAL +#define ARMA_USE_THREAD_LOCAL + +#if (defined(ARMA_RNG_ALT) || defined(ARMA_DONT_USE_CXX11_RNG)) + #undef ARMA_USE_CXX11_RNG +#endif + +#if defined(ARMA_DONT_USE_THREAD_LOCAL) + #undef ARMA_USE_THREAD_LOCAL #endif -// NOTE: mt19937_64_instance_warmup is used as a workaround +// NOTE: ARMA_WARMUP_PRODUCER enables a workaround // NOTE: for thread_local issue on macOS 11 and/or AppleClang 12.0 // NOTE: see https://gitlab.com/conradsnicta/armadillo-code/-/issues/173 // NOTE: if this workaround causes problems, please report it and -// NOTE: disable the workaround by uncommenting the code block below: +// NOTE: disable the workaround by commenting out the code block below: -// #if defined(__APPLE__) || defined(__apple_build_version__) -// #if !defined(ARMA_DONT_DISABLE_EXTERN_RNG) -// #undef ARMA_USE_EXTERN_RNG -// #endif -// #endif +#if defined(__APPLE__) || defined(__apple_build_version__) + #undef ARMA_WARMUP_PRODUCER + #define ARMA_WARMUP_PRODUCER +#endif +#if defined(ARMA_DONT_WARMUP_PRODUCER) + #undef ARMA_WARMUP_PRODUCER +#endif // NOTE: workaround for another thread_local issue on macOS // NOTE: where GCC (not Clang) may not have support for thread_local #if (defined(__APPLE__) && defined(__GNUG__) && !defined(__clang__)) - #if !defined(ARMA_DONT_DISABLE_EXTERN_RNG) - #undef ARMA_USE_EXTERN_RNG - #endif + #undef ARMA_USE_THREAD_LOCAL #endif +// NOTE: disable use of thread_local on MinGW et al; +// NOTE: i don't have the patience to keep looking into these broken platforms - -#if defined(ARMA_USE_EXTERN_RNG) - extern thread_local std::mt19937_64 mt19937_64_instance; - - #if defined(__APPLE__) || defined(__apple_build_version__) - namespace - { - struct mt19937_64_instance_warmup - { - inline mt19937_64_instance_warmup() - { - typename std::mt19937_64::result_type junk = mt19937_64_instance(); - arma_ignore(junk); - } - }; - - static mt19937_64_instance_warmup mt19937_64_instance_warmup_run; - } - #endif +#if (defined(__MINGW32__) || defined(__MINGW64__) || defined(__CYGWIN__) || defined(__MSYS__) || defined(__MSYS2__)) + #undef ARMA_USE_THREAD_LOCAL #endif +#if defined(ARMA_FORCE_USE_THREAD_LOCAL) + #undef ARMA_USE_THREAD_LOCAL + #define ARMA_USE_THREAD_LOCAL +#endif + +#if (!defined(ARMA_USE_THREAD_LOCAL)) + #undef ARMA_GUARD_PRODUCER + #define ARMA_GUARD_PRODUCER +#endif + +#if (defined(ARMA_DONT_GUARD_PRODUCER) || defined(ARMA_DONT_USE_STD_MUTEX)) + #undef ARMA_GUARD_PRODUCER +#endif class arma_rng @@ -77,7 +85,7 @@ class arma_rng #if defined(ARMA_RNG_ALT) typedef arma_rng_alt::seed_type seed_type; - #elif defined(ARMA_USE_EXTERN_RNG) + #elif defined(ARMA_USE_CXX11_RNG) typedef std::mt19937_64::result_type seed_type; #else typedef arma_rng_cxx03::seed_type seed_type; @@ -85,12 +93,24 @@ class arma_rng #if defined(ARMA_RNG_ALT) static constexpr int rng_method = 2; - #elif defined(ARMA_USE_EXTERN_RNG) + #elif defined(ARMA_USE_CXX11_RNG) static constexpr int rng_method = 1; #else static constexpr int rng_method = 0; #endif + #if defined(ARMA_USE_CXX11_RNG) + inline static std::mt19937_64& get_producer(); + inline static void warmup_producer(std::mt19937_64& producer); + + inline static void lock_producer(); + inline static void unlock_producer(); + + #if defined(ARMA_GUARD_PRODUCER) + inline static std::mutex& get_producer_mutex(); + #endif + #endif + inline static void set_seed(const seed_type val); inline static void set_seed_random(); @@ -102,6 +122,101 @@ class arma_rng +#if defined(ARMA_USE_CXX11_RNG) + +inline +std::mt19937_64& +arma_rng::get_producer() + { + #if defined(ARMA_USE_THREAD_LOCAL) + + // use a thread-safe RNG, with each thread having its own unique starting seed + + static std::atomic mt19937_64_producer_counter(0); + + static thread_local std::mt19937_64 mt19937_64_producer( std::mt19937_64::default_seed + mt19937_64_producer_counter++ ); + + arma_rng::warmup_producer(mt19937_64_producer); + + #else + + // use a plain RNG in case we don't have thread_local + + static std::mt19937_64 mt19937_64_producer( std::mt19937_64::default_seed ); + + arma_rng::warmup_producer(mt19937_64_producer); + + #endif + + return mt19937_64_producer; + } + + +inline +void +arma_rng::warmup_producer(std::mt19937_64& producer) + { + #if defined(ARMA_WARMUP_PRODUCER) + + static std::atomic_flag warmup_done = ATOMIC_FLAG_INIT; // init to false + + if(warmup_done.test_and_set() == false) + { + typename std::mt19937_64::result_type junk = producer(); + + arma_ignore(junk); + } + + #else + + arma_ignore(producer); + + #endif + } + + +inline +void +arma_rng::lock_producer() + { + #if defined(ARMA_GUARD_PRODUCER) + + std::mutex& producer_mutex = arma_rng::get_producer_mutex(); + + producer_mutex.lock(); + + #endif + } + + +inline +void +arma_rng::unlock_producer() + { + #if defined(ARMA_GUARD_PRODUCER) + + std::mutex& producer_mutex = arma_rng::get_producer_mutex(); + + producer_mutex.unlock(); + + #endif + } + + +#if defined(ARMA_GUARD_PRODUCER) + inline + std::mutex& + arma_rng::get_producer_mutex() + { + static std::mutex producer_mutex; + + return producer_mutex; + } +#endif + +#endif + + inline void arma_rng::set_seed(const arma_rng::seed_type val) @@ -110,9 +225,11 @@ arma_rng::set_seed(const arma_rng::seed_type val) { arma_rng_alt::set_seed(val); } - #elif defined(ARMA_USE_EXTERN_RNG) + #elif defined(ARMA_USE_CXX11_RNG) { - mt19937_64_instance.seed(val); + arma_rng::lock_producer(); + arma_rng::get_producer().seed(val); + arma_rng::unlock_producer(); } #else { @@ -141,7 +258,7 @@ arma_rng::set_seed_random() if(rd.entropy() > double(0)) { seed1 = static_cast( rd() ); } - if(seed1 != seed_type(0)) { have_seed = true; } + have_seed = (seed1 != seed_type(0)); } catch(...) {} @@ -162,12 +279,9 @@ arma_rng::set_seed_random() if(f.good()) { f.read((char*)(&(tmp.b[0])), sizeof(seed_type)); } - if(f.good()) - { - seed2 = tmp.a; + if(f.good()) { seed2 = tmp.a; } - if(seed2 != seed_type(0)) { have_seed = true; } - } + have_seed = (seed2 != seed_type(0)); } catch(...) {} } @@ -199,7 +313,7 @@ arma_rng::set_seed_random() } } - arma_rng::set_seed( seed1 + seed2 + seed3 + seed4 ); + arma_rng::set_seed(seed1 + seed2 + seed3 + seed4); } @@ -218,11 +332,17 @@ struct arma_rng::randi { return eT( arma_rng_alt::randi_val() ); } - #elif defined(ARMA_USE_EXTERN_RNG) + #elif defined(ARMA_USE_CXX11_RNG) { constexpr double scale = double(std::numeric_limits::max()) / double(std::mt19937_64::max()); - return eT( double(mt19937_64_instance()) * scale ); + arma_rng::lock_producer(); + + const eT out = eT(double(arma_rng::get_producer()()) * scale); + + arma_rng::unlock_producer(); + + return out; } #else { @@ -241,7 +361,7 @@ struct arma_rng::randi { return arma_rng_alt::randi_max_val(); } - #elif defined(ARMA_USE_EXTERN_RNG) + #elif defined(ARMA_USE_CXX11_RNG) { return std::numeric_limits::max(); } @@ -262,11 +382,17 @@ struct arma_rng::randi { arma_rng_alt::randi_fill(mem, N, a, b); } - #elif defined(ARMA_USE_EXTERN_RNG) + #elif defined(ARMA_USE_CXX11_RNG) { std::uniform_int_distribution local_i_distr(a, b); - for(uword i=0; i local_u_distr; - for(uword i=0; i < N; ++i) { mem[i] = eT( local_u_distr(mt19937_64_instance) ); } + std::mt19937_64& producer = arma_rng::get_producer(); + + arma_rng::lock_producer(); + + for(uword i=0; i < N; ++i) { mem[i] = eT( local_u_distr(producer) ); } + + arma_rng::unlock_producer(); } #else { @@ -358,11 +496,17 @@ struct arma_rng::randu for(uword i=0; i < N; ++i) { mem[i] = eT( arma_rng_alt::randu_val() * r + a ); } } - #elif defined(ARMA_USE_EXTERN_RNG) + #elif defined(ARMA_USE_CXX11_RNG) { std::uniform_real_distribution local_u_distr(a,b); - for(uword i=0; i < N; ++i) { mem[i] = eT( local_u_distr(mt19937_64_instance) ); } + std::mt19937_64& producer = arma_rng::get_producer(); + + arma_rng::lock_producer(); + + for(uword i=0; i < N; ++i) { mem[i] = eT( local_u_distr(producer) ); } + + arma_rng::unlock_producer(); } #else { @@ -396,12 +540,18 @@ struct arma_rng::randu< std::complex > return std::complex(a, b); } - #elif defined(ARMA_USE_EXTERN_RNG) + #elif defined(ARMA_USE_CXX11_RNG) { std::uniform_real_distribution local_u_distr; - const T a = T( local_u_distr(mt19937_64_instance) ); - const T b = T( local_u_distr(mt19937_64_instance) ); + std::mt19937_64& producer = arma_rng::get_producer(); + + arma_rng::lock_producer(); + + const T a = T( local_u_distr(producer) ); + const T b = T( local_u_distr(producer) ); + + arma_rng::unlock_producer(); return std::complex(a, b); } @@ -431,17 +581,23 @@ struct arma_rng::randu< std::complex > mem[i] = std::complex(a, b); } } - #elif defined(ARMA_USE_EXTERN_RNG) + #elif defined(ARMA_USE_CXX11_RNG) { std::uniform_real_distribution local_u_distr; + std::mt19937_64& producer = arma_rng::get_producer(); + + arma_rng::lock_producer(); + for(uword i=0; i < N; ++i) { - const T a = T( local_u_distr(mt19937_64_instance) ); - const T b = T( local_u_distr(mt19937_64_instance) ); + const T a = T( local_u_distr(producer) ); + const T b = T( local_u_distr(producer) ); mem[i] = std::complex(a, b); } + + arma_rng::unlock_producer(); } #else { @@ -491,17 +647,23 @@ struct arma_rng::randu< std::complex > mem[i] = std::complex(tmp1, tmp2); } } - #elif defined(ARMA_USE_EXTERN_RNG) + #elif defined(ARMA_USE_CXX11_RNG) { std::uniform_real_distribution local_u_distr(a,b); + std::mt19937_64& producer = arma_rng::get_producer(); + + arma_rng::lock_producer(); + for(uword i=0; i < N; ++i) { - const T tmp1 = T( local_u_distr(mt19937_64_instance) ); - const T tmp2 = T( local_u_distr(mt19937_64_instance) ); + const T tmp1 = T( local_u_distr(producer) ); + const T tmp2 = T( local_u_distr(producer) ); mem[i] = std::complex(tmp1, tmp2); } + + arma_rng::unlock_producer(); } #else { @@ -552,11 +714,17 @@ struct arma_rng::randn { return eT( arma_rng_alt::randn_val() ); } - #elif defined(ARMA_USE_EXTERN_RNG) + #elif defined(ARMA_USE_CXX11_RNG) { std::normal_distribution local_n_distr; - return eT( local_n_distr(mt19937_64_instance) ); + arma_rng::lock_producer(); + + const eT out = eT( local_n_distr(arma_rng::get_producer()) ); + + arma_rng::unlock_producer(); + + return out; } #else { @@ -575,12 +743,18 @@ struct arma_rng::randn { arma_rng_alt::randn_dual_val(out1, out2); } - #elif defined(ARMA_USE_EXTERN_RNG) + #elif defined(ARMA_USE_CXX11_RNG) { std::normal_distribution local_n_distr; - out1 = eT( local_n_distr(mt19937_64_instance) ); - out2 = eT( local_n_distr(mt19937_64_instance) ); + std::mt19937_64& producer = arma_rng::get_producer(); + + arma_rng::lock_producer(); + + out1 = eT( local_n_distr(producer) ); + out2 = eT( local_n_distr(producer) ); + + arma_rng::unlock_producer(); } #else { @@ -605,11 +779,17 @@ struct arma_rng::randn if(i < N) { mem[i] = eT( arma_rng_alt::randn_val() ); } } - #elif defined(ARMA_USE_EXTERN_RNG) + #elif defined(ARMA_USE_CXX11_RNG) { std::normal_distribution local_n_distr; - for(uword i=0; i < N; ++i) { mem[i] = eT( local_n_distr(mt19937_64_instance) ); } + std::mt19937_64& producer = arma_rng::get_producer(); + + arma_rng::lock_producer(); + + for(uword i=0; i < N; ++i) { mem[i] = eT( local_n_distr(producer) ); } + + arma_rng::unlock_producer(); } #else { @@ -657,11 +837,17 @@ struct arma_rng::randn mem[i] = (val_i * sd) + mu; } } - #elif defined(ARMA_USE_EXTERN_RNG) + #elif defined(ARMA_USE_CXX11_RNG) { std::normal_distribution local_n_distr(mu, sd); - for(uword i=0; i < N; ++i) { mem[i] = eT( local_n_distr(mt19937_64_instance) ); } + std::mt19937_64& producer = arma_rng::get_producer(); + + arma_rng::lock_producer(); + + for(uword i=0; i < N; ++i) { mem[i] = eT( local_n_distr(producer) ); } + + arma_rng::unlock_producer(); } #else { @@ -741,17 +927,23 @@ struct arma_rng::randn< std::complex > { for(uword i=0; i < N; ++i) { mem[i] = std::complex( arma_rng::randn< std::complex >() ); } } - #elif defined(ARMA_USE_EXTERN_RNG) + #elif defined(ARMA_USE_CXX11_RNG) { std::normal_distribution local_n_distr; + std::mt19937_64& producer = arma_rng::get_producer(); + + arma_rng::lock_producer(); + for(uword i=0; i < N; ++i) { - const T a = T( local_n_distr(mt19937_64_instance) ); - const T b = T( local_n_distr(mt19937_64_instance) ); + const T a = T( local_n_distr(producer) ); + const T b = T( local_n_distr(producer) ); mem[i] = std::complex(a,b); } + + arma_rng::unlock_producer(); } #else { @@ -818,11 +1010,17 @@ struct arma_rng::randg void fill(eT* mem, const uword N, const double a, const double b) { - #if defined(ARMA_USE_EXTERN_RNG) + #if defined(ARMA_USE_CXX11_RNG) { std::gamma_distribution local_g_distr(a,b); - for(uword i=0; i #include #include +#include #include "armadillo_bits/config.hpp" @@ -29,15 +30,11 @@ #include "armadillo_bits/typedef_elem.hpp" #include "armadillo_bits/include_superlu.hpp" - -#if defined(ARMA_USE_EXTERN_RNG) - #include - - namespace arma - { - thread_local std::mt19937_64 mt19937_64_instance; - } -#endif +namespace arma + { + // kept for compatibility with programs compiled with older versions of Armadillo + thread_local std::mt19937_64 mt19937_64_instance; + } namespace arma {