1use crate::{Error, Result};
18
19const SQRT_2PI: f64 = 2.506_628_274_631_000_2;
21
22#[derive(Debug, Clone, PartialEq)]
24pub enum Bandwidth {
25 Silverman,
31 Fixed(Vec<f64>),
33}
34
35#[derive(Debug, Clone, PartialEq)]
37pub struct KdeSampler {
38 samples: Vec<Vec<f64>>,
39 widths: Vec<f64>,
40}
41
42impl KdeSampler {
43 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 pub fn n_samples(&self) -> usize {
117 self.samples.len()
118 }
119
120 pub fn dims(&self) -> usize {
122 self.widths.len()
123 }
124
125 pub fn bandwidths(&self) -> &[f64] {
127 &self.widths
128 }
129
130 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 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 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 assert!(KdeSampler::fit(&vec![vec![3.0]; 8], Bandwidth::Silverman).is_err());
224 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 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 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 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}