You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

训练SSD后如何计算Precision、Recall及F1分数

SSD-300-TensorFlow 计算Precision、Recall、F1分数操作指南

前置准备

你当前使用的Colab环境需要先挂载存储了checkpoint和代码的谷歌云盘,执行以下命令完成挂载:

from google.colab import drive
drive.mount('/content/drive')
# 挂载后进入你的项目根目录,示例:
%cd /content/drive/MyDrive/你的项目文件夹路径

步骤1:加载训练完成的Checkpoint

首先构建和训练时一致的SSD-300模型结构,再从logs文件夹加载权重:

import tensorflow as tf
# 导入你的SSD-300模型构建代码
from your_model_file import ssd_300

# 构建和训练参数一致的模型
model = ssd_300.build_model(num_classes=你的类别数, input_shape=(300,300,3))
# 加载checkpoint
checkpoint = tf.train.Checkpoint(model=model)
checkpoint.restore(tf.train.latest_checkpoint('./logs')).expect_partial()

步骤2:导入已有的指标实现

你项目中tf_extended/metrics.py已有的Precision、Recall是原生TensorFlow实现,无需依赖scikit-learn即可直接调用:

from tf_extended.metrics import Precision, Recall

步骤3:准备验证数据集

将验证集的标注格式处理为和模型输出匹配的格式:标注需要包含真实框坐标、真实类别,模型输出需要包含预测框坐标、预测置信度、预测类别,和训练阶段的数据格式保持一致即可。

步骤4:批量计算指标并生成F1分数

遍历整个验证集完成指标统计,最后通过Precision和Recall计算F1:

# 初始化指标实例,如需按单类别计算可指定class_id参数
precision_obj = Precision(num_classes=你的类别数)
recall_obj = Recall(num_classes=你的类别数)

# 遍历验证集所有批次
for batch_imgs, batch_gt_boxes, batch_gt_labels in val_dataset:
    # 推理阶段关闭训练模式
    pred_boxes, pred_scores, pred_labels = model(batch_imgs, training=False)
    # 更新指标统计状态
    precision_obj.update_state(batch_gt_boxes, batch_gt_labels, pred_boxes, pred_labels, pred_scores)
    recall_obj.update_state(batch_gt_boxes, batch_gt_labels, pred_boxes, pred_labels, pred_scores)

# 提取最终指标结果
final_precision = precision_obj.result().numpy()
final_recall = recall_obj.result().numpy()
# 计算F1,避免除零报错
if final_precision + final_recall == 0:
    final_f1 = 0.0
else:
    final_f1 = 2 * (final_precision * final_recall) / (final_precision + final_recall)

# 输出结果
print(f"精确率: {final_precision:.4f}")
print(f"召回率: {final_recall:.4f}")
print(f"F1分数: {final_f1:.4f}")

Colab环境注意事项

  • 所有路径注意和云盘挂载后的实际路径匹配,避免出现文件找不到的错误
  • 验证集较大时建议分批次加载,不要一次性读入全部数据导致Colab内存不足崩溃
  • 指标支持按类别单独统计,只需修改初始化时的class_id参数即可输出对应类别的Precision、Recall、F1

内容的提问来源于stack exchange,提问作者EverydayDeveloper

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.01 23:36:03