Skip to content

Commit 656d4e3

Browse files
authored
Merge pull request #3402 from stan-dev/fix/gq-chain-id-rng
fix standalone gqs rng seed
2 parents d39d26b + b8e24a7 commit 656d4e3

2 files changed

Lines changed: 91 additions & 5 deletions

File tree

src/stan/services/sample/standalone_gqs.hpp

Lines changed: 11 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -30,13 +30,15 @@ namespace services {
3030
* @param[in, out] interrupt called every iteration
3131
* @param[in, out] logger logger to which to write warning and error messages
3232
* @param[in, out] sample_writer writer to which draws are written
33+
* @param[in] chain chain id to advance the pseudo random number generator
3334
* @return error code
3435
*/
3536
template <class Model>
3637
int standalone_generate(const Model &model, const Eigen::MatrixXd &draws,
3738
unsigned int seed, callbacks::interrupt &interrupt,
3839
callbacks::logger &logger,
39-
callbacks::writer &sample_writer) {
40+
callbacks::writer &sample_writer,
41+
unsigned int chain = 1) {
4042
if (draws.size() == 0) {
4143
logger.error("Empty set of draws from fitted model.");
4244
return error_codes::DATAERR;
@@ -63,7 +65,7 @@ int standalone_generate(const Model &model, const Eigen::MatrixXd &draws,
6365
util::gq_writer writer(sample_writer, logger, p_names.size());
6466
writer.write_gq_names(model);
6567

66-
stan::rng_t rng = util::create_rng(seed, 1);
68+
stan::rng_t rng = util::create_rng(seed, chain);
6769

6870
std::vector<double> unconstrained_params_r;
6971
std::vector<double> row(draws.cols());
@@ -115,17 +117,21 @@ int standalone_generate(const Model &model, const Eigen::MatrixXd &draws,
115117
* @param[in, out] logger logger to which to write warning and error messages
116118
* @param[in, out] sample_writers A vector of writers to which draws for each
117119
* chain are written
120+
* @param[in] init_chain_id first chain id. The pseudo random number generator
121+
* will advance for each chain by an integer sequence from `init_chain_id` to
122+
* `init_chain_id + num_chains - 1`
118123
* @return error code
119124
*/
120125
template <typename Model, typename SampleWriter>
121126
int standalone_generate(const Model &model, const int num_chains,
122127
const std::vector<Eigen::MatrixXd> &draws,
123128
unsigned int seed, callbacks::interrupt &interrupt,
124129
callbacks::logger &logger,
125-
std::vector<SampleWriter> &sample_writers) {
130+
std::vector<SampleWriter> &sample_writers,
131+
unsigned int init_chain_id = 1) {
126132
if (num_chains == 1) {
127133
return standalone_generate(model, draws[0], seed, interrupt, logger,
128-
sample_writers[0]);
134+
sample_writers[0], init_chain_id);
129135
}
130136

131137
std::vector<std::string> p_names;
@@ -157,7 +163,7 @@ int standalone_generate(const Model &model, const int num_chains,
157163
}
158164
writers.emplace_back(sample_writers[i], logger, p_names.size());
159165
writers[i].write_gq_names(model);
160-
rngs.emplace_back(util::create_rng(seed, i + 1));
166+
rngs.emplace_back(util::create_rng(seed, init_chain_id + i));
161167
}
162168
bool error_any = false;
163169
try {

src/test/unit/services/sample/standalone_gqs_parallel_test.cpp

Lines changed: 80 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -78,3 +78,83 @@ TEST_F(ServicesStandaloneGQ, genDraws_bernoulli) {
7878
match_csv_columns(bern_csv.samples, sample_ss[i].str(), 1000, 1, 8);
7979
}
8080
}
81+
82+
namespace {
83+
84+
Eigen::MatrixXd bernoulli_fit_draws() {
85+
std::stringstream out;
86+
std::ifstream csv_stream;
87+
csv_stream.open("src/test/test-models/good/services/bernoulli_fit.csv");
88+
stan::io::stan_csv bern_csv
89+
= stan::io::stan_csv_reader::parse(csv_stream, &out);
90+
csv_stream.close();
91+
return bern_csv.samples.middleCols<1>(7);
92+
}
93+
94+
// drop the timing-dependent " Elapsed Time" line, which the test writers
95+
// emit without a comment prefix
96+
std::string data_lines(const std::string& csv) {
97+
std::stringstream in(csv), out;
98+
std::string line;
99+
while (std::getline(in, line)) {
100+
if (!line.empty() && line[0] == '#')
101+
continue;
102+
if (line.find("Elapsed Time") != std::string::npos)
103+
continue;
104+
out << line << "\n";
105+
}
106+
return out.str();
107+
}
108+
109+
using test_writer
110+
= stan::callbacks::unique_stream_writer<std::stringstream, deleter_noop>;
111+
112+
} // namespace
113+
114+
// The chain id must reach the RNG: two runs of the same draws that differ
115+
// only in chain id must not share a random number stream.
116+
TEST_F(ServicesStandaloneGQ, genDraws_bernoulli_chain_id_rng) {
117+
Eigen::MatrixXd draws = bernoulli_fit_draws();
118+
auto gq = [&](unsigned int chain) {
119+
std::stringstream ss;
120+
test_writer writer(std::unique_ptr<std::stringstream, deleter_noop>(&ss),
121+
"");
122+
EXPECT_EQ(stan::services::standalone_generate(
123+
model, draws, 12345, interrupt, logger, writer, chain),
124+
stan::services::error_codes::OK);
125+
return data_lines(ss.str());
126+
};
127+
std::string chain_1 = gq(1);
128+
EXPECT_NE(chain_1, gq(2));
129+
EXPECT_EQ(chain_1, gq(1)); // reproducible given the same chain id
130+
}
131+
132+
// init_chain_id must offset the per-chain streams: chains started at 3 must
133+
// match chains 3 and 4 of a run started at 1.
134+
TEST_F(ServicesStandaloneGQ, genDraws_bernoulli_init_chain_id_offset) {
135+
Eigen::MatrixXd draws = bernoulli_fit_draws();
136+
auto gq = [&](int n_chains, unsigned int init_chain_id) {
137+
std::vector<std::stringstream> ss(n_chains);
138+
std::vector<test_writer> writers;
139+
writers.reserve(n_chains);
140+
std::vector<Eigen::MatrixXd> draws_vec;
141+
for (int i = 0; i < n_chains; i++) {
142+
writers.emplace_back(
143+
std::unique_ptr<std::stringstream, deleter_noop>(&ss[i]), "");
144+
draws_vec.push_back(draws);
145+
}
146+
EXPECT_EQ(stan::services::standalone_generate(model, n_chains, draws_vec,
147+
12345, interrupt, logger,
148+
writers, init_chain_id),
149+
stan::services::error_codes::OK);
150+
std::vector<std::string> out;
151+
for (int i = 0; i < n_chains; i++)
152+
out.push_back(data_lines(ss[i].str()));
153+
return out;
154+
};
155+
std::vector<std::string> from_1 = gq(4, 1);
156+
std::vector<std::string> from_3 = gq(2, 3);
157+
EXPECT_EQ(from_3[0], from_1[2]);
158+
EXPECT_EQ(from_3[1], from_1[3]);
159+
EXPECT_NE(from_1[0], from_1[1]);
160+
}

0 commit comments

Comments
 (0)