从零训练自己的验证码识别模型:数据生成到部署

验证码(CAPTCHA)是互联网最常见的人机验证手段之一,从早期的字符扭曲验证码到后来的点选、滑动、行为验证,技术形态不断演进。其中字符型验证码因实现简单、部署成本低,至今仍被大量网站使用。对于需要自动化业务流程的场景,训练一个专属的验证码识别模型,比调用第三方接口更灵活、成本更低。

本文将带你完整走通字符验证码识别的全链路:从合成训练数据、构建数据集、设计模型、训练调优,到最终封装成可调用的 API 服务。全程基于开源工具链,代码可复现,适合作为 CV 入门工程实践。

一、整体技术路线

字符验证码识别本质是不定长序列识别问题。一张图片包含 4-6 个字符,位置不固定、有干扰线和噪点,无法用普通单分类模型解决。

工业界主流方案是 CNN(特征提取)+ RNN/Transformer(序列建模)+ CTC Loss(对齐损失) 的端到端架构,不需要提前分割字符,输入整张图片直接输出字符序列。

我们的技术栈:

  • 数据生成:Python captcha 库 + OpenCV 数据增强
  • 深度学习框架:PyTorch
  • 模型骨干:ResNet18 + BiLSTM + CTC
  • 部署导出:ONNX 格式
  • 推理服务:FastAPI + Uvicorn

二、第一步:批量生成验证码数据集

训练模型的第一步是数据。真实验证码标注成本高,而字符验证码的样式规则明确,用合成数据完全可以训练出可用模型。

2.1 安装依赖

复制代码
pip install captcha opencv-python numpy pillow torch torchvision

2.2 生成脚本

我们生成常见的4 位数字 + 字母验证码,加入干扰线、噪点、字符扭曲,模拟真实网站的风格。

复制代码
import os
import random
import string
from captcha.image import ImageCaptcha
import numpy as np
import cv2

# 字符集:数字 + 大小写字母(也可去掉易混淆的 0/O、1/l)
CHARSET = string.digits + string.ascii_uppercase
NUM_CLASSES = len(CHARSET)
CHAR2IDX = {c: i for i, c in enumerate(CHARSET)}
IDX2CHAR = {i: c for i, c in enumerate(CHARSET)}

# 验证码参数
WIDTH, HEIGHT = 160, 60
CHAR_LEN = 4  # 4位字符

def generate_one_captcha():
    """生成单张验证码图片和对应标签"""
    label = ''.join(random.choices(CHARSET, k=CHAR_LEN))
    generator = ImageCaptcha(width=WIDTH, height=HEIGHT,
                             font_sizes=[40, 45, 50])
    img = generator.generate_image(label)
    
    # 转OpenCV格式,方便后续增强
    img_cv = cv2.cvtColor(np.array(img), cv2.COLOR_RGB2BGR)
    return img_cv, label

def generate_dataset(output_dir, count=10000):
    """批量生成数据集"""
    os.makedirs(f"{output_dir}/images", exist_ok=True)
    label_file = open(f"{output_dir}/labels.txt", "w", encoding="utf-8")
    
    for i in range(count):
        img, label = generate_one_captcha()
        filename = f"{i:06d}.png"
        cv2.imwrite(f"{output_dir}/images/{filename}", img)
        label_file.write(f"{filename}\t{label}\n")
        
        if (i+1) % 1000 == 0:
            print(f"已生成 {i+1}/{count} 张")
    
    label_file.close()
    print(f"数据集生成完成,共 {count} 张,保存在 {output_dir}")

# 生成训练集1万张、测试集1千张
generate_dataset("data/train", 10000)
generate_dataset("data/test", 1000)

2.3 数据增强策略

仅靠生成的数据容易过拟合,需要加入随机扰动提升泛化能力:

  • 随机亮度、对比度调整
  • 高斯噪声、椒盐噪点
  • 随机平移、轻微缩放
  • 增加 / 减少干扰线数量

建议在训练时用 torchvision.transforms 在线增强,避免磁盘占用。

三、第二步:构建数据加载器

3.1 数据集类

复制代码
import torch
from torch.utils.data import Dataset, DataLoader
from torchvision import transforms
from PIL import Image

class CaptchaDataset(Dataset):
    def __init__(self, label_path, img_dir, transform=None):
        self.data = []
        with open(label_path, "r", encoding="utf-8") as f:
            for line in f:
                fname, label = line.strip().split("\t")
                self.data.append((fname, label))
        self.img_dir = img_dir
        self.transform = transform or transforms.ToTensor()

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

    def __getitem__(self, idx):
        fname, label = self.data[idx]
        img = Image.open(f"{self.img_dir}/{fname}").convert("L")  # 转灰度
        img = self.transform(img)
        
        # 标签转编码
        target = [CHAR2IDX[c] for c in label]
        return img, torch.tensor(target, dtype=torch.long)

3.2 CTC 所需的 target 处理

CTC Loss 需要知道每个样本的真实长度。训练时我们用 target_lengths 记录,这里所有样本都是 4 位,直接填充即可。

