sklearn-rs 0.1.0

A scikit-learn inspired machine learning library in Rust
Documentation
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)
}

/// 计算 R² 分数
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)
}