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

如何修改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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.19 12:19:51