Skip to main content

nucleide_vr_tools/
kde.rs

1//! Gaussian kernel-density source sampling (KDSource-class, clean-room).
2//!
3//! [`KdeSampler`] fits an axis-aligned Gaussian KDE over caller-supplied
4//! particle vectors (energy/position/direction rows, e.g. projected from
5//! MCPL records by the caller) and resamples synthetic particles with the
6//! same density to boost downstream statistics. The KDE equations
7//! (Gaussian product kernel, Silverman bandwidth rule) are textbook
8//! statistics implemented here from scratch; no upstream code or data is
9//! read, ported, or vendored.
10//!
11//! Draws are deterministic in the house style ([`crate::sampling`]:
12//! caller-supplied randoms, no RNG inside): [`KdeSampler::draw`] takes one
13//! uniform `u` selecting the kernel centre plus one standard normal per
14//! dimension. MCPL projection stays caller-side (a ~10-line map from
15//! particle fields to rows); this module never depends on `mcpl-io`.
16
17use crate::{Error, Result};
18
19/// sqrt(2π): Gaussian normalization factor (no std const exists).
20const SQRT_2PI: f64 = 2.506_628_274_631_000_2;
21
22/// Bandwidth rule for [`KdeSampler::fit`].
23#[derive(Debug, Clone, PartialEq)]
24pub enum Bandwidth {
25    /// Silverman's rule per dimension,
26    /// `h = σ (4 / (d + 2) / n)^(1 / (d + 4))` with the unbiased
27    /// per-dimension standard deviation `σ`. Zero-variance dimensions are
28    /// an [`Error`] (a zero bandwidth is a delta spike, never a density);
29    /// use [`Bandwidth::Fixed`] with an explicit width instead.
30    Silverman,
31    /// Caller-supplied per-dimension widths (all finite and positive).
32    Fixed(Vec<f64>),
33}
34
35/// Axis-aligned Gaussian KDE over `n` samples in `d` dimensions.
36#[derive(Debug, Clone, PartialEq)]
37pub struct KdeSampler {
38    samples: Vec<Vec<f64>>,
39    widths: Vec<f64>,
40}
41
42impl KdeSampler {
43    /// Fit a KDE over `samples` (non-empty, rectangular, all finite).
44    pub fn fit(samples: &[Vec<f64>], bandwidth: Bandwidth) -> Result<Self> {
45        if samples.is_empty() {
46            return Err(Error::EmptyPdf);
47        }
48        let dims = samples[0].len();
49        if dims == 0 {
50            return Err(Error::LengthMismatch {
51                expected: 1,
52                got: 0,
53            });
54        }
55        for (i, row) in samples.iter().enumerate() {
56            if row.len() != dims {
57                return Err(Error::LengthMismatch {
58                    expected: dims,
59                    got: row.len(),
60                });
61            }
62            if row.iter().any(|v| !v.is_finite()) {
63                return Err(Error::NonFiniteTally {
64                    field: "kde sample",
65                    index: i,
66                });
67            }
68        }
69        let n = samples.len() as f64;
70        let d = dims as f64;
71        let widths = match bandwidth {
72            Bandwidth::Fixed(w) => {
73                if w.len() != dims {
74                    return Err(Error::LengthMismatch {
75                        expected: dims,
76                        got: w.len(),
77                    });
78                }
79                if w.iter().any(|v| !v.is_finite() || *v <= 0.0) {
80                    return Err(Error::NonFiniteTally {
81                        field: "kde bandwidth",
82                        index: w
83                            .iter()
84                            .position(|v| !v.is_finite() || *v <= 0.0)
85                            .unwrap_or(0),
86                    });
87                }
88                w
89            }
90            Bandwidth::Silverman => {
91                let factor = (4.0 / (d + 2.0) / n).powf(1.0 / (d + 4.0));
92                let mut widths = Vec::with_capacity(dims);
93                for j in 0..dims {
94                    let mean = samples.iter().map(|row| row[j]).sum::<f64>() / n;
95                    let var = samples
96                        .iter()
97                        .map(|row| (row[j] - mean).powi(2))
98                        .sum::<f64>()
99                        / (n - 1.0).max(1.0);
100                    let sigma = var.sqrt();
101                    if sigma <= 0.0 {
102                        return Err(Error::ZeroVarianceDim { dim: j });
103                    }
104                    widths.push(sigma * factor);
105                }
106                widths
107            }
108        };
109        Ok(Self {
110            samples: samples.to_vec(),
111            widths,
112        })
113    }
114
115    /// Number of fitted samples.
116    pub fn n_samples(&self) -> usize {
117        self.samples.len()
118    }
119
120    /// Sample dimensionality.
121    pub fn dims(&self) -> usize {
122        self.widths.len()
123    }
124
125    /// Fitted per-dimension bandwidths.
126    pub fn bandwidths(&self) -> &[f64] {
127        &self.widths
128    }
129
130    /// KDE density at `point` (Gaussian product kernel, length-checked).
131    pub fn pdf(&self, point: &[f64]) -> Result<f64> {
132        if point.len() != self.dims() {
133            return Err(Error::LengthMismatch {
134                expected: self.dims(),
135                got: point.len(),
136            });
137        }
138        if point.iter().any(|v| !v.is_finite()) {
139            return Err(Error::NonFiniteTally {
140                field: "kde point",
141                index: point.iter().position(|v| !v.is_finite()).unwrap_or(0),
142            });
143        }
144        let norm = self.widths.iter().map(|h| h * SQRT_2PI).product::<f64>();
145        let mut density = 0.0;
146        for row in &self.samples {
147            let mut z2 = 0.0;
148            for ((x, c), h) in point.iter().zip(row).zip(&self.widths) {
149                let z = (x - c) / h;
150                z2 += z * z;
151            }
152            density += (-0.5 * z2).exp();
153        }
154        Ok(density / norm / self.samples.len() as f64)
155    }
156
157    /// Resample one synthetic particle: `u` in `[0, 1)` selects kernel
158    /// centre `floor(u * n)`; `normals` (one finite standard normal per
159    /// dimension) perturbs it by `width * z` per dimension.
160    pub fn draw(&self, u: f64, normals: &[f64]) -> Result<Vec<f64>> {
161        if !(0.0..1.0).contains(&u) {
162            return Err(Error::BadDraw { value: u });
163        }
164        if normals.len() != self.dims() {
165            return Err(Error::LengthMismatch {
166                expected: self.dims(),
167                got: normals.len(),
168            });
169        }
170        if normals.iter().any(|v| !v.is_finite()) {
171            return Err(Error::NonFiniteTally {
172                field: "kde normal",
173                index: normals.iter().position(|v| !v.is_finite()).unwrap_or(0),
174            });
175        }
176        let centre = &self.samples[(u * self.samples.len() as f64) as usize];
177        Ok(centre
178            .iter()
179            .zip(&self.widths)
180            .zip(normals)
181            .map(|((c, h), z)| c + h * z)
182            .collect())
183    }
184}
185
186#[cfg(test)]
187mod tests {
188    use super::*;
189
190    /// Deterministic test stream (LCG + Box-Muller); test-only, never shipped.
191    struct TestRng(u64);
192
193    impl TestRng {
194        fn uniform(&mut self) -> f64 {
195            self.0 = self
196                .0
197                .wrapping_mul(6364136223846793005)
198                .wrapping_add(1442695040888963407);
199            ((self.0 >> 11) as f64) / ((1u64 << 53) as f64)
200        }
201
202        fn normal(&mut self) -> f64 {
203            let (u1, u2) = (self.uniform().max(1e-300), self.uniform());
204            (-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos()
205        }
206    }
207
208    fn gaussian_samples(mean: f64, sigma: f64, n: usize) -> Vec<Vec<f64>> {
209        let mut rng = TestRng(0x1234_5678_9abc_def0);
210        (0..n).map(|_| vec![mean + sigma * rng.normal()]).collect()
211    }
212
213    #[test]
214    fn input_errors_are_loud() {
215        assert!(KdeSampler::fit(&[], Bandwidth::Silverman).is_err());
216        assert!(KdeSampler::fit(&[vec![]], Bandwidth::Silverman).is_err());
217        assert!(KdeSampler::fit(&[vec![1.0], vec![1.0, 2.0]], Bandwidth::Silverman).is_err());
218        assert!(KdeSampler::fit(&[vec![f64::NAN]], Bandwidth::Silverman).is_err());
219        assert!(KdeSampler::fit(&[vec![1.0]], Bandwidth::Fixed(vec![])).is_err());
220        assert!(KdeSampler::fit(&[vec![1.0]], Bandwidth::Fixed(vec![0.0])).is_err());
221        assert!(KdeSampler::fit(&[vec![1.0]], Bandwidth::Fixed(vec![-2.0])).is_err());
222        // Zero variance under Silverman is a delta spike, not a density.
223        assert!(KdeSampler::fit(&vec![vec![3.0]; 8], Bandwidth::Silverman).is_err());
224        // ...but an explicit fixed width is the caller's choice.
225        assert!(KdeSampler::fit(&vec![vec![3.0]; 8], Bandwidth::Fixed(vec![0.5])).is_ok());
226    }
227
228    #[test]
229    fn draw_is_exact_and_deterministic() {
230        let kde = KdeSampler::fit(
231            &[vec![1.0, 2.0], vec![3.0, 4.0]],
232            Bandwidth::Fixed(vec![0.5, 2.0]),
233        )
234        .unwrap();
235        // u = 0.75 selects centre 1; perturbation is width * z per dim.
236        let got = kde.draw(0.75, &[1.0, -0.5]).unwrap();
237        assert_eq!(got, vec![3.5, 3.0]);
238        assert_eq!(kde.draw(0.75, &[1.0, -0.5]).unwrap(), got);
239        assert!(kde.draw(1.0, &[0.0, 0.0]).is_err());
240        assert!(kde.draw(-0.1, &[0.0, 0.0]).is_err());
241        assert!(kde.draw(0.5, &[0.0]).is_err());
242        assert!(kde.draw(0.5, &[0.0, f64::INFINITY]).is_err());
243    }
244
245    #[test]
246    fn gaussian_recovery_and_normalization() {
247        // 20000 draws from N(5, 2²): KDE mean within 5 SE, pdf at the mode
248        // within 2% of the closed form, integral within 1e-3 of 1.
249        let samples = gaussian_samples(5.0, 2.0, 20_000);
250        let kde = KdeSampler::fit(&samples, Bandwidth::Silverman).unwrap();
251        let h = kde.bandwidths()[0];
252        assert!(h > 0.0 && h < 1.0, "silverman width {h}");
253        let mut draws = Vec::with_capacity(4096);
254        let mut rng = TestRng(0xabcd);
255        for _ in 0..4096 {
256            draws.push(kde.draw(rng.uniform(), &[rng.normal()]).unwrap()[0]);
257        }
258        let mean = draws.iter().sum::<f64>() / draws.len() as f64;
259        let se = 2.0 / (draws.len() as f64).sqrt();
260        assert!((mean - 5.0).abs() < 5.0 * se, "mean {mean}");
261        let closed = (-0.5 * ((5.0f64 - 5.0) / 2.0).powi(2)).exp() / (2.0 * SQRT_2PI);
262        let got = kde.pdf(&[5.0]).unwrap();
263        assert!(
264            (got - closed).abs() / closed < 0.02,
265            "pdf {got} vs {closed}"
266        );
267        // Trapezoid integral over ±8σ.
268        let (lo, hi, m) = (-11.0, 21.0, 2048);
269        let mut area = 0.0;
270        let mut prev = kde.pdf(&[lo]).unwrap();
271        for i in 1..=m {
272            let x = lo + (hi - lo) * i as f64 / m as f64;
273            let cur = kde.pdf(&[x]).unwrap();
274            area += 0.5 * (prev + cur) * (hi - lo) / m as f64;
275            prev = cur;
276        }
277        assert!((area - 1.0).abs() < 1e-3, "integral {area}");
278    }
279
280    #[test]
281    fn bandwidth_is_deterministic() {
282        let samples = gaussian_samples(0.0, 1.0, 512);
283        let a = KdeSampler::fit(&samples, Bandwidth::Silverman).unwrap();
284        let b = KdeSampler::fit(&samples, Bandwidth::Silverman).unwrap();
285        assert_eq!(a.bandwidths(), b.bandwidths());
286    }
287}