☰
手写文本识别实战:Transformer+CTC完整链路解析
2026/10/11 22:55:21 网站建设 项目流程

简介:基于Transformer架构的手写文本识别项目源码与解析,面向具备一定深度学习基础、希望掌握序列识别与注意力机制的算法工程师和学生。系统不依赖字符分割,采用编码器-解码器完成端到端手写笔迹识别,核心是利用多头自注意力捕捉字形长距离依赖与二维空间结构。包体共18个文件、132KB,以Python脚本为主,涵盖数据增强、特征提取、模型训练与评估模块,另有Jupyter Notebook演示、阅读说明和备份文件,目录简洁便于对照调试。项目在IAM英文手写库和CASIA-HWDB中文数据集上分别达到94.7%和91.2%的行级准确率,较LSTM-CTC错误率降低23.6%,并配有数据增强工具、评估指标和可视化分析组件。已有133人下载学习,适合用来复现实验、改造自有手写识别任务或作为Transformer视觉文本识别项目的代码参考。

1. 手写文本识别为什么值得试试 Transformer:先把场景想清楚再动手

手写文本识别(HTR)这几年最大的变化,是 Transformer 不只在 NLP 里发光,还把手写字识别这个老任务拉回了聚光灯下。和印刷体 OCR 不同,手写笔迹没有统一模板,倾斜、粘连、连笔、涂抹叠在一起,传统的单字切分加分类思路基本被连笔和重叠击穿;Transformer 的自注意力能把整行的序列依赖一次建模清楚,这也是这套源码最核心的立足点。这份资源把数据预处理、模型实现、CTC 训练和解码推理串成了完整链路,适合已经会基础 CNN 分类、想把手写识别做成项目的开发者,也适合正在做文档数字化和表单识别的从业者直接照抄改造。

2. 手写识别要怎么换输入思路:预处理、标注编码与序列化的关键点

拿到手写识别资源包,第一步不是急着跑模型,而是先理解手写图像和印刷体图像在输入上的本质差异。这一章我把资源包里的数据管线拆开,讲清楚为什么切分思路在这里走不通,以及预处理和标签编码里哪些细节直接决定后续模型能不能收敛。

2.1 手写文本的难点与传统切分思路的死穴

手写文本和印刷体最明显的区别是没有标准字模。同一个字在不同人笔下,形态差异比印刷体大得多,而且往往带着连笔。很多手写行里,前一个字的最后一笔直接连到后一个字的起笔,字符之间根本不存在清晰的分割线。传统流程“先分割单字,再对每个单字做分类”在这种场景下非常吃亏:分割阶段一旦出错,后面的分类再准也挽救不回来,切割错误会直接变成识别错误的一部分。那些模糊、粘连、重叠的笔迹根本不存在稳妥的切分边界,这是切分思路的先天天花板。

后来常见做法是 CNN + BiLSTM + CTC。这类结构把整行图像直接送进网络,不需要切分,由 CNN 抽取局部视觉特征,BiLSTM 在时间方向建模上下文。它比单字分类稳健不少,但也有两个明显弱点:第一,BiLSTM 是顺序建模,序列越长,信息从尾部传回头部需要跨越的步数越多,梯度衰减问题无法回避;第二,它对横向长距离依赖的建模能力有限,而手写体里“一个字受前后字形态影响”恰恰是长距离关联。Transformer 把序列中任意两个位置直接相连,信息传递路径大幅缩短,这也是它在手写识别任务上比循环结构更能打的根本原因。

我一般会把这个环节拆成两条线:图像预处理线负责把任意尺寸的手写行图变成固定尺寸输入张量,标签编码线负责把字符串变成可供 CTC 训练的数字序列。这两条线最终由 CTC 在训练阶段做对齐,所以预处理质量几乎决定了模型能不能爬到可用精度。

2.2 图像预处理:高度固定、宽度动态、倾斜校正怎么做

预处理的第一步是灰度化和去噪。二值化不是必需的,直接拿灰度图进模型反而更稳,因为不少手写笔画的墨迹很浅,二值化阈值稍微选偏就会把细笔画切断,等于人为制造缺陷。去噪我一般用 3×3 或 5×5 中值滤波就够了,扫描件背景灰度低时再补一次自适应对比度拉伸。

