use crate::variogram::VariogramModelFamily;
#[derive(Clone, Debug)]
pub struct CrossVariogramBin {
pub lags: Vec<f64>,
pub semivariances: Vec<f64>,
pub counts: Vec<usize>,
pub bin_size: f64,
pub max_distance: f64,
}
impl CrossVariogramBin {
pub fn n_lags(&self) -> usize {
self.lags.len()
}
pub fn max_semivariance(&self) -> f64 {
self.semivariances
.iter()
.copied()
.fold(f64::NEG_INFINITY, f64::max)
}
pub fn mean_pairs_per_lag(&self) -> f64 {
if self.counts.is_empty() {
0.0
} else {
self.counts.iter().map(|&c| c as f64).sum::<f64>() / self.counts.len() as f64
}
}
}
#[derive(Clone, Debug)]
pub struct CrossVariogramModel {
pub nugget: f64,
pub sill: f64,
pub range: f64,
pub family: VariogramModelFamily,
pub wrss: f64,
pub condition_number: f64,
pub primary_var: String,
pub auxiliary_var: String,
}
impl CrossVariogramModel {
pub fn evaluate(&self, distance: f64) -> f64 {
if distance <= 0.0 {
return self.nugget;
}
let h_normalized = distance / self.range;
let model_value = match self.family {
VariogramModelFamily::Spherical => {
if h_normalized >= 1.0 {
1.0
} else {
1.5 * h_normalized - 0.5 * h_normalized.powi(3)
}
}
VariogramModelFamily::Exponential => {
1.0 - (-3.0 * h_normalized).exp()
}
VariogramModelFamily::Gaussian => {
1.0 - (-3.0 * h_normalized * h_normalized).exp()
}
};
self.nugget + (self.sill - self.nugget) * model_value
}
}
pub fn compute_cross_variogram(
primary: &[(f64, f64, f64)],
auxiliary: &[(f64, f64, f64)],
max_distance: f64,
bin_size: f64,
) -> Result<CrossVariogramBin, String> {
if primary.is_empty() {
return Err("Cannot compute cross-variogram: no sample locations".to_string());
}
if primary.len() != auxiliary.len() {
return Err(
"Primary and auxiliary arrays must have equal length for co-located sampling"
.to_string(),
);
}
if primary.len() < 2 {
return Err("Cannot compute cross-variogram: need at least 2 samples".to_string());
}
if bin_size <= 0.0 {
return Err("bin_size must be positive".to_string());
}
if max_distance <= 0.0 {
return Err("max_distance must be positive".to_string());
}
let n_lags = ((max_distance / bin_size).ceil() as usize).max(1);
let mut lag_sums = vec![0.0; n_lags];
let mut lag_counts = vec![0usize; n_lags];
for i in 0..primary.len() {
for j in (i + 1)..primary.len() {
let (x1, y1, z1) = primary[i];
let (x2, y2, z2) = primary[j];
let (_, _, w1) = auxiliary[i];
let (_, _, w2) = auxiliary[j];
let dx = x2 - x1;
let dy = y2 - y1;
let distance = (dx * dx + dy * dy).sqrt();
if distance <= max_distance {
let bin_idx = ((distance / bin_size).floor() as usize).min(n_lags - 1);
let product = (z1 - z2) * (w1 - w2);
lag_sums[bin_idx] += product;
lag_counts[bin_idx] += 1;
}
}
}
let mut lags = Vec::new();
let mut semivariances = Vec::new();
for bin_idx in 0..n_lags {
if lag_counts[bin_idx] > 0 {
let lag_center = ((bin_idx as f64) + 0.5) * bin_size;
let mean_product = lag_sums[bin_idx] / lag_counts[bin_idx] as f64;
let semivariance = mean_product / 2.0;
lags.push(lag_center);
semivariances.push(semivariance);
}
}
if lags.is_empty() {
return Err("No point pairs found within max_distance".to_string());
}
Ok(CrossVariogramBin {
lags,
semivariances,
counts: lag_counts.into_iter().filter(|&c| c > 0).collect(),
bin_size,
max_distance,
})
}
pub fn fit_cross_variogram_model(
cross_vgram: &CrossVariogramBin,
family: VariogramModelFamily,
primary_var: &str,
auxiliary_var: &str,
) -> Result<CrossVariogramModel, String> {
if cross_vgram.n_lags() < 2 {
return Err("Need at least 2 lags for model fitting".to_string());
}
let max_semi = cross_vgram
.semivariances
.iter()
.copied()
.fold(f64::NEG_INFINITY, f64::max);
let min_semi = cross_vgram
.semivariances
.iter()
.copied()
.fold(f64::INFINITY, f64::min);
let nugget = min_semi.max(0.0);
let sill = max_semi;
let target = nugget + 0.95 * (sill - nugget);
let mut range = cross_vgram.max_distance;
for (lag, semi) in cross_vgram.lags.iter().zip(&cross_vgram.semivariances) {
if semi >= &target {
range = lag * 1.2; break;
}
}
Ok(CrossVariogramModel {
nugget,
sill,
range: range.max(cross_vgram.bin_size), family,
wrss: 0.0, condition_number: 1.0, primary_var: primary_var.to_string(),
auxiliary_var: auxiliary_var.to_string(),
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_cross_variogram_basic() {
let primary = vec![
(0.0, 0.0, 10.0),
(1.0, 0.0, 12.0),
(0.0, 1.0, 11.0),
(1.0, 1.0, 13.0),
];
let auxiliary = vec![
(0.0, 0.0, 20.0),
(1.0, 0.0, 22.0),
(0.0, 1.0, 21.0),
(1.0, 1.0, 23.0),
];
let result = compute_cross_variogram(&primary, &auxiliary, 2.0, 0.5);
assert!(result.is_ok());
let cvgram = result.unwrap();
assert!(!cvgram.lags.is_empty());
assert_eq!(cvgram.n_lags(), cvgram.semivariances.len());
assert!(cvgram.max_semivariance() > 0.0);
}
#[test]
fn test_cross_variogram_mismatched_lengths() {
let primary = vec![(0.0, 0.0, 10.0), (1.0, 0.0, 12.0)];
let auxiliary = vec![(0.0, 0.0, 20.0)];
let result = compute_cross_variogram(&primary, &auxiliary, 2.0, 0.5);
assert!(result.is_err());
}
#[test]
fn test_cross_variogram_empty() {
let primary: Vec<(f64, f64, f64)> = vec![];
let auxiliary: Vec<(f64, f64, f64)> = vec![];
let result = compute_cross_variogram(&primary, &auxiliary, 2.0, 0.5);
assert!(result.is_err());
}
#[test]
fn test_cross_variogram_model_evaluation() {
let model = CrossVariogramModel {
nugget: 1.0,
sill: 5.0,
range: 10.0,
family: VariogramModelFamily::Exponential,
wrss: 0.0,
condition_number: 1.0,
primary_var: "Z".to_string(),
auxiliary_var: "Y".to_string(),
};
assert_eq!(model.evaluate(0.0), 1.0);
let far_eval = model.evaluate(100.0);
assert!(far_eval > 4.9);
let d1 = model.evaluate(1.0);
let d5 = model.evaluate(5.0);
let d10 = model.evaluate(10.0);
assert!(d1 < d5);
assert!(d5 < d10);
}
#[test]
fn test_fit_cross_variogram_model() {
let cvgram = CrossVariogramBin {
lags: vec![1.0, 2.0, 3.0, 4.0],
semivariances: vec![0.5, 1.0, 1.5, 1.8],
counts: vec![10, 8, 6, 4],
bin_size: 1.0,
max_distance: 4.0,
};
let result = fit_cross_variogram_model(
&cvgram,
VariogramModelFamily::Exponential,
"primary",
"auxiliary",
);
assert!(result.is_ok());
let model = result.unwrap();
assert!(model.nugget >= 0.0);
assert!(model.sill > model.nugget);
assert!(model.range > 0.0);
}
}