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

自定义数据集适配Keras CV RetinaNet:数据加载与准备问询

自定义数据集训练Keras CV RetinaNet的数据准备方案

一、明确标注格式

labelImg导出的标注通常是PASCAL VOC XML或YOLO txt格式,先确认你使用的类型,下面以最常用的PASCAL VOC XML为例说明。

二、整理数据集目录结构

建议按如下结构存放数据,便于后续批量处理:

dataset/
├── train/
│   ├── images/
│   │   ├── img1.jpg
│   │   ├── img2.png
│   │   └── ...
│   └── annotations/
│       ├── img1.xml
│       ├── img2.xml
│       └── ...
└── val/
    ├── images/
    └── annotations/

提前手动将数据集按8:2左右的比例划分成训练集和验证集。

三、解析标注并构建训练数据集

Keras CV的RetinaNet要求输入为(图像, 标签)结构,其中标签需包含bounding_boxes(格式[xmin, ymin, xmax, ymax])和classes(从0开始的整数类别ID)。

1. 编写XML标注解析函数

import xml.etree.ElementTree as ET
import tensorflow as tf
import os

def parse_voc_xml(xml_path):
    tree = ET.parse(xml_path)
    root = tree.getroot()
    
    # 获取原始图像尺寸
    size = root.find('size')
    img_width = int(size.find('width').text)
    img_height = int(size.find('height').text)
    
    boxes = []
    classes = []
    # 替换为你的类别映射,比如{'cat':0, 'dog':1}
    class_mapping = {'your_class_1':0, 'your_class_2':1}
    
    for obj in root.findall('object'):
        cls = obj.find('name').text
        if cls not in class_mapping:
            continue
        classes.append(class_mapping[cls])
        
        bndbox = obj.find('bndbox')
        xmin = float(bndbox.find('xmin').text)
        ymin = float(bndbox.find('ymin').text)
        xmax = float(bndbox.find('xmax').text)
        ymax = float(bndbox.find('ymax').text)
        boxes.append([xmin, ymin, xmax, ymax])
    
    return boxes, classes

2. 编写图像加载与预处理函数

def load_and_preprocess_data(img_path, xml_path, target_size=(640, 640)):
    # 加载并预处理图像
    img = tf.io.read_file(img_path)
    img = tf.image.decode_jpeg(img, channels=3)  # 若为png格式替换为decode_png
    img = tf.image.resize(img, target_size)
    img = tf.cast(img, tf.float32) / 255.0  # 归一化到0-1区间
    
    # 解析标注(用tf.py_function兼容Python代码)
    boxes, classes = tf.py_function(
        func=parse_voc_xml,
        inp=[xml_path],
        Tout=[tf.float32, tf.int32]
    )
    # 调整形状避免tf.data形状不匹配问题
    boxes = tf.reshape(boxes, (-1, 4))
    classes = tf.reshape(classes, (-1,))
    
    # 构建RetinaNet要求的标签格式
    labels = {
        'bounding_boxes': boxes,
        'classes': classes
    }
    return img, labels

3. 构建tf.data.Dataset

def create_dataset(img_dir, xml_dir, batch_size=8, target_size=(640, 640)):
    # 获取所有图像和对应标注的路径
    img_paths = [os.path.join(img_dir, fname) for fname in os.listdir(img_dir) if fname.endswith(('.jpg', '.png'))]
    xml_paths = [os.path.join(xml_dir, fname.replace('.jpg', '.xml').replace('.png', '.xml')) for fname in os.listdir(img_dir) if fname.endswith(('.jpg', '.png'))]
    
    # 构建数据集并优化加载流程
    dataset = tf.data.Dataset.from_tensor_slices((img_paths, xml_paths))
    dataset = dataset.map(
        lambda x, y: load_and_preprocess_data(x, y, target_size),
        num_parallel_calls=tf.data.AUTOTUNE
    )
    dataset = dataset.shuffle(buffer_size=100)
    dataset = dataset.batch(batch_size)
    dataset = dataset.prefetch(tf.data.AUTOTUNE)
    
    return dataset

# 生成训练和验证数据集
train_dataset = create_dataset('dataset/train/images', 'dataset/train/annotations')
val_dataset = create_dataset('dataset/val/images', 'dataset/val/annotations')

四、接入RetinaNet训练流程

拿到train_dataset和val_dataset后,直接替换教程中tfds.load返回的数据集即可开始训练:

import keras_cv

# 初始化RetinaNet(根据你的类别总数调整num_classes)
retinanet = keras_cv.models.RetinaNet(
    num_classes=2,
    backbone=keras_cv.models.ResNet50Backbone(),
    bounding_box_format='xyxy'  # 与我们的边界框格式保持一致
)

# 编译模型
retinanet.compile(
    classification_loss='focal',
    box_loss='smoothl1',
    optimizer=tf.keras.optimizers.SGD(learning_rate=0.01, momentum=0.9)
)

# 启动训练
retinanet.fit(
    train_dataset,
    validation_data=val_dataset,
    epochs=10
)

补充说明

  • 若使用YOLO格式标注:YOLO标注为[类别ID, x_center, y_center, width, height](相对坐标),需转换为xyxy绝对坐标,公式为:
    xmin = (x_center - width/2)*img_width
    ymin = (y_center - height/2)*img_height
    xmax = (x_center + width/2)*img_width
    ymax = (y_center + height/2)*img_height
  • 图像尺寸建议统一为640x640、512x512等RetinaNet常用尺寸,避免训练报错。
  • 大数据集可添加dataset.cache()缓存数据,进一步提升加载效率。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 10:05:31