简介:这是一份面向文本检测与OCR比赛场景的PyTorch实现资源,基于TextBoxes++改进模型,重点解决天池MTWI多目标文本识别中倾斜、任意形状文本的检测难题,适合正在备战OCR类竞赛或研究场景文本检测的开发者参考。压缩包共36个文件,以Python源码(.py)为主,另含少量编译缓存(.pyc)、Shell训练脚本、效果示例图、Jupyter演示与说明文档;整体仅2.38MB,轻量易部署。代码按model、data、train、eval、utils等模块组织,覆盖数据预处理、多尺度预测、Smooth L1损失设计、训练优化与MTWI指标评估全流程。已有98人学习,对希望完整复现比赛方案、快速搭建PyTorch文本检测基线,或借鉴赛题策略(如数据增强、模型融合)的读者来说,是一份结构紧凑、可直接运行的实战源码包。
1. 拆开这个 TextBoxes++ 的 PyTorch 源码包之前,先明确它解决什么问题
拿到TextBoxes++的pytorch版本,在天池mtwi比赛上进行应用.zip这个压缩包时,我第一反应是去看它是不是又一个“只给训练代码不给推理脚本”的半成品。解压之后发现结构比预想完整:train.py、eval_mtwi.py、test_mtwi.py、demo_multi.py都在,还有一份README.md和两个测试图片。这基本就是一套能在天池 MTWI 数据集上直接跑起来的完整 PyTorch 工程。
为什么 MTWI 比赛适合用 TextBoxes++?因为 MTWI 的数据里大量出现倾斜、弯曲、任意方向排列的文本,普通目标检测的矩形框(axis-aligned box)会把背景和邻近文本一并框进来,导致文字区域定位不干净,而 TextBoxes++ 在 SSD 的多尺度预测框架上引入了旋转矩形框(oriented box),让检测头可以回归每个文本行的中心点、宽高和旋转角度。对你来说是比赛场景,对工程场景来说则是广告文字识别、票据信息抽取、街景文本定位的前置检测模块。
这篇内容会按“模型原理 → 数据处理 → 训练调参 → 评测推理 → 实用技巧”的顺序把这套源码拆开。我会结合代码结构里的实际文件来讲,而不是泛泛介绍论文。你如果准备复现这个模型或者参加类似文本检测比赛,可以直接把后文的命令和参数拿去对照使用。
2. TextBoxes++ 模型架构拆解:从 TextBoxes 到任意方向文本检测
2.1 网络主干与默认框设计
TextBoxes++ 的主干网络沿用 SSD 的风格,基础层通常是 VGG16 的卷积部分(去掉全连接层),后面接额外卷积层来产生多尺度特征图。这个项目的model/modules.py和layers/functions基本是按照 SSD 的 Pytorch 实现方式改造的。关键点在于:普通 SSD 的默认框只有(cx, cy, w, h),TextBoxes 增加了一个(d, a)即文本线段的宽度和高度(可理解为垂直方向偏移),TextBoxes++ 在回归层上进一步输出了(cx, cy, w, h, angle)五元组。
看代码时我习惯先看config.py里的v2配置,因为默认框的aspect_ratio直接决定了检测召回的上限。这套代码里默认框设置了[1, 2, 3, 5, 1/2, 1/3, 1/5]这样的长条比例,这是文本检测和通用目标检测很不一样的地方——文字通常是细长条,长宽比超过 5 也很常见。如果你发现某些长文本行检测不到,优先检查aspect_ratio列表里是否包含足够大的值,比如7或10。
# config.py 片段(简化) v2 = { 'feature_maps': [38, 19, 10, 5, 3, 1], 'min_dim': 320, 'steps': [8, 16, 32, 64, 100, 320], 'aspect_ratios': [[1, 2, 3, 5, 1/2, 1/3, 1/5], [1, 2, 3, 5, 1/2, 1/3, 1/5], [1, 2, 3, 5, 1/2, 1/3, 1/5], [1, 2, 3, 5, 1/2, 1/3, 1/5], [1, 2, 3, 5, 1/2, 1/3, 1/5], [1, 2, 3, 5, 1/2, 1/3, 1/5]], }这段配置里feature_maps对应输入 320x320 时各卷积层的输出尺寸,steps是原图与特征图之间的缩放步长。注意这里min_dim是 320,意味着训练时图片会被缩放到 320x320。这个尺寸在 MTWI 这种高分辨率图像上会影响小字体的检测,后面训练部分我会建议改成 512 或 640,代价是 GPU 显存占用上升。
2.2 旋转框回归与损失函数
TextBoxes++ 的检测头同时输出分类置信度和旋转框回归值。分类部分和 SSD 一致,使用交叉熵;框回归部分不是简单的 Smooth L1,而是对旋转角度做了特殊处理。论文里把角度参数化成了两个元素:cos(2θ)和sin(2θ),目的是避免角度回归在 0° 和 180° 边界上的不连续问题。在layers/modules/box_utils.py里,你会看到类似encode函数中有对angle分量做cos、sin映射的代码。
损失函数采用多任务加权和:
# train.py 中 loss 计算示意 conf_loss = F.cross_entropy(conf_pred, conf_t, reduction='sum') loc_loss = smooth_l1_loss(loc_pred, loc_t, sigma=1.0) angle_loss = smooth_l1_loss(angle_pred, angle_t, sigma=1.0) loss = conf_loss + loc_loss + angle_losssmooth_l1_loss就是 Smooth L1 的实现,公式为:当|x| < 1时是0.5 * x^2,否则是|x| - 0.5。相比 L2 损失,它对离群点更不敏感,训练过程中不容易因为某一帧的标注框偏差产生大幅梯度。这里sigma=1.0控制平滑范围,sigma越大,损失对小误差的敏感度越高,训练初期容易震荡;sigma越小,对大误差的惩罚越平缓,收敛更稳定。我一般保持 1.0,只有在 loss 出现 NaN 时才考虑调大。
值得注意的是,这个项目的functions.py里可能同时包含match函数,负责把预测框和真实框做 IoU 匹配。文本检测里默认框和真实旋转框的交并比计算比普通矩形框复杂,代码里通常会先用最小外接矩形近似,或者直接把旋转框拆成四边形来计算多边形 IoU。你在训练时如果发现很多 anchor 没有被匹配到(正样本过少),可以适当降低overlap_threshold,默认通常是0.5,可以改成0.4来增加正样本量。
2.3 关键代码模块:box_utils 与 config.py
utils/box_utils.py是这个源码包的核心工具,里面至少包含decode、encode、nms这几个函数。decode把网络输出的 offsets 转换成最终的旋转框坐标,nms对重叠框做非极大值抑制。文本场景里 NMS 的 IoU 阈值很关键,通用目标检测常用0.45,但文本行之间经常有多个小框覆盖同一个长文本,阈值设太低会删掉有效框,设太高会输出大量重复框。我的做法是先用0.5做一次粗筛,再用0.3对角度相近的框做一次细筛。
config.py还有一个容易被忽略的配置项是max_num_text,它控制单张图片预测的最大文本框数量。在 MTWI 测试集上,如果一张图包含几十行文本,这个值设小了会被截断,设大了会拖慢 NMS 速度。代码里默认可能是 100,我建议根据实际数据统计设为 200 左右。另外score_threshold通常在0.5~0.7之间,情景区别很大:街景牌匾检测可以设低些(0.4),而票据扫描文本可以设高些(0.6),具体以验证集 F1 为准。
3. MTWI 数据集处理与数据增强配置
3.1 数据集结构解析与加载器实现
MTWI(Multi-Target Web Image)数据集来自天池比赛,标注格式是 XML 或 JSON,每个文本区域用四个点描述四边形顶点坐标。这个源码包里的data/mtwi2018.py就是专门解析这种格式的 PyTorch Dataset 类。使用它的第一步是确认目录结构:
MTWI2018/ ├── train/ │ ├── image1.jpg │ └── ... ├── train_txt/ │ ├── image1.txt │ └── ... ├── val/ │ └── ... └── val_txt/ └── ...打开mtwi2018.py会看到大约这样的加载逻辑:
# data/mtwi2018.py class MTWIDataset(Dataset): def __init__(self, root, transform=None, target_transform=None): self.images = sorted(glob.glob(root + '/*.jpg')) self.txts = sorted(glob.glob(root + '_txt/*.txt')) self.transform = transform def __getitem__(self, idx): img = cv2.imread(self.images[idx]) h, w = img.shape[:2] boxes = [] with open(self.txts[idx], 'r') as f: for line in f.readlines(): parts = line.strip().split(',') # 格式: x1,y1,x2,y2,x3,y3,x4,y4,text,ignore quad = [float(x) for x in parts[:8]] # 转换成旋转框 (cx, cy, w, h, angle) 或直接保存四边形 boxes.append(quad) return img, np.array(boxes, dtype=np.float32)这里要注意_txt目录名和root拼接逻辑。比赛数据里每个文件名的编号是对应的,如果出现错位,多半是glob排序的问题。字符串排序默认按字典序,10会排在9前面,所以sorted前需要做 key 转换:key=lambda x: int(x.split('/')[-1].split('.')[0])。这种小坑通常会让你的训练集和验证集错乱,训练指标看着正常但实际结果很差。
另外代码里target_transform会把四边形标注转换成(cx, cy, w, h, angle)格式,转换时用cv2.minAreaRect求最小外接矩形。这样做的缺点是:对 U 型或 S 型文本,四边形的最小外接矩形会包含很多背景。如果比赛数据里有大量弯曲文本,建议不要直接转旋转框,而是保留四边形,把回归头改成四边距离预测(类似 EAST),但那就是另一套代码了。当前这套代码只支持旋转框,你要有预期。
3.2 数据增强与预处理流程
utils/augmentations.py里实现了多种数据增强方法,包括随机裁剪、颜色扰动、旋转、缩放、翻转。文本检测训练时最需要注意的是“随机裁剪不能把标注框切掉一半”。常见做法是:
# utils/augmentations.py 中的随机裁剪逻辑 def random_crop(image, boxes, labels, max_trials=50): h, w = image.shape[:2] for _ in range(max_trials): min_iou = random.choice([0.1, 0.3, 0.5, 0.7, 0.9]) if min_iou >= 1.0: return image, boxes, labels for _ in range(50): nh = random.randint(0.5 * h, h) nw = random.randint(0.5 * w, w) x = random.randint(0, w - nw) y = random.randint(0, h - nh) crop = image[y:y+nh, x:x+nw] # 计算裁剪框与每个标注框的IoU,保留IoU大于阈值的框 new_boxes = [] for box in boxes: if iou(crop_box, box) >= min_iou: new_boxes.append(box - [x, y]) if len(new_boxes) > 0: return crop, np.array(new_boxes), labels return image, boxes, labels这段代码每次随机设定一个最低 IoU,然后尝试找到一块区域让至少一个文本框与裁剪框的 IoU 高于该阈值。这样既能做数据扩充,又不会把训练目标切得七零八落。如果训练时大量正样本框被裁剪得太小,可以检查这个函数里的min_iou取值区间,适当提高下限到0.3以上。
图片最终要缩放成config.py中的min_dim,但文本检测对高分辨率尤其敏感,直接压到 320 会丢失小字细节。源码包的data/mtwi2018.py里通常用cv2.resize连续插值,比赛场景下我建议你改成cv2.INTER_AREA做缩小、cv2.INTER_CUBIC做放大,这样小字边缘更锐利。另外颜色通道注意不要转成灰度,直接用 RGB,因为文本颜色本身是分类的重要特征。
3.3 配置文件 config.py 参数调整
config.py是比赛调参的核心文件。除了前面提到的aspect_ratios、feature_maps,还有几组参数直接决定训练效果:
| 参数名 | 默认值 | 作用 | 建议调整方向 |
|---|---|---|---|
lr | 1e-3 | 初始学习率 | 微调阶段降到1e-4 |
batch_size | 16 | 每批样本数 | 显存不够时减半并同步调小学习率 |
weight_decay | 5e-4 | 权重衰减系数 | 过拟合时增大到1e-3 |
num_workers | 4 | 数据加载线程数 | Windows 下建议设为 0,否则会报错 |
save_folder | weights/ | 模型保存路径 | 确保目录存在 |
angle_loss_weight | 1.0 | 角度损失的权重 | 角度偏差大时增大到2.0 |
调试时先跑通小数据子集,把config.py中的train_sets和val_sets指到同一个小目录,再逐步扩大。很多参赛者第一次训练直接全量数据,6 小时之后才发现代码 bug,这是最浪费时间的做法。我第一次跑这个源码时,直接用 100 张图训练 10 轮,确认 loss 能从 20 降到 5 左右,才敢全量训练。
4. 训练实战:优化器、学习率调整与训练脚本
4.1 训练流程与代码走读
train.py是完整的训练入口,它的逻辑和 SSD 的训练脚本很相似:
python train.py \ --dataset_root ./data/MTWI2018 \ --config ./config.py \ --batch_size 16 \ --num_workers 4 \ --start_iter 0 \ --lr 1e-3 \ --save_folder ./weights \ --resume ./weights/ssd300_epoch_100.pth脚本内部流程是:初始化模型 → 加载预训练权重 → 创建数据加载器 → 循环迭代 → 前向传播 → 计算损失 → 反向传播更新 → 周期性保存。这里面比较关键的是预训练权重。TextBoxes++ 的主干通常用 VGG16 在 ImageNet 上预训练过的权重,如果不开--resume,代码默认从随机初始化开始,那收敛速度会慢好几倍。
# train.py 中优化器设置 optimizer = optim.SGD(model.parameters(), lr=args.lr, momentum=0.9, weight_decay=config.weight_decay) scheduler = optim.lr_scheduler.MultiStepLR( optimizer, milestones=[100, 150, 200], gamma=0.1)这里用的是 SGD + Momentum,而不是 Adam。在文本检测这类密集预测任务上,SGD 的泛化能力通常优于 Adam,尤其是经过长时间训练后。milestones=[100,150,200]表示在第 100、150、200 轮迭代时把学习率乘以 0.1。如果你改用 Adam,建议初始学习率降到1e-4,因为 Adam 自适应的学习率步长在1e-3下容易震荡。
4.2 损失计算与梯度稳定性
训练过程中你会在终端看到类似iter 500 || Loss: 8.234 || Conf Loss: 5.678 || Loc Loss: 1.987 || Angle Loss: 0.569的输出。这份源码把三个损失分量分开了,这对定位问题很有帮助。如果Conf Loss居高不下,说明正负样本不平衡严重,检查是否限制负样本比例。常见的 SSD 实现会做 hard negative mining,让负样本与正样本的比例不超过 3:1。如果代码里没实现,你在train.py里找找是否有对conf_loss做top_k截断。
如果Angle Loss很大且不下降,先检查角度标注是否正确。MTWI 的四边形标注里,四个点的顺序不统一,有的按顺时针给,有的按逆时针给,直接把四个点传给cv2.minAreaRect可能得到错误角度。在mtwi2018.py里可以做一次顶点排序——先计算中心点,再按 atan2 角度排序,保证四边形顶点是顺时针排列。这个 bug 非常隐蔽,我的第二个训练失败教训就在这里。
训练时捕捉梯度爆炸的方法也很简单:
# 在 loss.backward() 之后、step() 之前插入 torch.nn.utils.clip_grad_norm_(model.parameters(), 10.0)这里把梯度范数裁剪到 10.0。如果训练前几个 iter 就出现NaN,通常是学习率过大或 backbone 权重初始化异常。把lr从1e-3降到1e-4先试跑 20 iter,如果还在 NaN,再检查输入图片里是否有全黑或全白图像导致 BN 统计量异常,过滤掉这类样本即可。
4.3 训练过程监控与常见问题
我实际训练这套模型时,遇到最多的问题是cv2.error: ... assertion failed。原因通常是augmentations.py里对图像做旋转或缩放时,边界框坐标越界。解决办法是在数据加载器的collate_fn或__getitem__末尾加一步过滤:
def filter_out_of_bound(bboxes, w, h): # 保留完全在图像内的框 mask = (bboxes[:, 0] >= 0) & (bboxes[:, 1] >= 0) & \ (bboxes[:, 2] <= w) & (bboxes[:, 3] <= h) return bboxes[mask]另外显存不足也很常见。如果你的显卡只有 8G 显存,把batch_size降到 4,同时把图片 resize 到 384x384,再配合torch.cuda.amp.autocast()做半精度训练,能省约 40% 显存。源码包没有提供混合精度代码,我自己加的时候注意在forward前后使用autocast和GradScaler。半精度下Smooth L1可能出现梯度下的损失较小,但最终检测框精度略降,所以最好只对 backbone 之外的层做混合精度。
5. 评估与推理应用:eval_mtwi.py使用与 demo 结果验证
5.1 评测指标与评估脚本使用
比赛的评价标准是端到端文本检测的 F1-score,具体调用方式是:
python eval_mtwi.py --trained_model ./weights/textboxes_pp_epoch_200.pth \ --config ./config.py \ --dataset_root ./data/MTWI2018 \ --eval_set valeval_mtwi.py会遍历验证集,对每张图执行前向推理,然后计算预测框与真实框的 IoU,匹配成功的框超过某个阈值就算检测正确。在 MTWI 比赛里,IoU 阈值通常设为 0.5,并且判对要求文本位置匹配且分类正确。代码输出会包括Precision、Recall和F1,你直接看 F1 即可。
这里有个容易忽略的点:推理时图像的尺寸必须与训练时一致。eval_mtwi.py内部只会调用config.py中设置的min_dim,并不会自动做多尺度测试。如果你训练时用 512,评估时用 320,F1 会非常惨。先检查eval_mtwi.py里是否有target_size变量,没有的话就让它从config里读取。
5.2 单图与多图推理 demo
源码包里提供了demo_mtwi.py和demo_multi.py两个推理脚本。单图推理的命令:
python demo_mtwi.py --trained_model ./weights/textboxes_pp_epoch_200.pth --image ./3.png --output ./result.jpgdemo_multi.py则是遍历一个目录下的所有图片。打开demo_mtwi.py看推理流程,核心代码比较直接:
# demo_mtwi.py 推理片段 def detect_bboxes(net, img, score_thresh=0.5): h, w = img.shape[:2] scale = 320.0 / max(h, w) # 保持长宽比缩放 new_w, new_h = int(w * scale), int(h * scale) resized = cv2.resize(img, (new_w, new_h)) # pad 到 320x320 canvas = np.zeros((320, 320, 3), dtype=np.uint8) canvas[:new_h, :new_w] = resized x = torch.from_numpy(canvas).permute(2, 0, 1).float().unsqueeze(0) with torch.no_grad(): boxes = net(x) # 输出已经是最终过滤后的旋转框 return boxes这里的缩放逻辑比较粗糙:先把长边缩放到 320,再补零到正方形。这样做的好处是避免了拉伸变形,但补零区域会让模型产生无意义的检测。如果你看到结果里有贴着右下边缘的误检框,多半就是 padding 引入的。
5.3 实际部署时的边界情况与调优技巧
最后一个值得展开的技巧是“解耦缩放与 pad,用多尺度推理提升召回”。在比赛场景中,文本行长度分布极宽,有的大标语横幅占满整张图,有的小字只有几十像素。固定缩放只能兼顾其中一类,我常用的做法是用demo_mtwi.py里的模型,在不同图片尺度上各推理一次,然后合并结果:
def multi_scale_infer(net, img, scales=[0.5, 1.0, 2.0]): h, w = img.shape[:2] all_boxes = [] for s in scales: nh, nw = int(h * s), int(w * s) resized = cv2.resize(img, (nw, nh)) # 把 resized 缩放到网络输入尺寸,再跑推理 boxes = run_net(net, resized) # 把坐标还原到原图尺度 boxes[:, ::2] /= s # x 坐标除 scale boxes[:, 1::2] /= s # y 坐标除 scale all_boxes.append(boxes) # 合并方式:对同一位置重叠的框取平均,或者直接 NMS merged = nms(np.vstack(all_boxes), thresh=0.5) return merged多尺度推理的好处是既能抓到小字(放大尺度)又能抓全大字(缩小尺度),坏处是推理时间翻倍。在比赛验证阶段可以用,但正式提交时如果时间有限,通常只用s=1.5一个尺度。另外,合并后的 NMS 阈值要适当调低(0.3),因为同一文本行在不同尺度下可能会各自输出一个高置信度框。
关于角度修正还有一个实用技巧:TextBoxes++ 输出的角度范围是[-π/4, π/4],如果目标文本竖排,模型容易输出相反方向。你可以在后处理里判断w < h时交换宽高并把角度加上π/2,这样可视化结果更符合直觉。对于竖排文本占比高的验证集,这个修正能让 recall 提高 1~2 个百分点。
如果你要在 CPU 或边缘设备部署,记得把模型转成 TorchScript 或 ONNX。转换前需要去掉代码里的动态循环匹配部分,固定输入尺寸为 320x320 或 512x512,否则 ONNX 导出会因为nms中的while循环失败。这个源码包里的nms是继承自 SSD 的普通 NMS,并不适合直接导出。建议只导出模型卷积部分,用 Python 做 NMS,这样避免了算子兼容问题。
本文还有配套的精品资源,点击获取