use rand::rngs::SmallRng;
use rand::SeedableRng;
use rayon::prelude::*;
use super::model::McmcModel;
pub struct McmcConfig {
pub n_samples: usize,
pub warmup: usize,
pub thin: usize,
pub seed: u64,
}
pub fn run_mcmc<M: McmcModel>(model: &M, config: &McmcConfig) -> M::Result {
let total = config.warmup + config.n_samples * config.thin;
let mut rng = SmallRng::seed_from_u64(config.seed);
let mut state = model.init(&mut rng);
let mut samples = Vec::with_capacity(config.n_samples);
for iter in 0..total {
model.sweep(&mut state, &mut rng);
if iter >= config.warmup && (iter - config.warmup).is_multiple_of(config.thin) {
samples.push(model.collect(&state));
}
}
model.summarize(samples)
}
pub fn run_mcmc_parallel<M: McmcModel + Sync>(
model: &M,
config: &McmcConfig,
n_chains: usize,
) -> Vec<M::Result>
where
M::State: Send,
M::Result: Send,
{
(0..n_chains)
.into_par_iter()
.map(|i| {
let chain_config = McmcConfig {
n_samples: config.n_samples,
warmup: config.warmup,
thin: config.thin,
seed: config.seed.wrapping_add(i as u64),
};
run_mcmc(model, &chain_config)
})
.collect()
}