#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum SpatialWeightsMode {
Queen,
Rook,
KNearest,
DistanceBand,
}
impl SpatialWeightsMode {
pub fn as_str(&self) -> &str {
match self {
Self::Queen => "queen",
Self::Rook => "rook",
Self::KNearest => "k_nearest",
Self::DistanceBand => "distance_band",
}
}
pub fn from_str(s: &str) -> Option<Self> {
match s.to_ascii_lowercase().as_str() {
"queen" => Some(Self::Queen),
"rook" => Some(Self::Rook),
"k_nearest" => Some(Self::KNearest),
"distance_band" => Some(Self::DistanceBand),
_ => None,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum IslandPolicy {
DropWithWarning,
KeepZeroWeight,
Error,
}
impl IslandPolicy {
pub fn as_str(&self) -> &str {
match self {
Self::DropWithWarning => "drop_with_warning",
Self::KeepZeroWeight => "keep_zero_weight",
Self::Error => "error",
}
}
pub fn from_str(s: &str) -> Option<Self> {
match s.to_ascii_lowercase().as_str() {
"drop_with_warning" => Some(Self::DropWithWarning),
"keep_zero_weight" => Some(Self::KeepZeroWeight),
"error" => Some(Self::Error),
_ => None,
}
}
}
#[derive(Clone, Debug)]
pub struct SpatialWeightsDiagnostics {
pub n_features: usize,
pub n_islands: usize,
pub neighbor_count_min: usize,
pub neighbor_count_mean: f64,
pub neighbor_count_max: usize,
pub connected_component_count: usize,
pub row_standardized: bool,
pub dropped_feature_count: usize,
}
#[derive(Clone, Debug)]
pub struct SpatialWeightsGraph {
pub neighbors: Vec<Vec<(usize, f64)>>,
pub diagnostics: SpatialWeightsDiagnostics,
pub warnings: Vec<String>,
}
impl SpatialWeightsGraph {
pub fn n_features(&self) -> usize {
self.neighbors.len()
}
pub fn n_islands(&self) -> usize {
self.diagnostics.n_islands
}
pub fn is_row_standardized(&self) -> bool {
self.diagnostics.row_standardized
}
pub fn warnings(&self) -> &[String] {
&self.warnings
}
}
pub fn normal_cdf(x: f64) -> f64 {
let z = x.abs();
let t = 1.0 / (1.0 + 0.231_641_9 * z);
let poly = t
* (0.319_381_530
+ t * (-0.356_563_782
+ t * (1.781_477_937 + t * (-1.821_255_978 + t * 1.330_274_429))));
let pdf = (-0.5 * z * z).exp() / (2.0 * std::f64::consts::PI).sqrt();
let cdf = 1.0 - pdf * poly;
if x >= 0.0 { cdf } else { 1.0 - cdf }
}
pub fn two_tailed_normal_p(z: f64) -> f64 {
(2.0 * (1.0 - normal_cdf(z.abs()))).clamp(0.0, 1.0)
}
pub fn connected_components(neighbors: &[Vec<(usize, f64)>]) -> usize {
let n = neighbors.len();
let mut undirected = vec![Vec::<usize>::new(); n];
for (i, row) in neighbors.iter().enumerate() {
for (j, _) in row {
undirected[i].push(*j);
undirected[*j].push(i);
}
}
let mut visited = vec![false; n];
let mut components = 0usize;
for start in 0..n {
if visited[start] {
continue;
}
components += 1;
let mut stack = vec![start];
visited[start] = true;
while let Some(node) = stack.pop() {
for next in &undirected[node] {
if !visited[*next] {
visited[*next] = true;
stack.push(*next);
}
}
}
}
components
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_spatial_weights_mode_parsing() {
assert_eq!(SpatialWeightsMode::from_str("queen"), Some(SpatialWeightsMode::Queen));
assert_eq!(SpatialWeightsMode::from_str("ROOK"), Some(SpatialWeightsMode::Rook));
assert_eq!(SpatialWeightsMode::from_str("k_nearest"), Some(SpatialWeightsMode::KNearest));
assert_eq!(SpatialWeightsMode::from_str("distance_band"), Some(SpatialWeightsMode::DistanceBand));
assert_eq!(SpatialWeightsMode::from_str("invalid"), None);
}
#[test]
fn test_island_policy_parsing() {
assert_eq!(IslandPolicy::from_str("drop_with_warning"), Some(IslandPolicy::DropWithWarning));
assert_eq!(IslandPolicy::from_str("KEEP_ZERO_WEIGHT"), Some(IslandPolicy::KeepZeroWeight));
assert_eq!(IslandPolicy::from_str("error"), Some(IslandPolicy::Error));
assert_eq!(IslandPolicy::from_str("invalid"), None);
}
#[test]
fn test_normal_cdf() {
assert!((normal_cdf(0.0) - 0.5).abs() < 0.001);
assert!(normal_cdf(2.0) > 0.95);
assert!(normal_cdf(-2.0) < 0.05);
}
#[test]
fn test_two_tailed_p() {
let p = two_tailed_normal_p(2.0);
assert!(p > 0.0 && p < 0.05);
}
#[test]
fn test_connected_components_simple() {
let neighbors = vec![
vec![(1, 1.0)],
vec![(0, 1.0), (2, 1.0)],
vec![(1, 1.0)],
];
assert_eq!(connected_components(&neighbors), 1);
}
#[test]
fn test_connected_components_multiple() {
let neighbors = vec![
vec![(1, 1.0)],
vec![(0, 1.0)],
vec![],
];
assert_eq!(connected_components(&neighbors), 2);
}
}