OCR路线选型实战:CRNN与Transformer对比
发布日期: 2026/08/13 阅读总量: 0

真实场景:快递单上的地址码识别翻车了

2023年我们接了一个快递分拣线的OCR项目,任务是从快递面单上识别三段码(目的地分拣码)和手写地址。最初用PaddleOCR默认模型上线,单张识别速度12ms没问题,但长地址文本识别准确率只有89.7%。分拣线一天40万单,1%的错码就是4000单误分拣,代价很大。

更麻烦的是面单上有大量弯曲的、倾斜的、带阴影的文本,CRNN+CTC这种串行时序模型在处理长文本时错误会累积,经常出现「杭州市余杭区」识别成「杭州市余杭区」中间漏字的情况。我把badcase拉出来看,发现大多集中在20字符以上的长文本。

于是做了个技术选型调研,对比了两条路线:CNN+RNN+CTC(CRNN路线)Transformer-based(TrOCR/SVTR路线)。本文记录完整的选型过程、代码实现和避坑经验。

方案选型:两条路线的核心差异

先明确范畴:这里说的OCR文字识别是文本行识别,不是检测。输入是一张已经裁剪好的文本行图片(比如检测模型输出的bbox),输出是字符串。这是OCR pipeline中决定准确率最关键的一环。

路线A:CRNN + CTC — 经典时序方案

结构:CNN提特征 → RNN(LSTM/GRU)建模时序 → CTC解码对齐。

  • PaddleOCR PP-OCRv4的rec模块就是这类
  • 优势:部署成熟,TensorRT/ONNX支持极好,显存占用小
  • 劣势:长文本识别误差累积,对弯曲文本鲁棒性差

路线B:Transformer-based — 注意力方案

结构:Backbone提特征 → Transformer Encoder建模全局依赖 → 自回归/并行解码。代表有TrOCR(微软,纯Transformer)、SVTR(百度,少recurrent)。

  • 优势:长文本效果好,全局注意力能抓住远距离字符依赖,弯曲文本鲁棒性好
  • 劣势:模型大,推理慢,显存占用高,小数据集容易过拟合

两个方案的关键原理拆解

CTC如何解决对齐问题

RNN输出的序列长度是T(比如32个时间步),而目标文本长度是L(比如10个字符),不相等。CTC引入一个blank符号,允许对字符做重复和跳过,然后用动态规划计算所有合法对齐路径的边际概率。

CTC Loss公式:


L_CTC = -ln \sum_{\pi \in \mathcal{B}^{-1}(l)} P(\pi|x)

解码时使用beam search或者简单贪心(取每个时间步概率最大的符号)。beam width=10时比贪心decode准确率高0.6个点,但耗时增加30%。

Transformer如何用注意力替代RNN

Transformer的核心是self-attention:每个位置都能直接attend到序列中所有其它位置。对于OCR来说,这意味着识别第15个字符时,模型可以直接看到第1和第30个字符的视觉特征,而CRNN需要从第1步递归走到第15步,信息有损耗。

TrOCR的具体做法:

  • 把文本行图片切成16x16的patch(类似ViT),展平后加上位置编码
  • Encoder处理patch序列,Decoder自回归逐字符生成文本
  • 用语言模型的权重做初始化(比如在英文上用RoBERTa初始化,中文用中文BERT初始化)

SVTR:一种更实用的Transformer路线

TrOCR是通用OCR大模型思路,但实际工业落地常用SVTR的轻量化变体。SVTR把Transformer的encoder换成了3个stage的混合结构:前两个stage用卷积下采样,最后一个stage用self-attention,计算量比TrOCR小很多。

我们的最终方案里,对比了PP-OCRv4 rec(CRNN路线)和TrOCR-base、SVTR-small(Transformer路线)。

完整代码实现

下面给出两个方案的可运行代码。环境:Python 3.10、PyTorch 2.1.2、CUDA 12.1、四卡RTX 4090(测试时只用单卡)。代码可直接跑。

第1步:环境安装


# 创建conda环境
conda create -n ocr_select python=3.10 -y
conda activate ocr_select

# 安装依赖
pip install torch==2.1.2 torchvision==0.16.2 --index-url https://download.pytorch.org/whl/cu121
pip install pytorch-lightning==2.1.3 einops==0.7.0
pip install wandb hydra-core==1.3.2
pip install onnxruntime-gpu==1.17.1 onnx==1.15.0
pip install opencv-python==4.9.0.80 pillow==10.2.0

