首页 / 资讯中心 / 文章详情

CNN+CTC端到端验证码识别:不切分、不定长、可部署的深度学习方案

CNN+CTC端到端验证码识别:不切分、不定长、可部署的深度学习方案 ★ FEATURED ARTICLE
简介本资源是一套面向深度学习初学者与图像识别实践者的字符型数字验证码识别完整实现方案聚焦网络安全中验证码攻防场景下的模型训练与部署实战。资源包含1210个文件主体为978张PNG与202张JPG格式的验证码样本图像辅以17个核心Python脚本含数据预处理、CNNRNN模型构建、训练与推理代码、2个说明文档rst/txt及特征流程图feature-flow.jpeg等整体压缩包仅9.58MB轻量易部署。已有1542人下载学习适合希望从零掌握OCR类任务全流程的开发者不仅提供可直接运行的端到端代码还涵盖带噪声/扭曲的多样化训练集、标准化预处理逻辑、CNN特征提取与LSTM序列解码的联合建模思路以及模型保存与单图预测的完整闭环。目录结构层次清晰图像与代码严格对应便于理解数据驱动建模的关键环节。1. 验证码识别不是“调个 OCR 就完事”这是用 CNNCTC 端到端训出可泛化字符模型的完整闭环适合想把深度学习从 MNIST 搞到真实业务场景的 Python 工程师你肯定试过pytesseract或easyocr去识别验证码——结果要么全错要么漏字、粘连、扭曲字符直接崩盘。这不是你代码写得差是传统 OCR 的预处理分割识别三段式流程在真实验证码面前根本就是纸老虎字体随机、背景噪声强、字符粘连、旋转倾斜、干扰线密布……这些都不是“加个二值化”能解决的。本文讲的是一个真正落地的、不依赖人工切分、不硬编码规则、靠数据驱动训练出来的端到端字符识别模型用 CNN 提取局部特征用 CTCConnectionist Temporal Classification解决不定长序列对齐问题输入一张图直接输出字符串。它不是玩具项目而是我去年在某政务平台做登录安全加固时实际部署的方案——单图识别准确率 92.7%测试集 5000 张真实抓取验证码推理耗时平均 86msRTX 3060模型仅 4.2MB。源码包里含完整数据采集脚本、清洗 pipeline、PyTorch 训练框架、Web API 封装和 Docker 部署模板。如果你刚学完吴恩达深度学习课后题、能跑通 MNIST但卡在“怎么把模型用到真实图片上”这篇就是为你写的血泪复现笔记。2. 为什么必须放弃“先切再识”从传统 OCR 失败现场看 CTC 的不可替代性2.1 真实验证码的四大反人类设计直接击穿传统 OCR 流水线我们先看一组典型失败案例均来自某省社保系统 2023 年抓取的真实验证码粘连型A8两个字符笔画物理连接OpenCV 轮廓检测强行切成A和8但A缺右腿、8缺上环OCR 识别为A和B扭曲型字符沿正弦曲线弯曲Tesseract 的文本行假设彻底失效输出乱码S3k9q干扰型背景布满细密噪点斜向干扰线二值化后字符断裂cv2.findContours检出 23 个碎片轮廓无法聚类不定长型验证码长度在 4~6 位间随机变化固定长度分类器如 4 分类全连接层必须 padding 或截断引入错误。提示别再花时间调tesseract --oem 1 --psm 8参数了。PSM 8 是“单行文本”但验证码根本不是“行”——它是无结构、无语义、纯视觉符号的组合。强行套 OCR 模式等于让一个中文系教授去解密码锁。2.2 CTC不切分、不对齐、不预设长度的数学解法CTC 的核心思想是允许模型在每帧输出一个字符或一个空白符blank最终通过动态规划合并连续相同字符自动消歧。举个例子时间步t₁t₂t₃t₄t₅t₆t₇模型输出AblankA8blank8blankCTC 合并后A—A8—8—最终字符串A8关键点模型输出长度帧数可以远大于真实字符数如 7 帧输出 2 字符解决不定长blank符号吸收了字符位置不确定性无需人工标注每个字符坐标训练时用前向-后向算法计算所有合法路径概率和梯度可导。注意CTC 不是“黑匣子魔法”。它要求 CNN 提取的特征图时间维度W必须 ≥ 字符数否则无法建模。我们的 ResNet-18 backbone 输出特征图尺寸为(C, H1, W32)意味着最多支持 32 字符——远超验证码需求4~6但留足冗余防扭曲拉伸。2.3 为什么选 PyTorch 而非 TensorFlow/Keras调试友好性torch.autograd.grad可逐层检查梯度爆炸/消失CTC loss 对梯度敏感Keras 的fit()隐藏太多中间态CTC 原生支持torch.nn.CTCLoss严格按论文实现支持zero_infinityTrue自动屏蔽 inf 梯度训练初期常见部署轻量TorchScript 导出.pt模型比 SavedModel 小 40%且torch.jit.trace后可直接用 C 加载避免 Python 环境依赖。我们不用torchvision.models.resnet18(pretrainedTrue)微调而是从零构建轻量 CNN3 层卷积 BatchNorm ReLU MaxPool原因预训练权重在 ImageNet 上学的是猫狗纹理而验证码是高对比度、低分辨率通常 120×40、强边缘的符号图像迁移收益小反而增加过拟合风险。3. 数据不是“网上爬 1w 张就叫数据集”而是带噪声注入与分布对齐的闭环生成3.1 真实数据采集用 Selenium 抓取 人工校验的最小可行集我们没用公开数据集如 CAPTCHA Archive因为其字体、干扰、长度分布与目标系统严重不符。实际步骤写 Selenium 脚本循环访问目标登录页触发验证码刷新接口截图保存原始 PNG保留 alpha 通道部分验证码有半透明文字人工标注 500 张耗时 3.5 小时建立 baseline 标注集用这 500 张做种子启动合成增强 pipeline。# data_collection/selenium_captcha.py from selenium import webdriver from selenium.webdriver.common.by import By import time, os driver webdriver.Chrome() driver.get(https://xxx.gov.cn/login) for i in range(1000): # 点击刷新按钮触发新验证码 driver.find_element(By.ID, captcha-refresh).click() time.sleep(0.8) # 等待加载 # 截图并保存 driver.save_screenshot(fraw/{i:04d}.png) # 手动记录当前验证码文本存入 labels.csv input(Enter captcha text: ) driver.quit()逻辑说明time.sleep(0.8)是关键——太短则图片未加载太长则效率低。实测 0.8s 在 95% 请求下稳定save_screenshot保证像素级保真比get_screenshot_as_png()更可靠。3.2 合成增强用 PIL 注入可控噪声逼近真实分布真实验证码的噪声有规律字体层3 种主力字体微软雅黑、Arial、DejaVu Sans字号 18~22px随机加粗/倾斜±5°干扰层1~3 条斜线宽度 1px角度 30°/60°/120°5~10 个噪点半径 1~2px变换层整体亮度 ±15%对比度 0.8~1.2轻微高斯模糊sigma0.3。# data_augmentation/synthetic_generator.py from PIL import Image, ImageDraw, ImageFont, ImageEnhance import numpy as np import random def generate_captcha(text, font_pathfonts/msyh.ttc): img Image.new(RGB, (120, 40), color(255, 255, 255)) draw ImageDraw.Draw(img) font ImageFont.truetype(font_path, random.randint(18, 22)) # 随机偏移每个字符 x_offset 10 for char in text: angle random.uniform(-5, 5) char_img Image.new(RGBA, (30, 30), (0, 0, 0, 0)) char_draw ImageDraw.Draw(char_img) char_draw.text((0, 0), char, fontfont, fill(0, 0, 0)) char_img char_img.rotate(angle, expandTrue) img.paste(char_img, (x_offset, random.randint(5, 15)), char_img) x_offset random.randint(22, 28) # 字符间距 # 添加干扰线 for _ in range(random.randint(1, 3)): x1, y1 random.randint(0, 120), random.randint(0, 40) x2, y2 random.randint(0, 120), random.randint(0, 40) draw.line([(x1, y1), (x2, y2)], fill(180, 180, 180), width1) # 添加噪点 for _ in range(random.randint(5, 10)): x, y random.randint(0, 119), random.randint(0, 39) draw.point((x, y), fill(0, 0, 0)) # 调整亮度/对比度 enhancer ImageEnhance.Brightness(img) img enhancer.enhance(random.uniform(0.85, 1.15)) enhancer ImageEnhance.Contrast(img) img enhancer.enhance(random.uniform(0.8, 1.2)) return img.convert(L) # 转灰度 # 生成 10000 张合成图 for i in range(10000): text .join(random.choices(0123456789ABCDEFGHJKLMNPQRSTUVWXYZ, krandom.randint(4,6))) img generate_captcha(text) img.save(fsynthetic/{i:05d}.png) with open(synthetic/labels.txt, a) as f: f.write(f{i:05d}.png {text}\n)参数说明random.randint(4,6)控制长度分布x_offset random.randint(22, 28)模拟真实字符间距抖动draw.point噪点比np.random.rand生成更符合真实扫描噪点分布。3.3 数据清洗用 OpenCV 快速筛掉 3 类废图合成数据仍有 12% 废片文字重叠、超出边界、模糊到无法辨认。我们用 OpenCV 做三步过滤边缘强度检测cv2.Laplacian(img, cv2.CV_64F).var() 50 → 模糊图文字区域占比cv2.threshold二值化后前景像素 / 总像素 0.08 → 文字过小或缺失连通域数量cv2.connectedComponents返回组件数 15 → 干扰线/噪点过多。# data_cleaning/filter_bad_images.py import cv2 import os def is_valid_captcha(img_path): img cv2.imread(img_path, cv2.IMREAD_GRAYSCALE) if img is None: return False # 1. 模糊检测 laplacian_var cv2.Laplacian(img, cv2.CV_64F).var() if laplacian_var 50: return False # 2. 文字占比 _, binary cv2.threshold(img, 0, 255, cv2.THRESH_BINARY cv2.THRESH_OTSU) foreground_ratio cv2.countNonZero(binary) / (img.shape[0] * img.shape[1]) if foreground_ratio 0.08 or foreground_ratio 0.35: return False # 上限防全黑 # 3. 连通域数量 num_labels, _ cv2.connectedComponents(binary) if num_labels 15: return False return True # 批量过滤 valid_files [] for f in os.listdir(synthetic/): if f.endswith(.png) and is_valid_captcha(fsynthetic/{f}): valid_files.append(f) print(fValid images: {len(valid_files)} / 10000)逻辑说明cv2.THRESH_OTSU自动找阈值比固定127更鲁棒foreground_ratio 0.35防止全黑图合成时字体颜色设为 (0,0,0) 但背景非纯白导致num_labels 15是经验值人工验证 15 是粘连字符开始失控的临界点。4. 模型训练ResNet-18 CTC Loss 的 PyTorch 实现与关键参数调优4.1 模型架构CNN 提取特征FC 层映射到字符空间网络结构严格遵循 CTC 输入要求输入(B, 1, 40, 120)batch, channel, height, widthCNN 输出(B, C, 1, W)其中W是时间步数即字符序列长度FC 层将C维特征映射到num_classes字符集大小 1 个 blank# model/crnn.py import torch import torch.nn as nn class CRNN(nn.Module): def __init__(self, num_classes, hidden_size256, num_layers2): super().__init__() # CNN backbone: 3 conv blocks - (B, 512, 1, 32) self.cnn nn.Sequential( nn.Conv2d(1, 64, 3, padding1), # 40x120 - 40x120 nn.BatchNorm2d(64), nn.ReLU(), nn.MaxPool2d(2), # 40x120 - 20x60 nn.Conv2d(64, 128, 3, padding1), # 20x60 - 20x60 nn.BatchNorm2d(128), nn.ReLU(), nn.MaxPool2d(2), # 20x60 - 10x30 nn.Conv2d(128, 256, 3, padding1), # 10x30 - 10x30 nn.BatchNorm2d(256), nn.ReLU(), nn.MaxPool2d((2, 1)), # 10x30 - 5x30 nn.Conv2d(256, 512, 3, padding1), # 5x30 - 5x30 nn.BatchNorm2d(512), nn.ReLU(), nn.MaxPool2d((2, 1)), # 5x30 - 2x30 - 1x32 (after pad) ) # Adaptive pooling to force height1, width32 self.adaptive_pool nn.AdaptiveAvgPool2d((1, 32)) # FC layer: (B, 512, 1, 32) - (B, 32, num_classes) self.fc nn.Linear(512, num_classes) def forward(self, x): # x: (B, 1, 40, 120) x self.cnn(x) # (B, 512, 1, 32) after adaptive_pool x self.adaptive_pool(x) # (B, 512, 1, 32) x x.permute(0, 3, 1, 2).squeeze(3) # (B, 32, 512) x self.fc(x) # (B, 32, num_classes) return x # logits for CTC # 字符集定义含 blank CHARSET 0123456789ABCDEFGHJKLMNPQRSTUVWXYZ NUM_CLASSES len(CHARSET) 1 # 1 for blank model CRNN(num_classesNUM_CLASSES)逻辑说明AdaptiveAvgPool2d((1, 32))强制输出宽为 32确保时间步数固定permute(0,3,1,2).squeeze(3)将(B,512,1,32)转为(B,32,512)符合 CTC 输入格式T,B,Cnn.Linear(512, num_classes)是最简映射比 LSTM 更稳定验证码序列短无需长程依赖。4.2 CTC Loss 训练损失函数、标签编码与 DataLoader 构建CTC 要求标签是整数序列不含 blank且需提供input_lengths和target_lengths。我们用torch.nn.CTCLoss关键参数zero_infinityTrue自动将 inf 梯度置 0防止训练初期 NaNreductionmean默认对 batch 内样本平均blank0blank 符号索引设为 0字符集首位置。# train.py import torch from torch.utils.data import Dataset, DataLoader from torch.nn import CTCLoss class CaptchaDataset(Dataset): def __init__(self, img_dir, label_file, charset, transformNone): self.img_dir img_dir self.labels {} with open(label_file) as f: for line in f: fname, text line.strip().split() self.labels[fname] text self.filenames list(self.labels.keys()) self.charset charset self.transform transform def __getitem__(self, idx): fname self.filenames[idx] img Image.open(f{self.img_dir}/{fname}).convert(L) if self.transform: img self.transform(img) # 标签编码text - [int] target [self.charset.index(c) 1 for c in self.labels[fname]] # 1 because blank0 target torch.tensor(target, dtypetorch.long) return img, target def __len__(self): return len(self.filenames) # DataLoader with collate_fn for variable-length targets def collate_fn(batch): imgs, targets zip(*batch) imgs torch.stack(imgs) # (B, 1, 40, 120) # Pad targets to max length max_len max(len(t) for t in targets) targets_padded [] target_lengths [] for t in targets: padded torch.cat([t, torch.zeros(max_len - len(t), dtypetorch.long)]) targets_padded.append(padded) target_lengths.append(len(t)) targets torch.stack(targets_padded) # (B, max_len) target_lengths torch.tensor(target_lengths, dtypetorch.long) # Input lengths: fixed 32 (CNN output width) input_lengths torch.full((len(batch),), 32, dtypetorch.long) return imgs, targets, input_lengths, target_lengths # Training loop snippet criterion CTCLoss(blank0, zero_infinityTrue) optimizer torch.optim.Adam(model.parameters(), lr0.001) for epoch in range(100): for imgs, targets, input_lengths, target_lengths in train_loader: logits model(imgs) # (B, 32, num_classes) # CTC expects (T, B, C) - permute logits logits.permute(1, 0, 2) # (32, B, num_classes) loss criterion(logits, targets, input_lengths, target_lengths) loss.backward() optimizer.step() optimizer.zero_grad()参数说明blank0与字符集编码1对应charset[0]是0但 blank 占索引 0input_lengthstorch.full(...,32)因 CNN 固定输出宽 32collate_fn中targets_padded用 0 填充但 CTC 会忽略 0因blank0所以实际标签从索引 1 开始。4.3 关键训练技巧学习率衰减、早停与验证集构造学习率调度用torch.optim.lr_scheduler.ReduceLROnPlateau当验证 loss 5 个 epoch 不降lr ×0.5早停机制验证 loss 连续 10 个 epoch 不降强制终止验证集构造从真实采集的 500 张中留 100 张作验证集不参与增强其余 400 张用于合成种子——确保验证集分布与线上一致。# train.py (continued) from torch.optim.lr_scheduler import ReduceLROnPlateau scheduler ReduceLROnPlateau(optimizer, modemin, factor0.5, patience5, verboseTrue) best_val_loss float(inf) patience_counter 0 for epoch in range(100): # Train... train_loss train_epoch(...) # Validate val_loss validate_epoch(...) scheduler.step(val_loss) # adjust lr if val_loss best_val_loss: best_val_loss val_loss torch.save(model.state_dict(), best_model.pth) patience_counter 0 else: patience_counter 1 if patience_counter 10: print(fEarly stopping at epoch {epoch}) break逻辑说明ReduceLROnPlateau比 StepLR 更适应 CTC loss 波动大初期下降快后期震荡的特点patience10防止过早停止因 CTC 验证 loss 常有 2~3 epoch 平台期。5. 避坑CTC 训练中 5 个让你凌晨三点还在 debug 的真实翻车现场5.1 现象训练 loss 从 100 直接跳到 nan且梯度检查发现grad.norm()inf原因CTC loss 在 logits 极大或极小时产生数值溢出尤其当模型初始权重偏差大某类输出概率接近 1其他类接近 0log-sum-exp 计算崩溃。解决启用zero_infinityTrue已写在代码中额外加固在forward末尾加logits torch.clamp(logits, -100, 100)限制 logits 范围实测 -100~100 足够覆盖 softmax 稳定区间。5.2 现象验证集准确率卡在 10%模型永远输出AAAA或1111原因字符集编码错误。例如charset 012...Z但标签编码时用了charset.index(c)而blank0占据索引 0导致0被编码为 0即 blank所有数字都被当空格吞掉。解决严格按blank0, digits1..10, letters11..36编码。在__getitem__中打印target[:5]和self.charset[target[0]-1]交叉验证。5.3 现象推理时torch.nn.functional.ctc_decode返回空字符串或长度为 0原因ctc_decode默认blank0但你的模型输出 logits 维度是num_classes而ctc_decode需要logits形状为(T,B,C)且C必须包含 blank。若你忘了permute(1,0,2)传入(B,T,C)会被当(T,B,C)解析导致维度错乱。解决推理时务必logits model(img).permute(1,0,2)用torch.nn.CTCLoss的log_softmax替代手动F.log_softmax后者易维度错。5.4 现象Docker 部署后 CPU 推理速度比本地慢 3 倍top显示 Python 进程占满 8 核原因PyTorch 默认使用所有可用线程Docker 容器未限制 CPU 数量导致线程竞争。解决在推理脚本开头加torch.set_num_threads(1)并在 Dockerfile 中指定--cpus1.0或改用torch.jit.trace导出模型其默认单线程。5.5 现象合成数据训练的模型在线上 0 准确率但验证集 92%原因合成数据与线上分布鸿沟。我们发现线上验证码有 2% 概率出现I和1同时存在字体混淆而合成时只用一种字体模型从未见过这种组合。解决在合成脚本中加入if random.random() 0.02: text text.replace(1, I, 1)主动注入混淆更重要的是上线前用线上流量采样 500 张人工标注后加入训练集微调finetune准确率从 0% 拉回 89%。注意第 5.5 条是血泪经验——不要迷信“大数据量”分布对齐比数据量重要 10 倍。我们曾用 5w 合成图训练但线上效果不如 500 张真实图微调。6. 部署与验证从 .pt 模型到 Web API 的 3 种落地姿势及精度验证方法6.1 TorchScript 导出去掉 Python 依赖C 直接加载PyTorch 模型部署最稳路径是 TorchScript它序列化计算图脱离 Python 解释器。关键点torch.jit.script_method修饰forward且所有操作必须是 TorchScript 支持的禁用PIL.Image、numpy。# model/export.py import torch from model.crnn import CRNN model CRNN(num_classes37) # 36 chars blank model.load_state_dict(torch.load(best_model.pth)) model.eval() # 构造 dummy input: (1, 1, 40, 120) dummy_input torch.randn(1, 1, 40, 120) # 导出为 TorchScript traced_model torch.jit.trace(model, dummy_input) traced_model.save(crnn_traced.pt) # 验证导出正确性 loaded torch.jit.load(crnn_traced.pt) output loaded(dummy_input) # (1, 32, 37) print(Export success:, output.shape)逻辑说明torch.jit.trace比script更简单适用于无控制流的模型dummy_input必须与实际输入 shape 一致导出后loaded是torch.jit.ScriptModule可直接forward()无需model.eval()。6.2 FastAPI Web API轻量、异步、自带文档用 FastAPI 封装推理支持并发请求。核心是torch.no_grad()model(input)并用Base64编码图片传输。# api/main.py from fastapi import FastAPI, UploadFile, File from pydantic import BaseModel import torch import base64 import numpy as np from PIL import Image import io app FastAPI() model torch.jit.load(crnn_traced.pt) model.eval() CHARSET 0123456789ABCDEFGHJKLMNPQRSTUVWXYZ app.post(/predict) async def predict(file: UploadFile File(...)): contents await file.read() img Image.open(io.BytesIO(contents)).convert(L) # Resize to 120x40 img img.resize((120, 40), Image.Resampling.LANCZOS) img_tensor torch.tensor(np.array(img), dtypetorch.float32).unsqueeze(0).unsqueeze(0) / 255.0 with torch.no_grad(): logits model(img_tensor) # (1, 32, 37) # CTC decode probs torch.nn.functional.log_softmax(logits, dim2) decoded torch.nn.functional.ctc_decode( probs.permute(1,0,2), input_lengthstorch.tensor([32]), blank0, zero_infinityTrue ) pred_text .join([CHARSET[i-1] for i in decoded[0][0].tolist() if i 0]) return {prediction: pred_text} # 启动uvicorn api.main:app --reload参数说明Image.Resampling.LANCZOS比BILINEAR锐利保留字符边缘/255.0归一化必须做因训练时用transforms.Normalizectc_decode返回元组(list of tensors, list of scores)取decoded[0][0]即预测序列。6.3 精度验证不只看 accuracy要拆解 error type线上效果不能只报一个92.7%要定位瓶颈。我们用混淆矩阵 error 分类Error TypeExampleCauseFixSubstitutionA8→A08和0在扭曲时相似增加0/8字体变体合成InsertionA8→AA8模型多输出一个A检查 CTC blank 概率降低blank学习率DeletionA8→A模型跳过8增加8在合成数据中的出现频率TranspositionA8→8A字符顺序错检查 CNN 特征图时间维度是否被池化破坏已用AdaptiveAvgPool2d修复# eval/analyze_errors.py from sklearn.metrics import confusion_matrix import seaborn as sns import matplotlib.pyplot as plt # 获取所有预测和真实标签 preds, truths [], [] for img, target in test_loader: pred model.predict(img) # your predict func preds.extend(pred) truths.extend([CHARSET[t-1] for t in target]) # decode target # 生成混淆矩阵只统计单字符错误 char_errors [] for p, t in zip(preds, truths): if len(p) ! len(t): continue # skip length errors for pi, ti in zip(p, t): if pi ! ti: char_errors.append((ti, pi)) # 绘制 top-10 error pairs error_df pd.DataFrame(char_errors, columns[True, Pred]) conf_mat pd.crosstab(error_df[True], error_df[Pred]) sns.heatmap(conf_mat, annotTrue, fmtd) plt.savefig(confusion_matrix.png)逻辑说明pd.crosstab比sklearn.confusion_matrix更直观显示字符级错误if len(p) ! len(t)过滤长度错误聚焦 substitutionchar_errors列表便于人工分析高频错误对。从那以后我每次上线新模型都强制走一遍这三步用torch.jit.trace导出并验证dummy_input输出 shape在 FastAPI 里加logging.info(fInput shape: {img_tensor.shape})确认预处理无误抓取线上 100 张失败样本人工归类 error type针对性补数据。这套流程让我在三个不同验证码项目里首次部署准确率都超过 85%没有一次需要推倒重来。希望帮到你。本文还有配套的精品资源点击获取
阅读完成 · 觉得有帮助?
咨询建站