如何修改Python脚本将XML中的多边形标注转为TFRecord格式?
修改XML转TFRecord脚本以支持多边形标注
要适配多边形标注并用于TensorFlow目标检测/分割模型训练,需要分两步修改脚本:提取XML中的多边形点并计算边界框,然后在TFRecord中添加多边形特征字段。以下是修改后的完整代码及说明:
1. 修改xml_to_csv函数:提取多边形点并计算边界框
该函数需要从XML中读取多边形坐标,计算出对应的边界框(Faster R-CNN仍依赖边界框生成候选框),并将多边形点序列化为字符串存入CSV:
import glob import pandas as pd import xml.etree.ElementTree as ET def xml_to_csv(path): xml_list = [] for xml_file in glob.glob(path + '/*.xml'): tree = ET.parse(xml_file) root = tree.getroot() img_width = int(root.find('size')[0].text) img_height = int(root.find('size')[1].text) filename = root.find('filename').text for member in root.findall('object'): class_name = member[0].text # 定位XML中的多边形标签(根据你的实际结构调整标签名) polygon = member.find('polygon') if not polygon: # 兼容原有边界框标注(可选) xmin = int(member[4][0].text) ymin = int(member[4][1].text) xmax = int(member[4][2].text) ymax = int(member[4][3].text) xs_str = f"{xmin},{xmax}" ys_str = f"{ymin},{ymax}" else: # 提取多边形点(适配<point x="..." y="..."/>格式) points = polygon.findall('point') xs = [] ys = [] for point in points: x = int(point.attrib['x']) y = int(point.attrib['y']) xs.append(str(x)) ys.append(str(y)) # 如果你的XML是<x1><y1><x2><y2>格式,替换为以下代码: # xs = [str(int(polygon.find(f'x{i}').text)) for i in range(1, len(polygon)//2 +1)] # ys = [str(int(polygon.find(f'y{i}').text)) for i in range(1, len(polygon)//2 +1)] xs_str = ','.join(xs) ys_str = ','.join(ys) # 从多边形点计算边界框 x_coords = list(map(int, xs)) y_coords = list(map(int, ys)) xmin = min(x_coords) xmax = max(x_coords) ymin = min(y_coords) ymax = max(y_coords) # 新增多边形点字段 value = ( filename, img_width, img_height, class_name, xmin, ymin, xmax, ymax, xs_str, ys_str ) xml_list.append(value) # 更新CSV列名,包含多边形点 column_name = ['filename', 'width', 'height', 'class', 'xmin', 'ymin', 'xmax', 'ymax', 'polygon_x', 'polygon_y'] xml_df = pd.DataFrame(xml_list, columns=column_name) return xml_df
2. 修改create_tf_example函数:添加多边形特征到TFRecord
该函数需要解析CSV中的多边形点,归一化后存入TFRecord,同时保留原有边界框字段(确保Faster R-CNN/Mask R-CNN兼容):
import tensorflow as tf import os import io from PIL import Image # 确保导入dataset_util(来自TensorFlow Object Detection API) from object_detection.utils import dataset_util def create_tf_example(group, path): with tf.gfile.GFile(os.path.join(path, '{}'.format(group.filename)), 'rb') as fid: encoded_jpg = fid.read() encoded_jpg_io = io.BytesIO(encoded_jpg) image = Image.open(encoded_jpg_io) width, height = image.size filename = group.filename.encode('utf8') image_format = b'jpg' xmins = [] xmaxs = [] ymins = [] ymaxs = [] classes_text = [] classes = [] # 初始化多边形点列表 polygon_xs = [] polygon_ys = [] for index, row in group.object.iterrows(): # 归一化边界框坐标 xmins.append(row['xmin'] / width) xmaxs.append(row['xmax'] / width) ymins.append(row['ymin'] / height) ymaxs.append(row['ymax'] / height) # 解析并归一化多边形点 xs = list(map(float, row['polygon_x'].split(','))) ys = list(map(float, row['polygon_y'].split(','))) normalized_xs = [x / width for x in xs] normalized_ys = [y / height for y in ys] polygon_xs.extend(normalized_xs) polygon_ys.extend(normalized_ys) # 处理类别 classes_text.append(row['class'].encode('utf8')) classes.append(class_text_to_int(row['class'])) tf_example = tf.train.Example(features=tf.train.Features(feature={ 'image/height': dataset_util.int64_feature(height), 'image/width': dataset_util.int64_feature(width), 'image/filename': dataset_util.bytes_feature(filename), 'image/source_id': dataset_util.bytes_feature(filename), 'image/encoded': dataset_util.bytes_feature(encoded_jpg), 'image/format': dataset_util.bytes_feature(image_format), # 保留边界框字段(Faster R-CNN必需) 'image/object/bbox/xmin': dataset_util.float_list_feature(xmins), 'image/object/bbox/xmax': dataset_util.float_list_feature(xmaxs), 'image/object/bbox/ymin': dataset_util.float_list_feature(ymins), 'image/object/bbox/ymax': dataset_util.float_list_feature(ymaxs), 'image/object/class/text': dataset_util.bytes_list_feature(classes_text), 'image/object/class/label': dataset_util.int64_list_feature(classes), # 添加多边形分割特征(适配Mask R-CNN等支持实例分割的模型) 'image/object/segmentation/polygon/x': dataset_util.float_list_feature(polygon_xs), 'image/object/segmentation/polygon/y': dataset_util.float_list_feature(polygon_ys), })) return tf_example
关键说明
- XML结构适配:根据你的数据集XML实际标签结构调整多边形提取逻辑,比如如果多边形用
<points>标签存储为空格分隔的字符串(如"10 20 30 40..."),需要修改解析代码。 - 模型兼容性:如果仅训练Faster R-CNN目标检测模型,多边形字段可选,但保留边界框字段即可;若训练Mask R-CNN等实例分割模型,必须添加多边形/掩码字段。
- 归一化:所有坐标需转换为相对于图像宽高的0-1浮点数,符合TensorFlow Object Detection API的要求。
- 语法修正:原脚本中
int(label_map)member[4][3].text存在语法错误,已修正为正确的边界框值提取逻辑。
内容的提问来源于stack exchange,提问作者Siddhi Powar
相关产品推荐
相关产品推荐

