use crate::{GeostatError, GeostatResult};
use crate::variogram::VariogramModel;
use super::KrigingResult;
use nalgebra as na;
use rayon::prelude::*;
#[derive(Clone, Debug)]
pub struct SimpleKriging {
training_coords: Vec<(f64, f64)>,
training_values: Vec<f64>,
variogram: VariogramModel,
known_mean: f64,
}
impl SimpleKriging {
pub fn new(
training_coords: Vec<(f64, f64)>,
training_values: Vec<f64>,
variogram: VariogramModel,
known_mean: f64,
) -> GeostatResult<Self> {
if training_coords.len() != training_values.len() {
return Err(GeostatError::InvalidParameters(
"Training coordinates and values must have the same length".to_string(),
));
}
if training_coords.len() < 3 {
return Err(GeostatError::InsufficientData(
"At least 3 training points required".to_string(),
));
}
Ok(SimpleKriging {
training_coords,
training_values,
variogram,
known_mean,
})
}
pub fn predict(&self, target_x: f64, target_y: f64) -> GeostatResult<KrigingResult> {
let target = (target_x, target_y);
let n = self.training_coords.len();
let mut gamma = na::DMatrix::zeros(n + 1, n + 1);
for i in 0..n {
for j in 0..n {
let dist = Self::distance(self.training_coords[i], self.training_coords[j]);
gamma[(i, j)] = self.variogram.evaluate(dist);
}
gamma[(i, n)] = 1.0;
gamma[(n, i)] = 1.0;
}
gamma[(n, n)] = 0.0;
let mut rhs = na::DVector::zeros(n + 1);
for i in 0..n {
let dist = Self::distance(self.training_coords[i], target);
rhs[i] = self.variogram.evaluate(dist);
}
rhs[n] = 1.0;
let weights = match gamma.clone().lu().solve(&rhs) {
Some(w) => w,
None => {
let svd = gamma.svd(true, true);
svd.solve(&rhs, 1e-10).map_err(|_| {
GeostatError::KrigingSolveFailed(
"Failed to solve kriging system".to_string(),
)
})?
}
};
let kriging_weights: Vec<f64> = weights.iter().take(n).copied().collect();
let sill = self.variogram.total_sill();
let mut kriging_variance = sill;
for i in 0..n {
let dist = Self::distance(self.training_coords[i], target);
kriging_variance -= kriging_weights[i] * self.variogram.evaluate(dist);
}
kriging_variance = kriging_variance.max(0.0);
let mut prediction = self.known_mean;
for i in 0..n {
prediction += kriging_weights[i] * (self.training_values[i] - self.known_mean);
}
let std_error = kriging_variance.sqrt();
let z_critical = 1.96; let ci_lower = prediction - z_critical * std_error;
let ci_upper = prediction + z_critical * std_error;
Ok(KrigingResult {
prediction,
variance: kriging_variance,
std_error,
ci_lower,
ci_upper,
})
}
fn distance(p1: (f64, f64), p2: (f64, f64)) -> f64 {
((p1.0 - p2.0).powi(2) + (p1.1 - p2.1).powi(2)).sqrt()
}
pub fn predict_batch(&self, targets: &[(f64, f64)]) -> GeostatResult<Vec<KrigingResult>> {
targets
.par_iter()
.map(|&(x, y)| self.predict(x, y))
.collect()
}
pub fn known_mean(&self) -> f64 {
self.known_mean
}
pub fn n_training(&self) -> usize {
self.training_coords.len()
}
pub fn variogram(&self) -> &VariogramModel {
&self.variogram
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::variogram::{VariogramModel, VariogramModelFamily};
#[test]
fn test_simple_kriging_creation() {
let coords = vec![(0.0, 0.0), (1.0, 0.0), (0.0, 1.0)];
let values = vec![10.0, 12.0, 11.0];
let variogram = VariogramModel {
family: VariogramModelFamily::Spherical,
nugget: 0.5,
partial_sill: 1.0,
range: 50.0,
wrss: 0.0,
condition_number: 1.0,
};
let sk = SimpleKriging::new(coords, values, variogram, 11.0);
assert!(sk.is_ok());
let sk = sk.unwrap();
assert_eq!(sk.known_mean(), 11.0);
assert_eq!(sk.n_training(), 3);
}
#[test]
fn test_simple_kriging_insufficient_points() {
let coords = vec![(0.0, 0.0), (1.0, 0.0)];
let values = vec![10.0, 12.0];
let variogram = VariogramModel {
family: VariogramModelFamily::Spherical,
nugget: 0.5,
partial_sill: 1.0,
range: 50.0,
wrss: 0.0,
condition_number: 1.0,
};
let sk = SimpleKriging::new(coords, values, variogram, 11.0);
assert!(sk.is_err());
}
#[test]
fn test_simple_kriging_mismatch() {
let coords = vec![(0.0, 0.0), (1.0, 0.0), (0.0, 1.0)];
let values = vec![10.0, 12.0];
let variogram = VariogramModel {
family: VariogramModelFamily::Spherical,
nugget: 0.5,
partial_sill: 1.0,
range: 50.0,
wrss: 0.0,
condition_number: 1.0,
};
let sk = SimpleKriging::new(coords, values, variogram, 11.0);
assert!(sk.is_err());
}
#[test]
fn test_simple_kriging_prediction() {
let coords = vec![(0.0, 0.0), (1.0, 0.0), (0.0, 1.0)];
let values = vec![10.0, 12.0, 11.0];
let variogram = VariogramModel {
family: VariogramModelFamily::Spherical,
nugget: 0.5,
partial_sill: 1.0,
range: 50.0,
wrss: 0.0,
condition_number: 1.0,
};
let sk = SimpleKriging::new(coords, values, variogram, 11.0).unwrap();
let result = sk.predict(0.5, 0.5).unwrap();
assert!(result.prediction.is_finite());
assert!(result.variance.is_finite());
assert!(result.variance >= 0.0);
assert!(result.ci_upper >= result.prediction);
assert!(result.ci_lower <= result.prediction);
}
#[test]
fn test_simple_kriging_batch_prediction() {
let coords = vec![(0.0, 0.0), (1.0, 0.0), (0.0, 1.0)];
let values = vec![10.0, 12.0, 11.0];
let variogram = VariogramModel {
family: VariogramModelFamily::Spherical,
nugget: 0.5,
partial_sill: 1.0,
range: 50.0,
wrss: 0.0,
condition_number: 1.0,
};
let sk = SimpleKriging::new(coords, values, variogram, 11.0).unwrap();
let targets = vec![(0.5, 0.5), (0.3, 0.7), (0.8, 0.2)];
let results = sk.predict_batch(&targets).unwrap();
assert_eq!(results.len(), 3);
for result in &results {
assert!(result.prediction.is_finite());
assert!(result.variance.is_finite());
assert!(result.variance >= 0.0);
}
}
#[test]
fn test_simple_kriging_at_data_point() {
let coords = vec![(0.0, 0.0), (1.0, 0.0), (0.0, 1.0)];
let values = vec![10.0, 12.0, 11.0];
let variogram = VariogramModel {
family: VariogramModelFamily::Spherical,
nugget: 0.5,
partial_sill: 1.0,
range: 50.0,
wrss: 0.0,
condition_number: 1.0,
};
let sk = SimpleKriging::new(coords.clone(), values.clone(), variogram, 11.0).unwrap();
let result = sk.predict(coords[0].0, coords[0].1).unwrap();
assert!(result.prediction.is_finite());
assert!(result.variance.is_finite());
assert!(result.variance >= 0.0);
}
#[test]
fn test_simple_kriging_different_means() {
let coords = vec![(0.0, 0.0), (1.0, 0.0), (0.0, 1.0)];
let values = vec![100.0, 120.0, 110.0];
let variogram = VariogramModel {
family: VariogramModelFamily::Spherical,
nugget: 0.5,
partial_sill: 1.0,
range: 50.0,
wrss: 0.0,
condition_number: 1.0,
};
let sk1 = SimpleKriging::new(coords.clone(), values.clone(), variogram.clone(), 110.0).unwrap();
let sk2 = SimpleKriging::new(coords, values, variogram, 100.0).unwrap();
let result1 = sk1.predict(0.5, 0.5).unwrap();
let result2 = sk2.predict(0.5, 0.5).unwrap();
assert!(result1.prediction.is_finite());
assert!(result2.prediction.is_finite());
assert!(result1.variance.is_finite());
assert!(result2.variance.is_finite());
}
}