# 5.2. 分类指标
分类指标放在 `rustyml::metrics` 里,和[回归指标](./5.1._回归指标.md)、[聚类指标](./5.3._聚类指标.md)在同一个模块。本页提到的每个函数和类型都从该模块平铺重导出,所以 `use rustyml::metrics::{accuracy, ConfusionMatrix, roc_auc};` 就是你需要的全部引入。这套接口按输入类型分成几类,这一点常让从 scikit-learn 转过来的人感到困惑:有的入口收硬类别标签,有的收决策阈值,还有的收原始概率或分数。把标签表示搞错是使用这个模块时最常见的错误。
## 5.2.1. 模块的约定与错误信号
[第 2 章](../Chapter-02/2.0._经典机器学习.md)的估计器返回 `Result<_, Error>`,本模块的函数不是这样,它们在前置条件不满足时直接 **panic**。这是有意为之:这里的分类函数都是纯粹的 `array -> scalar` 代码,只依赖 `ndarray` 和 `ahash`。遇到形状不匹配时,本模块的行为和 `ndarray` 本身一致,直接 panic,而不是返回 crate 的[错误类型](../Chapter-01/1.6._错误处理.md)。
每个函数都会检查两条前置条件:长度相等、输入非空。不满足时会以 `dimension mismatch: expected N, found M` 或 `input is empty: ...` 中止,这些措辞是刻意和 crate 的 `Error` 变体保持一致的。指标函数不是修复坏输入的地方:如果 `y_true` 和 `y_pred` 长度不同,说明上游有 bug,panic 会在调用点把它暴露出来,而不是悄悄返回一个具有误导性的 `0.0`。
参数顺序永远是 `(y_true, y_pred)`,真值在前。对 `accuracy` 这类对称指标来说顺序无所谓,但 `ConfusionMatrix::new` 和 `roc_auc` 靠参数顺序判断哪个数组是真值,所以养成把 `y_true` 放在前面的习惯。
模块一共用到三种标签表示,列在下表中。编译器会强制你遵守,但 panic 信息不会替你解释这套设计。
| 入口 | `y_true` 元素类型 | 预测 / 分数类型 | 适用范围 |
| --- | --- | --- | --- |
| `accuracy` | `f64`(离散标签) | `f64`(离散标签) | 二分类或多分类 |
| `ConfusionMatrix::new` | `f64` 硬标签 | `f64` 硬标签 | 仅二分类 |
| `ConfusionMatrix::new_with_labels` | `f64`,显式给定标签对 | `f64`,显式给定标签对 | 仅二分类 |
| `roc_auc`、`roc_curve`、`average_precision`、`precision_recall_curve` | `bool` | `f64` 分数 | 仅二分类 |
| `MulticlassConfusionMatrix::new` | `usize` | `usize` | 多分类 |
| `log_loss`、`top_k_accuracy` | `usize` | `f64` 概率矩阵 | 多分类 |
| `cohen_kappa` | `usize` | `usize` | 多分类 |
## 5.2.2. 准确率,以及它在类别不均衡时为何会骗人
自由函数 `accuracy(&y_true, &y_pred)` 返回完全匹配标签的占比,逐对以 `f64::EPSILON` 为容差比较。这意味着它是为**以 `f64` 存储的离散类别标签**设计的,比如 `0.0`、`1.0`、`2.0`,对二分类和多分类同样适用。比较是对称的,交换参数不影响结果。
不要把概率喂给 `accuracy`。它不做任何阈值处理,所以 `0.87` 和 `1.0` 会被算作不匹配。请自己先把概率阈值化。`ConfusionMatrix` 同样不会替你做这件事,它只收 `0.0`/`1.0` 硬标签,遇到别的值就会 panic。
准确率有一个实实在在的弱点:面对不均衡数据,它给出的汇总很糟糕。针对这种情况,模块提供了两个更诚实的指标。设想一个筛查问题:100 个样本,只有 5 个正例。一个对所有样本都预测负类的模型能拿到 95% 的准确率,却一个真正要紧的病例都没抓到。
```rust
use rustyml::metrics::{accuracy, ConfusionMatrix};
use ndarray::Array1;
fn main() {
// 不均衡问题:100 个样本,只有 5 个正例。
let mut truth = vec![0.0f64; 100];
for t in truth.iter_mut().take(5) {
*t = 1.0;
}
let y_true = Array1::from(truth);
// 一个永远预测多数(负)类的模型。
let y_pred = Array1::from(vec![0.0f64; 100]);
println!("accuracy: {:.3}", accuracy(&y_true, &y_pred)); // ~0.95
let cm = ConfusionMatrix::new(&y_true, &y_pred);
println!("recall: {:.3}", cm.recall()); // 0.0(什么都没抓到)
println!("balanced accuracy: {:.3}", cm.balanced_accuracy()); // 0.5(相当于瞎猜)
println!("MCC: {:.3}", cm.mcc()); // 0.0(毫无相关性)
}
```
`balanced_accuracy` 取召回率和特异度的平均,所以无论类别怎么倾斜,一个只会预测多数类的模型都被钉在 0.5。马修斯相关系数(`mcc`)更进一步,把混淆矩阵的四个格子折算成 `[-1, 1]` 区间内的单个相关值。对这个退化模型它读出 0,因为预测和真值之间根本没有任何关联。这两个数能告诉你 95% 的准确率是不是真的好。
## 5.2.3. 二分类混淆矩阵
`ConfusionMatrix` 是一个小巧的 `Copy` 结构体,持有四个计数:真正例、假正例、真负例、假负例,它暴露的每个标量都由这四个数派生。`ConfusionMatrix::new(&y_true, &y_pred)` 接收两个 `f64` **硬标签**数组,每个元素必须恰好是 `0.0` 或 `1.0`。它不做任何二值化。概率、无界的 decision function 得分,或者 `-1`/`+1` 标签,都会让构造函数 panic,而不是被悄悄转换。
scikit-learn 的 `confusion_matrix` 出于同样的理由也要求硬标签。请在调用前自己把分数阈值化,这也逼着你把切分点说清楚。
早先版本的 `ConfusionMatrix::new` 会把两个参数都用硬编码的 `0.5` 二值化。这会悄悄糟蹋概率形式的真值,还会在一个毫无意义的位置切断无界得分,并把 `NaN` 算成负例。如果有代码依赖那个旧行为,请在调用前显式加上 `mapv(|p| if p >= 0.5 { 1.0 } else { 0.0 })` 来还原它。
若标签用的是另一对取值,比如外部间隔分类器给出的 `-1`/`+1`,就用 `ConfusionMatrix::new_with_labels(&y_true, &y_pred, negative_label, positive_label)`,它是 scikit-learn `labels=[neg, pos]` 的二分类形式。RustyML 自带的 [`SVC` 与 `LinearSVC`](../Chapter-02/2.5._支持向量机.md) 预测出的就是 `0.0`/`1.0`,普通的 `new` 已经够用。两个参数的存储类型彼此独立,拿一个数组配一个视图完全没问题:`ConfusionMatrix::new(&y_test, &model.predict(&x)?.view())` 能编译通过。
```rust
use ndarray::array;
use rustyml::metrics::ConfusionMatrix;
fn main() {
// 硬 0/1 标签:5 个真正例,3 个真负例。
let y_true = array![1.0, 1.0, 1.0, 1.0, 1.0, 0.0, 0.0, 0.0];
let y_pred = array![1.0, 1.0, 1.0, 0.0, 0.0, 0.0, 0.0, 1.0];
let cm = ConfusionMatrix::new(&y_true, &y_pred);
let (tp, fp, tn, fn_) = cm.get_counts();
println!("TP={tp} FP={fp} TN={tn} FN={fn_}"); // TP=3 FP=1 TN=2 FN=2
println!("accuracy {:.3}", cm.accuracy());
println!("precision {:.3}", cm.precision()); // 3/4
println!("recall {:.3}", cm.recall()); // 3/5
println!("specificity {:.3}", cm.specificity()); // 2/3
println!("f1 {:.3}", cm.f1_score());
print!("{}", cm.summary());
}
```
`get_counts()` 返回原始的 `(tp, fp, tn, fn)` 元组。派生访问器都返回 `f64`:`accuracy`、`error_rate`(正好是 `1 - accuracy`)、`precision`、`recall`、`specificity`、`f1_score`、`mcc` 和 `balanced_accuracy`。每个指标对分母为零都有自己的约定,而且是刻意选择的:
- `precision` 和 `recall` 在分母为空时返回 `0.0`(没有任何正预测,或没有任何实际正例)。
- `specificity` 在没有实际负例时返回 `1.0`,这个 0/0 的情况按惯例算作"没什么可错的"。
- `mcc` 在任一边际和为零时返回 `0.0`,因为此时该系数确实无定义。
这些约定和多分类矩阵的按类约定一致,两者由此保持统一。
`summary()` 把矩阵和全部八个派生指标渲染成一张格式化表格,每个指标保留四位小数。它是给日志和 notebook 看的,不要用来解析,把它当作面向人的输出:
```text
Confusion Matrix:
+-----------------+--------------------+--------------------+
| | Predicted Positive | Predicted Negative |
+-----------------+--------------------+--------------------+
| Actual Positive | TP: 3 | FN: 2 |
| Actual Negative | FP: 1 | TN: 2 |
+-----------------+--------------------+--------------------+
Performance Metrics:
- Accuracy: 0.6250
- Balanced Accuracy: 0.6333
...
```
## 5.2.4. 精确率、召回率,以及二者的取舍
精确率和召回率回答的是不同的问题,优化哪一个是领域决策,不是统计问题。精确率(`TP / (TP + FP)`)衡量的是标出来的样本里有多少是真的。召回率(`TP / (TP + FN)`)衡量的是真实的病例里模型抓到了多少。这两个指标互相拉扯,因为都取决于决策阈值。阈值调低会标出更多样本:召回率上升,因为漏掉的真实病例变少,但精确率下降,因为更多的标记是误报。阈值调高则会让取舍反过来。
在**医疗筛查**中,一次假阴性可能致命,你会为此调向高召回率,接受那些后续复检会过滤掉的误报。在**欺诈审核**中,每标出一笔都要占用分析师的时间、还会惹恼正当客户,所以精确率更重要,你宁可漏掉一些欺诈,也不愿让团队被大量误报淹没。没有哪个阈值在抽象意义上是对的,只有和这两类错误代价相匹配的那个阈值才是对的。
`f1_score` 把精确率和召回率合成它们的**调和**平均,`2PR / (P + R)`。选调和平均而不是算术平均正是关键所在。精确率 1.0、召回率 0.0 的算术平均是好看的 0.5。调和平均在这里却是 0.0,因为它由较小的那个输入主导。只有当精确率和召回率都高时 F1 才会给分类器高分,这正好适合那种既不能误报、也不能漏检的分类器。当精确率和召回率都为 0 时,crate 返回 `0.0`,而不是去做除以 0 的运算。
F1 对两类错误一视同仁。当两类错误代价不等时,请分开报告精确率和召回率,有意地挑选阈值,而不是一味追逐最高的 F1 分数。
精确率、召回率和 F1 都是**混淆矩阵上的方法**:`cm.precision()`、`cm.recall()`、`cm.f1_score()`。这里没有 scikit-learn 那种 `precision_score(y_true, y_pred)` 式的自由函数。矩阵只遍历一次数据,之后从四个整数里回答每一个问题。换作按指标各建一个自由函数,则会在每次调用时重新遍历整个数组。移植 scikit-learn 脚本,就是把一串 `*_score` 调用收进一个 `ConfusionMatrix`,再从它身上读数字。
还有两个单数值汇总同样作为矩阵上的方法存在。`cm.balanced_accuracy()` 是两个类别召回率的平均,它是[准确率在不均衡数据下会骗人](#522-准确率以及它在类别不均衡时为何会骗人)的诚实对照物。一个九比一的分类器,准确率读作 0.9,在这里读作 0.5,也就是抛硬币应得的分数。`cm.mcc()` 是马修斯相关系数,衡量真实标注与预测标注之间的相关性,从 -1、经过代表随机的 0、到 +1。它只有在矩阵四个格子都好看时才会升高,这让它成为这几个数字里最难吹嘘的一个。
```rust
use ndarray::array;
use rustyml::metrics::ConfusionMatrix;
fn main() {
let y_true = array![0.0, 1.0, 1.0, 0.0, 1.0, 1.0, 0.0, 0.0];
let y_pred = array![0.0, 1.0, 0.0, 0.0, 1.0, 1.0, 1.0, 0.0];
// 一次遍历建好矩阵,之后每个指标都只是读取这些计数。
let cm = ConfusionMatrix::new(&y_true, &y_pred);
println!("precision {:.3}", cm.precision());
println!("recall {:.3}", cm.recall());
println!("f1 {:.3}", cm.f1_score());
println!("balanced accuracy {:.3}", cm.balanced_accuracy());
println!("mcc {:.3}", cm.mcc());
// TP=3, FP=1, TN=3, FN=1:precision = 3/4,recall = 3/4,于是 F1 也是 3/4。
let (tp, fp, tn, fn_) = cm.get_counts();
assert_eq!((tp, fp, tn, fn_), (3, 1, 3, 1));
assert!((cm.f1_score() - 0.75).abs() < 1e-12);
}
```
## 5.2.5. 免阈值评估:ROC AUC 与 PR 曲线
以上内容都是在单一固定阈值下算出来的。有时你想评估的是模型产生的*排序*本身,而不关心最终会在哪里切分,ROC AUC 和精确率-召回率曲线给出的正是这种视角。这些函数接收 `bool` 标签(`true` 表示正类)和 `f64` 分数,分数可以是概率,也可以是任何单调的决策值。
`roc_auc(&labels, &scores)` 返回 ROC 曲线下面积,用 Mann-Whitney U 统计量计算。这给了它一个干净的解释:随机取一个正样本、随机取一个负样本,前者得分高于后者的概率。1.0 是完美排序,0.5 是抛硬币。低于 0.5 说明模型把顺序排反了。分数相同的样本取平均秩。因此一个所有分数都相等的模型,会稳稳落在 0.5,而不是由数组顺序决定。
AUC 不需要阈值,也不依赖类别比例。它的短板出现在严重不均衡的数据上:庞大的负类会让曲线看着漂亮,而任何可用阈值下的精确率其实都很差。这也是你还要看平均精度的原因。
`average_precision(&labels, &scores)` 是精确率-召回率曲线下的面积,按召回率增量的精确率加权求和得到。在不均衡问题上,它给出的是更诚实的头条数字,因为它的基准线是正例率而不是固定的 0.5。它不会因为负类好对付而白占便宜。
```rust
use ndarray::array;
use rustyml::metrics::{average_precision, precision_recall_curve, roc_auc, roc_curve};
fn main() {
// 6 个样本的分数,`true` 标记正类。
let labels = array![true, false, true, false, true, false];
let scores = array![0.95, 0.4, 0.7, 0.3, 0.6, 0.2];
println!("ROC AUC: {:.3}", roc_auc(&labels, &scores));
println!("avg precision {:.3}", average_precision(&labels, &scores));
// 完整扫描:(fpr, tpr, thresholds),三者等长,从 (0,0) 原点开始。
let (fpr, tpr, thresholds) = roc_curve(&labels, &scores);
println!("ROC points: {}", fpr.len());
println!("first tpr={:.2} fpr={:.2}", tpr[0], fpr[0]);
assert_eq!(thresholds[0], f64::INFINITY); // 原点不把任何样本判为正
// precision/recall 在 thresholds 之外多带一个收尾点,
// 且按阈值递增 / recall 递减的顺序排列。
let (precision, recall, pr_thresholds) = precision_recall_curve(&labels, &scores);
assert_eq!(precision.len(), pr_thresholds.len() + 1);
assert_eq!(precision[precision.len() - 1], 1.0); // 收尾点:recall 0,precision 1
assert_eq!(recall[recall.len() - 1], 0.0);
}
```
`roc_curve` 返回 `(fpr, tpr, thresholds)`,是三个等长的 `Array1<f64>` 数组。每个不同的分数对应一个点,按分数递减排列,前面再加上 `(0, 0)` 原点。这个原点的阈值是 **`f64::INFINITY`**,唯一一个不把任何样本判为正的取值。和有限的 `max_score + 1.0` 不同,就算分数很大,无穷大也始终能与最高的那个真实阈值区分开(`1e17 + 1.0 == 1e17`)。这与 scikit-learn 一致。
还剩一处差异,而且是有意为之。RustyML 始终返回完整的扫描结果,因此会保留每一个共线的中间点。scikit-learn 默认的 `drop_intermediate=True` 则会丢掉这些点。RustyML 的点数因此可能更多,但曲线本身以及 `roc_auc` 完全一致。用梯形法则对这条曲线积分,能精确还原出 `roc_auc`。
`precision_recall_curve` 返回的 `(precision, recall, thresholds)` 顺序恰好**相反**。阈值*递增*,因此 recall 沿数组*递减*。末尾那个 `(precision = 1, recall = 0)` 的收尾点落在低 recall 的一端,它本就该在那儿。`precision` 和 `recall` 因此比 `thresholds` **多一个元素**。zip 这几个数组时要留意这个差一。scikit-learn 对这两个曲线函数做的是同样的顺序区分,输出逐元素一致。
早先的版本把收尾点接在了*高* recall 的一端,导致 recall 在两个方向上都不单调。依赖某个特定朝向的代码,需要对照现在的行为重新核对。
全部四个排序函数都会以 panic 拒绝 `NaN` 分数(`scores must not contain NaN`)。这不是死板。`f64::total_cmp` 会把 `NaN` 排成最极端的值,一个漏进来的 `NaN` 会被悄悄当成最有把握的预测,从而污染排序。`roc_auc` 和 `roc_curve` 还要求至少一个正标签*且*至少一个负标签。`average_precision` 和 `precision_recall_curve` 要求至少一个正标签。单一类别的输入对应一条退化曲线,所以它会 panic,而不是返回一个毫无意义的数。
这些函数只支持二分类。crate 里没有 one-vs-rest 或宏平均的多分类 AUC。要在多分类问题上得到按类的 ROC,就自己把每个类二值化,对每个类各调一次 `roc_auc`。
## 5.2.6. 多分类混淆矩阵与平均方式
超过两个类别时,请使用 `MulticlassConfusionMatrix`,它的真值和预测都接收 `usize` 标签。它的类别轴是两个输入中出现过的所有标签的**排序并集**,所以哪怕某个类只出现在预测里(一个凭空冒出来的类),也照样有一行一列。`matrix()` 把完整的 `K x K` 计数网格作为 `ArrayView2<usize>` 暴露出来,行按真实类别索引,列按预测类别索引。`labels()` 给出每个索引对应的标签,`n_classes()` 给出维度。
```rust
use ndarray::array;
use rustyml::metrics::{Average, MulticlassConfusionMatrix};
fn main() {
let y_true = array![0usize, 1, 2, 2, 1, 0, 2];
let y_pred = array![0usize, 2, 2, 2, 1, 0, 1];
let cm = MulticlassConfusionMatrix::new(&y_true, &y_pred);
println!("classes: {:?}", cm.labels()); // [0, 1, 2]
println!("support: {:?}", cm.support()); // 每个类的真实样本数
println!("accuracy: {:.3}", cm.accuracy());
println!("recall: {:?}", cm.per_class_recall());
// 聚合策略是显式参数,不是藏起来的默认值。
println!("macro F1: {:.3}", cm.f1(Average::Macro));
println!("micro F1: {:.3}", cm.f1(Average::Micro));
println!("weighted F1: {:.3}", cm.f1(Average::Weighted));
// 想要按类的数字,读 per_class_* 即可,它们按标签顺序返回 Vec<f64>。
println!("per-class F1: {:?}", cm.per_class_f1());
print!("{}", cm.summary()); // 计数网格 + 按类报告
}
```
按类的视图 `per_class_precision`、`per_class_recall`、`per_class_f1` 按标签顺序返回 `Vec<f64>`,沿用和二分类矩阵相同的分母为零约定:某个类从未被预测或从未为真时取 `0.0`。`support()` 返回每个类的真值样本数,加权平均正是用的这个数。
聚合版的 `precision`、`recall`、`f1` 方法各自接收一个 `Average` 参数,宏、微、加权之分正是在这里发挥价值:
- **`Average::Macro`** 是各类分数的不加权平均。每个类不论大小都同等计入,所以你在意的稀有类不会被常见类淹没。在不均衡的多分类问题上,通常应该报告这个数字。
- **`Average::Weighted`** 用每个类的 support 给该类分数加权,衡量的是模型在一个典型样本上的表现,比宏平均更贴近准确率。
- **`Average::Micro`** 先把所有类的计数汇总,再计算指标。这个类型只支持单标签分类:每个样本恰好一个预测类。在这种情况下,微精确率、微召回率和微 F1 会全部塌缩到同一个值:准确率。实现里 `Micro` 直接返回准确率。在单标签场景下,微 F1 和准确率按定义就是同一个数。报告出不同的值,就说明哪里错了。
这三种就是全部选项,没有对应 scikit-learn `average="binary"` 的那一个。想要单个类的 one-vs-rest 数字,就按标签位置从 `per_class_precision`、`per_class_recall` 或 `per_class_f1` 里取,`cm.labels()` 给出的正是这个顺序。
`summary()` 先打印计数网格,再接一份 scikit-learn 风格的按类报告,末尾带 `macro avg` 和 `weighted avg` 两行。它会按你手上的标签和计数自动调整表格大小。这也是这个类型唯一的报告入口,没有单独的自由函数版本。
## 5.2.7. 基于概率的指标与一致性指标
还有三个函数补全这个模块,它们接收概率或成对的标注,而不是混淆矩阵。
`log_loss(&y_true, &y_prob)` 计算多分类交叉熵:`y_true` 存每个样本的真实类别索引,类型是 `usize`。`y_prob` 是一个 `Array2<f64>`,每行一个样本、每列一个类别。只有分配给真实类别的那个概率参与计分。每行在打分前都会重新归一化到和为 1,因此本来不是归一化分布的行,也能被一致地处理。之后被选中的概率会被夹到远离 0 和 1 的范围,好让对数保持有限。一个自信却错误的预测,因此得到的是一个大但有限的惩罚,而不是 `+inf`。
越低越好。
`top_k_accuracy(&y_true, &y_prob, k)` 在样本的真实类别落在概率最高的 `k` 个类别之中时,把它算作正确。如果严格比它更可能的类别不足 `k` 个,该类就并列进 top-`k` 集合,所以边界处的并列算在样本这一边。当把"正确答案落在前 5 个类别之内"当作合理标准时,就报这个指标。遇到 `k == 0`、标签超出概率列的范围,或 `y_prob` 含 `NaN` 时,它会 panic。真实类别的概率若为 `NaN`,会让 `p > true_prob` 这个比较失效,从而把样本误算成命中。
`cohen_kappa(&y_true, &y_pred)` 衡量两份标注之间经随机校正后的一致性。公式是 `(p_o - p_e) / (1 - p_e)`。其中 `p_o` 是观测到的一致性(即准确率),`p_e` 是仅凭边际标签频率就能预期到的一致性。它从 -1、经过代表随机水平的 0、到 1(完全一致)。它告诉你,准确率是否真的比一个按类别频率比例瞎猜的模型更好,这在倾斜数据上是比原始准确率更尖锐的问题。
```rust
use ndarray::array;
use rustyml::metrics::{cohen_kappa, log_loss, top_k_accuracy};
fn main() {
let y_true = array![0usize, 1, 2];
// 第 i 行 = 样本 i 的预测类别分布。
let y_prob = array![
[0.8, 0.1, 0.1],
[0.1, 0.7, 0.2],
[0.2, 0.2, 0.6],
];
println!("log loss: {:.3}", log_loss(&y_true, &y_prob)); // 越低越好
println!("top-2 acc: {:.3}", top_k_accuracy(&y_true, &y_prob, 2));
// cohen_kappa 比较的是两份硬标注,不是概率。
let y_pred = array![0usize, 1, 1];
println!("kappa: {:.3}", cohen_kappa(&y_true, &y_pred));
}
```
## 5.2.8. 端到端:评估一个逻辑回归分类器
本节把这些指标和第 2 章的[逻辑回归](../Chapter-02/2.2._逻辑回归.md)模型放到一起。`LogisticRegression::predict` 返回的已经是 `Array1<f64>` 形式的硬标签 `{0.0, 1.0}`,正是 `ConfusionMatrix` 和 `accuracy` 想要的。唯一还需要转换的,是 `roc_auc` 需要的 `bool` 标签数组,配合 `predict_proba` 得到的分数使用。
```rust
use ndarray::{array, Array1};
use rustyml::machine_learning::LogisticRegression;
use rustyml::metrics::{accuracy, roc_auc, ConfusionMatrix};
fn main() {
// 特征空间为二维的两簇、分得很开的数据。
let x_train = array![
[1.0, 1.0], [1.5, 2.0], [2.0, 1.5],
[6.0, 5.0], [5.5, 6.5], [6.5, 5.5]
];
let y_train = array![0.0, 0.0, 0.0, 1.0, 1.0, 1.0];
let mut model = LogisticRegression::new(true, 0.5, 500, 1e-6).unwrap();
model.fit(&x_train, &y_train).unwrap();
// 带已知标签的留出测试集。
let x_test = array![[1.2, 1.4], [2.2, 1.8], [5.8, 6.0], [6.2, 5.2]];
let y_test = array![0.0, 0.0, 1.0, 1.0];
// predict -> 硬标签 {0.0, 1.0},可直接用于基于标签的指标。
let y_pred: Array1<f64> = model.predict(&x_test).unwrap();
println!("accuracy: {:.3}", accuracy(&y_test, &y_pred));
let cm = ConfusionMatrix::new(&y_test, &y_pred);
print!("{}", cm.summary());
// 从原始概率得到排序质量:bool 标签 + f64 分数。
let scores = model.predict_proba(&x_test).unwrap();
let labels = y_test.mapv(|v| v >= 0.5);
println!("ROC AUC: {:.3}", roc_auc(&labels, &scores));
}
```
这份数据干净可分,模型对测试集分类完美,所以每个指标都读作 1.0。这是一次有用的健全性检查,能确认流水线接得对,但真正有意思的决策发生在别处。要研究一个实际分类器的精确率/召回率取舍,请把 `predict_proba` 的输出喂给 5.2.5 节的 `roc_curve` 和 `precision_recall_curve`。把整条曲线扫一遍,而不是死守模型内置的 0.5 阈值。
这就是只报告一个准确率数字、和报告所选工作点连同它所接受的误差,这两者之间的区别。后一种报告才是能通过评审的那种。当标签是字符串或类别而不是 `0.0`/`1.0` 时,先用[标签编码](../Chapter-04/4.3._标签编码.md)转换它们。这样它们才会落到这些指标期望的 `f64` 或 `usize` 形式。