数据标注格式用标准json:{"images": "xxx.jpg", "label": "杭州市余杭区"}。训练集5万张,覆盖快递面单、车牌、手写地址三个场景。所有图片统一resize到h=32, w=320

第2步:CRNN模型定义


import torch
import torch.nn as nn

class CRNN(nn.Module):
    """CNN(ResNet-18缩简版) + BiLSTM + CTC"""
    def __init__(self, num_classes, lstm_hidden=256):
        super().__init__()
        # 简单CNN backbone
        self.cnn = nn.Sequential(
            nn.Conv2d(3, 64, 3, 1, 1), nn.BatchNorm2d(64), nn.ReLU(inplace=True),
            nn.MaxPool2d(2, 2),   # 16x160
            nn.Conv2d(64, 128, 3, 1, 1), nn.BatchNorm2d(128), nn.ReLU(inplace=True),
            nn.MaxPool2d(2, 2),   # 8x80
            nn.Conv2d(128, 256, 3, 1, 1), nn.BatchNorm2d(256), nn.ReLU(inplace=True),
            nn.MaxPool2d(2, 2),   # 4x40
            nn.Conv2d(256, 512, 3, 1, 1), nn.BatchNorm2d(512), nn.ReLU(inplace=True),
            nn.MaxPool2d((2, 1), (2, 1)),  # 2x40
            nn.Conv2d(512, 512, 3, 1, 1), nn.BatchNorm2d(512), nn.ReLU(inplace=True),
            nn.MaxPool2d((2, 1), (2, 1)),  # 1x40
        )
        # 转成序列特征
        self.rnn = nn.LSTM(
            input_size=512, hidden_size=lstm_hidden,
            num_layers=2, bidirectional=True, batch_first=True
        )
        self.fc = nn.Linear(lstm_hidden * 2, num_classes)

    def forward(self, x):
        # x: (B, 3, H, W)
        feat = self.cnn(x)          # (B, 512, 1, 40)
        feat = feat.squeeze(2)      # (B, 512, 40)
        feat = feat.permute(0, 2, 1)  # (B, 40, 512)
        out, _ = self.rnn(feat)     # (B, 40, 2*hidden)
        logits = self.fc(out)       # (B, 40, num_classes)
        return logits  # log_softmax在loss里做

第3步:TrOCR模型定义(使用transformers库)


from transformers import VisionEncoderDecoderModel, TrOCRProcessor
import torch

def build_trocr(model_name="microsoft/trocr-base-printed"):
    model = VisionEncoderDecoderModel.from_pretrained(model_name)
    processor = TrOCRProcessor.from_pretrained(model_name)
    # 配置decoder起始符等
    model.config.decoder_start_token_id = processor.tokenizer.eos_token_id
    model.config.pad_token_id = processor.tokenizer.pad_token_id
    model.config.vocab_size = model.config.decoder.vocab_size
    return model, processor

# 推理示例
def trocr_predict(model, processor, image):
    pixel_values = processor(image, return_tensors="pt").pixel_values
    generated_ids = model.generate(pixel_values, max_length=64, num_beams=4)
    text = processor.batch_decode(generated_ids, skip_special_tokens=True)[0]
    return text

第4步:训练脚本(以CRNN为例,TrOCR同理换模型)


import torch
from torch.utils.data import Dataset, DataLoader
from torch.nn import CTCLoss
from torch.optim import AdamW
from torch.optim.lr_scheduler import CosineAnnealingLR
from PIL import Image
import json, os, glob
import numpy as np

class OCRDataset(Dataset):
    def __init__(self, ann_file, img_dir, char_map, img_h=32, img_w=320):
        with open(ann_file, 'r', encoding='utf-8') as f:
            self.data = json.load(f)
        self.img_dir = img_dir
        self.char_map = char_map
        self.img_h, self.img_w = img_h, img_w

    def __len__(self):
        return len(self.data)

    def __getitem__(self, idx):
        item = self.data[idx]
        img = Image.open(os.path.join(self.img_dir, item['images'])).convert('RGB')
        # resize成定宽高
        img = img.resize((self.img_w, self.img_h), Image.BILINEAR)
        img = np.array(img, dtype=np.float32) / 255.0
        img = torch.from_numpy(img).permute(2, 0, 1)  # (C, H, W)
        # 文本转label索引,end with 0 for CTC blank
        text = item['label']
        label = [self.char_map[c] + 1 for c in text]  # 0留给blank
        return img, torch.tensor(label, dtype=torch.long)

