use scirs2_core::ndarray_ext::{Array1, Array2, ArrayView1, ArrayView2};
use sklears_core::{
error::{Result as SklResult, SklearsError},
traits::{Estimator, Fit, Predict, Untrained},
types::Float,
};
use std::collections::{HashMap, HashSet};
#[derive(Debug, Clone)]
pub struct CoTraining<S = Untrained> {
state: S,
view1_features: Vec<usize>,
view2_features: Vec<usize>,
p: usize,
n: usize,
max_iter: usize,
verbose: bool,
confidence_threshold: f64,
}
impl CoTraining<Untrained> {
pub fn new() -> Self {
Self {
state: Untrained,
view1_features: Vec::new(),
view2_features: Vec::new(),
p: 1,
n: 1,
max_iter: 30,
verbose: false,
confidence_threshold: 0.5,
}
}
pub fn view1_features(mut self, features: Vec<usize>) -> Self {
self.view1_features = features;
self
}
pub fn view2_features(mut self, features: Vec<usize>) -> Self {
self.view2_features = features;
self
}
pub fn p(mut self, p: usize) -> Self {
self.p = p;
self
}
pub fn n(mut self, n: usize) -> Self {
self.n = n;
self
}
pub fn max_iter(mut self, max_iter: usize) -> Self {
self.max_iter = max_iter;
self
}
pub fn verbose(mut self, verbose: bool) -> Self {
self.verbose = verbose;
self
}
pub fn confidence_threshold(mut self, threshold: f64) -> Self {
self.confidence_threshold = threshold;
self
}
fn extract_view(&self, X: &Array2<f64>, view_features: &[usize]) -> SklResult<Array2<f64>> {
if view_features.is_empty() {
return Err(SklearsError::InvalidInput(
"View features cannot be empty".to_string(),
));
}
let n_samples = X.nrows();
let n_features = view_features.len();
let mut view_X = Array2::zeros((n_samples, n_features));
for (new_j, &old_j) in view_features.iter().enumerate() {
if old_j >= X.ncols() {
return Err(SklearsError::InvalidInput(format!(
"Feature index {} out of bounds",
old_j
)));
}
for i in 0..n_samples {
view_X[[i, new_j]] = X[[i, old_j]];
}
}
Ok(view_X)
}
fn simple_classifier_predict(
&self,
X_train: &Array2<f64>,
y_train: &Array1<i32>,
X_test: &Array2<f64>,
classes: &[i32],
) -> (Array1<i32>, Array1<f64>) {
let n_test = X_test.nrows();
let mut predictions = Array1::zeros(n_test);
let mut confidences = Array1::zeros(n_test);
for i in 0..n_test {
let mut distances: Vec<(f64, i32)> = Vec::new();
for j in 0..X_train.nrows() {
let diff = &X_test.row(i) - &X_train.row(j);
let dist = diff.mapv(|x| x * x).sum().sqrt();
distances.push((dist, y_train[j]));
}
distances.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap_or(std::cmp::Ordering::Equal));
let k = distances.len().clamp(1, 5);
let mut class_votes: HashMap<i32, f64> = HashMap::new();
let mut total_weight = 0.0;
for &(dist, label) in distances.iter().take(k) {
let weight = if dist > 0.0 { 1.0 / (1.0 + dist) } else { 1.0 };
*class_votes.entry(label).or_insert(0.0) += weight;
total_weight += weight;
}
for (_, vote) in class_votes.iter_mut() {
*vote /= total_weight;
}
let (best_class, best_confidence) = class_votes
.iter()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap_or(std::cmp::Ordering::Equal))
.map(|(&class, &conf)| (class, conf))
.unwrap_or((classes[0], 0.0));
predictions[i] = best_class;
confidences[i] = best_confidence;
}
(predictions, confidences)
}
}
impl Default for CoTraining<Untrained> {
fn default() -> Self {
Self::new()
}
}
impl Estimator for CoTraining<Untrained> {
type Config = ();
type Error = SklearsError;
type Float = Float;
fn config(&self) -> &Self::Config {
&()
}
}
impl Fit<ArrayView2<'_, Float>, ArrayView1<'_, i32>> for CoTraining<Untrained> {
type Fitted = CoTraining<CoTrainingTrained>;
#[allow(non_snake_case)]
fn fit(self, X: &ArrayView2<'_, Float>, y: &ArrayView1<'_, i32>) -> SklResult<Self::Fitted> {
let X = X.to_owned();
let mut y = y.to_owned();
if self.view1_features.is_empty() || self.view2_features.is_empty() {
return Err(SklearsError::InvalidInput(
"Both views must have at least one feature".to_string(),
));
}
let overlap: HashSet<_> = self
.view1_features
.iter()
.filter(|f| self.view2_features.contains(f))
.collect();
if !overlap.is_empty() && self.verbose {
println!("Warning: Views have overlapping features: {:?}", overlap);
}
let mut labeled_mask = Array1::from_elem(y.len(), false);
let mut classes = HashSet::new();
for (i, &label) in y.iter().enumerate() {
if label != -1 {
labeled_mask[i] = true;
classes.insert(label);
}
}
if labeled_mask.iter().all(|&x| !x) {
return Err(SklearsError::InvalidInput(
"No labeled samples provided".to_string(),
));
}
let classes: Vec<i32> = classes.into_iter().collect();
if classes.len() != 2 {
return Err(SklearsError::InvalidInput(
"Co-training currently supports binary classification only".to_string(),
));
}
let X_view1 = self.extract_view(&X, &self.view1_features)?;
let X_view2 = self.extract_view(&X, &self.view2_features)?;
for iter in 0..self.max_iter {
let labeled_indices: Vec<usize> = labeled_mask
.iter()
.enumerate()
.filter(|(_, &is_labeled)| is_labeled)
.map(|(i, _)| i)
.collect();
let unlabeled_indices: Vec<usize> = labeled_mask
.iter()
.enumerate()
.filter(|(_, &is_labeled)| !is_labeled)
.map(|(i, _)| i)
.collect();
if unlabeled_indices.is_empty() {
if self.verbose {
println!("Iteration {}: All samples labeled", iter + 1);
}
break;
}
let X1_labeled = labeled_indices
.iter()
.map(|&i| X_view1.row(i).to_owned())
.collect::<Vec<_>>();
let X2_labeled = labeled_indices
.iter()
.map(|&i| X_view2.row(i).to_owned())
.collect::<Vec<_>>();
let y_labeled: Array1<i32> = labeled_indices.iter().map(|&i| y[i]).collect();
let X1_labeled = Array2::from_shape_vec(
(X1_labeled.len(), X_view1.ncols()),
X1_labeled.into_iter().flatten().collect(),
)
.map_err(|_| {
SklearsError::InvalidInput("Failed to create view1 training data".to_string())
})?;
let X2_labeled = Array2::from_shape_vec(
(X2_labeled.len(), X_view2.ncols()),
X2_labeled.into_iter().flatten().collect(),
)
.map_err(|_| {
SklearsError::InvalidInput("Failed to create view2 training data".to_string())
})?;
let X1_unlabeled = unlabeled_indices
.iter()
.map(|&i| X_view1.row(i).to_owned())
.collect::<Vec<_>>();
let X2_unlabeled = unlabeled_indices
.iter()
.map(|&i| X_view2.row(i).to_owned())
.collect::<Vec<_>>();
let X1_unlabeled = Array2::from_shape_vec(
(X1_unlabeled.len(), X_view1.ncols()),
X1_unlabeled.into_iter().flatten().collect(),
)
.map_err(|_| {
SklearsError::InvalidInput("Failed to create view1 unlabeled data".to_string())
})?;
let X2_unlabeled = Array2::from_shape_vec(
(X2_unlabeled.len(), X_view2.ncols()),
X2_unlabeled.into_iter().flatten().collect(),
)
.map_err(|_| {
SklearsError::InvalidInput("Failed to create view2 unlabeled data".to_string())
})?;
let (pred1, conf1) =
self.simple_classifier_predict(&X1_labeled, &y_labeled, &X2_unlabeled, &classes);
let (pred2, conf2) =
self.simple_classifier_predict(&X2_labeled, &y_labeled, &X1_unlabeled, &classes);
let mut added_any = false;
for &target_class in &classes {
let mut candidates1: Vec<(usize, f64)> = pred1
.iter()
.zip(conf1.iter())
.enumerate()
.filter(|(_, (&pred, &conf))| {
pred == target_class && conf >= self.confidence_threshold
})
.map(|(i, (_, &conf))| (i, conf))
.collect();
candidates1
.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
let add_count = if target_class == classes[0] {
self.p
} else {
self.n
};
for (candidate_idx, _) in candidates1.into_iter().take(add_count) {
let original_idx = unlabeled_indices[candidate_idx];
y[original_idx] = target_class;
labeled_mask[original_idx] = true;
added_any = true;
}
let mut candidates2: Vec<(usize, f64)> = pred2
.iter()
.zip(conf2.iter())
.enumerate()
.filter(|(_, (&pred, &conf))| {
pred == target_class && conf >= self.confidence_threshold
})
.map(|(i, (_, &conf))| (i, conf))
.collect();
candidates2
.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
for (candidate_idx, _) in candidates2.into_iter().take(add_count) {
let original_idx = unlabeled_indices[candidate_idx];
if !labeled_mask[original_idx] {
y[original_idx] = target_class;
labeled_mask[original_idx] = true;
added_any = true;
}
}
}
if !added_any {
if self.verbose {
println!("Iteration {}: No confident predictions, stopping", iter + 1);
}
break;
}
if self.verbose {
let n_labeled = labeled_mask.iter().filter(|&&x| x).count();
println!("Iteration {}: {} labeled samples", iter + 1, n_labeled);
}
}
Ok(CoTraining {
state: CoTrainingTrained {
X_train: X.clone(),
y_train: y,
classes: Array1::from(classes),
labeled_mask,
view1_features: self.view1_features.clone(),
view2_features: self.view2_features.clone(),
},
view1_features: self.view1_features,
view2_features: self.view2_features,
p: self.p,
n: self.n,
max_iter: self.max_iter,
verbose: self.verbose,
confidence_threshold: self.confidence_threshold,
})
}
}
impl CoTraining<CoTrainingTrained> {
fn extract_view(&self, X: &Array2<f64>, view_features: &[usize]) -> SklResult<Array2<f64>> {
if view_features.is_empty() {
return Err(SklearsError::InvalidInput(
"View features cannot be empty".to_string(),
));
}
let n_samples = X.nrows();
let n_features = view_features.len();
let mut view_X = Array2::zeros((n_samples, n_features));
for (new_j, &old_j) in view_features.iter().enumerate() {
if old_j >= X.ncols() {
return Err(SklearsError::InvalidInput(format!(
"Feature index {} out of bounds",
old_j
)));
}
for i in 0..n_samples {
view_X[[i, new_j]] = X[[i, old_j]];
}
}
Ok(view_X)
}
}
impl Predict<ArrayView2<'_, Float>, Array1<i32>> for CoTraining<CoTrainingTrained> {
#[allow(non_snake_case)]
fn predict(&self, X: &ArrayView2<'_, Float>) -> SklResult<Array1<i32>> {
let X = X.to_owned();
let n_test = X.nrows();
let mut predictions = Array1::zeros(n_test);
let labeled_indices: Vec<usize> = self
.state
.labeled_mask
.iter()
.enumerate()
.filter(|(_, &is_labeled)| is_labeled)
.map(|(i, _)| i)
.collect();
let X1_train = self.extract_view(&self.state.X_train, &self.state.view1_features)?;
let X1_labeled = labeled_indices
.iter()
.map(|&i| X1_train.row(i).to_owned())
.collect::<Vec<_>>();
let X1_labeled = Array2::from_shape_vec(
(X1_labeled.len(), X1_train.ncols()),
X1_labeled.into_iter().flatten().collect(),
)
.map_err(|_| {
SklearsError::InvalidInput("Failed to create view1 training data".to_string())
})?;
let y_labeled: Array1<i32> = labeled_indices
.iter()
.map(|&i| self.state.y_train[i])
.collect();
let mut all_features: Vec<usize> = self.state.view1_features.clone();
all_features.extend(&self.state.view2_features);
all_features.sort();
all_features.dedup();
let X_test_combined = self.extract_view(&X, &all_features)?;
for i in 0..n_test {
let mut min_dist = f64::INFINITY;
let mut best_label = 0;
for (j, &labeled_idx) in labeled_indices.iter().enumerate() {
let train_combined = self.extract_view(&self.state.X_train, &all_features)?;
let diff = &X_test_combined.row(i) - &train_combined.row(labeled_idx);
let dist = diff.mapv(|x| x * x).sum().sqrt();
if dist < min_dist {
min_dist = dist;
best_label = y_labeled[j];
}
}
predictions[i] = best_label;
}
Ok(predictions)
}
}
#[derive(Debug, Clone)]
pub struct CoTrainingTrained {
pub X_train: Array2<f64>,
pub y_train: Array1<i32>,
pub classes: Array1<i32>,
pub labeled_mask: Array1<bool>,
pub view1_features: Vec<usize>,
pub view2_features: Vec<usize>,
}
#[derive(Debug, Clone)]
pub struct MultiViewCoTraining<S = Untrained> {
state: S,
views: Vec<Vec<usize>>,
k_add: usize,
max_iter: usize,
confidence_threshold: f64,
selection_strategy: String,
verbose: bool,
}
impl MultiViewCoTraining<Untrained> {
pub fn new() -> Self {
Self {
state: Untrained,
views: Vec::new(),
k_add: 1,
max_iter: 30,
confidence_threshold: 0.6,
selection_strategy: "confidence".to_string(),
verbose: false,
}
}
pub fn views(mut self, views: Vec<Vec<usize>>) -> Self {
self.views = views;
self
}
pub fn k_add(mut self, k_add: usize) -> Self {
self.k_add = k_add;
self
}
pub fn max_iter(mut self, max_iter: usize) -> Self {
self.max_iter = max_iter;
self
}
pub fn confidence_threshold(mut self, threshold: f64) -> Self {
self.confidence_threshold = threshold;
self
}
pub fn selection_strategy(mut self, strategy: String) -> Self {
self.selection_strategy = strategy;
self
}
pub fn verbose(mut self, verbose: bool) -> Self {
self.verbose = verbose;
self
}
fn extract_view(&self, X: &Array2<f64>, view_features: &[usize]) -> SklResult<Array2<f64>> {
if view_features.is_empty() {
return Err(SklearsError::InvalidInput(
"View features cannot be empty".to_string(),
));
}
let n_samples = X.nrows();
let n_features = view_features.len();
let mut view_X = Array2::zeros((n_samples, n_features));
for (new_j, &old_j) in view_features.iter().enumerate() {
if old_j >= X.ncols() {
return Err(SklearsError::InvalidInput(format!(
"Feature index {} out of bounds",
old_j
)));
}
for i in 0..n_samples {
view_X[[i, new_j]] = X[[i, old_j]];
}
}
Ok(view_X)
}
fn train_view_classifier(
&self,
X_train: &Array2<f64>,
y_train: &Array1<i32>,
X_test: &Array2<f64>,
classes: &[i32],
) -> (Array1<i32>, Array1<f64>) {
let n_test = X_test.nrows();
let mut predictions = Array1::zeros(n_test);
let mut confidences = Array1::zeros(n_test);
for i in 0..n_test {
let mut distances: Vec<(f64, i32)> = Vec::new();
for j in 0..X_train.nrows() {
let diff = &X_test.row(i) - &X_train.row(j);
let dist = diff.mapv(|x| x * x).sum().sqrt();
distances.push((dist, y_train[j]));
}
distances.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap_or(std::cmp::Ordering::Equal));
let k = distances.len().clamp(3, 7);
let mut class_votes: HashMap<i32, f64> = HashMap::new();
let mut total_weight = 0.0;
for &(dist, label) in distances.iter().take(k) {
let weight = if dist > 0.0 { 1.0 / (1.0 + dist) } else { 1.0 };
*class_votes.entry(label).or_insert(0.0) += weight;
total_weight += weight;
}
for (_, vote) in class_votes.iter_mut() {
*vote /= total_weight;
}
let (best_class, best_confidence) = class_votes
.iter()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap_or(std::cmp::Ordering::Equal))
.map(|(&class, &conf)| (class, conf))
.unwrap_or((classes[0], 0.0));
predictions[i] = best_class;
confidences[i] = best_confidence;
}
(predictions, confidences)
}
fn select_confident_samples(
&self,
predictions: &Array1<i32>,
confidences: &Array1<f64>,
classes: &[i32],
) -> Vec<(usize, i32, f64)> {
let mut candidates = Vec::new();
for i in 0..predictions.len() {
if confidences[i] >= self.confidence_threshold {
candidates.push((i, predictions[i], confidences[i]));
}
}
candidates.sort_by(|a, b| b.2.partial_cmp(&a.2).unwrap_or(std::cmp::Ordering::Equal));
match self.selection_strategy.as_str() {
"confidence" => {
let mut selected = Vec::new();
for &class in classes {
let class_candidates: Vec<_> = candidates
.iter()
.filter(|(_, c, _)| *c == class)
.take(self.k_add)
.cloned()
.collect();
selected.extend(class_candidates);
}
selected
}
"diversity" => {
candidates
.into_iter()
.take(self.k_add * classes.len())
.collect()
}
_ => candidates
.into_iter()
.take(self.k_add * classes.len())
.collect(),
}
}
}
impl Default for MultiViewCoTraining<Untrained> {
fn default() -> Self {
Self::new()
}
}
impl Estimator for MultiViewCoTraining<Untrained> {
type Config = ();
type Error = SklearsError;
type Float = Float;
fn config(&self) -> &Self::Config {
&()
}
}
impl Fit<ArrayView2<'_, Float>, ArrayView1<'_, i32>> for MultiViewCoTraining<Untrained> {
type Fitted = MultiViewCoTraining<MultiViewCoTrainingTrained>;
#[allow(non_snake_case)]
fn fit(self, X: &ArrayView2<'_, Float>, y: &ArrayView1<'_, i32>) -> SklResult<Self::Fitted> {
let X = X.to_owned();
let mut y = y.to_owned();
if self.views.len() < 3 {
return Err(SklearsError::InvalidInput(
"Multi-view co-training requires at least 3 views".to_string(),
));
}
for (i, view) in self.views.iter().enumerate() {
if view.is_empty() {
return Err(SklearsError::InvalidInput(format!(
"View {} has no features",
i
)));
}
for &feature_idx in view {
if feature_idx >= X.ncols() {
return Err(SklearsError::InvalidInput(format!(
"Feature index {} out of bounds in view {}",
feature_idx, i
)));
}
}
}
let mut labeled_mask = Array1::from_elem(y.len(), false);
let mut classes = HashSet::new();
for (i, &label) in y.iter().enumerate() {
if label != -1 {
labeled_mask[i] = true;
classes.insert(label);
}
}
if labeled_mask.iter().all(|&x| !x) {
return Err(SklearsError::InvalidInput(
"No labeled samples provided".to_string(),
));
}
let classes: Vec<i32> = classes.into_iter().collect();
for iter in 0..self.max_iter {
let mut any_labels_added = false;
for view_idx in 0..self.views.len() {
let view = &self.views[view_idx];
let labeled_indices: Vec<usize> = labeled_mask
.iter()
.enumerate()
.filter(|(_, &is_labeled)| is_labeled)
.map(|(i, _)| i)
.collect();
if labeled_indices.is_empty() {
continue;
}
let X_view = self.extract_view(&X, view)?;
let X_labeled: Vec<Vec<f64>> = labeled_indices
.iter()
.map(|&i| X_view.row(i).to_vec())
.collect();
let y_labeled: Array1<i32> = labeled_indices.iter().map(|&i| y[i]).collect();
let X_labeled = Array2::from_shape_vec(
(X_labeled.len(), view.len()),
X_labeled.into_iter().flatten().collect(),
)
.map_err(|_| {
SklearsError::InvalidInput("Failed to create labeled training data".to_string())
})?;
let unlabeled_indices: Vec<usize> = labeled_mask
.iter()
.enumerate()
.filter(|(_, &is_labeled)| !is_labeled)
.map(|(i, _)| i)
.collect();
if unlabeled_indices.is_empty() {
continue; }
let mut all_predictions = Vec::new();
let mut all_confidences = Vec::new();
for other_view_idx in 0..self.views.len() {
if other_view_idx == view_idx {
continue; }
let other_view = &self.views[other_view_idx];
let X_other_view = self.extract_view(&X, other_view)?;
let X_other_labeled: Vec<Vec<f64>> = labeled_indices
.iter()
.map(|&i| X_other_view.row(i).to_vec())
.collect();
let X_other_labeled = Array2::from_shape_vec(
(X_other_labeled.len(), other_view.len()),
X_other_labeled.into_iter().flatten().collect(),
)
.map_err(|_| {
SklearsError::InvalidInput(
"Failed to create other view training data".to_string(),
)
})?;
let X_current_unlabeled: Vec<Vec<f64>> = unlabeled_indices
.iter()
.map(|&i| X_view.row(i).to_vec())
.collect();
let X_current_unlabeled = Array2::from_shape_vec(
(X_current_unlabeled.len(), view.len()),
X_current_unlabeled.into_iter().flatten().collect(),
)
.map_err(|_| {
SklearsError::InvalidInput(
"Failed to create current view unlabeled data".to_string(),
)
})?;
let (pred, conf) = self.train_view_classifier(
&X_other_labeled,
&y_labeled,
&X_current_unlabeled,
&classes,
);
all_predictions.push(pred);
all_confidences.push(conf);
}
if all_predictions.is_empty() {
continue;
}
let n_unlabeled = unlabeled_indices.len();
let mut final_predictions = Array1::zeros(n_unlabeled);
let mut final_confidences = Array1::zeros(n_unlabeled);
for i in 0..n_unlabeled {
let mut class_votes: HashMap<i32, f64> = HashMap::new();
let mut total_confidence = 0.0;
for (pred, conf) in all_predictions.iter().zip(all_confidences.iter()) {
let confidence = conf[i];
let prediction = pred[i];
*class_votes.entry(prediction).or_insert(0.0) += confidence;
total_confidence += confidence;
}
if total_confidence > 0.0 {
for (_, vote) in class_votes.iter_mut() {
*vote /= total_confidence;
}
let (best_class, best_confidence) = class_votes
.iter()
.max_by(|a, b| {
a.1.partial_cmp(b.1).unwrap_or(std::cmp::Ordering::Equal)
})
.map(|(&class, &conf)| (class, conf))
.unwrap_or((classes[0], 0.0));
final_predictions[i] = best_class;
final_confidences[i] = best_confidence;
}
}
let selected =
self.select_confident_samples(&final_predictions, &final_confidences, &classes);
for (unlabeled_idx, label, _confidence) in selected {
if unlabeled_idx < unlabeled_indices.len() {
let sample_idx = unlabeled_indices[unlabeled_idx];
y[sample_idx] = label;
labeled_mask[sample_idx] = true;
any_labels_added = true;
}
}
}
if !any_labels_added {
if self.verbose {
println!("Multi-view co-training converged at iteration {}", iter + 1);
}
break;
}
if self.verbose {
let n_labeled = labeled_mask.iter().filter(|&&x| x).count();
println!("Iteration {}: {} labeled samples", iter + 1, n_labeled);
}
}
Ok(MultiViewCoTraining {
state: MultiViewCoTrainingTrained {
X_train: X.clone(),
y_train: y,
classes: Array1::from(classes),
labeled_mask,
views: self.views.clone(),
},
views: self.views,
k_add: self.k_add,
max_iter: self.max_iter,
confidence_threshold: self.confidence_threshold,
selection_strategy: self.selection_strategy,
verbose: self.verbose,
})
}
}
impl MultiViewCoTraining<MultiViewCoTrainingTrained> {
fn extract_view(&self, X: &Array2<f64>, view_features: &[usize]) -> SklResult<Array2<f64>> {
if view_features.is_empty() {
return Err(SklearsError::InvalidInput(
"View features cannot be empty".to_string(),
));
}
let n_samples = X.nrows();
let n_features = view_features.len();
let mut view_X = Array2::zeros((n_samples, n_features));
for (new_j, &old_j) in view_features.iter().enumerate() {
if old_j >= X.ncols() {
return Err(SklearsError::InvalidInput(format!(
"Feature index {} out of bounds",
old_j
)));
}
for i in 0..n_samples {
view_X[[i, new_j]] = X[[i, old_j]];
}
}
Ok(view_X)
}
fn train_view_classifier(
&self,
X_train: &Array2<f64>,
y_train: &Array1<i32>,
X_test: &Array2<f64>,
classes: &[i32],
) -> (Array1<i32>, Array1<f64>) {
let n_test = X_test.nrows();
let mut predictions = Array1::zeros(n_test);
let mut confidences = Array1::zeros(n_test);
for i in 0..n_test {
let mut distances: Vec<(f64, i32)> = Vec::new();
for j in 0..X_train.nrows() {
let diff = &X_test.row(i) - &X_train.row(j);
let dist = diff.mapv(|x| x * x).sum().sqrt();
distances.push((dist, y_train[j]));
}
distances.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap_or(std::cmp::Ordering::Equal));
let k = distances.len().clamp(3, 7);
let mut class_votes: HashMap<i32, f64> = HashMap::new();
let mut total_weight = 0.0;
for &(dist, label) in distances.iter().take(k) {
let weight = if dist > 0.0 { 1.0 / (1.0 + dist) } else { 1.0 };
*class_votes.entry(label).or_insert(0.0) += weight;
total_weight += weight;
}
for (_, vote) in class_votes.iter_mut() {
*vote /= total_weight;
}
let (best_class, best_confidence) = class_votes
.iter()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap_or(std::cmp::Ordering::Equal))
.map(|(&class, &conf)| (class, conf))
.unwrap_or((classes[0], 0.0));
predictions[i] = best_class;
confidences[i] = best_confidence;
}
(predictions, confidences)
}
}
impl Predict<ArrayView2<'_, Float>, Array1<i32>>
for MultiViewCoTraining<MultiViewCoTrainingTrained>
{
#[allow(non_snake_case)]
fn predict(&self, X: &ArrayView2<'_, Float>) -> SklResult<Array1<i32>> {
let X = X.to_owned();
let n_test = X.nrows();
let mut predictions = Array1::zeros(n_test);
let labeled_indices: Vec<usize> = self
.state
.labeled_mask
.iter()
.enumerate()
.filter(|(_, &is_labeled)| is_labeled)
.map(|(i, _)| i)
.collect();
for i in 0..n_test {
let mut found_exact_match = false;
for j in 0..self.state.X_train.nrows() {
if i < self.state.X_train.nrows() {
let diff = &X.row(i) - &self.state.X_train.row(j);
let distance = diff.mapv(|x| x * x).sum().sqrt();
if distance < 1e-10 && i == j && self.state.labeled_mask[j] {
predictions[i] = self.state.y_train[j];
found_exact_match = true;
break;
}
}
}
if !found_exact_match {
let mut class_votes: HashMap<i32, f64> = HashMap::new();
let mut total_weight = 0.0;
for view in &self.state.views {
let X_view_train = self.extract_view(&self.state.X_train, view)?;
let X_view_test = self.extract_view(&X, view)?;
let X_labeled: Vec<Vec<f64>> = labeled_indices
.iter()
.map(|&idx| X_view_train.row(idx).to_vec())
.collect();
let y_labeled: Array1<i32> = labeled_indices
.iter()
.map(|&idx| self.state.y_train[idx])
.collect();
let X_labeled = Array2::from_shape_vec(
(X_labeled.len(), view.len()),
X_labeled.into_iter().flatten().collect(),
)
.map_err(|_| {
SklearsError::InvalidInput("Failed to create training data".to_string())
})?;
let test_sample = X_view_test
.row(i)
.to_owned()
.insert_axis(scirs2_core::ndarray::Axis(0));
let (view_predictions, view_confidences) = self.train_view_classifier(
&X_labeled,
&y_labeled,
&test_sample,
&self.state.classes.to_vec(),
);
let prediction = view_predictions[0];
let confidence = view_confidences[0];
*class_votes.entry(prediction).or_insert(0.0) += confidence;
total_weight += confidence;
}
let best_class = if total_weight > 0.0 {
class_votes
.iter()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap_or(std::cmp::Ordering::Equal))
.map(|(&class, _)| class)
.unwrap_or(self.state.classes[0])
} else {
self.state.classes[0]
};
predictions[i] = best_class;
}
}
Ok(predictions)
}
}
#[derive(Debug, Clone)]
pub struct MultiViewCoTrainingTrained {
pub X_train: Array2<f64>,
pub y_train: Array1<i32>,
pub classes: Array1<i32>,
pub labeled_mask: Array1<bool>,
pub views: Vec<Vec<usize>>,
}