1use crate::bayesian::rand_u01;
15
16fn 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#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
25pub enum DriftType {
26 Gbm,
28 Ou,
30}
31
32impl DriftType {
33 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#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
45pub enum Solver {
46 Euler,
48 Milstein,
50}
51
52impl Solver {
53 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#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
67pub struct SdeConfig {
68 pub x0: f64,
70 pub t_end: f64,
72 pub n_steps: usize,
74 pub n_paths: usize,
76 pub drift: DriftType,
78 pub mu: f64,
80 pub theta: f64,
82 pub sigma: f64,
84 pub solver: Solver,
86 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
107fn 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#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
117pub struct SdeResult {
118 pub mean: f64,
120 pub std: f64,
122 pub p05: f64,
124 pub p50: f64,
126 pub p95: f64,
128 pub min: f64,
130 pub max: f64,
132 pub n_paths: usize,
134 pub dt: f64,
136}
137
138#[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 let (_, diff) = drift_diffusion(x, cfg);
154 let sigma_prime = match cfg.drift {
155 DriftType::Gbm => cfg.sigma, DriftType::Ou => 0.0, };
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#[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 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#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
203pub struct MlMcResult {
204 pub mlmc_mean: f64,
206 pub fine_mean: f64,
208 pub coarse_mean: f64,
210 pub fine_std: f64,
212 pub n_paths: usize,
214 pub fine_steps: usize,
216 pub coarse_steps: usize,
218}
219
220fn 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)] use super::*;
254
255 #[test]
256 fn gbm_euler_matches_analytic_mean() {
257 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 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_x += cfg.mu * euler_x * dt + cfg.sigma * euler_x * dw;
310 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 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 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, 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}