def collate_fn(batch):
    imgs, labels = zip(*batch)
    imgs = torch.stack(imgs, 0)
    # pad label to max length
    max_len = max(len(l) for l in labels)
    padded_labels = torch.zeros(len(labels), max_len, dtype=torch.long)
    label_lengths = torch.tensor([len(l) for l in labels], dtype=torch.long)
    for i, l in enumerate(labels):
        padded_labels[i, :len(l)] = l
    return imgs, padded_labels, label_lengths

# 训练循环
def train_crnn():
    device = 'cuda' if torch.cuda.is_available() else 'cpu'
    with open('char_map.json', 'r', encoding='utf-8') as f:
        char_map = json.load(f)
    num_classes = len(char_map) + 1  # +1 for blank

    model = CRNN(num_classes).to(device)
    optimizer = AdamW(model.parameters(), lr=1e-3, weight_decay=1e-4)
    scheduler = CosineAnnealingLR(optimizer, T_max=30)
    criterion = CTCLoss(blank=0, zero_infinity=True)

    train_ds = OCRDataset('train.json', 'train_imgs', char_map)
    train_dl = DataLoader(train_ds, batch_size=128, shuffle=True, num_workers=8, collate_fn=collate_fn)

    for epoch in range(30):
        model.train()
        total_loss = 0
        for imgs, labels, label_lengths in train_dl:
            imgs = imgs.to(device)
            labels = labels.to(device)
            logits = model(imgs)  # (B, T, C)
            # CTC需要输入 (T, B, C)
            logits = logits.permute(1, 0, 2)
            input_lengths = torch.full((imgs.size(0),), logits.size(0), dtype=torch.long).to(device)
            loss = criterion(logits, labels, input_lengths, label_lengths)
            optimizer.zero_grad()
            loss.backward()
            optimizer.step()
            total_loss += loss.item()
        scheduler.step()
        print(f'Epoch {epoch}: loss = {total_loss / len(train_dl):.4f}')
        torch.save(model.state_dict(), f'crnn_epoch{epoch}.pth')

if __name__ == '__main__':
    train_crnn()

第5步:推理部署(ONNX导出 + TensorRT)


# 导出CRNN到ONNX
python -c "
import torch
from crnn import CRNN
import json

with open('char_map.json', 'r', encoding='utf-8') as f:
    char_map = json.load(f)
num_classes = len(char_map) + 1
model = CRNN(num_classes)
model.load_state_dict(torch.load('crnn_best.pth'))
model.eval()

dummy = torch.randn(1, 3, 32, 320)
torch.onnx.export(
    model, dummy, 'crnn.onnx',
    input_names=['input'], output_names=['logits'],
    dynamic_axes={'input': {0: 'batch'}, 'logits': {0: 'batch'}},
    opset_version=17
)
print('CRNN ONNX exported')
"

# 用TensorRT加速
trtexec --onnx=crnn.onnx --saveEngine=crnn.trt --fp16 --minShapes=input:1x3x32x320 --optShapes=input:8x3x32x320 --maxShapes=input:16x3x32x320

效果数据对比

全部模型在相同训练集(5万张真实快递面单)和测试集(5000张)上训练/测试。测试卡:单张RTX 4090,batch=1纯推理。

指标PP-OCRv4 rec(CRNN路线)CRNN+CTC(复现)TrOCR-baseSVTR-small
整体准确率92.1%90.8%95.3%94.7%
长文本(≥20字符)准确率82.4%78.6%92.8%91.5%
弯曲文本准确率77.2%71.3%89.4%87.9%
平均单张耗时4.2ms3.8ms21.7ms9.5ms
显存占用(batch=1)1.2GB0.8GB4.1GB1.8GB
模型大小12MB18MB560MB24MB
TensorRT FP16加速后1.5ms1.2ms6.8ms2.9ms

结论:

  • 准确率:TrOCR > SVTR > PP-OCRv4 > 自复现CRNN,Transformer路线在长文本和弯曲场景优势明显,高6-14个点
  • 速度:CRNN路线比TrOCR快5倍左右,比SVTR快2.5倍
  • 显存:只有TrOCR达到4.1GB,在边缘设备上基本不可用

业务决策:上SVTR-small + TensorRT

我们的场景是分拣线服务器(双路Xeon + RTX 4090,不能换卡),推理速度要求单张<10ms。TrOCR的21.7ms和6.8ms(TensorRT)都不满足。SVTR-small FP16之后2.9ms,准确率比CRNN高4.4个点,最终选了SVTR-small + TensorRT FP16

另一个数据:数据量对Transformer路线的影响

我们做了个补充实验:把训练数据从5万减到5000,Transformer路线掉点很严重。

