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_score、ConfusionMatrix::new 和 roc_auc 会因顺序不同给出不同结果,所以要养成习惯把顺序写对。第二,经典机器学习 里的估计器出错时返回 crate 的 Error,这里的函数不一样,它们在违反前置条件时会直接 panic 。长度不匹配或输入为空,总会触发 panic,而不是返回 Result。评分里出现 NaN,在大多数函数里同样会触发 panic,但 r2_score 和 explained_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_error、mean_absolute_error、对离群点稳健的 median_absolute_error、mean_absolute_percentage_error,还有两个衡量方差解释度的指标 r2_score 与 explained_variance_score。在它们之间做选择时要多加留意。r2_score 会让 NaN 一路传播,脏数据会大声暴露出来。explained_variance_score 则会悄悄跳过非有限样本,还会忽略一个恒定的预测偏差。这份便利很好用,直到它把真实问题盖住为止。
分类指标 是篇幅最大的一节,因为标签问题很少能压成单个数字。它覆盖二分类的 ConfusionMatrix。你需要用自己阈值化好的 0/1 硬标签来构建 ConfusionMatrix,喂别的值会让它 panic。准确率、精确率、召回率、特异度、F1、MCC 和平衡准确率,都从它的计数里推导出来。这一节还覆盖 MulticlassConfusionMatrix,支持通过 Average 枚举做宏平均、微平均、加权聚合。它还覆盖独立的 accuracy、roc_auc、average_precision、log_loss、cohen_kappa、top_k_accuracy 函数,以及扫描阈值的 roc_curve 与 precision_recall_curve。留意输入类型。有些函数接收 bool 标签配 f64 评分。另一些接收 usize 类别索引配一个概率矩阵。
聚类指标 一节沿着一条重要的界线分成两类。外部指标(adjusted_rand_index、normalized_mutual_info、adjusted_mutual_info,同质性、完整性、V-measure,以及 fowlkes_mallows_score)拿聚类结果去和真值标签比对。内部指标(silhouette_score、davies_bouldin_score、calinski_harabasz_score)则仅凭特征空间的几何结构给聚类打分,用在没有真值的场合。silhouette_score 接收一个来自 距离度量 的 DistanceCalculationMetric。这样一来,你就能用当初聚类时的同一种距离来做评估。
建议按顺序读这几节。它们都沿用上面立下的 (y_true, y_pred) 顺序与 panic 约定,5.1 为后面几节定下基调。先训练一个模型------用 经典机器学习 或 神经网络------再用 训练集与测试集划分 留出一份测试集,本章才能发挥最大价值。指标只有在模型训练时从未见过的数据上才有意义。唯一的硬性前提,是对 使用 ndarray 准备数据 有基本掌握,因为这里每个函数都吃进、也吐出 ndarray 类型。