倾斜校正在手写识别里非常关键。整行文本倾斜超过 15° 时,CNN 在横向滑窗时提取到的特征会变得很乱。常见做法是先通过图像投影找到文本行的上下边界,估计整体旋转角,再做仿射变换归零。霍夫变换检测最长直线也能估计角度,但扫描件里文本行经常带轻微弧度,直线检测容易误判,所以第一步只做全局小角度校正,局部弯曲交给网络自己学。

尺寸方面,我习惯把高度固定为 32 或 48,宽度按宽高比缩放后再做 padding,而不是直接拉伸到统一宽度。高度 32 时经过四次下采样刚好得到高度 1、宽度 W/4 的特征序列。padding 的目标宽度要算好,否则 CNN 下采样后序列长度会和预想的不一致。

import cv2 import numpy as np def preprocess_line(img, target_height=32, pad_value=255): # img: 灰度图 (H0, W0) h, w = img.shape[:2] scale = target_height / h new_w = int(round(w * scale)) # 缩小用 INTER_AREA,比 INTER_LINEAR 更能保留细笔画 resized = cv2.resize(img, (new_w, target_height), interpolation=cv2.INTER_AREA) # 宽度 padding 到 8 的倍数,方便 CNN 下采样后序列长度对齐 pad_w = (8 - new_w % 8) % 8 if pad_w > 0: resized = cv2.copyMakeBorder( resized, 0, 0, 0, pad_w, cv2.BORDER_CONSTANT, value=pad_value ) # 转成 (C, H, W),归一化到 [0, 1] x = resized.astype(np.float32) / 255.0 return np.expand_dims(x, axis=0) # (1, H, W)

这段代码里 INTER_AREA 用的是像素区域平均,图像缩小时能比线性插值保留更多笔画粗细信息,手写体很细的笔画在缩小时容易直接断掉,用 AREA 插值能明显缓解。padding 用 255 而不是 0,是因为手写数据基本都是白底黑字,填充白色才不会在图像边缘引入假的黑边。

2.3 字符表构建与标签编码:blank、UNK、PAD 三个位置要提前定

字符表是手写识别项目里最容易出疏漏的部分。开工前先统计标注集里的字符总量,确认是否包含数字、标点、大小写、连字符。随后要预留几个特殊 token:blank 在 CTC 中固定占 index 0,UNK 用于字符表之外的意外字符,PAD 只在需要定长张量时才会用。

chars = sorted(set(''.join(all_labels))) # CTC 要求 blank=0,所以从 2 开始给普通字符编号 vocab = {0: '<blank>', 1: '<unk>'} for i, c in enumerate(chars, start=2): vocab[i] = c char2idx = {v: k for k, v in vocab.items()} def encode_label(text, char2idx): ids = [char2idx.get(c, 1) for c in text] # 未登录字符给 <unk> return ids

CTC 的 blank 索引固定为 0,所以 0 号位必须留给 blank,1 号位留给 UNK。训练目标是变长整数序列,不需要按 batch 内最大长度做 padding,因为损失函数会依据 target_lengths 处理不同长度的目标。以前我总在这里搞混,后来习惯每次构造 Dataset 后都打印一个“字符 ID 解码回字符串”的往返结果做自检,保证 encode 和 decode 不会把字符表弄错位。

3. 模型主体实现:CNN 特征提取、Transformer Encoder 与 CTC 解码链路

这一章进入源码的核心部分。整套结构并不复杂:CNN Backbone 把二维图像压成一维时序特征,Transformer Encoder 在时序上建模依赖关系,最后接线性层映射到字符类别,训练用 CTC 损失,推理做去重解码。我拆这类源码时习惯先定位三个文件:网络结构定义、损失封装、解码脚本,这三个点理解透了,整份代码的骨架也就清晰了。

3.1 选型依据:Encoder + CTC 和 Encoder-Decoder 的分水岭

手写文本识别和机器翻译不一样。翻译必须逐词生成,因为目标语言有语序生成逻辑,而手写识别本质是“读出图像里的字符序列”,对齐关系隐藏在图像特征里,但每个字符并不依赖前一个字符的生成结果。所以 Encoder + CTC 是很多工程项目的首选:非自回归,一次前向推理就出完整序列,不用逐个字符循环;训练也稳定,没有 teacher forcing 和 exposure bias 的问题;显存占用低,适合小团队在单卡上做长文本。

Encoder-Decoder + Attention 也不是不能用,它更适合那些字符形态受上下文语义影响很大的场景,比如某些花体风格和特殊缩写。但代价是训练更慢,对数据量的要求更高。一句话判断:数据量小、字符表大、图像变形严重,先上 Encoder + CTC;数据充足且以后端语言模型辅助为目标,再考虑生成式方案。这份源码走的正是工程上最稳的 Encoder + CTC 路线。

