use scirs2_core::ndarray::{Array2, ArrayView2};
use sklears_core::{
error::{Result as SklResult, SklearsError},
types::Float,
};
use std::collections::HashMap;
#[derive(Debug, Clone)]
pub struct CombinationInfo {
pub combination: Vec<i32>,
pub frequency: usize,
pub relative_frequency: f64,
pub cardinality: usize,
}
#[derive(Debug, Clone)]
pub struct LabelAnalysisResults {
pub combinations: Vec<CombinationInfo>,
pub total_samples: usize,
pub unique_combinations: usize,
pub most_frequent: Option<CombinationInfo>,
pub least_frequent: Option<CombinationInfo>,
pub average_cardinality: f64,
pub cardinality_distribution: HashMap<usize, usize>,
}
pub fn analyze_combinations(y: &ArrayView2<'_, i32>) -> SklResult<LabelAnalysisResults> {
let (n_samples, n_labels) = y.dim();
if n_samples == 0 || n_labels == 0 {
return Err(SklearsError::InvalidInput(
"Input array must have at least one sample and one label".to_string(),
));
}
for sample_idx in 0..n_samples {
for label_idx in 0..n_labels {
let value = y[[sample_idx, label_idx]];
if value != 0 && value != 1 {
return Err(SklearsError::InvalidInput(format!(
"All label values must be 0 or 1, found: {}",
value
)));
}
}
}
let mut combination_counts: HashMap<Vec<i32>, usize> = HashMap::new();
let mut cardinality_distribution: HashMap<usize, usize> = HashMap::new();
let mut total_cardinality = 0;
for sample_idx in 0..n_samples {
let mut combination = Vec::new();
let mut cardinality = 0;
for label_idx in 0..n_labels {
let label_value = y[[sample_idx, label_idx]];
combination.push(label_value);
if label_value == 1 {
cardinality += 1;
}
}
*combination_counts.entry(combination).or_insert(0) += 1;
*cardinality_distribution.entry(cardinality).or_insert(0) += 1;
total_cardinality += cardinality;
}
let mut combinations: Vec<CombinationInfo> = combination_counts
.into_iter()
.map(|(combination, frequency)| {
let cardinality = combination.iter().sum::<i32>() as usize;
CombinationInfo {
combination,
frequency,
relative_frequency: frequency as f64 / n_samples as f64,
cardinality,
}
})
.collect();
combinations.sort_by(|a, b| b.frequency.cmp(&a.frequency));
let most_frequent = combinations.first().cloned();
let least_frequent = combinations.last().cloned();
let unique_combinations = combinations.len();
let average_cardinality = total_cardinality as f64 / n_samples as f64;
Ok(LabelAnalysisResults {
combinations,
total_samples: n_samples,
unique_combinations,
most_frequent,
least_frequent,
average_cardinality,
cardinality_distribution,
})
}
pub fn label_cooccurrence_matrix(y: &ArrayView2<'_, i32>) -> SklResult<Array2<usize>> {
let (n_samples, n_labels) = y.dim();
if n_samples == 0 || n_labels == 0 {
return Err(SklearsError::InvalidInput(
"Input array must have at least one sample and one label".to_string(),
));
}
let mut cooccurrence = Array2::<usize>::zeros((n_labels, n_labels));
for sample_idx in 0..n_samples {
for i in 0..n_labels {
for j in 0..n_labels {
if y[[sample_idx, i]] == 1 && y[[sample_idx, j]] == 1 {
cooccurrence[[i, j]] += 1;
}
}
}
}
Ok(cooccurrence)
}
pub fn label_correlation_matrix(y: &ArrayView2<'_, i32>) -> SklResult<Array2<f64>> {
let (n_samples, n_labels) = y.dim();
if n_samples == 0 || n_labels == 0 {
return Err(SklearsError::InvalidInput(
"Input array must have at least one sample and one label".to_string(),
));
}
let mut correlation = Array2::<Float>::zeros((n_labels, n_labels));
let mut means = vec![0.0; n_labels];
for j in 0..n_labels {
let mut sum = 0.0;
for i in 0..n_samples {
sum += y[[i, j]] as f64;
}
means[j] = sum / n_samples as f64;
}
for i in 0..n_labels {
for j in 0..n_labels {
if i == j {
correlation[[i, j]] = 1.0;
} else {
let mut numerator = 0.0;
let mut sum_sq_i = 0.0;
let mut sum_sq_j = 0.0;
for sample_idx in 0..n_samples {
let val_i = y[[sample_idx, i]] as f64 - means[i];
let val_j = y[[sample_idx, j]] as f64 - means[j];
numerator += val_i * val_j;
sum_sq_i += val_i * val_i;
sum_sq_j += val_j * val_j;
}
let denominator = (sum_sq_i * sum_sq_j).sqrt();
if denominator > 1e-10 {
correlation[[i, j]] = numerator / denominator;
} else {
correlation[[i, j]] = 0.0;
}
}
}
}
Ok(correlation)
}
pub fn get_rare_combinations(
results: &LabelAnalysisResults,
threshold: usize,
) -> Vec<CombinationInfo> {
results
.combinations
.iter()
.filter(|combo| combo.frequency <= threshold)
.cloned()
.collect()
}
pub fn get_combinations_by_cardinality(
results: &LabelAnalysisResults,
cardinality: usize,
) -> Vec<CombinationInfo> {
results
.combinations
.iter()
.filter(|combo| combo.cardinality == cardinality)
.cloned()
.collect()
}
pub fn find_singleton_labels(y: &ArrayView2<'_, i32>) -> SklResult<Vec<(usize, usize, f64)>> {
let (n_samples, n_labels) = y.dim();
if n_samples == 0 || n_labels == 0 {
return Err(SklearsError::InvalidInput(
"Input array must have at least one sample and one label".to_string(),
));
}
let mut singleton_counts = vec![0; n_labels];
for sample_idx in 0..n_samples {
let active_labels: Vec<usize> = (0..n_labels)
.filter(|&label_idx| y[[sample_idx, label_idx]] == 1)
.collect();
if active_labels.len() == 1 {
singleton_counts[active_labels[0]] += 1;
}
}
let results = singleton_counts
.into_iter()
.enumerate()
.map(|(label_idx, count)| {
let percentage = count as f64 / n_samples as f64 * 100.0;
(label_idx, count, percentage)
})
.collect();
Ok(results)
}
pub fn label_frequency_distribution(
y: &ArrayView2<'_, i32>,
) -> SklResult<Vec<(usize, usize, f64)>> {
let (n_samples, n_labels) = y.dim();
if n_samples == 0 || n_labels == 0 {
return Err(SklearsError::InvalidInput(
"Input array must have at least one sample and one label".to_string(),
));
}
let mut frequencies = vec![0; n_labels];
for sample_idx in 0..n_samples {
for label_idx in 0..n_labels {
if y[[sample_idx, label_idx]] == 1 {
frequencies[label_idx] += 1;
}
}
}
let results = frequencies
.into_iter()
.enumerate()
.map(|(label_idx, frequency)| {
let percentage = frequency as f64 / n_samples as f64 * 100.0;
(label_idx, frequency, percentage)
})
.collect();
Ok(results)
}