use crate::feature_selection::mrmr::mrmr_error::MrmrError;
use crate::mrmr::mrmr_result::MrmrResult;
use crate::mrmr::mrmr_utils;
use deep_causality_num::{Float, FloatOption};
use deep_causality_tensor::CausalTensor;
use std::collections::HashSet;
#[cfg(feature = "parallel")]
use rayon::prelude::*;
pub fn mrmr_features_selector<T, F>(
tensor: &CausalTensor<T>,
num_features: usize,
target_col: usize,
) -> Result<MrmrResult, MrmrError>
where
T: FloatOption<F>,
F: Float,
{
let shape = tensor.shape();
if shape.len() != 2 {
return Err(MrmrError::InvalidInput(
"Input tensor must be 2-dimensional".to_string(),
));
}
let n_rows = shape[0];
let n_cols = shape[1];
if n_rows < 3 {
return Err(MrmrError::SampleTooSmall(3));
}
if num_features == 0 || num_features >= n_cols {
return Err(MrmrError::InvalidInput(
"Invalid number of features requested".to_string(),
));
}
if target_col >= n_cols {
return Err(MrmrError::InvalidInput(
"Target column index out of bounds".to_string(),
));
}
let mut all_features: HashSet<usize> = (0..n_cols).collect();
all_features.remove(&target_col);
let mut selected_features_with_scores: Vec<(usize, f64)> = Vec::with_capacity(num_features);
#[cfg(feature = "parallel")]
let (first_feature, max_relevance) = {
let features: Vec<usize> = all_features.iter().copied().collect();
features
.into_par_iter()
.map(|feature_idx| {
let relevance = mrmr_utils::f_statistic(tensor, feature_idx, target_col)?;
if !relevance.is_finite() {
Err(MrmrError::FeatureScoreError(format!(
"Relevance score for feature {} is not finite: {}",
feature_idx, relevance
)))
} else {
Ok((feature_idx, relevance))
}
})
.reduce(
|| Ok((0, -1.0)),
|acc, res| {
let acc = acc?;
let res = res?;
if res.1 > acc.1 { Ok(res) } else { Ok(acc) }
},
)?
};
#[cfg(not(feature = "parallel"))]
let (first_feature, max_relevance) = {
let mut first_feature = 0;
let mut max_relevance = -1.0;
for &feature_idx in &all_features {
let relevance = mrmr_utils::f_statistic(tensor, feature_idx, target_col)?;
if !relevance.is_finite() {
return Err(MrmrError::FeatureScoreError(format!(
"Relevance score for feature {} is not finite: {}",
feature_idx, relevance
)));
}
if relevance > max_relevance {
max_relevance = relevance;
first_feature = feature_idx;
}
}
(first_feature, max_relevance)
};
if !max_relevance.is_finite() {
return Err(MrmrError::FeatureScoreError(format!(
"Initial max relevance score for feature {} is not finite: {}",
first_feature, max_relevance
)));
}
selected_features_with_scores.push((first_feature, max_relevance));
all_features.remove(&first_feature);
while selected_features_with_scores.len() < num_features {
#[cfg(feature = "parallel")]
let (best_feature, best_feature_score) = {
let features: Vec<usize> = all_features.iter().copied().collect();
features
.into_par_iter()
.map(|feature_idx| {
let relevance = mrmr_utils::f_statistic(tensor, feature_idx, target_col)?;
if !relevance.is_finite() {
return Err(MrmrError::FeatureScoreError(format!(
"Relevance score for feature {} is not finite: {}",
feature_idx, relevance
)));
}
let selected_indices: Vec<usize> = selected_features_with_scores
.iter()
.map(|(idx, _)| *idx)
.collect();
let redundancy: f64 = selected_indices
.par_iter()
.map(|&selected_idx| {
let (correlation, _) = mrmr_utils::pearson_correlation(tensor, feature_idx, selected_idx)?;
if !correlation.is_finite() {
return Err(MrmrError::FeatureScoreError(format!(
"Correlation for feature {} and selected feature {} is not finite: {}",
feature_idx, selected_idx, correlation
)))
}
Ok(correlation.abs())
})
.sum::<Result<f64, _>>()?;
let redundancy = redundancy / selected_indices.len() as f64;
if !redundancy.is_finite() {
return Err(MrmrError::FeatureScoreError(format!(
"Redundancy for feature {} is not finite: {}",
feature_idx, redundancy
)));
}
let mrmr_score = if redundancy == 0.0 {
if relevance == 0.0 {
return Err(MrmrError::FeatureScoreError(format!(
"mRMR score for feature {} is NaN (relevance {} / redundancy {}).",
feature_idx, relevance, redundancy
)));
} else {
return Err(MrmrError::FeatureScoreError(format!(
"mRMR score for feature {} is infinite (relevance {} / redundancy {}).",
feature_idx, relevance, redundancy
)));
}
} else {
relevance / redundancy
};
if !mrmr_score.is_finite() {
return Err(MrmrError::FeatureScoreError(format!(
"mRMR score for feature {} is not finite: {}",
feature_idx, mrmr_score
)));
}
Ok((feature_idx, mrmr_score))
})
.reduce(
|| Ok((0, -1.0)),
|acc, res| {
let acc = acc?;
let res = res?;
if res.1 > acc.1 {
Ok(res)
} else {
Ok(acc)
}
},
)?
};
#[cfg(not(feature = "parallel"))]
let (best_feature, best_feature_score) = {
let mut best_feature = 0;
let mut max_mrmr_score = -1.0;
for &feature_idx in &all_features {
let relevance = mrmr_utils::f_statistic(tensor, feature_idx, target_col)?;
if !relevance.is_finite() {
return Err(MrmrError::FeatureScoreError(format!(
"Relevance score for feature {} is not finite: {}",
feature_idx, relevance
)));
}
let mut redundancy = 0.0;
let selected_indices: Vec<usize> = selected_features_with_scores
.iter()
.map(|(idx, _)| *idx)
.collect();
for &selected_idx in &selected_indices {
let (correlation, _) =
mrmr_utils::pearson_correlation(tensor, feature_idx, selected_idx)?;
if !correlation.is_finite() {
return Err(MrmrError::FeatureScoreError(format!(
"Correlation for feature {} and selected feature {} is not finite: {}",
feature_idx, selected_idx, correlation
)));
}
redundancy += correlation.abs();
}
redundancy /= selected_indices.len() as f64;
if !redundancy.is_finite() {
return Err(MrmrError::FeatureScoreError(format!(
"Redundancy for feature {} is not finite: {}",
feature_idx, redundancy
)));
}
let mrmr_score = if redundancy == 0.0 {
if relevance == 0.0 {
return Err(MrmrError::FeatureScoreError(format!(
"mRMR score for feature {} is NaN (relevance {} / redundancy {}).",
feature_idx, relevance, redundancy
)));
} else {
return Err(MrmrError::FeatureScoreError(format!(
"mRMR score for feature {} is infinite (relevance {} / redundancy {}).",
feature_idx, relevance, redundancy
)));
}
} else {
relevance / redundancy
};
if mrmr_score > max_mrmr_score {
max_mrmr_score = mrmr_score;
best_feature = feature_idx;
}
}
(best_feature, max_mrmr_score)
};
selected_features_with_scores.push((best_feature, best_feature_score));
all_features.remove(&best_feature);
}
let max_score = selected_features_with_scores
.iter()
.map(|(_, score)| *score)
.fold(f64::MIN, |acc, score| acc.max(score));
if max_score > 0.0 {
for (_, score) in &mut selected_features_with_scores {
*score /= max_score;
}
}
Ok(MrmrResult::new(selected_features_with_scores))
}