1use crate::error::{InterpolateError, InterpolateResult};
36
37pub trait Variogram: Send + Sync {
44 fn gamma(&self, h: f64) -> f64;
46
47 fn is_bounded(&self) -> bool {
53 true
54 }
55
56 fn clone_box(&self) -> Box<dyn Variogram>;
58}
59
60#[derive(Debug, Clone, Copy)]
65pub struct SphericalVariogram {
66 pub nugget: f64,
68 pub sill: f64,
70 pub range: f64,
72}
73
74impl Variogram for SphericalVariogram {
75 fn gamma(&self, h: f64) -> f64 {
76 if h <= 0.0 {
77 return 0.0;
78 }
79 if h >= self.range {
80 return self.nugget + self.sill;
81 }
82 let u = h / self.range;
83 self.nugget + self.sill * (1.5 * u - 0.5 * u * u * u)
84 }
85
86 fn clone_box(&self) -> Box<dyn Variogram> {
87 Box::new(*self)
88 }
89}
90
91#[derive(Debug, Clone, Copy)]
95pub struct ExponentialVariogram {
96 pub nugget: f64,
97 pub sill: f64,
98 pub range: f64,
99}
100
101impl Variogram for ExponentialVariogram {
102 fn gamma(&self, h: f64) -> f64 {
103 if h <= 0.0 {
104 return 0.0;
105 }
106 self.nugget + self.sill * (1.0 - (-3.0 * h / self.range).exp())
107 }
108
109 fn clone_box(&self) -> Box<dyn Variogram> {
110 Box::new(*self)
111 }
112}
113
114#[derive(Debug, Clone, Copy)]
118pub struct GaussianVariogram {
119 pub nugget: f64,
120 pub sill: f64,
121 pub range: f64,
122}
123
124impl Variogram for GaussianVariogram {
125 fn gamma(&self, h: f64) -> f64 {
126 if h <= 0.0 {
127 return 0.0;
128 }
129 let u = h / self.range;
130 self.nugget + self.sill * (1.0 - (-3.0 * u * u).exp())
131 }
132
133 fn clone_box(&self) -> Box<dyn Variogram> {
134 Box::new(*self)
135 }
136}
137
138#[derive(Debug, Clone, Copy)]
142pub struct PowerVariogram {
143 pub nugget: f64,
144 pub slope: f64,
145 pub power: f64,
146}
147
148impl Variogram for PowerVariogram {
149 fn gamma(&self, h: f64) -> f64 {
150 if h <= 0.0 {
151 return 0.0;
152 }
153 self.nugget + self.slope * h.powf(self.power)
154 }
155
156 fn is_bounded(&self) -> bool {
157 false
158 }
159
160 fn clone_box(&self) -> Box<dyn Variogram> {
161 Box::new(*self)
162 }
163}
164
165fn lu_factor(mut a: Vec<f64>, n: usize) -> InterpolateResult<(Vec<f64>, Vec<usize>)> {
173 let mut piv: Vec<usize> = (0..n).collect();
174 for k in 0..n {
175 let mut max_val = a[k * n + k].abs();
177 let mut max_row = k;
178 for i in (k + 1)..n {
179 let v = a[i * n + k].abs();
180 if v > max_val {
181 max_val = v;
182 max_row = i;
183 }
184 }
185 if max_val < 1e-15 {
186 return Err(InterpolateError::ComputationError(
187 "Singular kriging matrix; add nugget > 0 or check data".into(),
188 ));
189 }
190 if max_row != k {
192 piv.swap(k, max_row);
193 for j in 0..n {
194 let tmp = a[k * n + j];
195 a[k * n + j] = a[max_row * n + j];
196 a[max_row * n + j] = tmp;
197 }
198 }
199 for i in (k + 1)..n {
201 a[i * n + k] /= a[k * n + k];
202 for j in (k + 1)..n {
203 let tmp = a[i * n + k] * a[k * n + j];
204 a[i * n + j] -= tmp;
205 }
206 }
207 }
208 Ok((a, piv))
209}
210
211fn lu_solve(lu: &[f64], piv: &[usize], b: &[f64], n: usize) -> Vec<f64> {
213 let mut x: Vec<f64> = (0..n).map(|i| b[piv[i]]).collect();
215 for i in 0..n {
217 for j in 0..i {
218 x[i] -= lu[i * n + j] * x[j];
219 }
220 }
221 for i in (0..n).rev() {
223 for j in (i + 1)..n {
224 x[i] -= lu[i * n + j] * x[j];
225 }
226 x[i] /= lu[i * n + i];
227 }
228 x
229}
230
231pub struct OrdinaryKriging {
248 pub points: Vec<Vec<f64>>,
250 pub values: Vec<f64>,
252 variogram: Box<dyn Variogram>,
254 lu_mat: Vec<f64>,
256 lu_piv: Vec<usize>,
258 c0: f64,
260 n: usize,
262 bounded: bool,
264}
265
266impl Clone for OrdinaryKriging {
267 fn clone(&self) -> Self {
268 Self {
269 points: self.points.clone(),
270 values: self.values.clone(),
271 variogram: self.variogram.clone_box(),
272 lu_mat: self.lu_mat.clone(),
273 lu_piv: self.lu_piv.clone(),
274 c0: self.c0,
275 n: self.n,
276 bounded: self.bounded,
277 }
278 }
279}
280
281impl std::fmt::Debug for OrdinaryKriging {
282 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
283 f.debug_struct("OrdinaryKriging")
284 .field("n", &self.n)
285 .finish()
286 }
287}
288
289impl OrdinaryKriging {
290 pub fn fit(
302 points: Vec<Vec<f64>>,
303 values: Vec<f64>,
304 variogram: Box<dyn Variogram>,
305 ) -> InterpolateResult<OrdinaryKriging> {
306 let n = points.len();
307 if n == 0 {
308 return Err(InterpolateError::InvalidInput {
309 message: "no data points".into(),
310 });
311 }
312 if values.len() != n {
313 return Err(InterpolateError::ShapeMismatch {
314 expected: format!("{}", n),
315 actual: format!("{}", values.len()),
316 object: "values".into(),
317 });
318 }
319
320 let bounded = variogram.is_bounded();
321
322 let c0 = if bounded {
325 variogram.gamma(1e12)
326 } else {
327 0.0 };
329
330 let m = n + 1;
332 let mut mat = vec![0.0_f64; m * m];
333 for i in 0..n {
334 for j in 0..n {
335 let h = euclidean_dist(&points[i], &points[j]);
336 let gamma = variogram.gamma(h);
337 mat[i * m + j] = if bounded { c0 - gamma } else { gamma };
338 }
339 mat[i * m + n] = 1.0;
341 mat[n * m + i] = 1.0;
342 }
343 mat[n * m + n] = 0.0;
345
346 let (lu_mat, lu_piv) = lu_factor(mat, m)?;
347
348 Ok(OrdinaryKriging {
349 points,
350 values,
351 variogram,
352 lu_mat,
353 lu_piv,
354 c0,
355 n,
356 bounded,
357 })
358 }
359
360 pub fn predict(&self, x: &[f64]) -> InterpolateResult<(f64, f64)> {
368 if !self.points.is_empty() && x.len() != self.points[0].len() {
369 return Err(InterpolateError::DimensionMismatch(format!(
370 "expected dim {}, got {}",
371 self.points[0].len(),
372 x.len()
373 )));
374 }
375
376 let m = self.n + 1;
377 let mut rhs = vec![0.0_f64; m];
378 for i in 0..self.n {
379 let h = euclidean_dist(x, &self.points[i]);
380 let gamma = self.variogram.gamma(h);
381 rhs[i] = if self.bounded { self.c0 - gamma } else { gamma };
382 }
383 rhs[self.n] = 1.0;
384
385 let sol = lu_solve(&self.lu_mat, &self.lu_piv, &rhs, m);
387
388 let estimate: f64 = (0..self.n).map(|i| sol[i] * self.values[i]).sum();
390
391 let rhs_dot_w: f64 = (0..self.n).map(|i| rhs[i] * sol[i]).sum();
393 let variance = if self.bounded {
394 (self.c0 - rhs_dot_w - sol[self.n]).max(0.0)
396 } else {
397 (rhs_dot_w + sol[self.n]).max(0.0)
399 };
400
401 Ok((estimate, variance))
402 }
403}
404
405fn euclidean_dist(a: &[f64], b: &[f64]) -> f64 {
410 a.iter()
411 .zip(b.iter())
412 .map(|(x, y)| (x - y) * (x - y))
413 .sum::<f64>()
414 .sqrt()
415}
416
417#[cfg(test)]
422mod tests {
423 use super::*;
424
425 fn make_1d_data() -> (Vec<Vec<f64>>, Vec<f64>) {
426 let xs = vec![0.0_f64, 1.0, 2.0, 3.0, 4.0];
427 let pts: Vec<Vec<f64>> = xs.iter().map(|&x| vec![x]).collect();
428 let vals: Vec<f64> = xs.iter().map(|&x| x * x).collect(); (pts, vals)
430 }
431
432 #[test]
433 fn test_spherical_kriging_interpolates_data() {
434 let (pts, vals) = make_1d_data();
435 let vgm = SphericalVariogram {
436 nugget: 0.0,
437 sill: 20.0,
438 range: 10.0,
439 };
440 let ok =
441 OrdinaryKriging::fit(pts.clone(), vals.clone(), Box::new(vgm)).expect("fit failed");
442
443 for (p, &v) in pts.iter().zip(vals.iter()) {
444 let (est, _var) = ok.predict(p).expect("predict failed");
445 assert!(
446 (est - v).abs() < 1e-6,
447 "spherical: at {:?} expected {} got {}",
448 p,
449 v,
450 est
451 );
452 }
453 }
454
455 #[test]
456 fn test_exponential_kriging_interpolates_data() {
457 let (pts, vals) = make_1d_data();
458 let vgm = ExponentialVariogram {
459 nugget: 0.0,
460 sill: 20.0,
461 range: 10.0,
462 };
463 let ok =
464 OrdinaryKriging::fit(pts.clone(), vals.clone(), Box::new(vgm)).expect("fit failed");
465
466 for (p, &v) in pts.iter().zip(vals.iter()) {
467 let (est, _) = ok.predict(p).expect("predict");
468 assert!((est - v).abs() < 1e-6, "exp: {:?} {} {}", p, v, est);
469 }
470 }
471
472 #[test]
473 fn test_gaussian_kriging_interpolates_data() {
474 let (pts, vals) = make_1d_data();
475 let vgm = GaussianVariogram {
476 nugget: 0.0,
477 sill: 20.0,
478 range: 10.0,
479 };
480 let ok =
481 OrdinaryKriging::fit(pts.clone(), vals.clone(), Box::new(vgm)).expect("fit failed");
482
483 for (p, &v) in pts.iter().zip(vals.iter()) {
484 let (est, _) = ok.predict(p).expect("predict");
485 assert!((est - v).abs() < 1e-6, "gauss: {:?} {} {}", p, v, est);
486 }
487 }
488
489 #[test]
490 fn test_power_variogram() {
491 let (pts, vals) = make_1d_data();
492 let vgm = PowerVariogram {
493 nugget: 0.0,
494 slope: 1.0,
495 power: 1.5,
496 };
497 let ok =
498 OrdinaryKriging::fit(pts.clone(), vals.clone(), Box::new(vgm)).expect("fit failed");
499
500 for (p, &v) in pts.iter().zip(vals.iter()) {
501 let (est, _) = ok.predict(p).expect("predict");
502 assert!((est - v).abs() < 1e-4, "power: {:?} {} {}", p, v, est);
503 }
504 }
505
506 #[test]
507 fn test_variance_is_nonnegative() {
508 let (pts, vals) = make_1d_data();
509 let vgm = SphericalVariogram {
510 nugget: 0.01,
511 sill: 20.0,
512 range: 10.0,
513 };
514 let ok = OrdinaryKriging::fit(pts, vals, Box::new(vgm)).expect("fit failed");
515 let test_pts = vec![vec![0.5_f64], vec![1.5], vec![2.5]];
516 for p in &test_pts {
517 let (_est, var) = ok.predict(p).expect("predict");
518 assert!(var >= 0.0, "variance negative at {:?}: {}", p, var);
519 }
520 }
521
522 #[test]
523 fn test_variogram_gamma_at_zero() {
524 let svgm = SphericalVariogram {
525 nugget: 0.1,
526 sill: 1.0,
527 range: 2.0,
528 };
529 assert_eq!(svgm.gamma(0.0), 0.0);
530 let evgm = ExponentialVariogram {
531 nugget: 0.1,
532 sill: 1.0,
533 range: 2.0,
534 };
535 assert_eq!(evgm.gamma(0.0), 0.0);
536 let gvgm = GaussianVariogram {
537 nugget: 0.1,
538 sill: 1.0,
539 range: 2.0,
540 };
541 assert_eq!(gvgm.gamma(0.0), 0.0);
542 let pvgm = PowerVariogram {
543 nugget: 0.1,
544 slope: 1.0,
545 power: 1.5,
546 };
547 assert_eq!(pvgm.gamma(0.0), 0.0);
548 }
549
550 #[test]
551 fn test_spherical_reaches_sill() {
552 let vgm = SphericalVariogram {
553 nugget: 0.0,
554 sill: 5.0,
555 range: 2.0,
556 };
557 let v = vgm.gamma(100.0);
558 assert!((v - 5.0).abs() < 1e-10, "should reach sill: {}", v);
559 }
560
561 #[test]
562 fn test_error_on_empty() {
563 let vgm = SphericalVariogram {
564 nugget: 0.0,
565 sill: 1.0,
566 range: 1.0,
567 };
568 let r = OrdinaryKriging::fit(vec![], vec![], Box::new(vgm));
569 assert!(r.is_err());
570 }
571
572 #[test]
573 fn test_error_on_dim_mismatch_predict() {
574 let pts = vec![vec![0.0_f64, 0.0], vec![1.0, 1.0]];
575 let vals = vec![0.0_f64, 1.0];
576 let vgm = GaussianVariogram {
577 nugget: 0.0,
578 sill: 1.0,
579 range: 5.0,
580 };
581 let ok = OrdinaryKriging::fit(pts, vals, Box::new(vgm)).expect("fit");
582 let r = ok.predict(&[0.5]); assert!(r.is_err());
584 }
585}