简介:面向需要将XML数据转换为深度学习训练格式的开发者,这份压缩包提供了两个轻量级Python脚本,专门解决从XML到CSV、再到TFRecord的格式转换问题,适用于TensorFlow模型训练前的数据预处理环节,尤其适合需要批量处理标注数据的中小型团队。压缩包内共2个脚本文件,均为py类型,整体大小仅3KB,不依赖重型框架,便于直接拷贝到项目中使用。其中一个脚本负责解析XML、提取节点并利用pandas生成CSV中间格式,另一个脚本则读取清洗后的CSV,通过TensorFlow的Example协议缓冲区将其转换为TFRecord二进制文件,衔接起数据清洗与模型训练两大流程。已有250人学习下载,脚本结构简洁、命名清晰,能显著减少手工编码与格式对接的重复劳动,帮助工程师快速搭建可复用的数据管线,让后续的训练数据读取更加高效。
1. 拿到 scripts(xml-csv-tfrecord).rar:XML 标注转 TFRecord 的中间层不能省
做目标检测训练,数据往往来自标注工具或老项目导出的 Pascal VOC XML,一张图对应一个 annotation 文件,<object>节点一个接一个。经常有人想写个 for 循环直接把 XML 喂给 TensorFlow,结果发现 XML 是树状结构、图片字节还得单独读,预处理代码越写越长。这个压缩包 scripts(xml-csv-tfrecord).rar 里是两个 Python 脚本:xml_to_csv.py 和 generate_tfrecord.py,把「XML → CSV → TFRecord」这条链串起来。它适合手里有 VOC 格式标注、想转成 TensorFlow 目标检测训练格式的人,也适合想搞明白数据流水线每一步在干什么的人。重点不是脚本多复杂,而是它帮你留住了中间检查数据的机会。
2. 先看懂两个脚本的分工:从 XML 到 TFRecord 的数据流与依赖清单
2.1 xml_to_csv.py:把树状 XML 拍平成表格行
XML 标注文件长这样:
<annotation> <filename>000001.jpg</filename> <size> <width>1920</width> <height>1080</height> <depth>3</depth> </size> <object> <name>cat</name> <bndbox> <xmin>100</xmin> <ymin>80</ymin> <xmax>320</xmax> <ymax>240</ymax> </bndbox> </object> <object> <name>dog</name> <bndbox> <xmin>400</xmin> <ymin>200</ymin> <xmax>800</xmax> <ymax>600</ymax> </bndbox> </object> </annotation>filename 是图片文件名,size 下面是宽高,每个 object 是一个标注框。对训练脚本来说,树状结构不方便直接喂给 Dataset API,因为你还要一层层 findall、迭代、取值,每个样本再组合多个 object。常见做法是先拍平成表格:一个框占一行,每行固定 8 列——filename、width、height、class、xmin、ymin、xmax、ymax。一张图有两个框,CSV 里就是连续两行。
列顺序不是随便定的:filename 在最前面,后面依次是宽高、类别、四个坐标。generate_tfrecord.py 里 groupby('filename') 依赖这一列,错位会导致分组错乱。我之前见过有人把 class 放到第八列,结果转换脚本按索引取值时全部错位,生成的 record 类别乱成一锅粥。所以用这套脚本时不要随意改表头顺序。
CSV 这种中间格式的好处是,pandas 可以直接读、可以筛、可以合并,中间检查数据分布非常方便。比如想看有多少张图没有标注框,用df.groupby('filename').size()数一下行数就出来了,比在 XML 里数 object 节点快得多。xml_to_csv.py 做的事情就是这个:输入一个存放 XML 的目录,输出一个 CSV。
2.2 generate_tfrecord.py:把一行记录打包成一个样本
CSV 只是中间层,TensorFlow 训练时不会直接拿 CSV 的行来喂模型,因为 CSV 里存的是文本路径和坐标,图片内容还得靠脚本去磁盘读。更常见的是把图片字节和标注框全部封装进 TFRecord 文件,每个样本对应一个序列化后的 tf.train.Example 协议缓冲区。
generate_tfrecord.py 就是干这个的:它读 CSV,把同属一张图片的多行归到一组,读入图片字节,按 TensorFlow 目标检测 API 约定的字段名写入 image/encoded、image/height、image/width、image/object/bbox/xmin 等特征,最终写成一个 .record 文件。这个文件对 TensorFlow 来说才是真正的高效数据源。
这些字段名不是随便起的。TensorFlow Object Detection API 的 data loader 会按名字找 image/encoded、image/object/bbox/xmin 这些特征,如果你改成 image/boxes_xmin,读数据时找不到对应字段,训练直接报 KeyError。所以脚本里的字段名必须严格沿用约定,别为了一时省事换个更短的名字。
TFRecord 不是玄学,说白了就是一个带长度前缀的二进制序列,Dataset API 用 TFRecordDataset 读取时,会自动处理缓冲、并行、预取。相比训练时每次从 CSV 现拼数据,TFRecord 避免了重复解析 XML 和重复组合字段的开销,而且单文件方便拷贝管理。几千张图的数据集,手工转换一次,后面训练阶段每秒读进来的样本量会稳定很多。
2.3 完整数据流与 Python 环境
两个脚本串起来的标准流程是:先跑 xml_to_csv.py,把所有 XML 转成 train.csv;中间用 pandas 或 Excel 打开检查有没有空行、标错类别;确认没问题后跑 generate_tfrecord.py,把 CSV 连同图片目录转成 train.record;最后在 TensorFlow 的 model_main_tf2.py 里指向这个 record 文件开始训练。如果要把数据分成 train/val,那就分别为两个目录执行一遍同样的操作,或者一次性转成一个大 CSV 再切分。
依赖环境如下表,装好之后基本不用额外配置:
| 依赖 | 用途 | 常见安装方式 |
|---|---|---|
| Python 3.7+ | 运行脚本 | Anaconda 或系统 Python |
| xml.etree.ElementTree | 解析 XML,Python 自带 | 无需安装 |
| pandas | 生成与合并 CSV | pip install pandas |
| tensorflow | 写 TFRecord、训练读取 | pip install tensorflow |
| object_detection utils(可选) | 复用 bytes_feature 等函数 | 下载 TF Models 仓库 |
我的习惯是先建一个干净的虚拟环境,再装 pandas 和 tensorflow,然后用 pip list 确认版本。装好后可以先跑两行环境检查:
python -c "import pandas; print(pandas.__version__)" python -c "import tensorflow as tf; print(tf.__version__)"第一条确认 pandas 可用,第二条确认 tensorflow 版本。tensorflow 版本差异最容易出问题,旧脚本里很多 tf.train.Example 的写法在 TF 2.x 依然保留,但如果用了 tf.contrib 或旧版 object_detection 工具类,需要提前改掉。这些在第 4 章里会提到具体改动点。
3. xml_to_csv.py 实操:解析节点、提取 bbox 与 CSV 路径的坑
3.1 用 ElementTree 遍历标注树:取什么、丢什么
xml_to_csv.py 的核心逻辑很直白,我一般会这样拆解:
import os import glob import pandas as pd import xml.etree.ElementTree as ET def xml_to_csv(xml_dir): xml_list = [] for xml_file in glob.glob(os.path.join(xml_dir, '*.xml')): tree = ET.parse(xml_file) root = tree.getroot() for member in root.findall('object'): value = ( root.find('filename').text, int(root.find('size/width').text), int(root.find('size/height').text), member.find('name').text, int(member.find('bndbox/xmin').text), int(member.find('bndbox/ymin').text), int(member.find('bndbox/xmax').text), int(member.find('bndbox/ymax').text), ) xml_list.append(value) column_name = ['filename', 'width', 'height', 'class', 'xmin', 'ymin', 'xmax', 'ymax'] df = pd.DataFrame(xml_list, columns=column_name) return df if __name__ == '__main__': df = xml_to_csv('annotations') df.to_csv('annotations.csv', index=False) print(df.head())运行逻辑是:glob 按通配符收集目录下所有 XML,ET.parse 把每个 XML 解析成 ElementTree,root.findall('object') 定位到所有标注框节点。每找到一个 object,就取出文件名、图片宽高、类别名和 bndbox 里的四个坐标,作为一个元组追加到列表。最后把整个列表套进 DataFrame,列名固定成上面 8 个,输出 CSV。
两个参数值得注意。第一是 xml_dir,传目录路径而不是单个文件,因为工具按目录批量处理;如果只有一个 XML 想试跑,可以先复制成一个单文件目录。第二是 filename 和坐标都用 int() 强转,这是为了后面 generate_tfrecord.py 能直接按数字处理。如果 XML 文本里带了缩进或换行,int() 会自动忽略前后的空白,但如果某个坐标被写成了小数比如 12.5,这里直接抛 ValueError。遇到这种标注要先修数据,不要跳过,不然坑会留到后面。
3.2 处理非 VOC 自定义 XML:命名空间与缺失节点的兼容
实际项目里很少拿到标准干净的 VOC 版本。最常见问题是 XML 根节点带了命名空间,导致findall('object')返回空列表,转出来的 CSV 只有表头没有数据。原因很简单:ElementTree 认为带命名空间的标签名是{http://...}object而不是object,直接查当然查不到。
我处理这种文件时,会先做一次去命名空间,把根节点下所有标签里的{命名空间}前缀剥掉再查:
for elem in root.iter(): if elem.tag.startswith('{'): elem.tag = elem.tag.split('}', 1)[1]这段代码放在 ET.parse 之后、findall 之前。做完后,原来的 findall('object') 才能命中。操作只影响内存中的树,不会改写磁盘上的 XML 文件。
另外,如果 XML 里有多个<name>或某个 object 缺 bndbox 子节点,int(None)会报错。稳妥做法是先判断子节点是否存在:
bbox = member.find('bndbox') if bbox is None: continue这个 continue 表示遇到残缺标注直接跳过该框,而不是让整个脚本崩掉。当然,跳过的框意味着标注数据缺失,事后要统计跳过数量,不能默默吞掉。我一般会在脚本里加一个计数器,最后打印出来:skip_count大于 0 就要回头查原始数据。
3.3 多个 XML 目录与合并 CSV:生成训练/测试两套 record
目标检测训练通常要分训练集、验证集。常见做法是维护 train_xml 和 val_xml 两个目录,分别跑一次 xml_to_csv.py,得到两个 CSV。这样分隔清晰,后面 generate_tfrecord.py 也按两套 CSV 各生成一个 record。
也有时候标完的数据只有一个大目录,train/val 的划分想放在 CSV 阶段做。这时可以先全部转成一个 all.csv,再用 pandas 按文件名哈希或随机数切分。我不太推荐随机切分后重写两个 CSV,而是直接做合并处理——比如标注工具分批导出,每个批次一个 CSV,想拼成一个:
import pandas as pd def combine_csvs(csv_paths, output_path): df_list = [pd.read_csv(p) for p in csv_paths] df = pd.concat(df_list, ignore_index=True) df.to_csv(output_path, index=False)pd.concat 时 ignore_index=True 是为了让合并后的行索引重新编号,否则后续按 filename groupby 会带出一堆旧索引,看着碍事。合并之后最好检查一下类别分布:df['class'].value_counts()。如果某个类别只有几条,训练时很可能因为样本太少导致这一个类别学不出来。
4. generate_tfrecord.py 实操:把 CSV 编码成 TFRecord Feature 的调用要点
4.1 TFRecord 与 tf.train.Example 的关系:为什么不直接写 JSON
TFRecord 文件里存的是一段段序列化后的二进制消息,每段消息是一个 tf.train.Example。Example 里所有字段都放进 features,features 是一个 map,key 是字符串特征名,value 是 BytesList、FloatList 或 Int64List 中的一种。看起来复杂,但实际就是为了让 TensorFlow 能够快速按名称取字段,不必解析整个文本文件。
为什么不直接写 JSON?因为 JSON 对每行都要做字符串解析,属性名重复冗余,数字和字符串混在一起,读取时还要按需转类型。TFRecord 是二进制编码,配合 tf.data 的并行读卡效率更高,尤其数据量到几千上万张图时差别很明显。这也是为什么 generate_tfrecord.py 要把 CSV 转成 TFRecord,而不是让训练脚本直接去吃 CSV。
目标检测最常用的字段约定来自 TensorFlow Models 仓库的 object_detection 接口。一张图对应一个 Example,里面至少有 encoded 图片字节、宽高、以及若干组归一化后的坐标框和类别标签。坐标要归一化到 0~1,这是训练的硬性要求。
4.2 逐行转换脚本的逻辑拆解
完整脚本的核心函数长这样:
import os import tensorflow as tf import pandas as pd from object_detection.utils import dataset_util VOC_TO_ID = {'cat': 1, 'dog': 2} def create_tf_example(group, image_path): with tf.io.gfile.GFile(image_path, 'rb') as fid: encoded_jpg = fid.read() width = int(group.iloc[0]['width']) height = int(group.iloc[0]['height']) xmins, ymins, xmaxs, ymaxs = [], [], [], [] classes_text, classes = [], [] for row in group.itertuples(): xmins.append(float(row.xmin) / width) ymins.append(float(row.ymin) / height) xmaxs.append(float(row.xmax) / width) ymaxs.append(float(row.ymax) / height) classes_text.append(row._3.encode('utf8')) classes.append(VOC_TO_ID[row._3]) tf_example = tf.train.Example(features=tf.train.Features(feature={ 'image/encoded': dataset_util.bytes_feature(encoded_jpg), 'image/height': dataset_util.int64_feature(height), 'image/width': dataset_util.int64_feature(width), 'image/object/bbox/xmin': dataset_util.float_list_feature(xmins), 'image/object/bbox/ymin': dataset_util.float_list_feature(ymins), 'image/object/bbox/xmax': dataset_util.float_list_feature(xmaxs), '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), })) return tf_example这段代码的逻辑是:group 是 CSV 里同一个 filename 对应的全部行,也就是一张图片的所有标注框。先从第一行取图片宽高,然后遍历 group 里每一行,把 bbox 的 xmin 除以 width、ymin 除以 height 归一化。类别同时写入 text 和 label 两个特征,text 是字符串用于可视化显示,label 是整数 id 用于 loss 计算。最后用 dataset_util 的辅助函数把 Python 列表包成 tf.train 对应的 Feature 类型。
这里我一般会用 object_detection.utils.dataset_util,它是 TensorFlow Models 仓库提供的辅助工具,省去手写 bytes_feature 的细节。如果不想依赖整个 object_detection 包,也可以自己实现:
def bytes_feature(value): return tf.train.Feature(bytes_list=tf.train.BytesList(value=[value]))这是一个很常用的替代写法。注意不要直接用 int64_feature 存坐标,坐标是浮点,存成整数会丢掉小数部分,训练时边界框会偏差很大。
然后看主循环:
def csv_to_tfrecord(csv_input, image_dir, output_path): writer = tf.io.TFRecordWriter(output_path) df = pd.read_csv(csv_input) grouped = df.groupby('filename') for filename, group in grouped: image_path = os.path.join(image_dir, filename) tf_example = create_tf_example(group, image_path) writer.write(tf_example.SerializeToString()) writer.close() print(f'Done. Total examples: {len(grouped)}')主循环的要点是按 filename 分组,每组写一个 Example。如果一个 filename 只出现一行,那就是单框图;出现多行就是多框图,全部框都固定进同一个 Example,而不是每个框单独写一个 Example。很多新手在这里翻车:不分组,遍历每一行都写一个 Example,最后一张图被拆成几个独立样本,训练时模型把同一个物体的几个框当成不同图片,坐标完全错乱。
writer.write 接收的是序列化后的字节,所以 create_tf_example 返回的 Example 要调用 SerializeToString()。TFRecordWriter 不需要手动 flush,close 时会自动写完剩余缓存。
4.3 运行命令与 label_map 的对应关系
把转换脚本接起来,标准命令是这样:
python xml_to_csv.py --xml_dir annotations --csv_path train.csv python generate_tfrecord.py --csv_path train.csv --image_dir images --output_path train.record如果你的脚本用 argparse 解析参数,通常会这样写 main:
if __name__ == '__main__': import argparse parser = argparse.ArgumentParser() parser.add_argument('--csv_path', required=True) parser.add_argument('--image_dir', required=True) parser.add_argument('--output_path', required=True) args = parser.parse_args() csv_to_tfrecord(args.csv_path, args.image_dir, args.output_path)命令行里的 image_dir 会跟 CSV 里的 filename 拼接,所以前面 xml_to_csv.py 生成的 CSV 里 filename 最好不要带子目录前缀,否则可能出现 images/train/img1.jpg 这种重复路径。如果 filename 本身带了路径,就在 xml_to_csv.py 阶段统一用os.path.basename()清理干净,这是第 5 章要展开的坑之一。
label_map 是另一个容易踩的点。TensorFlow 目标检测训练时,pipeline 配置文件里会指定 label_map_path,里面的 id 必须和 TFRecord 里 image/object/class/label 的整数一一对应。常见做法是让 generate_tfrecord.py 脚本里维护一个和 label_map 一致的字典:
item { id: 1 name: 'cat' } item { id: 2 name: 'dog' }如果脚本里字典和 label_map 不一致,比如脚本里 cat 是 1,pipeline 里 cat 是 2,训练时模型不会立刻报错,但 loss 和 mAP 永远对不上。我一般会在生成 TFRecord 后立刻跑一遍读回脚本,把 label 打出来核对。这个验证步骤放在第 6 章详细说。
5. 避坑与排查:XML 转 CSV、再转 TFRecord 的 5 个高频问题
5.1 filename 带路径或大小写不一致:图片读不到
现象:generate_tfrecord.py 运行时报NotFoundError: ...; No such file or directory,或者 TFRecord 生成后图片张数明显少于 CSV 里的 filename 数。
原因:xml_to_csv.py 直接取了 XML 里的<filename>原始文本。它可能是image/000001.jpg这样的相对路径,也可能是D:\datasets\000001.jpg这样的绝对路径;Windows 下还有反斜杠的问题。CSV 交给 generate_tfrecord.py 后,代码用os.path.join(image_dir, filename)拼接,于是变成了images/image/000001.jpg,自然找不到。
解决:在 xml_to_csv.py 里做一次归一化:
import os filename = os.path.basename(root.find('filename').text)这样不管原路径多长,CSV 里只保留纯文件名。如果你的图片目录里存在重名文件,那就要另外加一层结构,而不是用 basename 一刀切。另外注意大小写:img001.JPG和img001.jpg在 Linux 下是两个文件,生成 TFRecord 前先用 os.path.exists 检查一遍。
这段代码我一般会放在 3.1 节那个元组构造位置之前,先统一 filename,再进入后续取值逻辑。
5.2 XML 带命名空间导致找不到 object
现象:xml_to_csv.py 跑完,CSV 里除了表头一行都没有,但 XML 文件用文本编辑器打开明明有<object>标签。
原因:XML 根节点带有类似<annotation xmlns="...">的命名空间声明,ElementTree 会把 object 解析成{http://...}object。findall('object') 找不到这种带前缀的节点,返回空列表。
解决:按第 3.2 节的方法,在 findall 之前先剥掉命名空间前缀。我遇到这个问题时还会顺手打印一下 root.tag,看到类似{http://...}annotation就直接确认是命名空间问题。另外一种隐蔽情况是 XML 里根本没有 object 节点,而文件本身还存在,那就需要统计一下空标注文件的占比,别让脚本静默跳过。可以在循环里加个计数器:
if len(root.findall('object')) == 0: print(f'No object: {xml_file}')5.3 bbox 坐标读到 NaN 或字符串,CSV 类型不对
现象:转换时int(member.find('bndbox/xmin').text)抛TypeError,或者生成的 CSV 里 xmin 列是空值。更隐蔽的情况是 CSV 能生成,但 generate_tfrecord.py 里float(row.xmin)时报错。
原因:标注文件里某个框缺少 xmin 子节点,或者坐标值写成了12.0这种浮点文本。int('12.0')会抛 ValueError,int(None)会抛 TypeError。还有一个常见因素:XML 编辑器保存时在文本前后加了不可见字符,虽然 int() 能容忍空白,但 NaN 字符串不能。
解决:解析时加防御:
xmin_text = member.find('bndbox/xmin') if xmin_text is None: continue try: xmin = int(xmin_text.text) except ValueError: continue跳过之后要记录 filename 和缺失字段,最后打印出来人工核对。坐标是整数还是浮点要看标注工具:VOC 标准要求整数像素坐标,但有些工具导出的是带小数的归一化坐标,遇到这种就得区分处理,不要把小数直接 int() 截断。
5.4 坐标没归一化或归零
现象:TFRecord 生成成功,读回 bbox 全是 0,或者训练时 loss 不下降。
原因:常见是两处。一是 generate_tfrecord.py 里忘了除以 width/height,直接把像素坐标写进 float feature。TensorFlow Object Detection API 的默认数据增强和 loss 计算都假设坐标在 0~1 之间,绝对像素坐标进去会让 loss 变得巨大。二是用了 int64_feature 存坐标,浮点被截断成整数,0.123 变成 0,读回自然全是 0 或 1。
解决:严格按照 4.2 节的方式,xmin 除以 width,ymin 除以 height,并且坐标字段必须用 float_list_feature。生成后读回时注意观察数值范围:如果读回坐标全部小于 1 且分布合理,说明归一化正确。我还会顺手检查一组框的关系:xmax 是否大于 xmin、ymax 是否大于 ymin,如果出现颠倒,说明原始标注本身有问题。
5.5 多个 class 映射 id 错位
现象:训练能跑,但验证集的 mAP 在 0 附近震荡,打印预测框发现猫框上写的是 dog 标签。
原因:generate_tfrecord.py 里 VOC_TO_ID 字典和 label_map.pbtxt 不一致,或者两个目录(train/val)分别用了不同脚本版本。比如 val 的脚本里 dog 是 2,train 的脚本里 dog 是 3,同一类别的 label id 在两个 record 里不一样,模型训练时学到的类别语义就乱了。
解决:把 label_map 和字典都集中到一个统一的配置文件里,生成 train.record 和 val.record 时显式传入同一个 label_map 文件,并在脚本里读 label_map 自动构建字典,而不是手动写死 VOC_TO_ID。这样手动维护一份即可。生成完成后,进入第 6 章的读回验证,按类打印 text 字段和 integer label 字段。
6. 验证 TFRecord 没转坏:读回样本与标注框坐标复核
6.1 用 TFRecordDataset 读回样本
生成 TFRecord 只是第一步,真正坑的是转完之后没人检查。我有一个固定的验收动作:立刻用 tf.data 的 TFRecordDataset 读回几条记录,打印关键字段。脚本如下:
import tensorflow as tf def inspect_tfrecord(record_path, num_samples=3): dataset = tf.data.TFRecordDataset(record_path) for raw in dataset.take(num_samples): example = tf.train.Example() example.ParseFromString(raw.numpy()) f = example.features.feature print('height:', f['image/height'].int64_list.value) print('width:', f['image/width'].int64_list.value) print('class_text:', f['image/object/class/text'].bytes_list.value) print('xmin:', f['image/object/bbox/xmin'].float_list.value) print('ymax:', f['image/object/bbox/ymax'].float_list.value)这段代码从 record 文件里取 3 个样本,解析每个 Example 的 features,然后把宽高、类别文本、xmin 和 ymax 打出来。我建议检查四件事:宽高是否和原图一致;class_text 里的类别是否都在预期集合内;xmin/ymax 是否都在 0~1 之间;ymax 是否大于 ymin、xmax 是否大于 xmin。只要这四条通过,TFRecord 基本没转坏。
6.2 一张图确认标注框:和 bounding box 可视化交叉验证
读回数值只是逻辑上的验证,坐标有没有整体偏移、标注对象有没有张冠李戴,最好用一张图画出来肉眼确认。我会从 CSV 里随机挑一张图,读取同名的 TFRecord 样本,把归一化坐标还原成像素坐标,然后画矩形框:
from PIL import Image, ImageDraw import tensorflow as tf def draw_boxes_from_record(record_path, image_dir, filename, output='check.jpg'): dataset = tf.data.TFRecordDataset(record_path) for raw in dataset: example = tf.train.Example() example.ParseFromString(raw.numpy()) f = example.features.feature text_list = f['image/object/class/text'].bytes_list.value if not text_list: continue img = Image.open(f'{image_dir}/{filename}') draw = ImageDraw.Draw(img) h, w = img.size[1], img.size[0] for i, text in enumerate(text_list): xmin = f['image/object/bbox/xmin'].float_list.value[i] * w ymin = f['image/object/bbox/ymin'].float_list.value[i] * h xmax = f['image/object/bbox/xmax'].float_list.value[i] * w ymax = f['image/object/bbox/ymax'].float_list.value[i] * h draw.rectangle([xmin, ymin, xmax, ymax], outline='red') draw.text((xmin, ymin), text.decode()) img.save(output) break注意这段代码里我用 text_list 是否非空来判断样本,实际项目里我会直接从 CSV 侧找到 filename 对应的样本,再去 TFRecord 里定位索引。更稳的做法是在 generate_tfrecord.py 里为每个样本额外写一个 image/source_id 特征,验证时按 source_id 索引,而不是靠文件名匹配。
我第一次跑通这条流程时,就是信任脚本直接拿去训练,结果跑了两个 epoch 才发现 val 的类别 id 和 train 对不上,整个模型白训练。从那以后我每次转完 TFRecord 都强制走一遍这个验收动作,读回样本 + 画框抽查两件事做完才敢开始训练。希望帮到你。
本文还有配套的精品资源,点击获取