Skip to main content

wm_simulation/
sde.rs

1//! Stochastic differential equation solvers — Euler–Maruyama and Milstein.
2//!
3//! Supports two drift types:
4//! - **GBM**: `dX = μ·X·dt + σ·X·dW` (geometric Brownian motion, e.g. prices)
5//! - **OU**:   `dX = θ·(μ − X)·dt + σ·dW` (Ornstein–Uhlenbeck, mean reversion)
6//!
7//! Milstein adds the second-order correction `0.5·σ·σ'·X·(ΔW² − dt)`, which
8//! is non-zero only for GBM (the OU diffusion is constant).
9//!
10//! Also provides a two-level multilevel Monte Carlo (MLMC) extrapolation:
11//! `E ≈ E_fine + (E_fine − E_coarse)` — a cheap variance-reduction trick
12//! for terminal-statistic estimates.
13
14use crate::bayesian::rand_u01;
15
16/// Draw a standard normal via Box–Muller from the SplitMix64 state.
17fn randn(state: &mut u64) -> f64 {
18    let u1 = rand_u01(state).max(1e-12);
19    let u2 = rand_u01(state).max(1e-12);
20    (-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos()
21}
22
23/// SDE drift type.
24#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
25pub enum DriftType {
26    /// Geometric Brownian motion: `dX = μX dt + σX dW`.
27    Gbm,
28    /// Ornstein–Uhlenbeck: `dX = θ(μ − X) dt + σ dW`.
29    Ou,
30}
31
32impl DriftType {
33    /// Parse from the v26 tool's string names.
34    pub fn parse(s: &str) -> Result<Self, String> {
35        match s.to_ascii_lowercase().as_str() {
36            "gbm" | "geometric" => Ok(Self::Gbm),
37            "ou" | "ornstein" | "ornstein_uhlenbeck" | "mean_reversion" => Ok(Self::Ou),
38            other => Err(format!("unknown drift type '{other}' (expected gbm | ou)")),
39        }
40    }
41}
42
43/// Solver scheme.
44#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
45pub enum Solver {
46    /// Euler–Maruyama (strong order 0.5).
47    Euler,
48    /// Milstein (strong order 1.0) — identical to Euler for OU.
49    Milstein,
50}
51
52impl Solver {
53    /// Parse from the v26 tool's string names.
54    pub fn parse(s: &str) -> Result<Self, String> {
55        match s.to_ascii_lowercase().as_str() {
56            "euler" | "euler_maruyama" | "em" => Ok(Self::Euler),
57            "milstein" => Ok(Self::Milstein),
58            other => Err(format!(
59                "unknown solver '{other}' (expected euler | milstein)"
60            )),
61        }
62    }
63}
64
65/// SDE configuration.
66#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
67pub struct SdeConfig {
68    /// Initial value X(0).
69    pub x0: f64,
70    /// Terminal time T.
71    pub t_end: f64,
72    /// Number of time steps.
73    pub n_steps: usize,
74    /// Number of paths.
75    pub n_paths: usize,
76    /// Drift model.
77    pub drift: DriftType,
78    /// Drift coefficient μ (GBM) or mean-reversion level (OU).
79    pub mu: f64,
80    /// Mean-reversion strength θ (OU only).
81    pub theta: f64,
82    /// Diffusion coefficient σ.
83    pub sigma: f64,
84    /// Solver scheme.
85    pub solver: Solver,
86    /// PRNG seed.
87    pub seed: u64,
88}
89
90impl Default for SdeConfig {
91    fn default() -> Self {
92        Self {
93            x0: 100.0,
94            t_end: 1.0,
95            n_steps: 100,
96            n_paths: 1000,
97            drift: DriftType::Gbm,
98            mu: 0.05,
99            theta: 1.0,
100            sigma: 0.2,
101            solver: Solver::Euler,
102            seed: 42,
103        }
104    }
105}
106
107/// Drift and diffusion terms for the current state.
108fn drift_diffusion(x: f64, cfg: &SdeConfig) -> (f64, f64) {
109    match cfg.drift {
110        DriftType::Gbm => (cfg.mu * x, cfg.sigma * x),
111        DriftType::Ou => (cfg.theta * (cfg.mu - x), cfg.sigma),
112    }
113}
114
115/// Terminal statistics over all paths.
116#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
117pub struct SdeResult {
118    /// Mean terminal value.
119    pub mean: f64,
120    /// Standard deviation of terminal values.
121    pub std: f64,
122    /// 5th percentile of terminal values.
123    pub p05: f64,
124    /// Median (50th percentile).
125    pub p50: f64,
126    /// 95th percentile of terminal values.
127    pub p95: f64,
128    /// Min terminal value.
129    pub min: f64,
130    /// Max terminal value.
131    pub max: f64,
132    /// Number of paths simulated.
133    pub n_paths: usize,
134    /// Time step dt used.
135    pub dt: f64,
136}
137
138/// Solve the SDE and return terminal-value statistics.
139#[must_use]
140pub fn solve(cfg: &SdeConfig) -> SdeResult {
141    let dt = cfg.t_end / cfg.n_steps.max(1) as f64;
142    let sqrt_dt = dt.sqrt();
143    let mut rng = cfg.seed;
144    let mut terminals = Vec::with_capacity(cfg.n_paths);
145
146    for _ in 0..cfg.n_paths {
147        let mut x = cfg.x0;
148        for _ in 0..cfg.n_steps {
149            let (drift, diff) = drift_diffusion(x, cfg);
150            let dw = sqrt_dt * randn(&mut rng);
151            if cfg.solver == Solver::Milstein {
152                // Milstein correction: 0.5 · σ·σ' · (ΔW² − dt)
153                let (_, diff) = drift_diffusion(x, cfg);
154                let sigma_prime = match cfg.drift {
155                    DriftType::Gbm => cfg.sigma, // σ(x) = σx → σ' = σ
156                    DriftType::Ou => 0.0,        // σ(x) = σ → σ' = 0
157                };
158                let correction = 0.5 * diff * sigma_prime * (dw * dw - dt);
159                x += drift.mul_add(dt, diff * dw) + correction;
160            } else {
161                x += drift.mul_add(dt, diff * dw);
162            }
163        }
164        terminals.push(x);
165    }
166
167    stats(&terminals, cfg.n_paths, dt)
168}
169
170/// Two-level multilevel Monte Carlo estimate of the mean terminal value.
171///
172/// Uses the same seed for both levels so the coarse/fine paths share
173/// randomness (coupling), which makes the variance-reduction effective.
174#[must_use]
175pub fn solve_mlmc(cfg: &SdeConfig) -> MlMcResult {
176    let fine_steps = cfg.n_steps.max(2);
177    let coarse_steps = fine_steps / 2;
178
179    let fine = solve(&SdeConfig {
180        n_steps: fine_steps,
181        ..cfg.clone()
182    });
183    let coarse = solve(&SdeConfig {
184        n_steps: coarse_steps,
185        ..cfg.clone()
186    });
187
188    // E ≈ E_fine + (E_fine − E_coarse) — Richardson-style extrapolation
189    let mlmc_mean = fine.mean + (fine.mean - coarse.mean);
190    MlMcResult {
191        mlmc_mean,
192        fine_mean: fine.mean,
193        coarse_mean: coarse.mean,
194        fine_std: fine.std,
195        n_paths: cfg.n_paths,
196        fine_steps,
197        coarse_steps,
198    }
199}
200
201/// Result of a multilevel Monte Carlo run.
202#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
203pub struct MlMcResult {
204    /// MLMC-estimated mean terminal value.
205    pub mlmc_mean: f64,
206    /// Fine-level mean.
207    pub fine_mean: f64,
208    /// Coarse-level mean.
209    pub coarse_mean: f64,
210    /// Fine-level std.
211    pub fine_std: f64,
212    /// Paths per level.
213    pub n_paths: usize,
214    /// Fine steps.
215    pub fine_steps: usize,
216    /// Coarse steps.
217    pub coarse_steps: usize,
218}
219
220/// Percentile helper (nearest-rank).
221fn percentile(sorted: &[f64], q: f64) -> f64 {
222    if sorted.is_empty() {
223        return 0.0;
224    }
225    let idx = ((q * sorted.len() as f64).ceil() as usize)
226        .saturating_sub(1)
227        .min(sorted.len() - 1);
228    sorted[idx]
229}
230
231fn stats(terminals: &[f64], n_paths: usize, dt: f64) -> SdeResult {
232    let mean = terminals.iter().sum::<f64>() / terminals.len().max(1) as f64;
233    let var =
234        terminals.iter().map(|x| (x - mean).powi(2)).sum::<f64>() / terminals.len().max(1) as f64;
235    let mut sorted = terminals.to_vec();
236    sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
237    SdeResult {
238        mean,
239        std: var.sqrt(),
240        p05: percentile(&sorted, 0.05),
241        p50: percentile(&sorted, 0.5),
242        p95: percentile(&sorted, 0.95),
243        min: sorted.first().copied().unwrap_or(0.0),
244        max: sorted.last().copied().unwrap_or(0.0),
245        n_paths,
246        dt,
247    }
248}
249
250#[cfg(test)]
251mod tests {
252    #![allow(clippy::suboptimal_flops)] // test data expressions, not hot paths
253    use super::*;
254
255    #[test]
256    fn gbm_euler_matches_analytic_mean() {
257        // GBM: E[X_t] = X0 · e^(μt) — independent of σ
258        let cfg = SdeConfig {
259            x0: 100.0,
260            t_end: 1.0,
261            n_steps: 200,
262            n_paths: 20_000,
263            drift: DriftType::Gbm,
264            mu: 0.05,
265            sigma: 0.3,
266            solver: Solver::Euler,
267            seed: 42,
268            ..Default::default()
269        };
270        let r = solve(&cfg);
271        let analytic = 100.0 * (0.05_f64).exp();
272        assert!(
273            (r.mean - analytic).abs() / analytic < 0.02,
274            "euler mean {} vs analytic {}",
275            r.mean,
276            analytic
277        );
278        assert!(r.std > 0.0);
279    }
280
281    #[test]
282    fn milstein_reduces_pathwise_error() {
283        // Strong convergence: simulate one GBM path with both schemes using
284        // the SAME Brownian increments and compare to the exact solution
285        // X_t = X0·exp((μ − σ²/2)t + σ·W_t). Milstein (order 1.0) should be
286        // closer than Euler (order 0.5) on a coarse grid.
287        let cfg = SdeConfig {
288            x0: 100.0,
289            t_end: 1.0,
290            n_steps: 8,
291            n_paths: 1,
292            drift: DriftType::Gbm,
293            mu: 0.05,
294            sigma: 0.4,
295            seed: 7,
296            ..Default::default()
297        };
298        let dt = cfg.t_end / cfg.n_steps as f64;
299        let sqrt_dt = dt.sqrt();
300        let mut rng = cfg.seed;
301
302        let mut euler_x = cfg.x0;
303        let mut mil_x = cfg.x0;
304        let mut w = 0.0_f64;
305        for _ in 0..cfg.n_steps {
306            let dw = sqrt_dt * randn(&mut rng);
307            w += dw;
308            // Euler
309            euler_x += cfg.mu * euler_x * dt + cfg.sigma * euler_x * dw;
310            // Milstein
311            mil_x += cfg.mu * mil_x * dt
312                + cfg.sigma * mil_x * dw
313                + 0.5 * cfg.sigma * cfg.sigma * mil_x * (dw * dw - dt);
314        }
315        let exact =
316            cfg.x0 * ((cfg.mu - 0.5 * cfg.sigma * cfg.sigma) * cfg.t_end + cfg.sigma * w).exp();
317        let euler_err = (euler_x - exact).abs();
318        let mil_err = (mil_x - exact).abs();
319        assert!(
320            mil_err < euler_err,
321            "Milstein pathwise err {mil_err} should be smaller than Euler's {euler_err}"
322        );
323    }
324
325    #[test]
326    fn ou_reverts_to_mean() {
327        // OU: E[X_t] → μ as t → ∞. With θ=1, t=3, E ≈ μ + (x0−μ)e^(−3)
328        let cfg = SdeConfig {
329            x0: 0.0,
330            t_end: 3.0,
331            n_steps: 300,
332            n_paths: 20_000,
333            drift: DriftType::Ou,
334            mu: 5.0,
335            theta: 1.0,
336            sigma: 0.5,
337            solver: Solver::Euler,
338            seed: 3,
339        };
340        let r = solve(&cfg);
341        let analytic = 5.0 + (0.0 - 5.0) * (-3.0_f64).exp();
342        assert!(
343            (r.mean - analytic).abs() < 0.05,
344            "ou mean {} vs analytic {}",
345            r.mean,
346            analytic
347        );
348    }
349
350    #[test]
351    fn gbm_paths_never_negative_in_milstein_small_step() {
352        // Milstein with a small step on GBM should stay positive
353        let cfg = SdeConfig {
354            x0: 100.0,
355            t_end: 1.0,
356            n_steps: 500,
357            n_paths: 5000,
358            drift: DriftType::Gbm,
359            mu: 0.05,
360            sigma: 0.2,
361            solver: Solver::Milstein,
362            seed: 99,
363            ..Default::default()
364        };
365        let r = solve(&cfg);
366        assert!(
367            r.min > 0.0,
368            "GBM Milstein min should stay positive, got {}",
369            r.min
370        );
371    }
372
373    #[test]
374    fn mlmc_improves_estimate_on_coarse_grid() {
375        let base = SdeConfig {
376            x0: 100.0,
377            t_end: 1.0,
378            n_steps: 8, // coarse — big bias
379            n_paths: 10_000,
380            drift: DriftType::Gbm,
381            mu: 0.05,
382            sigma: 0.4,
383            seed: 11,
384            ..Default::default()
385        };
386        let analytic = 100.0 * (0.05_f64).exp();
387        let fine = solve(&SdeConfig { n_steps: 8, ..base });
388        let mlmc = solve_mlmc(&base);
389        let fine_err = (fine.mean - analytic).abs();
390        let mlmc_err = (mlmc.mlmc_mean - analytic).abs();
391        assert!(
392            mlmc_err <= fine_err + 1e-9,
393            "mlmc err {mlmc_err} should be <= fine err {fine_err}"
394        );
395    }
396
397    #[test]
398    fn drift_type_parsing() {
399        assert_eq!(DriftType::parse("gbm").unwrap(), DriftType::Gbm);
400        assert_eq!(DriftType::parse("ou").unwrap(), DriftType::Ou);
401        assert!(DriftType::parse("bogus").is_err());
402        assert_eq!(Solver::parse("euler").unwrap(), Solver::Euler);
403        assert_eq!(Solver::parse("milstein").unwrap(), Solver::Milstein);
404        assert!(Solver::parse("bogus").is_err());
405    }
406}