如何将自有目标检测数据集转换为TFDS专用COCO格式TFRecords
为Keras官方Retinanet实现生成匹配格式TFRecords的方法
完全可以手动生成符合项目输入要求的TFRecords文件,无需依赖tfds.load()加载公开数据集,具体操作流程如下:
1. 对齐目标数据结构
首先明确tfds加载的COCO数据集和通用COCO标注的结构差异:
生成自有数据集TFRecords时,单样本必须严格匹配以下字段定义、数据类型和维度:
image:JPEG编码后的图像字节串,shape为(),类型tf.stringimage/filename:图像原始文件名字符串,shape为(),类型tf.stringimage/id:图像全局唯一ID,shape为(),类型tf.int64image/height:图像原始像素高度,shape为(),类型tf.int64image/width:图像原始像素宽度,shape为(),类型tf.int64objects/bbox:归一化边界框坐标,格式为[ymin, xmin, ymax, xmax],所有值除以对应图像宽高映射到0-1区间,shape为[检测框数量, 4],类型tf.float32objects/label:检测框对应类别ID,shape为[检测框数量],类型tf.int64,注意类别ID从1开始计数,0默认保留给背景类objects/area:检测框面积,shape为[检测框数量],类型tf.int64objects/is_crowd:是否为拥挤样本标注,shape为[检测框数量],类型tf.bool,自有数据无该标注时可全部填False
2. TFRecords生成核心代码
遍历自有数据集的所有图像和标注,按照上述字段定义序列化后写入文件,核心实现参考:
import tensorflow as tf import cv2 import os def _bytes_feature(value): if isinstance(value, type(tf.constant(0))): value = value.numpy() return tf.train.Feature(bytes_list=tf.train.BytesList(value=[value])) def _int64_feature(value): return tf.train.Feature(int64_list=tf.train.Int64List(value=[value])) def _float_feature_list(value): return tf.train.Feature(float_list=tf.train.FloatList(value=value)) def serialize_sample(image_path, img_id, height, width, bboxes_px, labels, areas, is_crowd): # 读取并编码图像 img = cv2.imread(image_path) img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) _, img_encode = cv2.imencode('.jpg', img) img_bytes = img_encode.tobytes() # 转换bbox为归一化格式并展平 bbox_norm = [] for (ymin, xmin, ymax, xmax) in bboxes_px: bbox_norm.extend([ymin/height, xmin/width, ymax/height, xmax/width]) feature_map = { 'image': _bytes_feature(img_bytes), 'image/filename': _bytes_feature(os.path.basename(image_path).encode('utf-8')), 'image/id': _int64_feature(img_id), 'image/height': _int64_feature(height), 'image/width': _int64_feature(width), 'objects/bbox': _float_feature_list(bbox_norm), 'objects/label': tf.train.Feature(int64_list=tf.train.Int64List(value=labels)), 'objects/area': tf.train.Feature(int64_list=tf.train.Int64List(value=areas)), 'objects/is_crowd': tf.train.Feature(int64_list=tf.train.Int64List(value=is_crowd)) } return tf.train.Example(features=tf.train.Features(feature=feature_map)).SerializeToString() def write_tfrecord(sample_list, save_path): with tf.io.TFRecordWriter(save_path) as writer: for idx, sample in enumerate(sample_list): # 替换为自有数据集的标注读取逻辑,拿到每个样本的对应字段 serialized = serialize_sample( image_path=sample['path'], img_id=idx, height=sample['height'], width=sample['width'], bboxes_px=sample['bboxes'], # 输入为像素坐标下的[ymin,xmin,ymax,xmax]格式 labels=sample['labels'], areas=sample['areas'], is_crowd=sample.get('is_crowd', [0]*len(sample['labels'])) ) writer.write(serialized)
3. 加载生成的TFRecords
生成文件后直接用tf.data.TFRecordDataset读取,不要通过tfds加载,解析逻辑必须和写入时的字段定义完全对齐:
def parse_tfrecord(proto): feature_desc = { 'image': tf.io.FixedLenFeature([], tf.string), 'image/filename': tf.io.FixedLenFeature([], tf.string), 'image/id': tf.io.FixedLenFeature([], tf.int64), 'image/height': tf.io.FixedLenFeature([], tf.int64), 'image/width': tf.io.FixedLenFeature([], tf.int64), 'objects/bbox': tf.io.VarLenFeature(tf.float32), 'objects/label': tf.io.VarLenFeature(tf.int64), 'objects/area': tf.io.VarLenFeature(tf.int64), 'objects/is_crowd': tf.io.VarLenFeature(tf.int64), } sample = tf.io.parse_single_example(proto, feature_desc) # 解码图像、重构维度 sample['image'] = tf.io.decode_jpeg(sample['image'], channels=3) sample['objects/bbox'] = tf.reshape(tf.sparse.to_dense(sample['objects/bbox']), [-1, 4]) sample['objects/label'] = tf.sparse.to_dense(sample['objects/label']) sample['objects/area'] = tf.sparse.to_dense(sample['objects/area']) sample['objects/is_crowd'] = tf.cast(tf.sparse.to_dense(sample['objects/is_crowd']), tf.bool) # 输出格式对齐原Retinanet代码的输入要求即可 return sample # 数据集加载示例 train_ds = tf.data.TFRecordDataset("train.tfrecord").map(parse_tfrecord, num_parallel_calls=tf.data.AUTOTUNE) val_ds = tf.data.TFRecordDataset("val.tfrecord").map(parse_tfrecord, num_parallel_calls=tf.data.AUTOTUNE)
注意事项
- 边界框格式不要混淆:通用COCO标注默认是
[xmin, ymin, w, h]像素坐标,必须转换为归一化的[ymin, xmin, ymax, xmax]格式再写入,否则会导致IoU计算错误,模型完全无法收敛 - 变长检测框字段必须用
VarLenFeature解析:单张图像的检测框数量不固定,bbox、label、area、is_crowd字段不能使用FixedLenFeature,否则会触发解析报错 - 类别ID不要从0开始编号:Retinanet实现中0默认对应背景类,自有数据集的类别ID请从1开始顺延编号
内容的提问来源于stack exchange,提问作者Allan Santos
相关产品推荐
相关产品推荐

