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

从MNIST到手写汉字识别:CNN训练与避坑实战

从MNIST到手写汉字识别:CNN训练与避坑实战 ★ FEATURED ARTICLE
简介这份压缩包是面向深度学习与计算机视觉初学者的手写汉字识别测试资源聚焦利用深度卷积网络解决手写汉字自动识别问题适合正在学习模式识别、神经网络或相关交叉应用的开发者参考。包内仅含一个脚本文件压缩后大小约一千字节脚本虽精简却覆盖了数据加载与归一化、图像尺寸调整、卷积模型定义、损失函数与优化器选择、训练验证循环及模型保存等关键环节可帮助读者快速理清深度卷积网络处理图像数据的基本流程。已有177人浏览学习该资源。通过阅读脚本能直观看到如何将手写数字数据集风格的数据处理思路迁移到汉字样本上体会网络层设计对识别效果的影响也可为后续扩展更大规模训练集、调整网络结构或部署实际识别应用提供起点。压缩包尤其适合作为课程设计或入门项目的代码参考便于快速启动自己的手写汉字识别实验。1. chinese_test.zip 是什么手写汉字识别跟 MNIST 差在哪跑通了 MNIST 手写数字识别的朋友第一次把手写汉字抛给模型时通常会愣一下之前那套 90% 以上准确率的流程换到汉字上连 60% 都吃力。chinese_test.zip 这种手写汉字数据集本质上是把 MNIST 里的 0-9 换成了成百上千个汉字类别任务却从 10 选 1 变成了几百选 1——难的不只是数量汉字的笔画结构、形近字、书写自由度全都堆在一起。这里就沿着解包、预处理、训练、评估这条链路把能直接复现的代码和参数写出来也把最容易翻车的几个坑提前标记好。适合正在做单据识别、作业批改、手写输入相关需求或者想从 MNIST 进阶到真实中文场景的开发者和算法工程师。2. 从 zip 解包到可训练张量手写汉字数据的预处理全流程2.1 先看清目录和标签训练集/测试集怎么组织在手写汉字识别这个任务上数据集的组织方式比数字识别更值得花十分钟先确认。chinese_test.zip 虽带 test 字样但实际这类包通常是“训练集 测试集 说明”打包在一起解压后我一般先执行一条 find 命令看目录骨架再决定怎么写加载器mkdir -p data unzip chinese_test.zip -d data/ find data -maxdepth 2 -type d | head -20 # 期望看到 train/ 与 test/ 目录train 下每个子目录名是一个汉字类别注意 maxdepth 别太大类别多时目录树很深一次拉太多会把终端刷爆。常见做法是train/类别名/图片文件类别名用汉字本身命名比用编号命名友好得多后面做评估时可以直接读目录名还原真实标签。如果解出来不是目录结构而是序列化格式如 .mat、.npy、.pkl也不用慌多用一步读取转换即可。这种格式里 labels 通常是一维整数数组问题在于这些整数对应的是 Unicode 码点还是数据集的内部编号得单独打印几个值比对过再确认。我见过有人把码点直接当类别数传进模型最后num_classes设成了 19968训练直接爆显存。还有一类隐藏问题训练集和测试集的字符集可能不一致。测试时遇到训练里没出现过的字映射表里查不到推理程序直接抛 KeyError。稳妥的做法是构建映射时保留一个unknown槽位预测时输出置信度和对应的未知标记交给人工兜底。2.2 图片读入与统一尺寸灰度、缩放、归一化的边界条件手写汉字图片来源五花八门扫描件、手机拍摄、手写板导出的灰度图都有统一尺寸是训练前必须做的事。我一般把输入固定成 96x96比 MNIST 的 28x28 大得多因为汉字的笔画密度高28 像素会丢掉关键结构信息也不需要到 224那是为 ImageNet 大类目设计的对单字分类是浪费算力。“统一到多少”这件事没有标准答案但 64 以下基本不可用128 以上收益很小96 是性价比最高的中间值。读图时先转灰度再缩放。颜色信息对手写汉字几乎没有判别力反而增加过拟合和训练成本from PIL import Image import numpy as np def load_hanzi_image(path, size(96, 96)): img Image.open(path).convert(L) # 强制转灰度避免误入 RGB 分支 img img.resize(size, Image.LANCZOS) # 抗锯齿缩放笔画边缘更平滑 arr np.asarray(img, dtypenp.float32) arr arr / 255.0 # 归一化到 [0,1] return arr这里 resize 的插值方式值得说一句LANCZOS 在缩小时能保留较多笔画细节但也可能让噪声更明显如果发现训练集过度拟合、验证集掉点换成 BILINEAR 往往更稳。归一化直接除 255 在图像任务里足够不用像表格数据那样做标准化第一层 BN 会自己消化分布偏移。几万到几十万张图每个 epoch 都从磁盘读再 resize速度会拖垮训练。实操中我会在预处理后把结果缓存成 npy 文件一次转换终身使用cache np.memmap(train_cache.npy, dtypenp.float32, modew, shape(len(samples), 96, 96)) for i, path in enumerate(samples): cache[i] load_hanzi_image(path) cache.flush()memmap 的好处是数据太大时不用一次性全载入内存训练时按索引切片读取IO 压力也小。注意这个缓存文件不要放进 git它通常有几个 GB 大小。2.3 标签编码为什么别用自以为方便的 0~N 乱序编号血泪经验很多人在读目录时用 enumerate 给类别编号训练时看起来没事一旦要部署做推理逆映射就成了玄学。常见坑是两次运行 enumerate 的顺序不一致——文件系统返回顺序本身不保证稳定这就导致“训练时的 ID 3”和“部署时的 ID 3”不是同一个字。更稳的方式是用排序后的目录列表建映射。每个汉字对应一个稳定整数逆映射也不需要额外编码表import os, json def build_label_map(root): classes sorted(os.listdir(root)) # 固定顺序排除文件系统抖动 char_to_id {c: i for i, c in enumerate(classes)} id_to_char {i: c for c, i in char_to_id.items()} return char_to_id, id_to_char char_to_id, id_to_char build_label_map(data/train) with open(label_map.json, w, encodingutf-8) as f: json.dump({char_to_id: char_to_id, id_to_char: id_to_char}, f, ensure_asciiFalse)sorted 固定了遍历顺序class 名是汉字而不是目录编号时从 id_to_char 还原标签就是一次查表。映射表存成 json 随模型一起保存加载模型时必须带上否则预测结果根本没法解。后续要加新类别也在同一张表上追加并更新版本号避免线上模型和映射表错位。3. 用 CNN 做手写汉字识别模型结构、损失函数与训练参数3.1 从 MNIST 思路迁移过来要做哪些调整很多人把 LeNet 从 MNIST 直接搬到汉字上发现验证集准确率卡在 85% 上下不去。原因不是网络不行而是任务难度差了一个量级MNIST 是 10 选 1汉字常用字至少 3755 选 1GB2312 一级字表公开手写汉字数据集大多基于这个范围类间相似度又高“己/已/巳”“未/末”这类字形差异只有几像素。LeNet 的参数规模撑不起这种判别任务。我的调整思路有三层把卷积通道加宽64 → 128 → 256让浅层能容纳更多笔画组合特征每个卷积后接 BN加速收敛并稳住分布全连接前加 Dropout防止分类头在几万个样本上过拟合。深度上不着急堆到 ResNet 级别手写单字是结构相对简单的居中目标过深的网络很容易在小数据集上过拟合而且训练时间翻倍收益有限。损失函数用 CrossEntropyLoss 就够了它内部已经做了 softmax 和 log 运算不需要在模型输出层额外接 softmax。类别不均衡时给 loss 传入 weight 向量后面避坑章会展开。3.2 baseline一个能跑到 95% 以上的轻量 CNN 结构下面这个网络约 1.5M 参数在 96x96 输入上单卡能比较稳定地收敛我拿它作为手写汉字识别的 baseline比 LeNet 改版稳得多import torch.nn as nn class HanziCNN(nn.Module): def __init__(self, num_classes): super().__init__() self.features nn.Sequential( nn.Conv2d(1, 64, 3, padding1), nn.BatchNorm2d(64), nn.ReLU(), nn.Conv2d(64, 64, 3, padding1), nn.BatchNorm2d(64), nn.ReLU(), nn.MaxPool2d(2), # 96 - 48 nn.Conv2d(64, 128, 3, padding1), nn.BatchNorm2d(128), nn.ReLU(), nn.Conv2d(128, 128, 3, padding1), nn.BatchNorm2d(128), nn.ReLU(), nn.MaxPool2d(2), # 48 - 24 nn.Conv2d(128, 256, 3, padding1, groups64), nn.BatchNorm2d(256), nn.ReLU(), nn.AdaptiveAvgPool2d(1), # 24 - 1x1全局池化 ) self.head nn.Sequential( nn.Dropout(0.3), nn.Linear(256, num_classes) ) def forward(self, x): x self.features(x) return self.head(x.view(x.size(0), -1))第三组卷积用了 groups64 的深度可分离卷积把通道翻到 256 的同时控住参数量。全连接层直接到 num_classes如果类别数是 3755那这层约 96 万参数是模型里最大的一块。想进一步瘦身可以把 head 换成两层瓶颈全连接或者用 GlobalMaxPool 替代 AdaptiveAvgPool但后者对笔画响应更敏感训练波动也更大新手不建议首选。输入是 96x96 单通道所以第一层 Conv 的 in_channels1。如果你的数据是三通道读进来的要么在预处理里转灰度要么把第一层改成 3显存和训练时间都会涨一点。3.3 训练配置学习率、批次、早停与模型保存训练阶段我通常采用 AdamW 余弦退火 早停的组合对手写汉字这种中等规模分类任务比裸 SGD 收敛更稳。batch size 在 96x96 输入下单卡 128 到 256 之间合适大于 256 时 BN 的统计量会被稀释小类别学不动。下面是核心训练循环的简化版optimizer torch.optim.AdamW(model.parameters(), lr2e-3, weight_decay1e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max20) criterion nn.CrossEntropyLoss() for epoch in range(epochs): model.train() for images, labels in train_loader: images, labels images.cuda(), labels.cuda() optimizer.zero_grad() out model(images) loss criterion(out, labels) loss.backward() optimizer.step() scheduler.step() val_acc, val_loss evaluate(model, val_loader) if val_loss best_val_loss: best_val_loss val_loss torch.save({model: model.state_dict(), char_to_id: char_to_id, id_to_char: id_to_char}, best_hanzi.ckpt)几个关键参数的作用lr2e-3 配合 AdamW 不用手工衰减weight_decay1e-4 对大 FC 层有约束力能压住过拟合CosineAnnealing 的 T_max20等于让学习率走完半个余弦周期最后几个 epoch 会以极小学习率精修。早停条件是验证集 loss 连续 5 个 epoch 不下降就停防止在长尾类别上反复震荡。保存模型时把映射表一并存进去这是部署时最容易被忽略又最后悔的事。到这里模型能跑了但训练过程中有大量隐形坑下一章集中说。4. 避坑手写汉字识别项目里 5 个高频翻车点4.1 训练集 99%真实手写一测就只剩 70%现象模型在验证集上准得惊人一到真实手写板、手机拍的单字就原形毕露。原因公开数据集里的手写样本是“工整手写”都尽量居中、无背景、笔画完整真实场景有连笔、倾斜、曝光不均、边框干扰数据分布完全不同。解决训练结束后立刻收集 20~50 个真实样图做冒烟测试别只看验证集指标。上线阶段用真实样本做增强或微调比调网络结构有效得多。真实样本可以拍照、可以手写板导出关键是它们要来自目标任务的实际输入渠道而不是从测试集里挑的。4.2 长尾类别学不动loss 降不下来的元凶现象整体准确率尚可但低频字几乎永远预测错验证集 loss 在某个数值附近卡住不动。原因数据集类别分布幂律化高频字样本上千低频字样本可能只有几十个。CrossEntropy 对所有类别一视同仁长尾类的梯度被高频类淹没。解决在采样器上做文章。用 WeightedRandomSampler 按类别样本数倒数设权重让每个 epoch 里低频字被抽到更多次from torch.utils.data import WeightedRandomSampler counts np.bincount(labels, minlengthnum_classes) weights (1.0 / counts.astype(np.float32)) ** 0.5 weights weights / weights.sum() sampler WeightedRandomSampler(weights, num_sampleslen(labels), replacementTrue)权重取 0.5 次幂是为了松弛直接用倒数会让高频字欠拟合。num_samples 保持和数据集长度一致即可replacementTrue 允许重复采样低频字。如果显存够把采样器的效果和 loss 权重叠加使用长尾类别通常能再涨几个点。4.3 相似字卷成麻花未/末、己/已/巳在混淆矩阵里扎堆现象混淆矩阵里高频错误组合高度集中在笔画数接近、结构相似的汉字上模型对“横的长短差异”这种像素级差别分不开。原因数据里本身存在标注错误和书写模糊模型又在特征层面把相似结构映射到了相邻的向量空间二者叠加导致难例扎堆。解决先量化再决定策略。统计 Top-K 混淆对人工确认是标注噪声还是模型没学会。若是模型问题在 loss 里加 CenterLoss 拉紧类内距离或对相似字做局部笔画裁切的专门增强通常能拉回几个点。不要一上来就换大模型先看混淆对长什么样。4.4 显存够用但 batch 跑到一半 OOM现象训练进行到第几步就报 CUDA out of memory准确率还没起来。原因96x96 输入加 256 通道的中间特征图占显存不少全连接层的反向传播也要存梯度类别越多分类头的梯度矩越大显存峰值比预想高得多。解决batch size 从 256 降到 128配合梯度累积模拟更大的 batchaccum_steps 4 optimizer.zero_grad() for i, (images, labels) in enumerate(train_loader): loss criterion(model(images), labels) / accum_steps loss.backward() if (i 1) % accum_steps 0: optimizer.step() optimizer.zero_grad()这样既不牺牲等效 batch 的稳定性又能把显存峰值压下一截。注意 loss 要除以 accum_steps让梯度等效于原 batch 的平均。如果还是 OOM把输入尺寸临时缩到 80x80 跑一次实验确认收益再决定要不要换硬件。4.5 验证集准确率莫名掉 20%loss 却正常现象某一次重跑实验严格按验证集评估loss 正常acc 骤降甚至低到随机水平。原因多半是并行 DataLoader 打乱了数据而你把 image 和 label 分开两个 list 传递打乱后二者不再对齐。这种 bug 不会报错只会默默垃圾化训练结果最难查。解决在 Dataset 的__getitem__里同时返回 image 和 label用同一个索引取数永远不要分开传列表。数据加载器的 shuffle 影响的是样本顺序不该影响样本配对。写完加载器后先用一小批数据打印 image 和 label 对应的文件路径人工核对一次再开训。5. 评估与验证手写汉字模型怎么算“真的能用”5.1 按字拆准确率别被平均数骗了chinese_test.zip 这类数据集的测试集字频分布往往不均匀。整体准确率 95% 可能是高频字全对、低频字全错的结果这个指标放在业务里会掩盖严重问题。每类单独算准确率后画一条降序曲线比一个平均数有价值得多from sklearn.metrics import accuracy_score class_acc {} for char in test_classes: mask (test_labels char_to_id[char]) class_acc[char] accuracy_score(test_true[mask], test_pred[mask]) worst sorted(class_acc.items(), keylambda x: x[1])[:20] for char, acc in worst: print(f{char}: {acc:.3f})用目录名做标签的价值在这体现按字展示时能直接看出哪些字出问题而不是面对一堆整型 ID 无从下手。高频字要求 98% 以上低频字 80% 以上才敢往粗粒度场景推。如果某个中频字准确率只有 60%优先查训练样本量和标注噪声别急着调模型。5.2 混淆矩阵里找形近字模型是“没学会”还是“认不清”整体准确率之外形近字的混淆结构更值得研究。用 sklearn 的 confusion_matrix 输出 3755 类不可能逐格看我一般只挑对角线附近的高混淆对再用汉字的笔画数或部首做分组def top_confusions(conf, id_to_char, k30): np.fill_diagonal(conf, 0) flat conf.flatten() idx np.argsort(flat)[::-1][:k] for i in idx: r, c divmod(i, conf.shape[1]) print(f{id_to_char[r]} - {id_to_char[c]}: {flat[i]})如果混淆对高度集中在形近字上模型处在“特征模糊区”需要补难例或加 CenterLoss如果错误分散到无关字才要考虑结构性问题比如标签错乱、数据加载 bug。这一步相当于给模型做体检别等部署了才头疼。人工核对 Top-K 混淆对时把对应样本图打出来一起看能直接判断是标注错还是模型错。5.3 用真实手写板数据做冒烟测试评估集之外的第二道关公开测试集和真实场景总有 gap。一个成本不高但能救命的做法自己准备 20 张圆珠笔、铅笔、手写屏写出的单字照片训练完当天跑一次批量推理记录 Top-1 肉眼能认对几个、Top-5 才认对几个model.eval() with torch.no_grad(): for img_path in real_samples: img load_hanzi_image(img_path) out model(torch.Tensor(img).unsqueeze(0).unsqueeze(0).cuda()) top5 torch.topk(out, 5).indices.squeeze().cpu().numpy() print(img_path, [id_to_char[i] for i in top5])这里的 unsqueeze(0) 是为了加 batch 维和通道维load_hanzi_image 已经返回 HxW 的单通道数组。冒烟测试不用追求数量关键是样本要来自目标任务的真实输入渠道。我见过模型在数据集上 97%拍一张纸上的字就掉到 75% 的项目原因是训练集图片全经过二值化真实照片带灰度渐变和阴影网络直接懵了。冒烟测试结果记录下来每次改动模型后重跑一遍比盯着 tensorboard 更贴近业务。6. 进阶把识别器推进业务的数据增强与轻量化6.1 针对汉字的数据增强旋转别超过 10°糊一点会更强汉字是有严格方向性的文字。旋转超过 15° 的增强会主动制造“倒字”“斜字”样本把模型教糊涂。我常用的一组增强参数旋转 -8° 到 8°平移 ±10%缩放 0.9 到 1.1亮度扰动 0.8 到 1.2加少量高斯模糊import albumentations as A train_transform A.Compose([ A.Rotate(limit8, border_mode0), A.ShiftScaleRotate(shift_limit0.1, scale_limit(0.9, 1.1), rotate_limit0), A.RandomBrightnessContrast(brightness_limit0.2), A.GaussBlur(blur_limit(3, 5)), ])Rotate 的 border_mode0 填充黑色比默认反射填充更贴近真实纸面背景。高斯模糊虽然会降低训练集准确率但能显著提升对手机拍摄模糊图画的鲁棒性这是血泪经验换来的。增强只加在训练侧验证和推理保持原始输入。6.2 模型轻量化量化与端侧部署的取舍要把识别器塞进手机或嵌入式设备参数量和计算量就得压下来。常用路径是把 HanziCNN 的卷积宽度减半64→32256→128再配合 8-bit PTQ 量化参数量能降到原来的四分之一左右Top-1 掉点在 1~2 个点以内。不同 CPU、不同框架差距很大以下只是量级参考方案参数量单张推理延迟Top-1 掉点原版 HanziCNN基线基线基线宽度减半版本约 1/2约 0.7 倍-0.5% 左右宽度减半 8-bit 量化约 1/4约 0.4 倍-1.0% 到 -2.0%6.3 增量更新新字形来了怎么补手写识别没有“训完就完”一说。业务上线后新用户、新书写工具会不断产生模型没见过的新字形。正确姿势是把它们按置信度分级低置信度样本抽人确认确认后进重训练集用原模型权重继续训练几个 epoch学习率降到 1e-4 左右只微调分类头和最后两层卷积避免灾难性遗忘。把 chinese_test.zip 这类数据集真正吃透收获的不只是模型精度而是面对真实手写场景时判断“该调数据还是调模型”的直觉。我最深的教训是一个手写识别模型好不好用不看测试集刷到多少分而看你肯不肯挖混淆矩阵、敢不敢用真实写坏的字把它击穿。标签映射表和冒烟样本记得留好它们比一轮精调更能救项目于水火。希望帮到你。本文还有配套的精品资源点击获取
阅读完成 · 觉得有帮助?
咨询建站