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

新手求教:如何将自定义图像数据集导入TensorFlow用于目标检测?

如何用TensorFlow导入自定义目标检测数据集并训练模型

Hey there! 作为刚入坑TensorFlow和深度学习的新手,自定义数据集的导入确实是个容易卡壳的点——毕竟MNIST这类公开数据集都有现成的工具链,但自己的数据就得一步步搭流程。我来给你梳理一套清晰的实操步骤,帮你把同一物体的图像数据集用上:

第一步:整理数据集并完成标注

目标检测需要每个图像里物体的位置(边界框)和类别信息,先把基础工作做好:

  • 把所有图像归类到一个文件夹,比如./images,建议再拆成train(训练集)和val(验证集)子文件夹,比例大概8:2就行
  • 标注工具推荐用LabelImg,它简单直观,能生成PASCAL VOC格式的XML标注文件——每个图像对应一个XML,里面包含物体的xmin、ymin、xmax、ymax(边界框坐标)和类别名称
  • 因为你只有同一类物体,标注时统一设一个类别名,比如my_target_object就行

第二步:转换为TensorFlow友好的TFRecord格式

TensorFlow目标检测API最常用的是TFRecord格式,它能高效加载批量数据。你可以按下面的步骤来:

  1. 先搞定TensorFlow Object Detection API的环境:
    • 获取官方的模型工具集,配置好环境变量,编译相关依赖(按照官方文档的步骤来就行,确保Python能导入object_detection模块)
  2. 写个简单的转换脚本:
    核心是遍历你的图像和XML标注,把每个样本转换成TFRecord的Example格式。给你个简化版的代码片段参考:
    import tensorflow as tf
    from object_detection.utils import dataset_util
    import xml.etree.ElementTree as ET
    import os
    
    def xml_to_tfrecord(xml_dir, img_dir, output_path):
        writer = tf.io.TFRecordWriter(output_path)
        # 遍历所有XML标注文件
        for xml_filename in os.listdir(xml_dir):
            if not xml_filename.endswith('.xml'):
                continue
            # 解析XML
            tree = ET.parse(os.path.join(xml_dir, xml_filename))
            root = tree.getroot()
            # 获取对应图像路径并读取
            img_filename = root.find('filename').text
            img_path = os.path.join(img_dir, img_filename)
            with tf.io.gfile.GFile(img_path, 'rb') as fid:
                encoded_img = fid.read()
            # 获取图像尺寸
            img_width = int(root.find('size').find('width').text)
            img_height = int(root.find('size').find('height').text)
            # 整理边界框和类别信息
            xmins, xmaxs, ymins, ymaxs = [], [], [], []
            class_texts, class_ids = [], []
            for obj in root.findall('object'):
                # 把坐标归一化(除以图像宽高)
                xmin = float(obj.find('bndbox').find('xmin').text) / img_width
                xmax = float(obj.find('bndbox').find('xmax').text) / img_width
                ymin = float(obj.find('bndbox').find('ymin').text) / img_height
                ymax = float(obj.find('bndbox').find('ymax').text) / img_height
                xmins.append(xmin)
                xmaxs.append(xmax)
                ymins.append(ymin)
                ymaxs.append(ymax)
                # 类别信息
                class_name = obj.find('name').text
                class_texts.append(class_name.encode('utf8'))
                class_ids.append(1)  # 单类别直接设为1就行
            # 构建TFRecord Example
            example = tf.train.Example(features=tf.train.Features(feature={
                'image/encoded': dataset_util.bytes_feature(encoded_img),
                'image/format': dataset_util.bytes_feature(b'jpg'),  # PNG的话改成b'png'
                'image/object/bbox/xmin': dataset_util.float_list_feature(xmins),
                'image/object/bbox/xmax': dataset_util.float_list_feature(xmaxs),
                'image/object/bbox/ymin': dataset_util.float_list_feature(ymins),
                'image/object/bbox/ymax': dataset_util.float_list_feature(ymaxs),
                'image/object/class/text': dataset_util.bytes_list_feature(class_texts),
                'image/object/class/label': dataset_util.int64_list_feature(class_ids),
            }))
            writer.write(example.SerializeToString())
        writer.close()
    
    # 转换训练集和验证集
    xml_to_tfrecord('./annotations/train', './images/train', './train.record')
    xml_to_tfrecord('./annotations/val', './images/val', './val.record')
    

第三步:创建标签映射文件

需要一个label_map.pbtxt文件,告诉模型类别ID对应的名称,放在./annotations文件夹里,内容如下:

item {
  id: 1
  name: 'my_target_object'
}

第四步:配置训练Pipeline

从官方的samples/configs里选一个适合新手的预训练模型配置,比如ssd_mobilenet_v2_fpnlite_320x320_coco17_tpu-8.config(轻量、训练快),然后修改几个关键参数:

  • num_classes:改成你的类别数(这里是1)
  • fine_tune_checkpoint:设置预训练模型的 checkpoint 路径,比如./pretrained_model/ssd_mobilenet_v2_fpnlite_320x320_coco17_tpu-8/checkpoint/ckpt-0
  • train_input_reader下的input_path改成你的train.record路径,label_map_path改成./annotations/label_map.pbtxt
  • eval_input_reader同理,换成val.record的路径

第五步:启动训练

用官方提供的训练脚本启动训练:

python model_main_tf2.py --model_dir=./training_logs --pipeline_config_path=./ssd_mobilenet_v2_fpnlite_320x320_coco17_tpu-8.config

--model_dir是保存训练 checkpoint 和日志的文件夹,训练过程中可以用TensorBoard实时看进度:

tensorboard --logdir=./training_logs

新手小贴士

  • 如果数据集很小,不想折腾TFRecord,也可以用tf.data.Dataset直接加载图像和标注,比如用tf.io.read_file读图像,再解析XML或者CSV格式的标注,代码更灵活
  • 优先选轻量级预训练模型,训练速度快,方便调试
  • 一定要检查图像和标注的路径对应,别出现“找不到文件”的错误,这是新手常踩的坑

内容的提问来源于stack exchange,提问作者Mjd Al Mahasneh

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.19 04:31:53