Python / AI · 服务化与结课项目 · LESSON 26

评估、checkpoint 与实验记录

保存最佳模型、记录指标和比较实验,避免只凭最后一次训练结果下结论。

18 分钟evaluation · checkpoint · experiments

学习目标

本节把“训练结束打印一个 loss”升级为可审计的模型评估。你将会:

  • 在 eval 和 inference_mode 下计算正确的 loss、precision、recall、F1 与混淆矩阵。
  • 区分 batch 平均与按样本加权的指标,处理类别不平衡和阈值选择。
  • 保存包含模型、optimizer、epoch、配置、数据版本和指标的 checkpoint。
  • 记录实验上下文,验证加载 checkpoint 后结果一致,并避免用 test 集反复调参。

从 JS/TS 迁移的心智模型

JavaScript/TypeScript 项目会把构建产物和测试报告一起保存;AI 实验的产物还包括参数、预处理、数据切分和随机状态。一个 JSON 里写着 f1=0.91,并不能说明它来自哪个 checkpoint、哪套标签、哪次切分。评估函数要像 API 一样固定输入输出,实验记录要像审计日志一样能回溯。

TRANSLATION LENS 同一个意图,两种工程表达 窄屏可左右滑动查看完整代码
JS / TS
const best = metrics.f1 > bestF1;
if (best) await saveModel(model);
console.log({ epoch, metrics });
Python / PyTorch
metrics = evaluate(model, valid_loader)
if metrics["f1"] > best_f1:
  save_checkpoint(model, optimizer, metrics)
print({"epoch": epoch, **metrics})

评估模式与指标收集

评估前调用 model.eval(),再用 torch.inference_mode() 关闭梯度和 autograd 图。不要把 evaluate 放在训练模式下,否则 Dropout 和 BatchNorm 会让每次结果不同;不要只使用 loss,因为 loss 的尺度不一定对应业务风险。

对二分类,收集每个样本的 logits 或正类概率和真实标签,最后统一选择阈值。precision、recall、F1 来自全体样本的 TP、FP、FN;逐 batch 计算 F1 再平均,会让小 batch 和大 batch 权重错误。多分类要输出 macro、weighted 和每类指标,防止多数类掩盖稀有类。

示例一:返回可审核的预测结果

import torch

def collect_predictions(model, loader, device):
    model.eval()
    all_logits, all_labels, all_ids = [], [], []
    with torch.inference_mode():
        for batch in loader:
            features = batch["features"].to(device, dtype=torch.float32)
            labels = batch["labels"].to(device, dtype=torch.long)
            logits = model(features)
            if logits.ndim != 2 or logits.shape[0] != labels.shape[0]:
                raise ValueError("logits must be (batch, classes)")
            all_logits.append(logits.cpu())
            all_labels.append(labels.cpu())
            all_ids.extend(batch.get("sample_ids", []))
    if not all_logits:
        raise ValueError("empty evaluation loader")
    return torch.cat(all_logits), torch.cat(all_labels), all_ids

logits, labels, sample_ids = collect_predictions(model, loader, device)
probabilities = logits.softmax(dim=1)[:, 1]
print(tuple(logits.shape), tuple(labels.shape), len(sample_ids))

输出应显示 logits 的第一维等于评估样本数,labels shape 是该样本数,sample_ids 数量一致。先收集 CPU Tensor 再计算指标可以避免 GPU 显存随整个验证集增长;大数据集可分块写入临时文件,但必须保持 ID 与结果对应。

阈值、混淆矩阵与不平衡

二分类 logits 的 argmax 等价于概率阈值 0.5 只在特定编码下成立。唤醒任务可能宁愿增加 false negative,也不愿让设备频繁误唤醒;医疗或安全任务的代价又不同。用 validation 选择 threshold,记录目标约束,例如 recall 至少 0.90 或 false positive rate 不超过预算。test 只在 threshold 冻结后计算。

指标要和样本分组一起报告。整体 F1 之外,按 device_id、site、用户、时间窗口给出样本数和 recall;某组只有 2 条样本时标记不稳定。若标签延迟或有未标注样本,明确评估覆盖范围,不能把未知当负例后宣称精度。

示例二:从 logits 计算指标

import numpy as np
from sklearn.metrics import (
    average_precision_score,
    confusion_matrix,
    f1_score,
    precision_score,
    recall_score,
)

