【RustyML入门】5.0. 模型评估

5. 模型评估

训练好的模型,可信度不会超过你用来评判它的那个数字。指标选错,真实的失败就会被盖住。在一个 99% 都是负样本的欺诈数据集上,哪怕模型什么都没抓到,单看准确率也会显得很漂亮。本章讲 RustyML 的 metrics 模块。这些是把原始预测压成标量的评分函数,负责把预测结果转成你会写进报告的诊断数值。这里的一切都在 rustyml::metrics 下,由 metrics feature 控制,full 会一并打开它。这个模块做了扁平化重导出,所以 mean_squared_error 既能通过 metrics::mean_squared_error 取到,也能通过 metrics::regression::mean_squared_error 取到。想一次性全部引入作用域,写 use rustyml::prelude::metrics::*; 即可。

贯穿整个模块的有两条约定。第一,参数顺序是 (y_true, y_pred),真值在前。对 MSE、MAE、准确率这类对称指标来说,顺序不影响结果。但 r2_scoreConfusionMatrix::newroc_auc 会因顺序不同给出不同结果,所以要养成习惯把顺序写对。第二,经典机器学习 里的估计器出错时返回 crate 的 Error,这里的函数不一样,它们在违反前置条件时会直接 panic 。长度不匹配或输入为空,总会触发 panic,而不是返回 Result。评分里出现 NaN,在大多数函数里同样会触发 panic,但 r2_scoreexplained_variance_score 会把 NaN 当作数据来处理,回归指标一节有说明。这个模块是一组纯函数构成的轻量叶子。它的行为对齐 ndarray 自身处理维度不匹配的方式。这个取舍与 crate 其余部分有何不同,见 错误处理

rust 复制代码
use ndarray::array;
use rustyml::metrics::{accuracy, mean_squared_error, r2_score};

fn main() {
    // 回归:(y_true, y_pred),真值在前
    let y_true = array![3.0, -0.5, 2.0, 7.0];
    let y_pred = array![2.5, 0.0, 2.0, 8.0];
    println!("MSE = {:.4}", mean_squared_error(&y_true, &y_pred));
    println!("R^2 = {:.4}", r2_score(&y_true, &y_pred));

    // 分类:以 f64 存储的整数标签,计算精确匹配准确率
    let labels = array![0.0, 1.0, 1.0, 0.0];
    let preds = array![0.0, 1.0, 0.0, 0.0];
    println!("accuracy = {:.4}", accuracy(&labels, &preds));
}

回归指标 一节覆盖连续目标的评分:mean_squared_error 及其开方版 root_mean_squared_errormean_absolute_error、对离群点稳健的 median_absolute_errormean_absolute_percentage_error,还有两个衡量方差解释度的指标 r2_scoreexplained_variance_score。在它们之间做选择时要多加留意。r2_score 会让 NaN 一路传播,脏数据会大声暴露出来。explained_variance_score 则会悄悄跳过非有限样本,还会忽略一个恒定的预测偏差。这份便利很好用,直到它把真实问题盖住为止。

分类指标 是篇幅最大的一节,因为标签问题很少能压成单个数字。它覆盖二分类的 ConfusionMatrix。你需要用自己阈值化好的 0/1 硬标签来构建 ConfusionMatrix,喂别的值会让它 panic。准确率、精确率、召回率、特异度、F1、MCC 和平衡准确率,都从它的计数里推导出来。这一节还覆盖 MulticlassConfusionMatrix,支持通过 Average 枚举做宏平均、微平均、加权聚合。它还覆盖独立的 accuracyroc_aucaverage_precisionlog_losscohen_kappatop_k_accuracy 函数,以及扫描阈值的 roc_curveprecision_recall_curve。留意输入类型。有些函数接收 bool 标签配 f64 评分。另一些接收 usize 类别索引配一个概率矩阵。

聚类指标 一节沿着一条重要的界线分成两类。外部指标(adjusted_rand_indexnormalized_mutual_infoadjusted_mutual_info,同质性、完整性、V-measure,以及 fowlkes_mallows_score)拿聚类结果去和真值标签比对。内部指标(silhouette_scoredavies_bouldin_scorecalinski_harabasz_score)则仅凭特征空间的几何结构给聚类打分,用在没有真值的场合。silhouette_score 接收一个来自 距离度量DistanceCalculationMetric。这样一来,你就能用当初聚类时的同一种距离来做评估。

建议按顺序读这几节。它们都沿用上面立下的 (y_true, y_pred) 顺序与 panic 约定,5.1 为后面几节定下基调。先训练一个模型------用 经典机器学习神经网络------再用 训练集与测试集划分 留出一份测试集,本章才能发挥最大价值。指标只有在模型训练时从未见过的数据上才有意义。唯一的硬性前提,是对 使用 ndarray 准备数据 有基本掌握,因为这里每个函数都吃进、也吐出 ndarray 类型。

相关推荐
Evand J1 小时前
【MATLAB例程,PDR11】二维平面下的PDR(行人航位推算)步态检测与EKF融合定位,附代码的下载链接
开发语言·matlab·平面·pdr·行人导航·步态检测
唐青枫1 小时前
别只把 enum 当常量列表:Zig 枚举、状态机与 Tagged Union 实战
后端
watersink1 小时前
机器学习关联分析
人工智能·算法·机器学习
牛油果子哥q2 小时前
C++STL算法超全精讲:排序/查找/去重/遍历/最值/合并全套API、仿函数、Lambda适配、实战避坑
开发语言·c++·算法
2601_963870212 小时前
【计算机毕业设计】基于Spring Boot的非遗文创产品交易平台的设计与实现
java·spring boot·后端
lv__pf2 小时前
Spring之AOP底层源码解析(下)【TL spring14】
java·后端·spring
Xeon_CC2 小时前
用C++编写屏幕准星
开发语言·c++
孤存5222 小时前
c语言动态内存管理
c语言·开发语言·算法
脑海科技实验室2 小时前
CNS Neurosci. Ther.:亚临床抑郁症中凸显-默认模式网络动态变化:基于前聚类的共激活模式分析
机器学习·聚类·抑郁症