Skip to content
Merged
Show file tree
Hide file tree
Changes from 4 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 3 additions & 4 deletions examples/walnuts_api.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -63,7 +63,7 @@ int main() {

auto warmup_cfg = walnuts::WarmupConfigBuilder()
.min_max_iter(50, 2000)
.mass_converge_tol(2.0)
.mass_converge_tol(2.0)
.step_size_converge_tol(0.2)
.mass_init_count(4.0)
.build();
Expand All @@ -80,9 +80,8 @@ int main() {

// 2) SAMPLE =================================================================
// output sent to handlers
walnuts::WalnutsConfig config{std::move(init_cfg),
std::move(warmup_cfg),
std::move(sampling_cfg)};
walnuts::WalnutsConfig config{std::move(init_cfg), std::move(warmup_cfg),
std::move(sampling_cfg)};
walnuts::walnuts<std::mt19937_64>(seed, chain_handlers, global_handler,
interrupt_callback, logp_grad, config);

Expand Down
28 changes: 9 additions & 19 deletions include/walnuts/config.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -1046,8 +1046,7 @@ inline std::ostream& operator<<(std::ostream& out, const SamplingConfig& cfg) {
* configurations.
*/
class WalnutsConfig {
public:

public:
/**
* @brief Construct a Walnuts configuration given the component
* configurations.
Expand All @@ -1058,42 +1057,33 @@ class WalnutsConfig {
* @param[in] warmup The warmup configuration.
* @param[in] sampling The sampling configuration.
*/
WalnutsConfig(InitConfig init,
WarmupConfig warmup,
SamplingConfig sampling)
: init_(std::move(init)),
warmup_(std::move(warmup)),
sampling_(std::move(sampling)) {
};
WalnutsConfig(InitConfig init, WarmupConfig warmup, SamplingConfig sampling)
: init_(std::move(init)),
warmup_(std::move(warmup)),
sampling_(std::move(sampling)) {};

/**
* @brief Return the initialization configuration.
*
* @return The initialization configuration.
*/
const InitConfig& init() const noexcept {
return init_;
}
const InitConfig& init() const noexcept { return init_; }

/**
* @brief Return the warmup configuration.
*
* @return The warmup configuration.
*/
const WarmupConfig& warmup() const noexcept {
return warmup_;
}
const WarmupConfig& warmup() const noexcept { return warmup_; }

/**
* @brief Return the sampling configuration.
*
* @return The sampling configuration.
*/
const SamplingConfig& sampling() const noexcept {
return sampling_;
}
const SamplingConfig& sampling() const noexcept { return sampling_; }

private:
private:
/** The initialization configuration for all chains. */
InitConfig init_;

Expand Down
11 changes: 9 additions & 2 deletions include/walnuts/summary.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -479,8 +479,15 @@ inline Eigen::RowVectorXd sample_standard_deviation(const MC& chains) {
template <MarkovChainSequence MC>
inline Eigen::MatrixXd quantiles(const MC& chains,
const Eigen::VectorXd& probs) {
if (std::ranges::any_of(probs,
[](double p) { if (!(p >= 0)) return true; if (!(p <= 1)) return true; return false; })) {
if (std::ranges::any_of(probs, [](double p) {
if (!(p >= 0)) {
return true;
}
if (!(p <= 1)) {
return true;
}
return false;
})) {
throw std::invalid_argument("probs must be in [0, 1]");
}
const Eigen::Index N = static_cast<Eigen::Index>(chains.num_draws());
Expand Down
11 changes: 9 additions & 2 deletions include/walnuts/util.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -4,8 +4,8 @@
#include <cstddef>
#include <functional>
#include <limits>
#include <random>
#include <numeric>
#include <random>
#include <type_traits>

#include <Eigen/Dense>
Expand Down Expand Up @@ -58,7 +58,14 @@ enum class Direction {
Forward /**< Step forward in time. */
};

/**
* @brief A type definition for constructing `Direction::Backward` constants.
*/
using Backward_t = std::integral_constant<Direction, Direction::Backward>;

/**
* @brief A type definition for constructing `Direction::Forward` constants.
*/
using Forward_t = std::integral_constant<Direction, Direction::Forward>;

/**
Expand All @@ -80,7 +87,7 @@ class Random {
*
* @param[in,out] rng The base random number generator.
*/
explicit Random(RNG& rng)
explicit Random(RNG& rng) noexcept
: rng_(rng), unif_(0.0, 1.0), binary_(0.5), normal_(0.0, 1.0) {}

/**
Expand Down
17 changes: 9 additions & 8 deletions tests/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -15,25 +15,26 @@ include(GoogleTest)
## Source Coverage ##
##########################

function(add_test TEST_NAME TEST_SOURCE)
add_executable(${TEST_NAME} ${TEST_SOURCE})
target_link_libraries(${TEST_NAME}
function(add_test TEST_NAME)
add_executable(${TEST_NAME}_test ${TEST_NAME}_test.cpp)
target_link_libraries(${TEST_NAME}_test
PRIVATE
gtest_main
Eigen3::Eigen
walnuts
)
if(WALNUTS_COVERAGE)
target_compile_options(${TEST_NAME} PRIVATE
target_compile_options(${TEST_NAME}_test PRIVATE
-fprofile-instr-generate
-fcoverage-mapping
)
target_link_options(${TEST_NAME} PRIVATE
target_link_options(${TEST_NAME}_test PRIVATE
-fprofile-instr-generate
)
endif()
gtest_discover_tests(${TEST_NAME})
gtest_discover_tests(${TEST_NAME}_test)
endfunction()

add_test(summary_test summary_test.cpp)
add_test(config_test config_test.cpp)
add_test(summary)
add_test(config)
add_test(util)
118 changes: 56 additions & 62 deletions tests/config_test.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -5,8 +5,8 @@
#include <stdexcept>
#include <vector>

#include <walnuts.hpp>
#include <../tests/test_util.hpp>
#include <walnuts.hpp>

// class InitChainConfig ********************************************

Expand Down Expand Up @@ -59,10 +59,9 @@ TEST(InitChainConfig, MassIsReturnedByReference) {
EXPECT_EQ(&cfg.mass(), &cfg.mass());
}


// classes InitConfig and InitConfigBuilder *************************

// default InitConfig
// default InitConfig

TEST(InitConfigBuilder, DefaultsAreCorrect) {
Eigen::Index D_Eigen = Eigen::Index{2};
Expand All @@ -73,7 +72,7 @@ TEST(InitConfigBuilder, DefaultsAreCorrect) {
EXPECT_EQ(cfg.dims(), D);
for (std::size_t m = 0; m < M; ++m) {
EXPECT_DOUBLE_EQ(cfg.step_size(m), 0.1); // default step size 0.1
EXPECT_TRUE(cfg.position(m).isZero()); // default position 0
EXPECT_TRUE(cfg.position(m).isZero()); // default position 0
EXPECT_EQ(cfg.position(m).size(), D_Eigen);
EXPECT_TRUE(cfg.mass(m).isOnes()); // default mass I
EXPECT_EQ(cfg.mass(m).size(), D_Eigen);
Expand Down Expand Up @@ -302,14 +301,12 @@ TEST(InitConfigBuilder, VectorMassesSetsPerChain) {
vs[0] << 1.0, 2.0;
vs[1] << 3.0, 4.0;
vs[2] << 5.0, 6.0;
walnuts::InitConfig cfg =
walnuts::InitConfigBuilder(3, 2).masses(vs).build();
walnuts::InitConfig cfg = walnuts::InitConfigBuilder(3, 2).masses(vs).build();
for (std::size_t m = 0; m < 3; ++m) {
expect_near(cfg.mass(m), vs[m]);
}
}


TEST(InitConfigBuilder, VectorMassesThrowsOnWrongNumberOfChains) {
walnuts::InitConfigBuilder b(3, 2);
std::vector<Eigen::VectorXd> wrong_chains(2, Eigen::VectorXd::Ones(2));
Expand Down Expand Up @@ -432,7 +429,7 @@ TEST(InitConfig, InitChainConfigReturnsCorrectValues) {
expect_near(cc.mass(), mass);
}

// chaining
// chaining

TEST(InitConfigBuilder, MethodChainingReturnsBuilder) {
Eigen::VectorXd pos(2);
Expand Down Expand Up @@ -690,41 +687,40 @@ TEST(WarmupConfigBuilder, YieldPeriodThrowsOnZero) {

TEST(WarmupConfigBuilder, ChainReferenceIdentity) {
walnuts::WarmupConfigBuilder builder;
auto& builder_chain = builder
.min_max_iter(25, 200)
.step_size_converge_tol(0.05)
.mass_converge_tol(0.5)
.mass_init_count(2.0)
.mass_additive_smoothing(1e-4)
.max_macro_steps_target(10.0)
.step_accept_rate_target(0.75)
.step_learning_rate(0.1)
.step_gradient_decay(0.85)
.step_sq_gradient_decay(0.95)
.step_stabilization(1e-3)
.step_learn_rate_decay(0.6)
.publish_stride(2)
.yield_period(16);
auto& builder_chain = builder.min_max_iter(25, 200)
.step_size_converge_tol(0.05)
.mass_converge_tol(0.5)
.mass_init_count(2.0)
.mass_additive_smoothing(1e-4)
.max_macro_steps_target(10.0)
.step_accept_rate_target(0.75)
.step_learning_rate(0.1)
.step_gradient_decay(0.85)
.step_sq_gradient_decay(0.95)
.step_stabilization(1e-3)
.step_learn_rate_decay(0.6)
.publish_stride(2)
.yield_period(16);
EXPECT_EQ(&builder, &builder_chain);
}

TEST(WarmupConfigBuilder, FullChainProducesCorrectConfig) {
walnuts::WarmupConfigBuilder builder;
auto cfg = builder.min_max_iter(25, 200)
.step_size_converge_tol(0.05)
.mass_converge_tol(0.5)
.mass_init_count(2.0)
.mass_additive_smoothing(1e-4)
.max_macro_steps_target(10.0)
.step_accept_rate_target(0.75)
.step_learning_rate(0.1)
.step_gradient_decay(0.85)
.step_sq_gradient_decay(0.95)
.step_stabilization(1e-3)
.step_learn_rate_decay(0.6)
.publish_stride(2)
.yield_period(16)
.build();
.step_size_converge_tol(0.05)
.mass_converge_tol(0.5)
.mass_init_count(2.0)
.mass_additive_smoothing(1e-4)
.max_macro_steps_target(10.0)
.step_accept_rate_target(0.75)
.step_learning_rate(0.1)
.step_gradient_decay(0.85)
.step_sq_gradient_decay(0.95)
.step_stabilization(1e-3)
.step_learn_rate_decay(0.6)
.publish_stride(2)
.yield_period(16)
.build();
EXPECT_EQ(cfg.min_iter(), std::size_t{25});
EXPECT_EQ(cfg.max_iter(), std::size_t{200});
EXPECT_DOUBLE_EQ(cfg.step_size_converge_tol(), 0.05);
Expand Down Expand Up @@ -853,13 +849,13 @@ TEST(SamplingConfigBuilder, RhatConvergeTolThrowsOnBadValues) {

TEST(SamplingConfigBuilder, FullChainProducesCorrectConfig) {
walnuts::SamplingConfig cfg = walnuts::SamplingConfigBuilder()
.min_max_iter(25, 200)
.max_trajectory_doublings(8)
.max_step_halvings(3)
.max_hamiltonian_error(1.0)
.min_micro_steps(2)
.rhat_converge_tol(1.05)
.build();
.min_max_iter(25, 200)
.max_trajectory_doublings(8)
.max_step_halvings(3)
.max_hamiltonian_error(1.0)
.min_micro_steps(2)
.rhat_converge_tol(1.05)
.build();
EXPECT_EQ(cfg.min_iter(), std::size_t{25});
EXPECT_EQ(cfg.max_iter(), std::size_t{200});
EXPECT_EQ(cfg.max_trajectory_doublings(), std::size_t{8});
Expand All @@ -871,30 +867,28 @@ TEST(SamplingConfigBuilder, FullChainProducesCorrectConfig) {

TEST(SamplingConfigBuilder, ChainingReferenceEquality) {
walnuts::SamplingConfigBuilder builder = walnuts::SamplingConfigBuilder();
auto& builder_chain = builder
.min_max_iter(25, 200)
.max_trajectory_doublings(8)
.max_step_halvings(3)
.max_hamiltonian_error(1.0)
.min_micro_steps(2)
.rhat_converge_tol(1.05);
auto& builder_chain = builder.min_max_iter(25, 200)
.max_trajectory_doublings(8)
.max_step_halvings(3)
.max_hamiltonian_error(1.0)
.min_micro_steps(2)
.rhat_converge_tol(1.05);
EXPECT_EQ(&builder_chain, &builder);
}

// classes WalnutsConfig and WalnutsConfigBuilder *******************

TEST(WalnutsConfig, MembersAreIndependent) {
walnuts::WalnutsConfig cfg{
walnuts::InitConfigBuilder(2, 3).step_sizes(0.25).build(),
walnuts::InitConfigBuilder(2, 3).step_sizes(0.25).build(),
walnuts::WarmupConfigBuilder().min_max_iter(10, 200).build(),
walnuts::SamplingConfigBuilder().min_max_iter(5, 100).build()
};

EXPECT_EQ(cfg.warmup().min_iter(), std::size_t{10});
EXPECT_EQ(cfg.warmup().max_iter(), std::size_t{200});
EXPECT_EQ(cfg.sampling().min_iter(), std::size_t{5});
EXPECT_EQ(cfg.sampling().max_iter(), std::size_t{100});
EXPECT_EQ(cfg.init().num_chains(), std::size_t{2});
EXPECT_DOUBLE_EQ(cfg.init().step_size(0), 0.25);
EXPECT_DOUBLE_EQ(cfg.init().step_size(1), 0.25);
walnuts::SamplingConfigBuilder().min_max_iter(5, 100).build()};

EXPECT_EQ(cfg.warmup().min_iter(), std::size_t{10});
EXPECT_EQ(cfg.warmup().max_iter(), std::size_t{200});
EXPECT_EQ(cfg.sampling().min_iter(), std::size_t{5});
EXPECT_EQ(cfg.sampling().max_iter(), std::size_t{100});
EXPECT_EQ(cfg.init().num_chains(), std::size_t{2});
EXPECT_DOUBLE_EQ(cfg.init().step_size(0), 0.25);
EXPECT_DOUBLE_EQ(cfg.init().step_size(1), 0.25);
}
Loading
Loading