def binary_metrics(logits, labels, threshold=0.5):
    probabilities = logits.softmax(dim=1)[:, 1].numpy()
    truth = labels.numpy()
    predictions = (probabilities >= threshold).astype(np.int64)
    matrix = confusion_matrix(truth, predictions, labels=[0, 1])
    return {
        "threshold": float(threshold),
        "precision": float(precision_score(truth, predictions, zero_division=0)),
        "recall": float(recall_score(truth, predictions, zero_division=0)),
        "f1": float(f1_score(truth, predictions, zero_division=0)),
        "pr_auc": float(average_precision_score(truth, probabilities)),
        "confusion_matrix": matrix.tolist(),
        "rows": int(len(truth)),
    }

metrics = binary_metrics(logits, labels, threshold=0.65)
print(metrics)

运行结果的 confusion_matrix 是 TN、FP、FN、TP 的二维列表,metrics 中还应有 rows 和 threshold。若一类完全没有样本,precision/recall=0 只是计算定义,不代表模型合格;报告数据覆盖和告警。保存浮点指标时使用足够精度,展示时四舍五入,避免选择模型时因显示值相同而误判。

checkpoint:保存什么、何时保存

最小 checkpoint 应保存 model.state_dict;可继续训练的 checkpoint 还要保存 optimizer.state_dict、scheduler、epoch、best_metric、config、data_version、feature_names、class_mapping、random seed 和库版本。加载时使用 map_location,先重建同一 Module,再 load_state_dict;不要直接把完整 Python 对象反序列化成不可审计依赖。

保存策略通常有 latest 和 best 两个文件:latest 用于断点续训,best 按 validation F1 或业务指标更新。写文件时先写临时路径并原子替换,避免进程中断留下半个 artifact。checkpoint 不能包含用户原始音频、密钥或无关敏感数据。

示例三:按 validation 指标保存最佳权重

from pathlib import Path
import torch

def save_if_best(path, model, optimizer, epoch, metrics, config, data_version, best_f1):
    current = float(metrics["f1"])
    if current <= best_f1:
        return best_f1, False
    payload = {
        "model": model.state_dict(),
        "optimizer": optimizer.state_dict(),
        "epoch": epoch,
        "metrics": metrics,
        "config": config,
        "data_version": data_version,
        "best_f1": current,
    }
    temporary = Path(str(path) + ".tmp")
    torch.save(payload, temporary)
    temporary.replace(path)
    return current, True

best_f1, saved = save_if_best(
    "artifacts/best.pt", model, optimizer, epoch, metrics,
    {"batch_size": 32, "seed": 42}, "events-v3", best_f1,
)
print({"saved": saved, "best_f1": best_f1})

输出 saved=True 只表示 validation 指标超过历史 best;它不表示 test 已通过。若两个 epoch 指标相同,应固定 tie-break 规则,例如选择更早 epoch 或更低延迟。保存后立即重新加载并在同一 validation batch 上比较 logits,验证文件确实对应当前模型。

实验记录、运行与验证

每次 run 的记录至少包含 run_id、开始时间、git 或代码版本、data_version、split_id、配置、设备、参数量、训练/验证样本数、每 epoch 指标、阈值、耗时和失败信息。test 指标单独标记 final,不能混入训练曲线参与选择。

checkpoint = torch.load("artifacts/best.pt", map_location=device)
restored = EventClassifier().to(device)
restored.load_state_dict(checkpoint["model"])
restored.eval()

with torch.inference_mode():
    before = model(batch_features.to(device))
    after = restored(batch_features.to(device))
print({
    "checkpoint_epoch": checkpoint["epoch"],
    "data_version": checkpoint["data_version"],
    "max_logit_diff": float((before - after).abs().max()),
})
assert torch.allclose(before, after, atol=1e-6, rtol=1e-5)

这里的结果给出 max_logit_diff,能证明权重加载没有改变同一输入的前向结果。若差异很大,检查 Module 结构、参数 dtype/device、是否加载了错误文件和 batch 预处理;不要只看 checkpoint 文件存在。

