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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 18:40:03