Skip to main content

hpr_core/
random.rs

1//! Seeded pseudorandom numbers for turbulence, dispersions and Monte Carlo.
2//!
3//! [`SeededRng`] is **xoshiro256++**, seeded from one `u64` through **SplitMix64**, exactly as the
4//! authors' reference code recommends:
5//!
6//! - D. Blackman and S. Vigna, "Scrambled linear pseudorandom number generators", *ACM Trans.
7//!   Math. Softw.* 47(4), article 36 (2021), <https://doi.org/10.1145/3460772>; reference C code
8//!   (public domain) at <https://prng.di.unimi.it/>.
9//!
10//! The algorithm is part of the result: the same seed gives the same integer stream on every
11//! platform and in every release. Changing the generator, the seeding or the order of draws
12//! changes every seeded result, so it needs an ADR.
13//!
14//! Normal deviates use the polar method of G. Marsaglia and T. A. Bray, "A convenient method for
15//! generating normal variables", *SIAM Review* 6(3), 260–264 (1964): draw `u, v` uniform on
16//! `(−1, 1)` until `0 < s = u² + v² < 1`, then `u·√(−2 ln s / s)` and `v·√(−2 ln s / s)` are two
17//! independent standard normal deviates. The second one is kept for the next call. The method
18//! needs only `ln` and `sqrt`, so normal deviates are bit-identical on one platform, and across
19//! platforms wherever their math libraries' `ln` agree.
20
21use serde::{Deserialize, Serialize};
22
23use crate::error::CoreError;
24
25/// The SplitMix64 increment, `⌊2⁶⁴/φ⌋` (odd).
26const SPLITMIX64_GAMMA: u64 = 0x9e37_79b9_7f4a_7c15;
27
28/// One SplitMix64 step: advances `state` and returns the next output (Steele, Lea and Flood,
29/// "Fast splittable pseudorandom number generators", OOPSLA 2014; the constants are those of the
30/// public-domain reference at <https://prng.di.unimi.it/splitmix64.c>).
31fn splitmix64(state: &mut u64) -> u64 {
32    *state = state.wrapping_add(SPLITMIX64_GAMMA);
33    mix64(*state)
34}
35
36/// SplitMix64's output function: a bijection on 64-bit words that scrambles every input bit into
37/// every output bit.
38fn mix64(mut z: u64) -> u64 {
39    z = (z ^ (z >> 30)).wrapping_mul(0xbf58_476d_1ce4_e5b9);
40    z = (z ^ (z >> 27)).wrapping_mul(0x94d0_49bb_1331_11eb);
41    z ^ (z >> 31)
42}
43
44/// A seeded xoshiro256++ generator with a standard normal sampler.
45///
46/// It serializes as its 256-bit state, written as four hexadecimal strings (`"0x…"`, because JSON
47/// readers such as JavaScript's lose integers above 2⁵³), and the spare normal deviate, so a run
48/// can be checkpointed and resumed bit for bit. The all-zero state is rejected: xoshiro would stay
49/// at zero forever.
50#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
51#[serde(try_from = "SeededRngData", into = "SeededRngData")]
52pub struct SeededRng {
53    state: [u64; 4],
54    spare_normal: Option<f64>,
55}
56
57#[derive(Serialize, Deserialize)]
58#[serde(deny_unknown_fields)]
59struct SeededRngData {
60    state: [String; 4],
61    #[serde(default)]
62    spare_normal: Option<f64>,
63}
64
65/// Parses one state word written as `0x` and 1 to 16 hexadecimal digits.
66fn parse_state_word(text: &str) -> Result<u64, CoreError> {
67    text.strip_prefix("0x")
68        .filter(|digits| {
69            !digits.is_empty()
70                && digits.len() <= 16
71                && digits.bytes().all(|b| b.is_ascii_hexdigit())
72        })
73        .and_then(|digits| u64::from_str_radix(digits, 16).ok())
74        .ok_or(CoreError::InvalidRandomState)
75}
76
77impl TryFrom<SeededRngData> for SeededRng {
78    type Error = CoreError;
79
80    fn try_from(data: SeededRngData) -> Result<Self, CoreError> {
81        let mut state = [0; 4];
82        for (word, text) in state.iter_mut().zip(&data.state) {
83            *word = parse_state_word(text)?;
84        }
85        if state == [0; 4] {
86            return Err(CoreError::InvalidRandomState);
87        }
88        if let Some(spare) = data.spare_normal
89            && !spare.is_finite()
90        {
91            return Err(CoreError::Domain {
92                what: "spare normal deviate",
93                value: spare,
94            });
95        }
96        Ok(SeededRng {
97            state,
98            spare_normal: data.spare_normal,
99        })
100    }
101}
102
103impl From<SeededRng> for SeededRngData {
104    fn from(rng: SeededRng) -> Self {
105        SeededRngData {
106            state: rng.state.map(|word| format!("{word:#018x}")),
107            spare_normal: rng.spare_normal,
108        }
109    }
110}
111
112impl SeededRng {
113    /// Seeds the generator: the four state words are the first four SplitMix64 outputs starting
114    /// from `seed`. SplitMix64 never gives four zero words in a row, so every seed is valid.
115    pub fn seed_from_u64(seed: u64) -> Self {
116        let mut sm = seed;
117        let state = [
118            splitmix64(&mut sm),
119            splitmix64(&mut sm),
120            splitmix64(&mut sm),
121            splitmix64(&mut sm),
122        ];
123        SeededRng {
124            state,
125            spare_normal: None,
126        }
127    }
128
129    /// A generator for one stream of a seeded run, keyed by `keys`: a Monte Carlo sample's index
130    /// and the variable it draws, say. The stream depends only on `seed` and `keys`, so a sample
131    /// draws the same numbers whatever else the run draws, however many samples it has and
132    /// however they are shared among threads: the counter-based idea of J. K. Salmon, M. A.
133    /// Moraes, R. O. Dror and D. E. Shaw, "Parallel random numbers: as easy as 1, 2, 3", *Proc.
134    /// SC11* (2011), <https://doi.org/10.1145/2063384.2063405>.
135    ///
136    /// The key folds into one word with SplitMix64's output function `mix`, a bijection on 64-bit
137    /// words: `h₀ = seed`, `hᵢ = mix(hᵢ₋₁ ⊕ mix(kᵢ + γ))` with `γ` SplitMix64's increment, and the
138    /// generator is [`SeededRng::seed_from_u64`]`(hₙ)`. Two different keys give the same `hₙ` only
139    /// by a collision of 64-bit hashes, about one chance in 2⁶⁴ for a pair of streams; with no
140    /// keys the stream is `seed_from_u64(seed)`'s.
141    pub fn for_stream(seed: u64, keys: &[u64]) -> Self {
142        let hash = keys.iter().fold(seed, |h, &key| {
143            mix64(h ^ mix64(key.wrapping_add(SPLITMIX64_GAMMA)))
144        });
145        Self::seed_from_u64(hash)
146    }
147
148    /// The next 64 random bits (the xoshiro256++ output function and state transition).
149    pub fn next_u64(&mut self) -> u64 {
150        let s = &mut self.state;
151        let result = s[0].wrapping_add(s[3]).rotate_left(23).wrapping_add(s[0]);
152        let t = s[1] << 17;
153        s[2] ^= s[0];
154        s[3] ^= s[1];
155        s[1] ^= s[2];
156        s[0] ^= s[3];
157        s[2] ^= t;
158        s[3] = s[3].rotate_left(45);
159        result
160    }
161
162    /// A uniform deviate on `[0, 1)`: the top 53 bits of [`SeededRng::next_u64`] times `2⁻⁵³`,
163    /// so every value is a multiple of `2⁻⁵³`.
164    pub fn uniform(&mut self) -> f64 {
165        // 2⁻⁵³ as an exact f64 literal.
166        const SCALE: f64 = 1.0 / 9_007_199_254_740_992.0;
167        // Cast: a 53-bit integer converts to f64 exactly.
168        let top = (self.next_u64() >> 11) as f64;
169        top * SCALE
170    }
171
172    /// A standard normal deviate (mean 0, variance 1) by the Marsaglia–Bray polar method.
173    ///
174    /// Each accepted pair gives two deviates; the second is returned by the next call. On average
175    /// a pair costs `4/π ≈ 1.27` attempts, two uniforms each.
176    pub fn standard_normal(&mut self) -> f64 {
177        if let Some(spare) = self.spare_normal.take() {
178            return spare;
179        }
180        loop {
181            let u = 2.0 * self.uniform() - 1.0;
182            let v = 2.0 * self.uniform() - 1.0;
183            let s = u * u + v * v;
184            if s > 0.0 && s < 1.0 {
185                let factor = (-2.0 * s.ln() / s).sqrt();
186                self.spare_normal = Some(v * factor);
187                return u * factor;
188            }
189        }
190    }
191}
192
193#[cfg(test)]
194mod tests {
195    use rand_core::{Rng, SeedableRng};
196    use rand_xoshiro::{SplitMix64, Xoshiro256PlusPlus};
197
198    use super::*;
199
200    /// The generator and its seeding match the rust-random `rand_xoshiro` crate, an independent
201    /// implementation of the same public-domain reference code.
202    #[test]
203    fn stream_matches_the_rand_xoshiro_implementation() {
204        for seed in [0, 1, 42, 0x0123_4567_89ab_cdef, u64::MAX] {
205            let mut ours = SeededRng::seed_from_u64(seed);
206            let mut theirs = Xoshiro256PlusPlus::seed_from_u64(seed);
207            for _ in 0..10_000 {
208                assert_eq!(ours.next_u64(), theirs.next_u64(), "seed {seed}");
209            }
210        }
211    }
212
213    #[test]
214    fn seeding_is_four_splitmix64_outputs() {
215        let mut theirs = SplitMix64::seed_from_u64(7);
216        let mut sm = 7;
217        for _ in 0..100 {
218            assert_eq!(splitmix64(&mut sm), theirs.next_u64());
219        }
220    }
221
222    #[test]
223    fn a_stream_depends_only_on_its_seed_and_keys() {
224        // No keys: the seed's own stream.
225        assert_eq!(SeededRng::for_stream(42, &[]), SeededRng::seed_from_u64(42));
226        // The fold, written out: one key is mix(seed ^ mix(key + γ)), seeded by SplitMix64.
227        let mut sm = 7_u64;
228        let inner = splitmix64(&mut sm);
229        let mut outer = (42 ^ inner).wrapping_sub(SPLITMIX64_GAMMA);
230        assert_eq!(
231            SeededRng::for_stream(42, &[7]),
232            SeededRng::seed_from_u64(splitmix64(&mut outer))
233        );
234        // Repeatable, and different for a different seed, key, key order or key count.
235        let mut first = SeededRng::for_stream(1, &[3, 5]);
236        let mut again = SeededRng::for_stream(1, &[3, 5]);
237        let a: Vec<u64> = (0..4).map(|_| first.next_u64()).collect();
238        let b: Vec<u64> = (0..4).map(|_| again.next_u64()).collect();
239        assert_eq!(a, b);
240        let others = [
241            SeededRng::for_stream(2, &[3, 5]),
242            SeededRng::for_stream(1, &[3, 6]),
243            SeededRng::for_stream(1, &[5, 3]),
244            SeededRng::for_stream(1, &[3]),
245            SeededRng::for_stream(1, &[3, 5, 0]),
246        ];
247        for mut other in others {
248            assert_ne!(other.next_u64(), a[0]);
249        }
250    }
251
252    #[test]
253    fn neighbouring_streams_are_uncorrelated() {
254        // The first uniform of 20,000 streams keyed 0, 1, 2…: the mean and the lag-one
255        // correlation of independent uniforms, within four standard errors.
256        let n = 20_000_u64;
257        let u: Vec<f64> = (0..n)
258            .map(|k| SeededRng::for_stream(2026, &[k]).uniform())
259            .collect();
260        let count = n as f64;
261        let mean = u.iter().sum::<f64>() / count;
262        assert!(
263            (mean - 0.5).abs() < 4.0 * (1.0 / 12.0 / count).sqrt(),
264            "{mean}"
265        );
266        let lag: f64 = u
267            .windows(2)
268            .map(|w| (w[0] - 0.5) * (w[1] - 0.5))
269            .sum::<f64>()
270            / (count - 1.0)
271            * 12.0;
272        assert!(lag.abs() < 4.0 / (count - 1.0).sqrt(), "{lag}");
273    }
274
275    #[test]
276    fn uniform_is_in_the_half_open_unit_interval() {
277        let mut rng = SeededRng::seed_from_u64(3);
278        let n = 1_000_000;
279        let mut sum = 0.0;
280        for _ in 0..n {
281            let x = rng.uniform();
282            assert!((0.0..1.0).contains(&x));
283            sum += x;
284        }
285        // Mean of U(0,1) is 1/2 with standard error √(1/12/n) ≈ 2.9e-4; allow 5 standard errors.
286        let mean = sum / f64::from(n);
287        assert!((mean - 0.5).abs() < 5.0 * (1.0 / 12.0 / f64::from(n)).sqrt());
288    }
289
290    /// Mean, variance, and the empirical CDF at −2, −1, 0, 1 and 2 standard deviations, each
291    /// within 5 standard errors for 10⁶ draws. Φ values from M. Abramowitz and I. A. Stegun,
292    /// *Handbook of Mathematical Functions*, Table 26.1.
293    #[test]
294    fn standard_normal_has_the_normal_moments_and_cdf() {
295        let mut rng = SeededRng::seed_from_u64(2026);
296        let n = 1_000_000;
297        let points = [-2.0, -1.0, 0.0, 1.0, 2.0];
298        let phi = [
299            0.022_750_131_948_179,
300            0.158_655_253_931_457,
301            0.5,
302            0.841_344_746_068_543,
303            0.977_249_868_051_821,
304        ];
305        let mut below = [0_u32; 5];
306        let (mut sum, mut sum_sq, mut sum_4) = (0.0, 0.0, 0.0);
307        for _ in 0..n {
308            let x = rng.standard_normal();
309            sum += x;
310            sum_sq += x * x;
311            sum_4 += x * x * x * x;
312            for (count, &point) in below.iter_mut().zip(&points) {
313                if x <= point {
314                    *count += 1;
315                }
316            }
317        }
318        let nf = f64::from(n);
319        let mean = sum / nf;
320        let variance = sum_sq / nf - mean * mean;
321        assert!(mean.abs() < 5.0 / nf.sqrt(), "mean {mean}");
322        // Var of the sample variance of N(0,1) is 2/n.
323        assert!(
324            (variance - 1.0).abs() < 5.0 * (2.0 / nf).sqrt(),
325            "variance {variance}"
326        );
327        // E[x⁴] = 3 with Var = (105 − 9)/n.
328        assert!((sum_4 / nf - 3.0).abs() < 5.0 * (96.0 / nf).sqrt());
329        for ((&count, &p), &point) in below.iter().zip(&phi).zip(&points) {
330            let fraction = f64::from(count) / nf;
331            let standard_error = (p * (1.0 - p) / nf).sqrt();
332            assert!(
333                (fraction - p).abs() < 5.0 * standard_error,
334                "CDF at {point}: {fraction} vs {p}"
335            );
336        }
337    }
338
339    #[test]
340    fn same_seed_same_normals_and_serde_resumes_the_stream() {
341        let mut a = SeededRng::seed_from_u64(99);
342        let mut b = SeededRng::seed_from_u64(99);
343        for _ in 0..1001 {
344            assert_eq!(a.standard_normal().to_bits(), b.standard_normal().to_bits());
345        }
346        // Mid-pair: `a` holds a spare deviate, which the checkpoint must carry.
347        assert!(a.spare_normal.is_some());
348        let json = serde_json::to_string(&a).unwrap();
349        let mut resumed: SeededRng = serde_json::from_str(&json).unwrap();
350        for _ in 0..1000 {
351            assert_eq!(
352                a.standard_normal().to_bits(),
353                resumed.standard_normal().to_bits()
354            );
355        }
356    }
357
358    #[test]
359    fn state_serializes_as_hex_words_and_bad_states_are_rejected() {
360        let rng = SeededRng::seed_from_u64(1);
361        let value = serde_json::to_value(&rng).unwrap();
362        let words = value["state"].as_array().unwrap();
363        assert_eq!(words.len(), 4);
364        for (word, &expected) in words.iter().zip(&rng.state) {
365            let text = word.as_str().unwrap();
366            assert_eq!(text.len(), 18);
367            assert_eq!(parse_state_word(text).unwrap(), expected);
368        }
369        let zero = r#"{"state":["0x0","0x0","0x0","0x00"],"spare_normal":null}"#;
370        assert_eq!(
371            serde_json::from_str::<SeededRng>(zero)
372                .unwrap_err()
373                .to_string(),
374            CoreError::InvalidRandomState.to_string()
375        );
376        for bad in [
377            r#"{"state":[1,2,3,4]}"#,
378            r#"{"state":["1","0x2","0x3","0x4"]}"#,
379            r#"{"state":["0x","0x2","0x3","0x4"]}"#,
380            r#"{"state":["0x10000000000000000","0x2","0x3","0x4"]}"#,
381            r#"{"state":["0xg","0x2","0x3","0x4"]}"#,
382            r#"{"state":["0x+1","0x2","0x3","0x4"]}"#,
383        ] {
384            assert!(serde_json::from_str::<SeededRng>(bad).is_err(), "{bad}");
385        }
386        let short = r#"{"state":["0x1","0x2","0x3","0x4"]}"#;
387        assert_eq!(
388            serde_json::from_str::<SeededRng>(short).unwrap().state,
389            [1, 2, 3, 4]
390        );
391    }
392}