常见错误、排错与调试

  • validation 指标每次变化:确认 eval、inference_mode、DataLoader 顺序和随机增强关闭。
  • F1 比逐 batch 平均高或低很多:收集全体预测后计算,按样本加权 loss。
  • test 指标反复上升:检查是否用 test 选 threshold、epoch 或特征;重新划出 validation。
  • 加载后输出不同:比较 feature_names、标准化参数、class_mapping、state_dict keys、dtype 和 device。
  • checkpoint 很大或加载慢:检查是否误存优化器之外的缓存、完整数据和计算图;保存 state_dict 和必要元数据。
  • 最佳模型过拟合:比较 train/valid 曲线、分组指标和时间切分,增加正则化或更公平数据。
  • 评估超时或 OOM:减少评估 batch、按块收集 CPU 结果,记录峰值显存和每批耗时,不要降低指标覆盖而不记录。

练习与任务

为二分类事件模型写 evaluate_and_checkpoint:计算 validation 的 precision、recall、F1、PR-AUC、混淆矩阵和每设备 recall;达到更高 F1 时保存 best checkpoint;记录配置、data_version、split_id 和耗时;最后加载 checkpoint 验证 logits 一致。

01
TRY IT YOURSELF

评估与 checkpoint 练习

实现 collect_predictions、binary_metrics、save_if_best 和 load_and_verify;覆盖不平衡标签、阈值变化、空 loader、错误 logits shape、断点加载和重复运行结果。

给我一点提示

指标在全体预测上计算;best 只看 validation;torch.save 后重新 load_state_dict,并用 allclose 比较同一 batch 的 logits。

查看参考答案
model.eval()
with torch.inference_mode():
  logits = torch.cat([model(batch.to(device)) for batch in loader])
prob = logits.softmax(dim=1)[:, 1].cpu().numpy()
pred = (prob >= threshold).astype(int)
metrics = {"f1": f1_score(labels, pred, zero_division=0)}
if metrics["f1"] > best_f1:
  torch.save({"model": model.state_dict(), "metrics": metrics}, path)

完整答案

def evaluate_and_checkpoint(model, loader, device, path, best_f1, context):
    model.eval()
    logits_parts, label_parts = [], []
    with torch.inference_mode():
        for features, labels in loader:
            features = features.to(device, dtype=torch.float32)
            labels = labels.to(device, dtype=torch.long)
            logits = model(features)
            if logits.ndim != 2 or logits.shape[1] != 2:
                raise ValueError("expected binary logits (batch, 2)")
            logits_parts.append(logits.cpu())
            label_parts.append(labels.cpu())
    if not logits_parts:
        raise ValueError("validation loader is empty")

    logits = torch.cat(logits_parts)
    labels = torch.cat(label_parts)
    metrics = binary_metrics(logits, labels, threshold=context["threshold"])
    saved = False
    if metrics["f1"] > best_f1:
        torch.save({
            "model": model.state_dict(),
            "metrics": metrics,
            "context": context,
        }, path)
        best_f1 = metrics["f1"]
        saved = True
    return metrics, best_f1, saved

把 context 中的 threshold、data_version、split_id 和特征版本写入 JSON;再测试同一 checkpoint 在 CPU 和可用 GPU 上的输出误差。评估结果只有在指标、数据、权重和运行环境一一对应时,才是可以审核的模型结论。

本节结论

评估的交付物不是一个最高分,而是“哪个模型、在哪份数据、用什么阈值、为何被选中”的证据。下一节会把已选择的模型变成有资源预算的推理入口。

与同一 AI 项目主线的连接

DataLoader 提供有序的 batch,模型训练课程提供 train/eval 状态,本节把输出变成跨实验可比较的指标和 checkpoint。split manifest 与 data_version 保护评估不被泄漏;feature_names、标准化参数和 class_mapping 保护推理不被错位。把 validation 的 F1、recall、误报成本和评估耗时一并记录,才能决定模型是否值得进入服务,而不是只看一次漂亮的离线结果。

小结

评估要在正确模式下收集全体预测,再按业务代价计算指标;checkpoint 要保存权重、状态、配置和数据上下文;最佳模型由 validation 选择,test 只做最终验收;加载后用同一输入比较 logits。这样实验记录才从日志数字变成可复现、可回退、可审核的 AI 资产。

FURTHER READING

延伸阅读

先完成本节练习,再用这些资料查阅完整 API 和真实项目组织方式。

当前学习阶段服务化与结课项目
0/6

阶段共 6 节课,按顺序完成更容易建立完整的迁移模型。