复制代码
train_transform = transforms.Compose([
    transforms.Resize((60, 160)),
    transforms.RandomAffine(degrees=2, translate=(0.02, 0.02)),
    transforms.ColorJitter(brightness=0.3, contrast=0.3),
    transforms.ToTensor(),
])

train_set = CaptchaDataset("data/train/labels.txt", "data/train/images", train_transform)
test_set = CaptchaDataset("data/test/labels.txt", "data/test/images")

train_loader = DataLoader(train_set, batch_size=64, shuffle=True, num_workers=2)
test_loader = DataLoader(test_set, batch_size=128, shuffle=False, num_workers=2)

四、第三步:模型结构设计

4.1 整体架构

  • Backbone:ResNet18 提取视觉特征,输出特征图序列
  • 序列层:BiLSTM 捕捉字符间上下文依赖
  • 输出层:线性层映射到字符集维度 + softmax
  • 损失函数:CTCLoss(处理字符对齐问题)

4.2 模型代码

复制代码
import torch.nn as nn
import torch.nn.functional as F
from torchvision.models import resnet18

class CaptchaModel(nn.Module):
    def __init__(self, num_classes, hidden_size=128):
        super().__init__()
        
        # 用ResNet18做特征提取,去掉最后两层
        resnet = resnet18(weights=None)
        self.backbone = nn.Sequential(
            resnet.conv1,
            resnet.bn1,
            resnet.relu,
            resnet.maxpool,
            resnet.layer1,
            resnet.layer2,
            resnet.layer3,
            resnet.layer4,  # [B, 512, 2, 5]
        )
        
        # 特征维度压缩
        self.fc1 = nn.Linear(512 * 2, hidden_size)
        
        # 双向LSTM
        self.lstm = nn.LSTM(hidden_size, hidden_size, 
                            bidirectional=True, 
                            num_layers=2,
                            batch_first=True)
        
        # 输出层
        self.fc2 = nn.Linear(hidden_size * 2, num_classes + 1)  # +1是blank

    def forward(self, x):
        # x: [B, 1, 60, 160] -> 灰度图需要先转3通道
        x = x.repeat(1, 3, 1, 1)
        feat = self.backbone(x)  # [B, 512, H, W]
        
        B, C, H, W = feat.shape
        # 把高度维度压平,宽度作为序列长度
        feat = feat.permute(0, 3, 1, 2).contiguous()  # [B, W, C, H]
        feat = feat.view(B, W, -1)  # [B, seq_len, C*H]
        
        feat = self.fc1(feat)
        feat, _ = self.lstm(feat)
        out = self.fc2(feat)  # [B, seq_len, num_classes+1]
        
        # CTC需要 [T, B, C] 格式
        return out.permute(1, 0, 2).log_softmax(2)

4.3 为什么用 CTC Loss?

验证码中字符位置不固定,我们不知道第几个特征帧对应第几个字符。CTC(Connectionist Temporal Classification)可以自动对齐输入序列和输出标签,不需要字符级标注,非常适合 OCR、语音识别这类序列任务。

五、第四步:模型训练

5.1 训练循环

复制代码
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = CaptchaModel(NUM_CLASSES).to(device)
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
criterion = nn.CTCLoss(blank=NUM_CLASSES, zero_infinity=True)

def train_one_epoch():
    model.train()
    total_loss = 0
    for imgs, targets in train_loader:
        imgs = imgs.to(device)
        targets = targets.to(device)
        
        B = imgs.size(0)
        input_lengths = torch.full((B,), 5, dtype=torch.long).to(device)  # 序列长度=5
        target_lengths = torch.full((B,), CHAR_LEN, dtype=torch.long).to(device)
        
        optimizer.zero_grad()
        outputs = model(imgs)
        loss = criterion(outputs, targets, input_lengths, target_lengths)
        loss.backward()
        optimizer.step()
        
        total_loss += loss.item()
    return total_loss / len(train_loader)

5.2 解码与准确率计算

模型输出是概率序列,需要用贪心解码 或束搜索转成最终字符串。

复制代码
def decode(output):
    """CTC贪心解码:去重+去blank"""
    argmax = output.argmax(-1).permute(1, 0).cpu().numpy()  # [B, T]
    results = []
    for seq in argmax:
        chars = []
        prev = -1
        for idx in seq:
            if idx != prev and idx != NUM_CLASSES:  # 去掉blank和重复
                chars.append(IDX2CHAR[idx])
            prev = idx
        results.append(''.join(chars))
    return results

@torch.no_grad()
def evaluate():
    model.eval()
    correct = 0
    total = 0
    for imgs, targets in test_loader:
        imgs = imgs.to(device)
        outputs = model(imgs)
        preds = decode(outputs)
        gts = [''.join(IDX2CHAR[i.item()] for i in t) for t in targets]
        
        for p, g in zip(preds, gts):
            if p == g:
                correct += 1
            total += 1
    return correct / total

