Keras CV RetinaNet自定义数据加载:类别与Bounding Box张量转换问题
Keras CV RetinaNet自定义数据集加载修复方案
问题说明
我正在基于RetinaNet的目标检测教程加载自定义数据,教程要求Keras CV的Bounding Box必须封装为指定格式的字典:
Keras CV对边界框有预定义规范,需将边界框封装为如下格式的字典:
bounding_boxes = { # num_boxes 可以是不规则维度 'boxes': Tensor(shape=[batch, num_boxes, 4]), 'classes': Tensor(shape=[batch, num_boxes])}
我的数据包含图片和存储边界框坐标的XML文件,已经把边界框转成了教程推荐的xywh格式,但现在有两个核心问题:
- 如何为每个边界框关联对应的类别
- 如何将这个字典转换为有效的Tensor
原代码
import tensorflow as tf import xml.etree.ElementTree as et import os import numpy as np img_path = '/home/joaquin/TFM/Doom_KerasCV/IA_training_data_reduced_640/' img_list = [] xml_list = [] box_list = [] box_dict = {} img_norm = [] def list_creation (img_path): for subdir, dirs, files in os.walk(img_path): for file in files: if file.endswith('.png'): img_list.append(subdir+"/"+file) img_list.sort() if file.endswith('.xml'): xml_list.append(subdir+"/"+file) xml_list.sort() return img_list, xml_list def box_extraction (xml_list): for element in xml_list: root = et.parse(element) boxes = list() for box in root.findall('.//object'): label = box.find('name').text xmin = int(box.find('./bndbox/xmin').text) ymin = int(box.find('./bndbox/ymin').text) xmax = int(box.find('./bndbox/xmax').text) ymax = int(box.find('./bndbox/ymax').text) width = xmax - xmin height = ymax - ymin data = np.array([xmin,ymax,width,height]) box_dict = {'boxes':data,'classes':label} # boxes.append(data) box_list.append(box_dict) return box_list list_creation(img_path) boxes_dataset = tf.data.Dataset.from_tensor_slices(box_extraction(xml_list)) def loader (img_list): for image in img_list: img = tf.keras.utils.load_img(image) # loads the image # Normalizamos los pixeles de la imagen entre 0 y 1: img = tf.image.per_image_standardization(img) img = tf.keras.utils.img_to_array(img) # converts the image to numpy array img_norm.append(img) return img_norm img_dataset = tf.data.Dataset.from_tensor_slices(loader(img_list)) dataset = tf.data.Dataset.zip((img_dataset, boxes_dataset)) def get_dataset_partitions_tf(ds, ds_size, train_split=0.8, val_split=0.1, test_split=0.1, shuffle=True, shuffle_size=10): assert (train_split + test_split + val_split) == 1 if shuffle: ds = ds.shuffle(shuffle_size, seed=12) train_size = int(train_split * ds_size) val_size = int(val_split * ds_size) train_ds = ds.take(train_size) val_ds = ds.skip(train_size).take(val_size) test_ds = ds.skip(train_size).skip(val_size) return train_ds, val_ds, test_ds train,validation,test = get_dataset_partitions_tf(dataset, len(dataset))
修复后的代码及关键说明
import tensorflow as tf import xml.etree.ElementTree as et import os import numpy as np img_path = '/home/joaquin/TFM/Doom_KerasCV/IA_training_data_reduced_640/' # 1. 收集所有类别,建立文本到整数的映射 def collect_classes(xml_paths): classes = set() for xml_path in xml_paths: root = et.parse(xml_path) for obj in root.findall('.//object'): classes.add(obj.find('name').text) return {cls: idx for idx, cls in enumerate(sorted(classes))} # 2. 生成图片-标注对的列表(每张图片对应一组边界框和类别) def create_img_anno_pairs(img_path): img_list = [] xml_list = [] for subdir, _, files in os.walk(img_path): for file in files: if file.endswith('.png'): img_list.append(os.path.join(subdir, file)) elif file.endswith('.xml'): xml_list.append(os.path.join(subdir, file)) # 确保图片和XML按文件名匹配(假设命名一致) img_list.sort() xml_list.sort() class_map = collect_classes(xml_list) img_anno_pairs = [] for img_path, xml_path in zip(img_list, xml_list): # 加载并预处理图片 img = tf.keras.utils.load_img(img_path) img = tf.image.per_image_standardization(img) img = tf.keras.utils.img_to_array(img) # 解析XML获取边界框和类别 root = et.parse(xml_path) boxes = [] classes = [] for obj in root.findall('.//object'): label = obj.find('name').text xmin = int(obj.find('./bndbox/xmin').text) ymin = int(obj.find('./bndbox/ymin').text) xmax = int(obj.find('./bndbox/xmax').text) ymax = int(obj.find('./bndbox/ymax').text) width = xmax - xmin height = ymax - ymin # 注意:Keras CV的xywh格式是(x, y, width, height),此处修正原代码y坐标错误(原用ymax,应为ymin) boxes.append([xmin, ymin, width, height]) classes.append(class_map[label]) # 封装成Keras CV要求的字典格式 bounding_boxes = { 'boxes': tf.convert_to_tensor(boxes, dtype=tf.float32), 'classes': tf.convert_to_tensor(classes, dtype=tf.float32) } img_anno_pairs.append((img, bounding_boxes)) return img_anno_pairs, len(img_anno_pairs) # 3. 创建数据集(处理不规则维度) img_anno_pairs, dataset_size = create_img_anno_pairs(img_path) dataset = tf.data.Dataset.from_generator( lambda: (pair for pair in img_anno_pairs), output_signature=( tf.TensorSpec(shape=(None, None, 3), dtype=tf.float32), { 'boxes': tf.TensorSpec(shape=(None, 4), dtype=tf.float32), 'classes': tf.TensorSpec(shape=(None,), dtype=tf.float32) } ) ) # 4. 划分数据集 def get_dataset_partitions_tf(ds, ds_size, train_split=0.8, val_split=0.1, test_split=0.1, shuffle=True, shuffle_size=1000): assert (train_split + test_split + val_split) == 1 if shuffle: ds = ds.shuffle(shuffle_size, seed=12) train_size = int(train_split * ds_size) val_size = int(val_split * ds_size) train_ds = ds.take(train_size) val_ds = ds.skip(train_size).take(val_size) test_ds = ds.skip(train_size + val_size) return train_ds, val_ds, test_ds train, validation, test = get_dataset_partitions_tf(dataset, dataset_size)
关键修改点
- 类别编码:将文本类别转换为整数,符合模型训练的数值输入要求
- 标注结构修正:每张图片对应一个字典,包含该图片所有边界框和类别,而非每个边界框单独生成字典
- 数据集创建:使用
from_generator处理不规则维度(不同图片边界框数量不同),指定输出签名确保Tensor格式有效 - 边界框格式校验:修正原代码中y坐标的错误(原代码用了ymax,正确xywh的y应为ymin)
- 代码结构优化:合并图片加载和标注解析逻辑,确保图片与标注一一对应
内容的提问来源于stack exchange,提问作者Joaquin
相关产品推荐
相关产品推荐

