use ndarray::{ArrayBase, Ix1, Data};
use crate::error::{Result, SklearnError};
pub fn mean_squared_error<S1, S2>(
y_true: &ArrayBase<S1, Ix1>,
y_pred: &ArrayBase<S2, Ix1>,
) -> Result<f64>
where
S1: Data<Elem = f64>,
S2: Data<Elem = f64>,
{
if y_true.len() != y_pred.len() {
return Err(SklearnError::ShapeMismatch {
expected: format!("{}", y_true.len()),
actual: format!("{}", y_pred.len()),
});
}
let diff = y_true - y_pred;
let mse = diff.mapv(|x| x.powi(2)).mean().unwrap_or(0.0);
Ok(mse)
}
pub fn mean_absolute_error<S1, S2>(
y_true: &ArrayBase<S1, Ix1>,
y_pred: &ArrayBase<S2, Ix1>,
) -> Result<f64>
where
S1: Data<Elem = f64>,
S2: Data<Elem = f64>,
{
if y_true.len() != y_pred.len() {
return Err(SklearnError::ShapeMismatch {
expected: format!("{}", y_true.len()),
actual: format!("{}", y_pred.len()),
});
}
let diff = y_true - y_pred;
let mae = diff.mapv(|x| x.abs()).mean().unwrap_or(0.0);
Ok(mae)
}
pub fn r2_score<S1, S2>(
y_true: &ArrayBase<S1, Ix1>,
y_pred: &ArrayBase<S2, Ix1>,
) -> Result<f64>
where
S1: Data<Elem = f64>,
S2: Data<Elem = f64>,
{
if y_true.len() != y_pred.len() {
return Err(SklearnError::ShapeMismatch {
expected: format!("{}", y_true.len()),
actual: format!("{}", y_pred.len()),
});
}
let y_mean = y_true.mean().unwrap_or(0.0);
let total_sum_squares: f64 = y_true.mapv(|y| (y - y_mean).powi(2)).sum();
if total_sum_squares.abs() < 1e-10 {
return Err(SklearnError::NumericalError {
reason: "目标值的方差为零".to_string(),
});
}
let residual_sum_squares: f64 = y_true
.iter()
.zip(y_pred.iter())
.map(|(&y, &y_hat)| (y - y_hat).powi(2))
.sum();
let r2 = 1.0 - (residual_sum_squares / total_sum_squares);
Ok(r2)
}