use crate::bayesian::rand_u01;
fn randn(state: &mut u64) -> f64 {
let u1 = rand_u01(state).max(1e-12);
let u2 = rand_u01(state).max(1e-12);
(-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub enum DriftType {
Gbm,
Ou,
}
impl DriftType {
pub fn parse(s: &str) -> Result<Self, String> {
match s.to_ascii_lowercase().as_str() {
"gbm" | "geometric" => Ok(Self::Gbm),
"ou" | "ornstein" | "ornstein_uhlenbeck" | "mean_reversion" => Ok(Self::Ou),
other => Err(format!("unknown drift type '{other}' (expected gbm | ou)")),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub enum Solver {
Euler,
Milstein,
}
impl Solver {
pub fn parse(s: &str) -> Result<Self, String> {
match s.to_ascii_lowercase().as_str() {
"euler" | "euler_maruyama" | "em" => Ok(Self::Euler),
"milstein" => Ok(Self::Milstein),
other => Err(format!(
"unknown solver '{other}' (expected euler | milstein)"
)),
}
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct SdeConfig {
pub x0: f64,
pub t_end: f64,
pub n_steps: usize,
pub n_paths: usize,
pub drift: DriftType,
pub mu: f64,
pub theta: f64,
pub sigma: f64,
pub solver: Solver,
pub seed: u64,
}
impl Default for SdeConfig {
fn default() -> Self {
Self {
x0: 100.0,
t_end: 1.0,
n_steps: 100,
n_paths: 1000,
drift: DriftType::Gbm,
mu: 0.05,
theta: 1.0,
sigma: 0.2,
solver: Solver::Euler,
seed: 42,
}
}
}
fn drift_diffusion(x: f64, cfg: &SdeConfig) -> (f64, f64) {
match cfg.drift {
DriftType::Gbm => (cfg.mu * x, cfg.sigma * x),
DriftType::Ou => (cfg.theta * (cfg.mu - x), cfg.sigma),
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct SdeResult {
pub mean: f64,
pub std: f64,
pub p05: f64,
pub p50: f64,
pub p95: f64,
pub min: f64,
pub max: f64,
pub n_paths: usize,
pub dt: f64,
}
#[must_use]
pub fn solve(cfg: &SdeConfig) -> SdeResult {
let dt = cfg.t_end / cfg.n_steps.max(1) as f64;
let sqrt_dt = dt.sqrt();
let mut rng = cfg.seed;
let mut terminals = Vec::with_capacity(cfg.n_paths);
for _ in 0..cfg.n_paths {
let mut x = cfg.x0;
for _ in 0..cfg.n_steps {
let (drift, diff) = drift_diffusion(x, cfg);
let dw = sqrt_dt * randn(&mut rng);
if cfg.solver == Solver::Milstein {
let (_, diff) = drift_diffusion(x, cfg);
let sigma_prime = match cfg.drift {
DriftType::Gbm => cfg.sigma, DriftType::Ou => 0.0, };
let correction = 0.5 * diff * sigma_prime * (dw * dw - dt);
x += drift.mul_add(dt, diff * dw) + correction;
} else {
x += drift.mul_add(dt, diff * dw);
}
}
terminals.push(x);
}
stats(&terminals, cfg.n_paths, dt)
}
#[must_use]
pub fn solve_mlmc(cfg: &SdeConfig) -> MlMcResult {
let fine_steps = cfg.n_steps.max(2);
let coarse_steps = fine_steps / 2;
let fine = solve(&SdeConfig {
n_steps: fine_steps,
..cfg.clone()
});
let coarse = solve(&SdeConfig {
n_steps: coarse_steps,
..cfg.clone()
});
let mlmc_mean = fine.mean + (fine.mean - coarse.mean);
MlMcResult {
mlmc_mean,
fine_mean: fine.mean,
coarse_mean: coarse.mean,
fine_std: fine.std,
n_paths: cfg.n_paths,
fine_steps,
coarse_steps,
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct MlMcResult {
pub mlmc_mean: f64,
pub fine_mean: f64,
pub coarse_mean: f64,
pub fine_std: f64,
pub n_paths: usize,
pub fine_steps: usize,
pub coarse_steps: usize,
}
fn percentile(sorted: &[f64], q: f64) -> f64 {
if sorted.is_empty() {
return 0.0;
}
let idx = ((q * sorted.len() as f64).ceil() as usize)
.saturating_sub(1)
.min(sorted.len() - 1);
sorted[idx]
}
fn stats(terminals: &[f64], n_paths: usize, dt: f64) -> SdeResult {
let mean = terminals.iter().sum::<f64>() / terminals.len().max(1) as f64;
let var =
terminals.iter().map(|x| (x - mean).powi(2)).sum::<f64>() / terminals.len().max(1) as f64;
let mut sorted = terminals.to_vec();
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
SdeResult {
mean,
std: var.sqrt(),
p05: percentile(&sorted, 0.05),
p50: percentile(&sorted, 0.5),
p95: percentile(&sorted, 0.95),
min: sorted.first().copied().unwrap_or(0.0),
max: sorted.last().copied().unwrap_or(0.0),
n_paths,
dt,
}
}
#[cfg(test)]
mod tests {
#![allow(clippy::suboptimal_flops)] use super::*;
#[test]
fn gbm_euler_matches_analytic_mean() {
let cfg = SdeConfig {
x0: 100.0,
t_end: 1.0,
n_steps: 200,
n_paths: 20_000,
drift: DriftType::Gbm,
mu: 0.05,
sigma: 0.3,
solver: Solver::Euler,
seed: 42,
..Default::default()
};
let r = solve(&cfg);
let analytic = 100.0 * (0.05_f64).exp();
assert!(
(r.mean - analytic).abs() / analytic < 0.02,
"euler mean {} vs analytic {}",
r.mean,
analytic
);
assert!(r.std > 0.0);
}
#[test]
fn milstein_reduces_pathwise_error() {
let cfg = SdeConfig {
x0: 100.0,
t_end: 1.0,
n_steps: 8,
n_paths: 1,
drift: DriftType::Gbm,
mu: 0.05,
sigma: 0.4,
seed: 7,
..Default::default()
};
let dt = cfg.t_end / cfg.n_steps as f64;
let sqrt_dt = dt.sqrt();
let mut rng = cfg.seed;
let mut euler_x = cfg.x0;
let mut mil_x = cfg.x0;
let mut w = 0.0_f64;
for _ in 0..cfg.n_steps {
let dw = sqrt_dt * randn(&mut rng);
w += dw;
euler_x += cfg.mu * euler_x * dt + cfg.sigma * euler_x * dw;
mil_x += cfg.mu * mil_x * dt
+ cfg.sigma * mil_x * dw
+ 0.5 * cfg.sigma * cfg.sigma * mil_x * (dw * dw - dt);
}
let exact =
cfg.x0 * ((cfg.mu - 0.5 * cfg.sigma * cfg.sigma) * cfg.t_end + cfg.sigma * w).exp();
let euler_err = (euler_x - exact).abs();
let mil_err = (mil_x - exact).abs();
assert!(
mil_err < euler_err,
"Milstein pathwise err {mil_err} should be smaller than Euler's {euler_err}"
);
}
#[test]
fn ou_reverts_to_mean() {
let cfg = SdeConfig {
x0: 0.0,
t_end: 3.0,
n_steps: 300,
n_paths: 20_000,
drift: DriftType::Ou,
mu: 5.0,
theta: 1.0,
sigma: 0.5,
solver: Solver::Euler,
seed: 3,
};
let r = solve(&cfg);
let analytic = 5.0 + (0.0 - 5.0) * (-3.0_f64).exp();
assert!(
(r.mean - analytic).abs() < 0.05,
"ou mean {} vs analytic {}",
r.mean,
analytic
);
}
#[test]
fn gbm_paths_never_negative_in_milstein_small_step() {
let cfg = SdeConfig {
x0: 100.0,
t_end: 1.0,
n_steps: 500,
n_paths: 5000,
drift: DriftType::Gbm,
mu: 0.05,
sigma: 0.2,
solver: Solver::Milstein,
seed: 99,
..Default::default()
};
let r = solve(&cfg);
assert!(
r.min > 0.0,
"GBM Milstein min should stay positive, got {}",
r.min
);
}
#[test]
fn mlmc_improves_estimate_on_coarse_grid() {
let base = SdeConfig {
x0: 100.0,
t_end: 1.0,
n_steps: 8, n_paths: 10_000,
drift: DriftType::Gbm,
mu: 0.05,
sigma: 0.4,
seed: 11,
..Default::default()
};
let analytic = 100.0 * (0.05_f64).exp();
let fine = solve(&SdeConfig { n_steps: 8, ..base });
let mlmc = solve_mlmc(&base);
let fine_err = (fine.mean - analytic).abs();
let mlmc_err = (mlmc.mlmc_mean - analytic).abs();
assert!(
mlmc_err <= fine_err + 1e-9,
"mlmc err {mlmc_err} should be <= fine err {fine_err}"
);
}
#[test]
fn drift_type_parsing() {
assert_eq!(DriftType::parse("gbm").unwrap(), DriftType::Gbm);
assert_eq!(DriftType::parse("ou").unwrap(), DriftType::Ou);
assert!(DriftType::parse("bogus").is_err());
assert_eq!(Solver::parse("euler").unwrap(), Solver::Euler);
assert_eq!(Solver::parse("milstein").unwrap(), Solver::Milstein);
assert!(Solver::parse("bogus").is_err());
}
}