-
-
Notifications
You must be signed in to change notification settings - Fork 385
Expand file tree
/
Copy pathfullrank_test.cpp
More file actions
131 lines (113 loc) · 4.42 KB
/
Copy pathfullrank_test.cpp
File metadata and controls
131 lines (113 loc) · 4.42 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
#include <stan/services/experimental/advi/fullrank.hpp>
#include <gtest/gtest.h>
#include <stan/io/empty_var_context.hpp>
#include <test/test-models/good/services/test_lp.hpp>
#include <test/unit/services/instrumented_callbacks.hpp>
class ServicesExperimentalAdvi : public testing::Test {
public:
ServicesExperimentalAdvi() : model(context, 0, &model_log) {}
std::stringstream model_log;
stan::test::unit::instrumented_writer init, parameter, diagnostic;
stan::test::unit::instrumented_logger logger;
stan::io::empty_var_context context;
stan::test::unit::instrumented_interrupt interrupt;
stan_model model;
};
TEST_F(ServicesExperimentalAdvi, experimental_message) {
unsigned int seed = 0;
unsigned int chain = 1;
double init_radius = 0;
int grad_samples = 1;
int elbo_samples = 100;
int max_iterations = 10000;
double tol_rel_obj = 0.01;
double eta = 1.0;
bool adapt_engaged = true;
int adapt_iterations = 50;
int eval_elbo = 100;
int output_samples = 1000;
stan::services::experimental::advi ::fullrank(
model, context, seed, chain, init_radius, grad_samples, elbo_samples,
max_iterations, tol_rel_obj, eta, adapt_engaged, adapt_iterations,
eval_elbo, output_samples, interrupt, logger, init, parameter,
diagnostic);
EXPECT_GT(logger.call_count(), 0);
EXPECT_EQ(logger.call_count(), logger.call_count_info())
<< "all messages go to info";
EXPECT_EQ(1, logger.find_info("EXPERIMENTAL ALGORITHM"))
<< "Missing experimental algorithm message";
}
TEST_F(ServicesExperimentalAdvi, fullrank) {
unsigned int seed = 0;
unsigned int chain = 1;
double init_radius = 0;
int grad_samples = 1;
int elbo_samples = 100;
int max_iterations = 10000;
double tol_rel_obj = 0.01;
double eta = 1.0;
bool adapt_engaged = true;
int adapt_iterations = 50;
int eval_elbo = 100;
int output_samples = 1000;
int return_code = stan::services::experimental::advi ::fullrank(
model, context, seed, chain, init_radius, grad_samples, elbo_samples,
max_iterations, tol_rel_obj, eta, adapt_engaged, adapt_iterations,
eval_elbo, output_samples, interrupt, logger, init, parameter,
diagnostic);
EXPECT_EQ(0, return_code);
std::vector<std::vector<std::string> > parameter_names;
parameter_names = parameter.vector_string_values();
std::vector<std::vector<double> > parameter_values;
parameter_values = parameter.vector_double_values();
// Expectations of parameter parameter names.
ASSERT_EQ(8, parameter_names[0].size());
EXPECT_EQ("lp__", parameter_names[0][0]);
EXPECT_EQ("log_p__", parameter_names[0][1]);
EXPECT_EQ("log_g__", parameter_names[0][2]);
EXPECT_EQ("y.1", parameter_names[0][3]);
EXPECT_EQ("y.2", parameter_names[0][4]);
EXPECT_EQ("z.1", parameter_names[0][5]);
EXPECT_EQ("z.2", parameter_names[0][6]);
EXPECT_EQ("xgq", parameter_names[0][7]);
// Expect one name per parameter value.
EXPECT_EQ(parameter_names[0].size(), parameter_values[0].size());
ASSERT_EQ(1, init.vector_double_values().size());
ASSERT_EQ(2, init.vector_double_values().at(0).size());
std::vector<double> init_values = init.vector_double_values().at(0);
EXPECT_FLOAT_EQ(0, init_values[0]);
EXPECT_FLOAT_EQ(0, init_values[1]);
ASSERT_EQ(output_samples + 1, parameter.vector_double_values().size());
ASSERT_EQ(eval_elbo, diagnostic.vector_double_values().size());
EXPECT_EQ(0, interrupt.call_count());
}
TEST_F(ServicesExperimentalAdvi, fullrank_timing_info) {
unsigned int seed = 0;
unsigned int chain = 1;
double init_radius = 0;
int grad_samples = 1;
int elbo_samples = 100;
int max_iterations = 100;
double tol_rel_obj = 0.01;
double eta = 1.0;
bool adapt_engaged = true;
int adapt_iterations = 50;
int eval_elbo = 100;
int output_samples = 10;
stan::services::experimental::advi::fullrank(
model, context, seed, chain, init_radius, grad_samples, elbo_samples,
max_iterations, tol_rel_obj, eta, adapt_engaged, adapt_iterations,
eval_elbo, output_samples, interrupt, logger, init, parameter,
diagnostic);
EXPECT_TRUE(logger.find_info("Elapsed Time:") > 0)
<< "Should find 'Elapsed Time:' in logger info output";
bool found_in_parameter = false;
for (const auto& msg : parameter.string_values()) {
if (msg.find("Elapsed Time:") != std::string::npos) {
found_in_parameter = true;
break;
}
}
EXPECT_TRUE(found_in_parameter)
<< "Should find 'Elapsed Time:' in parameter_writer output";
}