资讯动态

MORAN 场景文本识别训练源码逐行精读:randomSequentialSampler 采样器与训练流程全解析

发布时间:2026/8/23 16:45:18 来源:尧图企业网站定制
MORAN 场景文本识别训练源码逐行精读randomSequentialSampler 采样器与训练流程全解析【免费下载链接】MORAN_v2MORAN: A Multi-Object Rectified Attention Network for Scene Text Recognition项目地址: https://gitcode.com/gh_mirrors/mo/MORAN_v2MORANMulti-Object Rectified Attention Network是一个带矫正机制的场景文本识别网络由 ResNet 风格骨干 MORN 矫正网络 ASRN 注意力识别网络组成。本文带你逐行精读 MORAN 的训练入口 main.py重点拆解自定义采样器 randomSequentialSampler 的设计动机与逐行实现并完整梳理从数据加载、模型组装到训练循环的全流程帮助新手快速上手训练自己的场景文本识别模型。 为什么需要精读 MORAN 的训练源码MORAN v2 在 IIIT5K 上达到 93.4% 的准确率详见 README.md其核心卖点是一次训练即可稳定收敛的矫正网络。对新手而言训练脚本看似简单但藏着两个容易被忽略的工程细节为什么不用 PyTorch 默认 shuffle——答案就在tools/dataset.py中的自定义采样器randomSequentialSampler为什么只支持单 GPU 训练——变长文本 batch 导致多卡切分困难理解这两点是读懂整个训练流程的钥匙。项目结构速览MORAN_v2/ ├── main.py # 训练/验证主入口 ├── train_MORAN.sh # 一键训练脚本 ├── demo.py # 推理演示 ├── models/ │ ├── moran.py # MORAN MORN(矫正) ASRN(识别) │ ├── morn.py # MORN 矫正网络 │ ├── asrn_res.py # ASRN 注意力识别网络 │ └── fracPickup.py # 训练时注意力扰动用 └── tools/ ├── dataset.py # LMDB 数据集 randomSequentialSampler └── utils.py # 标签编解码、loss 平均器一键启动train_MORAN.sh 里的关键参数训练只需修改数据集路径后执行sh train_MORAN.sh。脚本中几个值得注意的参数参数值说明--batchSize64每个 batch 的图片数--niter10训练轮数--lr1学习率作者提醒换任务请手动调低--adadelta-优化器选择 Adadelta--BidirDecoder-启用双向注意力解码器--valInterval1000每 1000 步验证一次--saveInterval40000每 40000 步存一次权重依赖环境见requirements.txt注意官方要求 PyTorch 0.3.1更高版本会显著拖慢训练pip install -r requirements.txt克隆仓库地址git clone https://link.gitcode.com/i/c536a3c63223d1e3a49299ccda5dbb98 核心精读randomSequentialSampler 采样器逐行解析先看它是如何被接管的在main.py构建 DataLoader 时有一处反常识的写法train_loader torch.utils.data.DataLoader( train_dataset, batch_sizeopt.batchSize, shuffleFalse, samplerdataset.randomSequentialSampler(train_dataset, opt.batchSize), num_workersint(opt.workers))shuffleFalse却又传入了自定义 sampler——打乱交给采样器DataLoader 自己不再洗牌。这是本文最值得深挖的设计。为什么要用随机起点 连续索引MORAN 的数据集是 LMDB 格式见tools/dataset.py的lmdbDataset。LMDB 按 key 顺序顺序读取时缓存命中率极高而纯随机索引访问会造成大量磁盘随机 I/O数据加载速度大幅下降。因此randomSequentialSampler的思路是每个 batch 随机选一个起点然后取一段连续索引——既保留了随机又保留了顺序读的效率。这是训练性能优化的经典技巧。逐行精读完整实现位于tools/dataset.pyclass randomSequentialSampler(sampler.Sampler): def __init__(self, data_source, batch_size): self.num_samples len(data_source) self.batch_size batch_size def __len__(self): return self.num_samples def __iter__(self): n_batch len(self) // self.batch_size tail len(self) % self.batch_size index torch.LongTensor(len(self)).fill_(0) for i in range(n_batch): random_start random.randint(0, len(self) - self.batch_size) batch_index random_start torch.arange(0, self.batch_size) index[i * self.batch_size:(i 1) * self.batch_size] batch_index # deal with tail if tail: random_start random.randint(0, len(self) - self.batch_size) tail_index random_start torch.arange(0, tail) index[(i 1) * self.batch_size:] tail_index return iter(index)逐段拆解__init__只记录样本总数与 batch 大小。没有保存任何预打乱顺序——顺序是每次__iter__时动态生成的所以每个 epoch 的采样顺序都不同。__len__返回num_samples注意它返回的是样本数而非 batch 数这让len(train_loader)与总 batch 数对齐DataLoader 内部会按 batch_size 整除训练循环中while i len(train_loader)正好跑完一个 epoch。n_batch与tail样本总数整除 batch 得到完整 batch 数余数 tail 单独处理保证一个 epoch 中每个样本恰好出现一次尾部样本靠最后一段补齐。random_start random.randint(0, len(self) - self.batch_size)起点上限是总数 - batch 大小确保random_start batch_size - 1不越界。batch_index random_start torch.arange(0, self.batch_size)起点加上[0, 1, ..., batch_size-1]得到一段连续索引。这是整个采样器的精髓一行。尾部处理同样随机取一个起点只取tail个连续索引填补最后一个不完整的 batch。返回iter(index)把一整个 epoch 的索引序列一次性生成DataLoader 按 batch 切块取数据。对比思考默认RandomSampler每次取一个完全随机的索引LMDB 随机读性能差而BatchSampler(RandomSampler)只是把随机索引分组组内依然不连续。randomSequentialSampler兼顾了两者这正是 MORAN 训练速度快的原因之一。数据流水线从 LMDB 到 TensorlmdbDataset 读取逻辑tools/dataset.py中lmdbDataset的读取规则样本按image-000000001/label-000000001这样的 key 存取num-samples存总数__getitem__中先index 1LMDB 从 1 开始编号图片转灰度后过resizeNormalize变换缩放到 200x64再ToTensor并做sub_(0.5).div_(0.5)归一化到 [-1, 1]标签会过滤掉字符表之外的字符追加结束符$若开启reverseTrue对应--BidirDecoder还会额外返回一份反转标签供双向解码器训练使用。两个容错细节值得学习图片损坏时返回self[index 1]标签为空时同样跳过——避免脏数据中断训练。多数据源合并main.py中把 NIPS 2014 与 CVPR 2016 两个训练集用ConcatDataset拼成一个数据集再交给同一个 DataLoader——采样器的连续索引依然在整个拼接数据集上生效。模型组装MORAN MORN ASRNmodels/moran.py只有 20 多行结构一目了然class MORAN(nn.Module): def __init__(self, ...): self.MORN MORN(nc, targetH, targetW, ...) # 矫正网络 self.ASRN ASRN(targetH, nc, nclass, nh, ...) # 识别网络 def forward(self, x, length, text, text_rev, testFalse): x_rectified self.MORN(x, test) # 先矫正 preds self.ASRN(x_rectified, ...) # 再识别 return predsMORNmodels/morn.py估计垂直方向偏移场用grid_sample把歪斜文本拉直。训练时有 50% 概率直接跳过矫正if np.random.random() 0.5这是一种正则化防止矫正网络过度拟合ASRNmodels/asrn_res.pyResNet 提特征 双向 LSTM 注意力解码器用fracPickup对注意力权重做训练扰动。训练主循环main.py 逐段精读环境与随机种子assert opt.ngpu 1, Multi-GPU training is not supported yet...单卡限制的原因注释写得很清楚batch 内文本长度不一多卡 DDP 无法均匀切分。随后固定random/numpy/torch三个种子保证可复现并打开cudnn.benchmark。张量复用与优化器main.py预先分配好固定大小的image/text/length张量每个 batch 用utils.loadData原地拷贝——省去每步重复申请显存的开销Variable 显式张量复用是 PyTorch 0.3 时代的典型写法。优化器支持四选一Adam / Adadelta / SGD / RMSprop默认 RMSprop。train_MORAN.sh选用的是 Adadelta学习率 1。trainBatch单步训练def trainBatch(): data train_iter.next() t, l converter.encode(cpu_texts, scannedTrue) # 文本 - 索引序列 长度 ... preds MORAN(image, length, text, text_rev) cost criterion(preds, text) # 交叉熵 MORAN.zero_grad() cost.backward() optimizer.step() return cost注意标签编码方式strLabelConverterForAttention.encode把一个 batch 的文本拼接成长度不等的索引序列length记录每条文本长度——注意力识别非 CTC必须知道每段边界。主循环验证、保存、训练三件事for epoch in range(opt.niter): train_iter iter(train_loader) while i len(train_loader): if i % opt.valInterval 0: # 定期验证 MORAN.eval() acc_tmp val(test_dataset, criterion) if acc_tmp acc: # 只留最优权重 torch.save(MORAN.state_dict(), {0}/{1}_{2}.pth.format(...)) if i % opt.saveInterval 0: # 定期存档 torch.save(MORAN.state_dict(), ...) MORAN.train() cost trainBatch() # 一次前向 反向 ...几个关键细节验证前后切换训练态验证前把全部参数requires_grad False并eval()验证后恢复train()——确保 BN/Dropout 行为与梯度开关正确双权重保存策略valInterval触发时保存精度最优权重文件名含精度值saveInterval触发时无条件存档文件名含 epoch 与步数兼顾最好的模型与可回溯的训练轨迹双向解码器的验证正向/反向各跑一次前向取两边平均置信度更高的一条作为最终预测main.py中按$截断、再反转还原这是 v2 稳定性的又一个小技巧。新手训练实操建议学习率要手动调低脚本里--lr 1是给 Adadelta 用的换成 Adam/SGD 时请调回常规量级数据必须是 LMDB 格式可用 ASTER 仓库提供的转换工具把自定义数据集转成 LMDB并写入num-samples⚡想提速理解本文的采样器设计后你可以把它替换为按 key 排序的索引生成器或在多进程 worker 中预取想直观看到矫正效果demo.py支持加载demo.pth做逐帧矫正可视化见 demo/ 目录下的演示图。总结MORAN 的训练代码虽短处处是工程巧思randomSequentialSampler用随机起点 连续索引化解 LMDB 随机读瓶颈主循环用最优权重 定期存档双保险保存模型MORN 的 50% 跳过矫正与双向解码器置信度择优共同换来 v2 的一次性稳定收敛。读懂这套流程你也就掌握了场景文本识别项目通用的训练骨架——数据管道、注意力编解码、变长文本处理缺一不可。【免费下载链接】MORAN_v2MORAN: A Multi-Object Rectified Attention Network for Scene Text Recognition项目地址: https://gitcode.com/gh_mirrors/mo/MORAN_v2创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价