3.2 从图像到特征序列:CNN Backbone 与位置编码

CNN Backbone 的职责是把高度为 32 的输入图像下采样成高度 1、宽度 W/4 的特征序列。卷积层和池化层的参数要严格匹配图像高度,否则特征序列长度会算错。

import torch import torch.nn as nn class ConvBackbone(nn.Module): def __init__(self, in_channels=1): super().__init__() # 高度 32 -> 16 -> 8 -> 4 -> 1,宽度保持 W -> W/2 -> W/4 self.conv = nn.Sequential( nn.Conv2d(in_channels, 32, kernel_size=3, padding=1), nn.BatchNorm2d(32), nn.ReLU(inplace=True), nn.MaxPool2d((2, 2)), nn.Conv2d(32, 64, kernel_size=3, padding=1), nn.BatchNorm2d(64), nn.ReLU(inplace=True), nn.MaxPool2d((2, 2)), nn.Conv2d(64, 128, kernel_size=3, padding=1), nn.BatchNorm2d(128), nn.ReLU(inplace=True), nn.MaxPool2d((2, 1)), # 宽度只缩一半,保留序列长度 nn.Conv2d(128, 256, kernel_size=3, padding=1), nn.BatchNorm2d(256), nn.ReLU(inplace=True), nn.MaxPool2d((2, 1)), ) def forward(self, x): # x: (B, 1, H, W) x = self.conv(x) # (B, 256, 1, W/4) return x.squeeze(2).permute(0, 2, 1) # (B, T, C),T = W/4

注意最后两个 MaxPool2d 用的是 (2, 1),宽度方向只缩一半,这是为了不让序列长度缩得太过分。如果四个池化都用 (2, 2),宽度会缩到 W/16,序列太短会导致每个时间步要承载过多字符信息,模型基本学不动。高度方向必须缩到 1,否则后面输入不了 Transformer Encoder。

Positional Encoding 这里采用经典的正余弦编码,加到特征序列上。CNN 已经做了局部感知,位置编码负责告诉 Transformer 每个时间步在整行里的相对位置。

class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len=512): super().__init__() pe = torch.zeros(max_len, d_model) position = torch.arange(0, max_len).unsqueeze(1) div_term = 10000 ** (torch.arange(0, d_model, 2) / d_model) pe[:, 0::2] = torch.sin(position / div_term) pe[:, 1::2] = torch.cos(position / div_term) self.register_buffer('pe', pe.unsqueeze(0)) def forward(self, x): # x: (B, T, C) return x + self.pe[:, :x.size(1)]

max_len 要大于训练集最长序列。比如训练样本最长宽度 256,下采样后 T=64,那 max_len 设 512 完全够用。d_model 必须和 CNN Backbone 的输出通道数一致,这里都是 256。

3.3 Transformer Encoder 堆叠与 CTC 解码

Encoder block 用标准的多头自注意力加前馈网络,每一层都接残差和 LayerNorm。

class TransformerEncoderBlock(nn.Module): def __init__(self, d_model=256, nhead=8, dim_feedforward=1024, dropout=0.1): super().__init__() self.self_attn = nn.MultiheadAttention( d_model, nhead, dropout=dropout, batch_first=True ) self.feed_forward = nn.Sequential( nn.Linear(d_model, dim_feedforward), nn.ReLU(), nn.Dropout(dropout), nn.Linear(dim_feedforward, d_model), ) self.norm1 = nn.LayerNorm(d_model) self.norm2 = nn.LayerNorm(d_model) def forward(self, x): attn_out, _ = self.self_attn(x, x, x) x = self.norm1(x + attn_out) ff_out = self.feed_forward(x) return self.norm2(x + ff_out) class HTRModel(nn.Module): def __init__(self, num_classes, d_model=256, nhead=8, num_layers=4): super().__init__() self.backbone = ConvBackbone() self.pos_encoder = PositionalEncoding(d_model) self.encoder_layers = nn.ModuleList([ TransformerEncoderBlock(d_model, nhead) for _ in range(num_layers) ]) self.fc = nn.Linear(d_model, num_classes) def forward(self, x): feat = self.backbone(x) # (B, T, C) feat = self.pos_encoder(feat) for layer in self.encoder_layers: feat = layer(feat) logits = self.fc(feat) # (B, T, num_classes) return logits

