Skip to content
Merged
Show file tree
Hide file tree
Changes from all 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
16 changes: 11 additions & 5 deletions src/stan/services/sample/standalone_gqs.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -30,13 +30,15 @@ namespace services {
* @param[in, out] interrupt called every iteration
* @param[in, out] logger logger to which to write warning and error messages
* @param[in, out] sample_writer writer to which draws are written
* @param[in] chain chain id to advance the pseudo random number generator
* @return error code
*/
template <class Model>
int standalone_generate(const Model &model, const Eigen::MatrixXd &draws,
unsigned int seed, callbacks::interrupt &interrupt,
callbacks::logger &logger,
callbacks::writer &sample_writer) {
callbacks::writer &sample_writer,
unsigned int chain = 1) {
if (draws.size() == 0) {
logger.error("Empty set of draws from fitted model.");
return error_codes::DATAERR;
Expand All @@ -63,7 +65,7 @@ int standalone_generate(const Model &model, const Eigen::MatrixXd &draws,
util::gq_writer writer(sample_writer, logger, p_names.size());
writer.write_gq_names(model);

stan::rng_t rng = util::create_rng(seed, 1);
stan::rng_t rng = util::create_rng(seed, chain);

std::vector<double> unconstrained_params_r;
std::vector<double> row(draws.cols());
Expand Down Expand Up @@ -115,17 +117,21 @@ int standalone_generate(const Model &model, const Eigen::MatrixXd &draws,
* @param[in, out] logger logger to which to write warning and error messages
* @param[in, out] sample_writers A vector of writers to which draws for each
* chain are written
* @param[in] init_chain_id first chain id. The pseudo random number generator
* will advance for each chain by an integer sequence from `init_chain_id` to
* `init_chain_id + num_chains - 1`
* @return error code
*/
template <typename Model, typename SampleWriter>
int standalone_generate(const Model &model, const int num_chains,
const std::vector<Eigen::MatrixXd> &draws,
unsigned int seed, callbacks::interrupt &interrupt,
callbacks::logger &logger,
std::vector<SampleWriter> &sample_writers) {
std::vector<SampleWriter> &sample_writers,
unsigned int init_chain_id = 1) {
if (num_chains == 1) {
return standalone_generate(model, draws[0], seed, interrupt, logger,
sample_writers[0]);
sample_writers[0], init_chain_id);
}

std::vector<std::string> p_names;
Expand Down Expand Up @@ -157,7 +163,7 @@ int standalone_generate(const Model &model, const int num_chains,
}
writers.emplace_back(sample_writers[i], logger, p_names.size());
writers[i].write_gq_names(model);
rngs.emplace_back(util::create_rng(seed, i + 1));
rngs.emplace_back(util::create_rng(seed, init_chain_id + i));
}
bool error_any = false;
try {
Expand Down
80 changes: 80 additions & 0 deletions src/test/unit/services/sample/standalone_gqs_parallel_test.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -78,3 +78,83 @@ TEST_F(ServicesStandaloneGQ, genDraws_bernoulli) {
match_csv_columns(bern_csv.samples, sample_ss[i].str(), 1000, 1, 8);
}
}

namespace {

Eigen::MatrixXd bernoulli_fit_draws() {
std::stringstream out;
std::ifstream csv_stream;
csv_stream.open("src/test/test-models/good/services/bernoulli_fit.csv");
stan::io::stan_csv bern_csv
= stan::io::stan_csv_reader::parse(csv_stream, &out);
csv_stream.close();
return bern_csv.samples.middleCols<1>(7);
}

// drop the timing-dependent " Elapsed Time" line, which the test writers
// emit without a comment prefix
std::string data_lines(const std::string& csv) {
std::stringstream in(csv), out;
std::string line;
while (std::getline(in, line)) {
if (!line.empty() && line[0] == '#')
continue;
if (line.find("Elapsed Time") != std::string::npos)
continue;
out << line << "\n";
}
return out.str();
}

using test_writer
= stan::callbacks::unique_stream_writer<std::stringstream, deleter_noop>;

} // namespace

// The chain id must reach the RNG: two runs of the same draws that differ
// only in chain id must not share a random number stream.
TEST_F(ServicesStandaloneGQ, genDraws_bernoulli_chain_id_rng) {
Eigen::MatrixXd draws = bernoulli_fit_draws();
auto gq = [&](unsigned int chain) {
std::stringstream ss;
test_writer writer(std::unique_ptr<std::stringstream, deleter_noop>(&ss),
"");
EXPECT_EQ(stan::services::standalone_generate(
model, draws, 12345, interrupt, logger, writer, chain),
stan::services::error_codes::OK);
return data_lines(ss.str());
};
std::string chain_1 = gq(1);
EXPECT_NE(chain_1, gq(2));
EXPECT_EQ(chain_1, gq(1)); // reproducible given the same chain id
}

// init_chain_id must offset the per-chain streams: chains started at 3 must
// match chains 3 and 4 of a run started at 1.
TEST_F(ServicesStandaloneGQ, genDraws_bernoulli_init_chain_id_offset) {
Eigen::MatrixXd draws = bernoulli_fit_draws();
auto gq = [&](int n_chains, unsigned int init_chain_id) {
std::vector<std::stringstream> ss(n_chains);
std::vector<test_writer> writers;
writers.reserve(n_chains);
std::vector<Eigen::MatrixXd> draws_vec;
for (int i = 0; i < n_chains; i++) {
writers.emplace_back(
std::unique_ptr<std::stringstream, deleter_noop>(&ss[i]), "");
draws_vec.push_back(draws);
}
EXPECT_EQ(stan::services::standalone_generate(model, n_chains, draws_vec,
12345, interrupt, logger,
writers, init_chain_id),
stan::services::error_codes::OK);
std::vector<std::string> out;
for (int i = 0; i < n_chains; i++)
out.push_back(data_lines(ss[i].str()));
return out;
};
std::vector<std::string> from_1 = gq(4, 1);
std::vector<std::string> from_3 = gq(2, 3);
EXPECT_EQ(from_3[0], from_1[2]);
EXPECT_EQ(from_3[1], from_1[3]);
EXPECT_NE(from_1[0], from_1[1]);
}
Loading