5.3 训练效果

  • 1 万张训练集,batch_size=64,在单张 GPU 上训练 30 轮左右
  • 简单验证码准确率可以达到 98%+
  • 干扰多、字符重叠严重的场景,准确率会降到 90% 左右,需要增加数据量和增强强度

六、第五步:模型优化技巧

如果准确率达不到预期,可以从以下方向优化:

  1. 扩大数据集:从 1 万增到 5 万 - 10 万张,覆盖更多字体、干扰样式
  2. 加入真实样本微调:用几十到几百张真实标注样本微调模型,泛化性会大幅提升
  3. 更换骨干网络:用更深的 ResNet34/50,或者替换为 MobileNetV3 轻量化
  4. 引入注意力机制:在 LSTM 后加 Attention 层,或者直接用 Transformer 替换 LSTM
  5. 多尺度训练:随机缩放图片宽度,增强对不同字符间距的适应性

七、第六步:模型导出与部署

训练好的模型需要导出成通用格式,才能在生产环境调用。

7.1 导出 ONNX

复制代码
model.eval()
dummy_input = torch.randn(1, 1, 60, 160).to(device)
torch.onnx.export(
    model, dummy_input, "captcha.onnx",
    input_names=["image"],
    output_names=["output"],
    dynamic_axes={"image": {0: "batch"}, "output": {1: "batch"}}
)
print("ONNX模型导出成功")

7.2 搭建推理 API

用 FastAPI 快速封装成 HTTP 服务:

复制代码
# app.py
import cv2
import numpy as np
import onnxruntime as ort
from fastapi import FastAPI, UploadFile, File
import uvicorn

app = FastAPI(title="验证码识别服务")
session = ort.InferenceSession("captcha.onnx", providers=["CPUExecutionProvider"])

CHARSET = "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZ"

def decode_output(output):
    argmax = output.argmax(-1)[0]
    chars = []
    prev = -1
    for idx in argmax:
        if idx != prev and idx != len(CHARSET):
            chars.append(CHARSET[idx])
        prev = idx
    return ''.join(chars)

@app.post("/recognize")
async def recognize(file: UploadFile = File(...)):
    # 读取图片
    content = await file.read()
    nparr = np.frombuffer(content, np.uint8)
    img = cv2.imdecode(nparr, cv2.IMREAD_GRAYSCALE)
    
    # 预处理
    img = cv2.resize(img, (160, 60))
    img = img.astype(np.float32) / 255.0
    img = img[np.newaxis, np.newaxis, :, :]
    
    # 推理
    output = session.run(["output"], {"image": img})[0]
    result = decode_output(output)
    
    return {"code": 0, "result": result}

if __name__ == "__main__":
    uvicorn.run(app, host="0.0.0.0", port=8000)

启动服务:

复制代码
python app.py

调用方式:

复制代码
curl -X POST -F "file=@captcha.png" http://localhost:8000/recognize

单张图片推理耗时在 CPU 上约 10-30ms,完全满足日常自动化需求。

八、总结与进阶方向

本文完整实现了数据生成 → 模型训练 → 部署上线的验证码识别全流程。这套方案对标准字符验证码效果显著,工程上可直接落地。

如果需要应对更复杂的验证码,可以继续深入:

  • 点选 / 文字验证码:改用目标检测模型(YOLO)定位字符,再做识别
  • 滑块验证码:用 CV 模板匹配或分割模型计算缺口位置
  • 对抗验证:加入对抗样本训练,提升模型鲁棒性
  • 联邦学习:多站点联合训练,保护数据隐私

验证码识别是一场攻防博弈,没有一劳永逸的方案。但掌握了从数据到部署的完整方法论后,面对新型验证方式也能快速迭代出解决方案。

相关推荐
wuyk5551 天前
Python 网络爬虫入门到实战 第 04 章:请求头、UA 伪装、超时、异常处理、基础反爬绕过
开发语言·爬虫·python
weixin_440401691 天前
质朴的爬虫+数据处理
爬虫·python·数据分析·pandas
沙漠之主2 天前
Python 教学设计资料:从入门到实战的完整课程方案
爬虫·python
深蓝电商API2 天前
用YOLO训练验证码目标检测模型的全流程
爬虫·yolo·目标检测·目标跟踪
wuyk5553 天前
Python网络爬虫入门到实战 第02章:HTTP/HTTPS协议超通俗精讲(GET/POST、请求头、响应码、爬虫核心基础)
爬虫·python·http
szial4 天前
网络爬虫与 CDP 实战(二):浏览器能访问,Python 却被拒绝?拆开登录态与 CSRF
爬虫·python·csrf
梅雅达编程笔记4 天前
04-Python CSV数据保存与翻页抓取
开发语言·爬虫·python·pandas·数据采集·csv
梅雅达编程笔记4 天前
03_网页表格数据抓取
爬虫·python·beautifulsoup·pandas·数据采集
傻啦嘿哟4 天前
爬虫代理IP池从0到1:构建高可用代理池,彻底解决IP被封问题
网络·爬虫·tcp/ip