这里有一个很容易踩的细节:PyTorch 的 MultiheadAttention 默认 batch_first=False,如果忘了设置,输入 (B, T, C) 会被误读成 (T, B, C),结果就是训练时数据维度对不上。我在代码里显式写上 batch_first=True,省去后面 permute 的麻烦。nhead=8 要求 d_model 能被 8 整除,256 正好。Encoder 层数 4 层在中等规模数据集上足够,加深到 6 层收益有限但显存上涨明显。

CTC 损失的计算方式比分类任务的交叉熵要稍微绕一点。关键在于 log_softmax 的维度顺序必须转成 (T, B, C)。

import torch.nn.functional as F def ctc_loss_batch(logits, targets, input_lengths, target_lengths): # logits: (B, T, C) log_probs = logits.permute(1, 0, 2).log_softmax(dim=-1) return F.ctc_loss( log_probs, targets, input_lengths, target_lengths, blank=0, reduction="mean", zero_infinity=True )

zero_infinity=True 的用途是:当某条序列的 CTC loss 因为概率对数为零而变成 -inf 时,把它当 0 处理,避免梯度爆掉。训练时 input_lengths 是每个样本的实际特征长度,target_lengths 是每个字符串的字符数,这两个张量必须由数据加载器额外提供。

推理阶段的 greedy 解码要处理两件事:去掉 blank 占位符,合并相邻重复字符。

def greedy_decode(logits, vocab): # logits: (1, T, C) preds = logits.argmax(dim=-1).squeeze(0).tolist() prev = None chars = [] for token in preds: if token != 0 and token != prev: chars.append(vocab[token]) prev = token return ''.join(chars)

CTC 允许同一个字符跨越多个时间步输出,所以连续帧概率最大的字符相同时只输出一次。如果漏掉“去重”这一步,解码结果会出现大量连串重复字符,后面第五章节会专门讲这个坑。

4. 训练策略与调参:从 loss 下降慢到 CER 收敛的四个实操步骤

模型结构搭好之后,训练策略决定这套源码能不能在你自己的数据集上复现出 GitHub README 里贴的效果。手写识别的数据标注贵,尤其“图像 + 文本”成对数据很难积累,所以增强、学习率、评估指标这几件事每个都得认真对待。

4.1 数据增强:不能伤到标签的轻量扰动

手写数据集如果只有几千张,Transformer 这种大参数模型很容易过拟合。数据增强在 HTR 里不是锦上添花,而是必需品。我常用一组固定组合:小幅仿射变换、弹性形变、亮度对比度扰动、随机擦除。

旋转角度控制在 ±3° 以内,这个幅度不会改变字符本身的拓扑结构,但足以让模型对轻微倾斜鲁棒。弹性形变的幅度也要压住,否则本来完整的单词会被扭曲成另一个形状,标签却没变,反而引入错误监督。随机擦除是模拟遮挡和墨迹污染,只擦背景或细线,避免把整字抹掉。

import random import cv2 def augment_line(img): h, w = img.shape[:2] # 随机旋转,角度控制很小 angle = random.uniform(-3, 3) M = cv2.getRotationMatrix2D((w / 2, h / 2), angle, 1.0) img = cv2.warpAffine(img, M, (w, h), borderValue=255) # 随机擦除一个小矩形 if random.random() < 0.3: x0 = random.randint(0, max(1, w - 20)) y0 = random.randint(0, max(1, h - 10)) img[y0:y0 + 10, x0:x0 + 20] = 255 # 亮度对比度扰动 img = cv2.convertScaleAbs( img, alpha=random.uniform(0.8, 1.2), beta=random.randint(-15, 15) ) return img

warpAffine 的 borderValue 用 255,否则旋转后四个角会出现黑色三角区,模型会把黑边误当成笔画特征。随机擦除的矩形宽 20、高 10,对手写行来说不至于覆盖掉一个完整字符,但又足够模拟局部噪声。

4.2 超参表与训练循环的基本盘

Transformer 对学习率敏感,没有预训练权重时尤其依赖 warmup。下面是这套源码里我觉得比较稳的一组默认参数。

参数建议值说明
图像高度32 或 48决定 CNN 下采样后的特征序列高度
最大宽度数据集中最长宽度超长样本切段或缩放处理
batch size32~64显存不足降到 16,配合梯度累积
基础学习率1e-4AdamW 下稳定起步
warmup steps2000前 2000 步线性升到目标学习率
学习率调度cosine 或 Noam后半段衰减到接近 0
字符表大小视数据集而定标点符号必须覆盖
dropout0.1数据量少时提到 0.3

