CPU Parallelism
In this section we demonstrate how to use the popular rayon crate to solve many ODEs in parallel on the CPU. The example is a simple population dynamics model: we solve the same ODE with different values of the growth parameter in parallel. To generate the growth parameters, we use rand_chacha::ChaCha12Rng to generate random values according to a log-normal distribution.
First, let's define the types and constants that we use. We use the nalgebra dense matrix type for the linear algebra backend, and define:
- the number of ODEs we want to solve in parallel (
N_SAMPLES), - the number of time points we want to solve for (
N_TIMES), - the final time point we want to solve to (
T_FINAL), - the carrying capacity parameter for the population dynamics model (
K), - an initial value for the population (
Y0), and - the random seed for the random number generator (
SEED).
type M = NalgebraMat<f64>;
type V = <M as MatrixCommon>::V;
const N_SAMPLES: usize = 1000;
const N_TIMES: usize = 101;
const T_FINAL: f64 = 20.0;
const K: f64 = 10.0;
const Y0: f64 = 0.1;
const SEED: u64 = 42;
Now we define an ensemble function that draws the random growth rates and solves the ensemble of ODEs in parallel. Once the ensemble has been solved, we reduce the results, also in parallel, to the 5%, 50%, and 95% quantiles at each evaluation time. The function returns the quantile bands; the evaluation times are passed in by the caller.
Both the solve and reduce stages run on the Rayon thread pool. A single random number generator cannot be shared across threads, but ChaCha12Rng offers independent streams from one seed, so each sample uses the stream matching its index.
Similarly, the diffsol Problem cannot be shared across threads as it needs to be mutated in order to set the growth rate for each sample. It also uses a RefCell to store operator statistics, so in Rust terms it is Send but not Sync (please raise an issue on the repo if Sync is required for your work). So that the Problem is not shared between worker threads, we will use map_init to give each thread its own problem, which will be reused across every sample that this particular thread works on.
fn ensemble(n_samples: usize, t_eval: &[f64]) -> Vec<[f64; 3]> {
let solutions: Vec<M> = (0..n_samples)
.into_par_iter()
.map_init(
|| {
OdeBuilder::<M>::new()
.p([1.0, K])
.rhs(|y, p, _t, dy| dy[0] = p[0] * y[0] * (1.0 - y[0] / p[1]))
.init(|_p, _t, y| y[0] = Y0, 1)
.build()
.unwrap()
},
|problem, i| {
let mut rng = ChaCha12Rng::seed_from_u64(SEED);
rng.set_stream(i as u64);
let r = LogNormal::new(0.5_f64.ln(), 0.3).unwrap().sample(&mut rng);
let p = V::from_vec(vec![r, K], *problem.eqn.context());
problem.eqn_mut().set_params(&p);
problem.tsit45().unwrap().solve_dense(t_eval).unwrap().0
},
)
.collect();
(0..t_eval.len())
.into_par_iter()
.map(|i| {
let mut ys: Vec<f64> = solutions.iter().map(|s| s.column(i)[0]).collect();
ys.sort_by(f64::total_cmp);
[quantile(&ys, 0.05), quantile(&ys, 0.5), quantile(&ys, 0.95)]
})
.collect()
}
/// Quantile of an ascending slice, linearly interpolating between order statistics.
fn quantile(sorted: &[f64], q: f64) -> f64 {
let pos = q * (sorted.len() - 1) as f64;
let (lo, hi) = (pos.floor() as usize, pos.ceil() as usize);
sorted[lo] + (sorted[hi] - sorted[lo]) * (pos - lo as f64)
}
Once we have called ensemble and obtained the results, we can plot the quantiles using Plotly. The following code creates a plot of the 5%, 50%, and 95% quantiles of the population dynamics model, shown below:
/// Plot the median with a shaded 5%-95% band. The band is drawn by filling the 95% trace
/// down to the 5% trace that precedes it.
fn plot_bands(t_eval: &[f64], bands: &[[f64; 3]]) -> Plot {
let t: Vec<f64> = t_eval.to_vec();
let lower: Vec<f64> = bands.iter().map(|b| b[0]).collect();
let median: Vec<f64> = bands.iter().map(|b| b[1]).collect();
let upper: Vec<f64> = bands.iter().map(|b| b[2]).collect();
let mut plot = Plot::new();
plot.add_trace(
Scatter::new(t.clone(), lower)
.mode(Mode::Lines)
.line(Line::new().width(0.0))
.name("5%"),
);
plot.add_trace(
Scatter::new(t.clone(), upper)
.mode(Mode::Lines)
.line(Line::new().width(0.0))
.fill(Fill::ToNextY)
.fill_color("rgba(31, 119, 180, 0.25)")
.name("95%"),
);
plot.add_trace(Scatter::new(t, median).mode(Mode::Lines).name("median"));
plot.set_layout(
Layout::new()
.x_axis(Axis::new().title("t"))
.y_axis(Axis::new().title("y")),
);
plot
}
Thread Scaling
Now we can examine how effective the parallelism is by varying the number of threads used in the rayon thread pool. The following code will run the ensemble with different numbers of threads, and record the time taken for each run. We will repeat each run a few times and then take the median to reduce benchmark noise.
fn thread_scaling(t_eval: &[f64]) -> (Vec<f64>, Vec<f64>, Vec<f64>) {
ensemble(BENCH_SAMPLES, t_eval); // warm up the pool and the allocator
let mut threads = Vec::new();
let mut elapsed = Vec::new();
for n in 1..=rayon::current_num_threads() {
let pool = rayon::ThreadPoolBuilder::new()
.num_threads(n)
.build()
.unwrap();
// use the median of a few runs to avoid outliers
let mut runs: Vec<f64> = (0..BENCH_REPEATS)
.map(|_| {
let start = Instant::now();
let bands = pool.install(|| ensemble(BENCH_SAMPLES, t_eval));
assert_eq!(bands.len(), t_eval.len());
start.elapsed().as_secs_f64()
})
.collect();
runs.sort_by(f64::total_cmp);
elapsed.push(runs[runs.len() / 2]);
threads.push(n as f64);
}
let speedup = elapsed.iter().map(|t| elapsed[0] / t).collect();
(threads, elapsed, speedup)
}
We use Plotly to plot the results and compare them against a reference line indicating ideal linear scaling.
Ideal speed-up would be linear with the number of threads, but several factors reduce the measured scaling below the ideal line:
- problem setup: each thread worker needs to build its own
OdeSolverProbleminmap_init. - reducing the solutions: This requires
N_TIMESsorts which areO(N log N)each. - memory bandwidth: here we need to allocate, write then read
N_SAMPLESsolution trajectories. - scheduling overhead: Rayon uses a thread pool and schedules work across it, this creates more work at higher thread counts.
- simple ODE: The logistic growth ODE is trivial to solve, so the linear portion is relatively cheap, increasing the weighting of the other factors above.