ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

水表识别双网络实战:TensorFlow定位与读数模型训练指南

水表识别双网络实战:TensorFlow定位与读数模型训练指南 简介本资源是一套基于深度学习的水表识别项目源码面向计算机视觉学习者与仪表识别方向的开发者采用“定位网络识别网络”两阶段方案解决水表读数自动提取问题。包内共55个文件以18个Python脚本、13个pyc编译文件、14张jpg样本图及少量xml、iml配置为主压缩包约144KB涵盖数据预处理、模型定义、训练与测试等模块目录按base、utils、summaries等分层组织结构清晰便于二次开发。项目包含WM_config、D_config等配置脚本与wm_train、d_train等训练入口配合image_preprocess等预处理工具可帮助读者理解从表盘定位到字符识别的完整流程并据此复现实验、调整网络结构或迁移到其他仪表识别场景。目前已有128人学习下载适合具备一定深度学习基础、希望积累视觉项目实战经验的人群参考。1. 水表识别为什么要拆成两个网络从一张表盘图说起水表读数识别这件事很多人第一反应是「上一个 OCR 不就完了」。真上手拍几张照片就会发现水表图像里既有指针表盘又有数字滚轮表盘区域在整张图里占比很小背景还经常是井盖、泥水、反光。直接把整图丢给识别网络模型大部分算力都浪费在找表盘上精度也上不去。这个项目WaterMeter-master的思路就是把任务拆成两步先用一个定位网络把表盘区域框出来再用一个识别网络对框出来的区域做读数识别。定位网络负责「表在哪」识别网络负责「读数是多少」两个网络各司其职。这套代码基于 TensorFlow 实现目录里能看到wm_model.py、wm_trainer.py、wm_train.py这一组是识别网络WM 即 Water Meterd_model.py、d_trainer.py、d_train.py这一组是定位网络D 即 Detection。数据放在newdatares下文件名里带着坐标标注比如1271_0706140506449)(3,33 6,69 138,27 139,62 )这种括号里的四组数字就是表盘四角坐标。适合谁看做过一点深度学习、想找一个完整双网络落地案例的工程师或者手头有仪表识别需求、想参考工程结构的人。不适合完全没碰过 TensorFlow 的纯新手直接照搬环境配置那关会卡住。2. 定位网络与识别网络的分工数据标注格式与模型结构2.1 从文件名解析标注坐标是怎么存的这个项目最特别的地方是标注直接写在文件名里没有单独的标注文件。看一个真实样本1271_0706140506449)(3,33 6,69 138,27 139,62 )[3,33 6,69 138,27 139,62 ].jpg拆开看结构1271_0706140506449是图像 ID后面跟的)(3,33 6,69 138,27 139,62 )是表盘四角坐标再后面的[3,33 6,69 138,27 139,62 ]是同样的坐标用方括号再存一遍。四组坐标按顺序是左上、右上、右下、左下每个坐标是x,y格式。这种命名方式的好处是不用维护额外的标注文件坏处是文件名超长、容易在传输中截断而且解析逻辑必须写得足够健壮。解析这类文件名的常见做法是用正则把括号里的坐标段抠出来import re def parse_annotation(filename): # 匹配方括号内的四组坐标格式如 [3,33 6,69 138,27 139,62 ] pattern r\[(\d,\d)\s(\d,\d)\s(\d,\d)\s(\d,\d)\s*\] match re.search(pattern, filename) if not match: return None points [] for group in match.groups(): x, y group.split(,) points.append((int(x), int(y))) # 返回顺序左上、右上、右下、左下 return points这段逻辑的关键点是用方括号而不是圆括号做匹配因为圆括号在文件名里出现了两次方括号只有一次能避免歧义。\s匹配坐标之间的空格\s*容忍末尾可能多出的空格。返回的四个点顺序固定后续做透视变换或者裁剪都依赖这个顺序。如果解析返回None说明这个文件名不符合规范训练时要跳过否则会污染数据。2.2 定位网络为什么用回归而不是检测框d_model.py里做的是坐标回归不是常见的 YOLO 那种检测框加分类。原因很直接水表表盘在图像里基本只有一个不需要处理多目标也不需要区分类别。回归四个角点坐标比检测框更精确因为表盘可能是倾斜的矩形框会带进大量背景。四个角点确定后可以做透视变换把倾斜的表盘摆正再送给识别网络。定位网络的输入是整张水表图输出是归一化后的八个值四个点的 x、y。损失函数一般用均方误差或者 Smooth L1后者对异常值更鲁棒。训练时要注意坐标归一化把像素坐标除以图像宽高落到 0 到 1 之间否则不同分辨率图像混在一起训练会震荡。2.3 识别网络的结构与输入处理wm_model.py是识别网络输入是定位网络裁剪并摆正后的表盘区域。识别网络要同时处理指针读数和数字滚轮常见做法是 CNN 提特征后接两个分支一个分支做数字分类滚轮上的数字一个分支做指针角度回归。项目里image_preprocess.py和image_pre.py负责预处理包括灰度化、二值化、尺寸归一化这些步骤。识别网络的训练数据来自newdatares里已经标注好的图像但要注意识别网络的输入应该是定位网络裁剪后的结果而不是原图。如果训练识别网络时直接用原图推理时又用裁剪图分布不一致会导致精度暴跌。我一般会先把定位网络训到收敛用它批量裁剪出表盘区域再用这些裁剪图训识别网络。3. 环境配置与训练流程从零跑通两个网络3.1 TensorFlow 环境与依赖版本这套代码用的是 TensorFlow 1.x 风格的 APITensorflowUtils.py里能看到tf.Session、tf.placeholder这些老接口。如果你装的是 TensorFlow 2.x直接跑会报一堆AttributeError。常见做法是建一个 Python 3.6 或 3.7 的虚拟环境装 TensorFlow 1.15conda create -n watermeter python3.7 conda activate watermeter pip install tensorflow1.15.0 pip install opencv-python numpy pillowtensorflow1.15.0是最后一个支持 1.x API 的版本能在 Python 3.7 上跑。如果你只有 Python 3.8 以上装 TensorFlow 2.x 后需要开兼容模式import tensorflow.compat.v1 as tf但项目里没有做这个适配得自己改。opencv-python用于图像预处理numpy和pillow是基础依赖。装完后跑python -c import tensorflow as tf; print(tf.__version__)确认版本是 1.15。3.2 数据准备与目录结构项目的数据放在newdatares目录下图像文件名带标注。训练前要确认几件事图像格式是 jpg文件名符合解析规则坐标在图像范围内。我一般会先跑一个检查脚本把解析失败的、坐标越界的、图像损坏的挑出来import os import cv2 from parse_annotation import parse_annotation # 前面写的解析函数 data_dir newdatares bad_files [] for fname in os.listdir(data_dir): if not fname.endswith(.jpg): continue points parse_annotation(fname) if points is None: bad_files.append((fname, 解析失败)) continue img cv2.imread(os.path.join(data_dir, fname)) if img is None: bad_files.append((fname, 图像损坏)) continue h, w img.shape[:2] for x, y in points: if x 0 or x w or y 0 or y h: bad_files.append((fname, 坐标越界)) break print(f问题文件数{len(bad_files)}) for f, reason in bad_files[:10]: print(f, reason)这段脚本先解析文件名拿坐标再读图像确认没损坏最后检查每个坐标是否落在图像范围内。bad_files里的文件在训练前要移走或者修正否则训练时读一张崩一次。坐标越界通常是标注时手滑或者图像被裁剪过但文件名没更新。3.3 定位网络训练参数与启动命令定位网络的训练入口是d_train.py配置在D_config.py。启动前先看一眼配置里的关键参数参数含义常见取值batch_size每批样本数8 或 16learning_rate学习率1e-4 到 1e-3num_epochs训练轮数100 到 200input_size输入图像尺寸224x224 或 256x256checkpoint_dir模型保存路径自定义启动命令python d_train.py --config D_config.py如果显存不够先把batch_size降到 4 或 2再把input_size从 256 降到 224。训练过程中看 loss 曲线定位网络的 loss 应该在前 20 个 epoch 快速下降之后缓慢收敛。如果 loss 一直不降检查坐标归一化有没有做或者学习率是不是太大导致震荡。3.4 识别网络训练与两阶段串联识别网络训练入口是wm_train.py配置在WM_config.py。关键区别是识别网络的输入不是原图而是定位网络裁剪后的表盘区域。所以流程是先训定位网络用它批量裁剪出表盘图存到一个新目录再用这些裁剪图训识别网络。# 第一步训练定位网络 python d_train.py # 第二步用定位网络裁剪表盘区域假设有推理脚本 python d_inference.py --input newdatares --output cropped_data # 第三步训练识别网络 python wm_train.py --data cropped_datad_inference.py项目里没有直接提供但可以根据d_model.py和test.py改一个。核心逻辑是加载定位网络权重对每张图预测四个角点做透视变换裁剪出表盘区域保存成新图。裁剪后的图命名可以保持原 ID方便和标注对应。识别网络训练时读cropped_data目录配置里的data_dir指向它。4. 避坑与排查训练不收敛、坐标解析失败、显存爆了怎么办4.1 现象定位网络 loss 降到某个值就不动了原因坐标归一化没做或者做了但范围不对。如果坐标是像素值直接喂给网络输出也是像素值loss 会很大且难收敛。正确做法是把坐标除以图像宽高落到 0 到 1。另外检查一下图像 resize 后坐标有没有同步缩放resize 了图但没缩放坐标标注就全错了。解决在数据加载环节统一做归一化resize 图像的同时按比例缩放坐标。归一化后的坐标在 0 到 1 之间网络输出也用 sigmoid 激活保证范围一致。4.2 现象文件名解析返回 None训练时频繁跳过样本原因文件名里的坐标格式不统一有的用方括号有的用圆括号有的坐标之间是逗号有的是空格。项目里的样本大部分是[x,y x,y x,y x,y ]格式但可能有少量历史文件格式不同。解决解析函数里加容错先尝试方括号失败再尝试圆括号再失败记录到日志人工检查。不要直接抛异常中断训练跳过并记录训练完统一处理。4.3 现象训练到一半显存爆了报 OOM原因batch_size 太大或者输入图像尺寸太大。定位网络输入是整图如果原图分辨率很高比如 2000x3000resize 到 256 后显存占用还好但如果没 resize 直接喂原图显存瞬间爆。解决在数据加载时强制 resize 到固定尺寸比如 256x256。batch_size 从 4 开始试能跑通再往上加。另外检查有没有在训练循环里累积张量没释放比如把 loss 存到列表里一直不清理。4.4 现象识别网络在裁剪图上精度很高但推理时用原图精度暴跌原因训练和推理的输入分布不一致。训练用的是定位网络裁剪后的表盘图推理时如果直接拿原图喂给识别网络表盘只占图像一小部分识别网络看到的和训练时完全不一样。解决推理时也要先过定位网络裁剪再送识别网络。两个网络必须串联使用不能跳过定位直接上识别。如果定位网络在某些图上框不准识别结果也会跟着错这是级联系统的固有风险。4.5 现象训练 loss 正常下降但验证集精度不涨原因过拟合或者验证集和训练集分布差异大。水表图像如果训练集都是白天拍的验证集有夜间拍的模型泛化跟不上。解决加数据增强随机亮度、对比度、轻微旋转。另外检查验证集里有没有训练集出现过的图像 ID如果有说明数据划分时泄漏了验证精度虚高。5. 进阶技巧用定位网络的输出做数据增强与模型验证定位网络训好之后它的输出不只是用来裁剪。我一般会拿它做两件事一是检查标注质量二是做识别网络的数据增强。检查标注质量的思路用定位网络对训练集做推理把预测的四个角点和文件名里的标注角点画在同一张图上肉眼比对。如果某个样本预测和标注偏差很大要么是标注错了要么是这张图太难模型没学会。这个步骤能揪出一批脏数据比单纯看 loss 曲线有用得多。import cv2 import numpy as np def visualize_compare(img_path, gt_points, pred_points, save_path): img cv2.imread(img_path) # 标注用绿色画 for i in range(4): cv2.line(img, gt_points[i], gt_points[(i1)%4], (0, 255, 0), 2) # 预测用红色画 for i in range(4): cv2.line(img, pred_points[i], pred_points[(i1)%4], (0, 0, 255), 2) cv2.imwrite(save_path, img)绿色是标注红色是预测重合度高说明模型学得好。偏差大的样本单独拎出来看是标注问题就修标注是模型问题就加进难例训练集。数据增强的思路定位网络预测的四个角点可以做轻微扰动模拟标注误差让识别网络对裁剪偏差更鲁棒。具体做法是在裁剪时把角点随机偏移几个像素裁剪出的表盘图会有轻微位移和旋转识别网络在这种扰动下训练推理时对定位误差的容忍度更高。扰动幅度一般控制在 5 到 10 个像素太大反而让识别网络学不到稳定特征。验证方法上我习惯把数据集按 8:1:1 划分训练、验证、测试测试集只在最后跑一次。测试时两个网络串联端到端算读数准确率而不是分别看定位精度和识别精度。因为实际使用中用户只关心最终读数对不对中间环节的指标再好串联起来错也是白搭。从那以后我每次做级联模型都强制走一遍端到端测试不再只看单模块指标。希望帮到你。本文还有配套的精品资源点击获取
返回列表