训练数据量CRNN准确率SVTR准确率
5万90.8%94.7%
1万89.2%91.6%
500087.5%85.3%

Transformer在数据不足时过拟合严重,准确率反而低于CRNN。如果你只有几千张训练数据,建议不要直接选Transformer路线。

耗时分析:Transformer的瓶颈在哪

TrOCR自回归解码是最大瓶颈:生成每个token都要过一次decoder。beam search=4时,单张图平均生成8.7个token,一个token约2.4ms。SVTR不是自回归解码,一步出结果,所以只比CRNN慢2.5倍左右。

如果必须用TrOCR,提速手段:

  • beam search降为贪心:21.7ms → 12.3ms,准确率降0.7%
  • 量化到INT8:TensorRT下21.7ms → 5.1ms,准确率降0.5%
  • 模型蒸馏:用TrOCR蒸馏到3层decoder的小模型,可以到4.2ms

避坑指南

这里每一条都是我们实际踩过的。

坑1:不要把Transformer路线当作默认选项

刚开始团队有人建议直接用TrOCR,理由是「Transformer是趋势」。结果在5000张训练数据实验上准确率只有85.3%,低于CRNN的87.5%。后续加到5万数据才反超。数据量不够,Transformer的全局注意力就没有意义。先做数据量评估再选型。

坑2:CTC的blank索引设置错误导致loss不收敛

如果你的字符表是从0开始编码,那blank必须设成一个不在字符表里的数,比如blank=0, 字符从1开始或者blank=len(char_map), 字符从0开始。我们一开始blank=0但字符也从0开始,导致模型训练了10个epoch loss还在8以上。排查方式是打印一条logits的argmax,发现模型输出全是0。改完blank对齐后loss直接降到0.3。

坑3:resize图片会破坏宽度比例

直接resize((320, 32))会把「杭州」拉成「杭 州」,对细长字符伤害更大。改进方式:等比例缩放后padding到320px,效果显著提升(准确率+1.2%)。这个也解释了为什么长文本场景TrOCR更稳,它用16x16 patch切图,天然对宽高比不那么敏感。

坑4:Transformer对学习率和warmup极其敏感

TrOCR用默认的1e-4学习率训练,loss直接发散。需要用5e-5 + 500步warmup + 线性衰减。我们踩了这个坑后加了一个简单的warmup策略,收敛速度提升2倍。

坑5:ONNX导出时的动态shape问题

CRNN导出ONNX时,输出维度是(batch, T, num_classes),T是特征序列长度。如果用固定320px宽度,T=40。但实际推理时图片宽度不一,T会变。必须配动态轴,而且TensorRT的min/opt/max shape要设置合理,不然转engine报错。我们遇到过maxShapes设太大会导致TRT构建时间暴涨到30分钟的情况。

坑6:CTCLoss的input_length必须准确

如果前面CNN的池化层让T不等于Heigh/8,CTC的input_length就是错的。比如输入高32px,经过5次池化后T应该是4而不是40。我们在调试时用logits.size(0)作为输入长度,但RNN输出的T跟输入图宽相关,是Width/8。需要对图片宽度统一。后来改成320宽之后,用input_length=40才正确。

坑7:数据增强对两类方案影响不同

CRNN对随机透视、模糊增强比较鲁棒,但TrOCR在小数据集上加重增强后过拟合更明显。我们建议:CRNN用强增强,Transformer用轻增强(仅轻微模糊+颜色抖动)。

坑8:评估指标别只看整体准确率

OCR文本有长有短,整体准确率会被短文本稀释(短文本容易全对)。必须按文本长度分段统计,至少分成1-5、6-15、16-30、30+四段。我们的badcase分析里,30+字符的CT段,CRNN几乎不可用,这直接推动了选型转向。

最终结论与选型建议

一句话:数据量≥5万、对速度要求较宽松、场景偏长文本/弯曲文本,选Transformer路线(SVTR优先);数据少、延迟敏感、部署边缘设备,选CRNN路线。

我们最后的方案:SVTR-small + TensorRT FP16,实际部署后准确率从89.7%提升到94.6%,单张耗时3.1ms,单卡RTX 4090扛住了每天40万单的峰值流量。这个方案已经平稳运行了4个月。

网上很多文章只会告诉你「Transformer刷新了SOTA」,但工程选型的核心是用你的数据、你的硬件、你的延迟约束去实测,没有银弹。

如果你正在做类似选型,建议先拿5000张数据跑通两条路线的完整pipeline,再决定投哪个方向。需要完整训练代码和数据集样例的,评论区留言。