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

RetinaNet模型训练后测试方法及训练流程疑问咨询

嘿,我来帮你梳理下这两个问题——我对keras-retinanet这套工具还算熟悉,刚好能给你点实用的建议:

解答

问题1:训练过程是否遗漏生成器步骤?

你完全没有遗漏必要步骤!

因为你使用的是retinanet-train csv这个命令模式,脚本内部已经自动帮你构建了CSV格式的数据生成器(对应keras_retinanet.preprocessing.csv_generator.CSVGenerator类),它会读取你提供的标注CSV和类别CSV,自动处理图片加载、数据增强、批量生成等逻辑。

而PascalVocGenerator只是针对VOC格式数据集的生成器实现,和你用的CSV模式是平行的不同数据入口,所以不用手动创建它,放心继续就行。

问题2:如何测试训练好的模型?

既然官方示例Notebook失效了,我给你两种实用的测试方案:

方案1:单张图片可视化测试

适合快速验证模型的预测效果,步骤如下:

  1. 加载训练好的模型
    注意训练保存的.h5模型包含自定义层,需要用库内的加载方法:
    from keras_retinanet.models import load_model
    # 加载模型时设置compile=False,避免因训练相关的损失函数报错
    model = load_model('path/to/your/trained_model.h5', compile=False)
    
  2. 预处理测试图片
    用库内工具统一处理图片格式:
    from keras_retinanet.utils.image import preprocess_image, resize_image
    import cv2
    import numpy as np
    
    # 读取图片(注意OpenCV默认BGR格式,转成RGB)
    image = cv2.imread('your_test_image.jpg')
    image_rgb = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
    # 预处理+缩放图片
    processed_img = preprocess_image(image_rgb)
    processed_img, scale_factor = resize_image(processed_img)
    
  3. 执行预测并筛选结果
    # 扩展维度适配模型输入
    inputs = np.expand_dims(processed_img, axis=0)
    # 预测得到框、置信度、类别
    boxes, scores, labels = model.predict_on_batch(inputs)
    # 还原框的坐标(因为之前缩放了图片)
    boxes /= scale_factor
    # 筛选置信度高于阈值的结果(比如0.5)
    valid_idx = scores[0] > 0.5
    valid_boxes = boxes[0][valid_idx]
    valid_scores = scores[0][valid_idx]
    valid_labels = labels[0][valid_idx]
    
  4. 可视化预测结果
    用库内工具或者OpenCV画框:
    from keras_retinanet.utils.visualization import draw_boxes
    # 你的类别列表,要和训练时的classes.csv顺序完全一致
    class_names = ["class1", "class2", ..., "class50"]
    # 在原图上画框
    result_img = draw_boxes(image_rgb, valid_boxes, valid_scores, valid_labels, class_names=class_names)
    # 保存或显示图片
    cv2.imwrite('prediction_result.jpg', cv2.cvtColor(result_img, cv2.COLOR_RGB2BGR))
    

方案2:批量测试并计算指标

如果需要评估模型在测试集上的整体性能(比如mAP),可以用官方提供的retinanet-eval脚本,命令格式和训练类似:

retinanet-eval csv <path_to_test_annotations.csv> <path_to_classes.csv> <path_to_trained_model.h5>

这个脚本会自动遍历测试集,计算各类别的AP、整体mAP等指标,输出到控制台。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.15 04:18:01