训练循环里最容易被忽略的是梯度裁剪。CTC 的反向传播在概率接近零时会产生巨大梯度,不做裁剪,模型可能在某个 batch 后直接炸掉。

optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4) scheduler = torch.optim.lr_scheduler.LambdaLR( optimizer, lr_lambda=lambda step: min(step / 2000, 1.0) ) for batch in train_loader: images, targets, target_lengths = batch logits = model(images) input_lengths = torch.full((images.size(0),), logits.size(1), dtype=torch.long) loss = ctc_loss_batch(logits, targets, input_lengths, target_lengths) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0) optimizer.step() scheduler.step()

warmup 调度器在前 2000 步把学习率从 0 线性升到 1e-4,后续步数由 cosine 调度接管。input_lengths 必须显式传入,因为 batch 内不同样本虽然宽度 padding 到同一长度,但实际有效宽度不同,CTC 需要根据真实长度计算。

4.3 验证用 CER 和 WER,别只盯着 loss

训练 loss 下降不代表识别正确率高,最好每 5 轮跑一次验证集,输出三个指标:训练 loss、验证集 CER、验证集 WER。

CER 是字符错误率,用编辑距离除以标签字符数。WER 是词错误率,把手写行按空格切词后计算。

import Levenshtein def compute_cer(pred_text, gt_text): d = Levenshtein.distance(pred_text, gt_text) return d / max(len(gt_text), 1)

Levenshtein 库可以直接 pip 安装,算编辑距离非常方便。分母加 max 是为了避免空标签时除零。保存 checkpoint 时只按验证集 CER 最低来存,不要只看训练 loss,否则十有八九会存下过拟合模型。

5. 避坑排查:手写识别项目容易翻车的五个典型位置

这一章把我在复现和改造类似项目过程中遇到的高频问题整理出来。每条都是“现象 → 原因 → 解决”的结构,照顺序排查基本能定位。

5.1 训练 loss 出现 NaN 或迅速爆炸

现象:训练前几百步 loss 正常,某一步突然变成 NaN,或者连续几个 batch 后 loss 抬升到十几甚至几十。多数情况下 batch size 没变,数据也没明显异常。

原因:最常见的是 CTC 的 log_softmax 输出接近零时取对数得到负无穷,反向传播再把这个负无穷的梯度放大;其次是学习率过大,Transformer 对学习率比 CNN 敏感;再次是梯度裁剪没开。

解决:在 log_softmax 之后做一次数值下限保护,给 log_probs 加一个 clip,比如log_probs = log_probs.clamp(min=-1e-10);同时开启clip_grad_norm_(model.parameters(), 5.0);学习率从 1e-4 起步,不要直接上 1e-3。另外检查 blank 索引是否真的为 0,字符表错位也会让 loss 在训练中途跳变。

5.2 验证结果全是空行

现象:训练 loss 正常下降,但验证集解码出来的文本全是空字符串,一行字符都识别不出来。

原因:模型把所有概率都压给了 blank 类。CTC 的训练机制允许模型通过输出大量 blank 来压缩 loss,当模型容量不足或训练不充分时,它倾向于选择这条“偷懒”路径。特别是手写数据里字符间距较宽时,blank 帧占比天然很高,模型更容易走偏。

解决:先确认字符表里 blank 索引真的是 0,class 数量和 vocab 对齐;再把训练轮数拉长,给模型足够时间学会输出非 blank 字符;检查验证集图像和训练集的预处理是否一致,比如训练时做强增强、验证时没做,本身是合理的,但尺寸缩放和 padding 必须一致。还可以把 blank 的初始化 bias 调小,比如线性层初始化时给 blank 类一个更小的权重。

5.3 训练集 CER 很低但验证集崩盘

现象:训练集 CER 已经降到 5% 以下,验证集 CER 却还在 40% 以上,明显是泛化失败。

原因:过拟合。手写数据集的笔迹风格非常分散,同一个字符在不同人手里差异很大。如果训练集只覆盖了一两种书写风格,模型很容易记住这些风格的局部纹理,而不是学到一个“字符是什么”。相比之下印刷体 OCR 的过拟合问题要轻得多,因为印刷体字形高度统一。

