自定义数据集适配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_widthymin = (y_center - height/2)*img_heightxmax = (x_center + width/2)*img_widthymax = (y_center + height/2)*img_height - 图像尺寸建议统一为640x640、512x512等RetinaNet常用尺寸,避免训练报错。
- 大数据集可添加
dataset.cache()缓存数据,进一步提升加载效率。
内容的提问来源于stack exchange,提问作者Joaquin
相关产品推荐
相关产品推荐