解决:数据增强的旋转角度和弹性形变再加大一点;尝试把 drop路径调高到 0.2~0.3;确认验证集和训练集确实来自不同书写者,如果验证集本身就来自同一批人,那 CER 40% 的问题可能另有原因。数据来源混杂时,建议按书写者做划分,而不是随机划分样本。

5.4 长文本后半段识别质量骤降

现象:短文本行(比如 8 到 10 个字符)识别效果不错,一行二三十个字符时,后半段错误率急剧上升,甚至后半段几乎全错。

原因:位置编码外推问题。训练时图像最长宽度有限,Transformer 的绝对位置编码只见过训练集内的长度范围,推理时图像宽度超过训练范围,后半段位置编码向量是模型从未见过的,注意力分布自然混乱。

解决:把训练数据里的长样本比例提上来,或者把最长训练宽度直接设成与推理一致。如果推理需要处理 512 像素宽的图,训练增强里就要保证有一部分样本接近这个宽度。另一个变通办法是把超长图切成两半,分别识别后再拼接,但这会引入跨切分处的字符粘连问题,要谨慎使用。

5.5 解码结果出现连串重复字符

现象:greedy 解码出来的文本里出现“sssee”这种连续重复,明显不该重复的位置重复了好几遍。

原因:相邻时间步输出了同一个字符但中间没有 blank 分隔。CTC 的合并规则是“先合并相邻重复字符,再删除 blank”,但模型对字符结束的边界判断不够明确时,会把同一个字符传播多个时间步。训练数据里如果存在大量连笔和重叠,这个问题会更严重。

解决:先用 greedy 解码把重复合并逻辑检查一遍,确认不是解码脚本写错。然后增大模型容量或加深 Encoder 层数,让模型有更强能力判断字符边界。也可以考虑在解码阶段引入 beam search,beam 会在多个候选路径里挑更合理的路径。注意,beam search 并不一定能消除所有重复,必要时结合语言模型做一次纠错。

6. 落地加速:ONNX 导出、动态宽度与 Beam Search 的实践取舍

模型训好只是第一步,把模型推到真实业务里还需要处理推理速度和部署依赖的问题。这三件事是我每次做 OCR 类项目都会顺手做掉的,性价比很高。

6.1 ONNX 导出:摆脱训练框架的推理方式

ONNX 导出可以保证模型脱离训练框架独立运行。导出时用动态轴支持宽度变化,宽度固定会强迫长图被 resize 变形,识别效果会明显下降。

import torch model.eval() dummy_input = torch.randn(1, 1, 32, 128) torch.onnx.export( model, dummy_input, "htr.onnx", input_names=["input"], output_names=["logits"], dynamic_axes={ "input": {0: "batch", 3: "width"}, "logits": {0: "batch", 1: "time"} }, opset_version=13 )

dynamic_axes 里把 batch 和 width 声明为动态,推理时就可以传入不一样宽度的图。时间步 T 会随宽度变化,所以 logits 的第二维也要标记为动态。opset_version 用 13 以上基本能覆盖常见的算子。

6.2 Beam Search:greedy 之外的性价比选择

Greedy 解码简单快,但只保留每一步最优,无法修正局部错误。字符级 beam search 每个时间步保留 top K 条候选路径,能挽回部分局部错误。beam=5 时 CER 通常能降 1 到 3 个百分点,但收益不是线性的,beam=10 之后提升很小,推理耗时却翻倍。没有语言模型时,beam 的提升主要来自它对重复输出的抑制和对低置信度区域的候选探索。

6.3 置信度阈值:解码阶段的一剂后悔药

这是我最常用的落地技巧。推理时对 logits 做 softmax,取每帧最大概率和对应类别,当置信度低于某个阈值时强制把该帧置为 blank。它相当于一个解码阶段的“后悔药”:模型信心不足时宁可跳过,也不要硬认,能明显减少低置信度区域的胡猜字符。阈值一般取 0.6 到 0.7,太高会把真实字符也抹成空白。

probs, pred_idx = logits.softmax(dim=-1).max(dim=-1) low_conf = probs < 0.6 pred_idx[low_conf] = 0 # 置为 blank

这份源码我前后拆了两遍,第一遍只顾跑通训练循环,结果在字符表对齐和动态宽度上各翻了一次车。从那以后,我每次开新数据集都强制自己先跑一遍预处理和字符映射的往返自检,再进训练。先确认“能吃进模型”再谈精度,这个顺序能省出大把调试时间。希望帮到你。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询