ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

基于TensorFlow+ResNet的EMNIST手写字母数字识别实战解析

基于TensorFlow+ResNet的EMNIST手写字母数字识别实战解析 简介基于Python实现的字母数字识别项目采用TensorFlow2与EMNIST数据集在tf2环境中对ResNet网络进行简化实现用于手写英文字母与数字的分类识别。项目面向初入计算机视觉的开发者及课程设计场景代码结构清晰便于二次修改与实验复现。资源共50个文件压缩包约104.79MB主要包含Python源码、checkpoint训练权重、h5模型文件、测试图片、训练过程可视化的png图像以及README和字符映射txt等辅助资料。其中源码覆盖数据加载、模型定义、训练与推理测试四个环节checkpoint便于断点续训或直接调用png图片直观展示不同训练轮次下的准确率变化已有280人学习下载。借助预训练权重与演示脚本学习者可快速加载模型验证效果也能结合源码理解ResNet在TensorFlow2中的搭建细节、训练策略和参数调优思路是图像识别入门与课设参考的实用素材。1. 从一份 TF2 手写识别工程看字母数字识别的正确打开方式手写字母数字识别不是“能跑通 CNN 就算完”的玩具。真正把 EMNIST 数据集训到能交付需要同时处理字符集映射、ResNet 残差结构、checkpoint 续训和真实图片预处理。这份基于 Python TensorFlow 2.1 的项目正好把这几块串在一起mytrain.py 负责训练model.py 实现 ResNettest.py 做批量评估demo.py 跑单图预测。它适合刚完成 python 入门、想往计算机视觉方向走的人也适合做课程设计时需要完整技术栈的读者。下面按数据准备、模型实现、训练检查点、推理落地四个层次拆开讲重点是可复现的参数和容易翻车的地方。2. EMNIST 字符集与 ResNet 选型为什么不是 MNIST 也不是 VGG2.1 EMNIST 的结构与 characters.txt 的类别对齐MNIST 只有 10 类数字对“字母数字识别”来说信息量不够。EMNIST 是 MNIST 的扩展版本把英文字母也纳入了同一个评估体系。常见子集有 byclass、bymerge、digits、letters 四类。这个项目的训练数据基于 EMNIST输出层维度则跟着characters.txt走所以不能直接把 EMNIST 原始类别数写死。子集类别数样本数说明byclass62814255数字 大小写字母0/O、1/l 混淆最严重bymerge47814255大小写合并字母数字混合识别更实用digits10280000只有数字效果等价于 MNISTletters26145600只有字母一般作为单独任务使用项目里的characters.txt决定了类别顺序。读取时常见做法是每行一个字符跳过空行最后得到一个字符串列表。这个列表的索引就是训练标签。写代码时要注意编码Windows 10 下用 UTF-8 读取最稳。def load_chars(pathcharacters.txt): chars [] with open(path, r, encodingutf-8) as f: for line in f: line line.strip() if line: chars.append(line[0]) # 一行一个字符取首个非空白字符 return chars这段代码的逻辑是把字符顺序转成 Python 列表例如第 0 行是0那么标签0对应字符0第 10 行是A标签A的索引就是10。参数上唯一需要注意的是line[0]只取第一个字符如果characters.txt里写了中文注释或空格会直接污染类别表。实际项目里我会先做一次可见字符校验过滤掉\ufeff这类 BOM 头。训练和推理必须共用同一个characters.txt。很多复现项目准确率崩掉不是因为模型差而是 demo.py 里重新手写了一个字符列表顺序和训练时不一致最后预测结果全错。这个映射是整条链路的“基础设施”比网络结构更容易被忽略。2.2 ResNet 相对 VGG 在小尺寸图像上的优势EMNIST 的图片是 28×28 灰度图图像内容简单不需要超大感受野。VGG 风格堆叠 3×3 卷积在 28×28 输入上也能跑但参数和计算量明显冗余。以一个 64 通道的 VGG block 为例两个 3×3 卷积需要的参数量是 3×3×64×64×2而残差结构里相同通道数的两个 3×3 卷积参数量相同但多了一条恒等连接。真正的差距在于梯度传播ResNet 的 shortcut 让回传信号可以直接流过深层避免十层以后梯度消失。对于 28×28 小图我的经验是不要在一开始就疯狂降采样。输入宽度只有 28如果第一个卷积就用 stride2特征图直接变成 14×14后续残差块再做一次 stride2只剩 7×7空间信息保留得太少。这个项目里 ResNet 的“简单实现”定位也印证了这一点残差块数量不必多通道数从 32 起步在后面某个阶段降一倍分辨率即可。def residual_block(x, filters, stride1): shortcut x if stride ! 1 or x.shape[-1] ! filters: shortcut tf.keras.layers.Conv2D( filters, 1, stridesstride, use_biasFalse)(x) shortcut tf.keras.layers.BatchNormalization()(shortcut) out tf.keras.layers.Conv2D( filters, 3, stridesstride, paddingsame, use_biasFalse)(x) out tf.keras.layers.BatchNormalization()(out) out tf.keras.layers.ReLU()(out) out tf.keras.layers.Conv2D( filters, 3, strides1, paddingsame, use_biasFalse)(out) out tf.keras.layers.BatchNormalization()(out) out tf.keras.layers.Add()([out, shortcut]) out tf.keras.layers.ReLU()(out) return outshortcut分支里的 1×1 卷积只在通道数或特征图尺寸变化时出现作用是把输入张量调整成残差分支输出一样的形状。use_biasFalse是因为后面紧跟 BatchNormalization偏置会被 BN 的均值减法抵消留着反而浪费参数。两个 3×3 卷积中间也夹了一层 BN 和 ReLU这是 ResNet 原论文的经典排列方式。2.3 类别合并对输出维度的影响如果characters.txt同时包含A和a模型需要区分大小写输出类别会变多。但 EMNIST byclass 里的C和c、O和0视觉差异极小强行让模型区分它们训练代价高且收获有限。更稳妥的做法是在生成训练标签前做类别合并把同义字符映射到同一个类别再重建characters.txt。合并后类别数可能从 62 降到 36 或 47输出层Dense的 units 也要跟着改否则sparse_categorical_crossentropy的标签索引会超过输出维度训练直接报错。这里有一个和 ResNet 选型相关的细节类别越多最后一个全连接层的参数量越大。假设残差块输出做 GlobalAveragePooling 后特征维度是 6462 类时全连接参数是 64×62而 10 类时只有 64×10。差异不大但当特征维度扩到 512 时影响就明显了。因此在小数据集上优先用较少的类别数配合合理的残差宽度而不是盲目堆深度。3. model.py 与 mytrain.py 的 ResNet 实现细节3.1 从残差块到完整模型的组装方式model.py的职责不是定义一个 50 层的 ResNet而是提供一个适配 28×28 灰度输入、输出类别数可配置的轻量残差网络。先写残差块再写网络装配函数。import tensorflow as tf from tensorflow.keras import layers, Model def build_resnet_for_emnist(num_classes): inp layers.Input(shape(28, 28, 1)) x layers.Conv2D(32, 3, paddingsame, use_biasFalse)(inp) x layers.BatchNormalization()(x) x layers.ReLU()(x) x residual_block(x, 32) x residual_block(x, 64, stride2) x residual_block(x, 64) x layers.GlobalAveragePooling2D()(x) x layers.Dense(num_classes, activationsoftmax)(x) return Model(inp, x)GlobalAveragePooling2D替代 Flatten Dense 的经典组合减少参数量也天然对空间过拟合有抑制作用。为什么第二层残差块要stride2因为经过第一层后特征图还是 28×28保留一部分位置信息到第二层降采样成 14×14后续就没有再缩小保证最后进入全局池化前仍有足够的空间特征。实际训练中这个结构在 EMNIST 上的拟合速度比 VGG-11 快很多。这里的num_classes必须和load_chars的长度一致。常见写法是chars load_chars(characters.txt) model build_resnet_for_emnist(len(chars))如果characters.txt少了某一行但标签文件还是原来的轻则准确率下降重则Dense输出维度小于最大标签索引训练直接崩溃。所以我习惯在训练脚本开头加一行断言。assert max_label len(chars), characters.txt 与标签不匹配3.2 mytrain.py 的训练超参与数据流水线mytrain.py是训练入口。TF2.1 环境下model.compile和model.fit是最省事的组合没必要自己写tf.GradientTape循环。项目日志里的 20epochs.png 和 80epochs.png 说明默认训练轮数至少在 80 以上常见搭配是批量大小 64、Adam 初始学习率 0.001。model.compile( optimizertf.keras.optimizers.Adam(learning_rate0.001), losssparse_categorical_crossentropy, metrics[accuracy] ) history model.fit( train_ds, # tf.data.Dataset epochs80, validation_dataval_ds, callbackscallbacks, )选sparse_categorical_crossentropy是因为标签是整数索引不是 one-hot 向量。如果标签是 one-hot就要改成categorical_crossentropy两者混用会导致 loss 一直停在 log 尺度下不去。数据显示流水线里shuffle的 buffer 至少应该大于单个类别样本数的量级否则连续的字符会被分成同一批模型看到的数据变换不足收敛更快但泛化变差。def prepare_dataset(images, labels, batch_size64): ds tf.data.Dataset.from_tensor_slices((images, labels)) ds ds.shuffle(10000) ds ds.batch(batch_size) ds ds.prefetch(tf.data.AUTOTUNE) return dsprefetch(AUTOTUNE)让 CPU 准备下一批数据的同时GPU 或训练循环处理当前批次避免数据加载成为瓶颈。EMNIST 图片只有 28×28预处理开销极小这个配置的实际收益可能不大但对于从 TFRecord 加载大图片的迁移场景这个习惯仍然值得保留。3.3 数据增强要不要加、怎么加TF2.1 自带的tf.keras.layers.experimental.preprocessing里已经有 RandomTranslation 之类但稳定性不如后来版本所以我通常直接在 Dataset 上做轻量变换避免把增强逻辑写进模型层导致推理时也要经过同一套预处理。def slight_distort(image, label): image tf.image.random_crop(image, size(24, 24, 1)) image tf.image.resize(image, (28, 28)) return image, label把 28×28 随机裁剪成 24×24 再放大回 28×28等效于模拟轻微偏移和缩放这是手写字符识别里性价比最高的增强。不推荐对字母和数字做水平翻转因为b翻转后看起来像d数字6翻转后接近9会引入标签噪声。最多做一点 15 度以内的随机旋转但 TF2.1 里tf.image.rot90只有固定 90 度只能自己写角度变换或放弃旋转。这个项目没有把增强写进模型说明作者保留了原始语义我倾向于认同这个选择。4. 训练曲线与 checkpoint 管理20 epochs 和 80 epochs 差在哪4.1 checkpoint 文件为什么拆成 .index 和>model.save_weights(checkpoints/pro1-10.ckpt)上面这个调用会把权重保存成pro1-10.ckpt.index和若干个data分片。下一次训练想续跑不要直接load_model因为这里保存的是权重而不是完整网络结构。正确恢复方式如下。model build_resnet_for_emnist(len(chars)) model.load_weights(checkpoints/pro1-10.ckpt)先重建模型再加载权重。这里有个陷阱如果characters.txt被改动过len(chars)变了Dense层权重形状不匹配load_weights会报错或者跳过某些层。所以凡是涉及类别数变化的调整都要重新训练不能直接拿旧 checkpoint 顶上去。4.2 20epochs.png 和 80epochs.png 反映了什么训练日志里两张曲线图分别对应 20 轮和 80 轮。20 轮时训练 loss 通常已经降到比较低验证准确率能达到 90% 上下但验证 loss 还在缓慢波动说明模型尚未完全收敛。80 轮时验证准确率会进入平台期靠学习率衰减继续压低 loss但这种压低的收益越来越小。训练轮次典型状态处理建议5训练 loss 快速下降验证准确率低于 70%观察学习率是否过大确认标签没有错位20验证准确率 90% 左右loss 仍缓慢下降不要急着停配合 ReduceLROnPlateau 继续训练80验证准确率进入平台期可能出现轻微过拟合使用 EarlyStopping 保存最佳权重只看准确率曲线会忽略一个重要事实手写字符识别里某些类别本身容易混淆比如0和O、1和l。如果验证集里这些样本的占比偏高准确率会呈现断崖式波动loss 曲线反而更平滑。所以我在判断是否收敛时优先盯val_loss而不是val_accuracy。4.3 用 ReduceLROnPlateau 和 EarlyStopping 控制训练节奏80 轮纯训练很容易在最后 20 轮出现无效震荡。常见做法是让学习率在验证 loss 停滞时自动减半。reduce_lr tf.keras.callbacks.ReduceLROnPlateau( monitorval_loss, factor0.5, patience3, min_lr1e-6 ) early_stop tf.keras.callbacks.EarlyStopping( monitorval_loss, patience8, restore_best_weightsTrue ) callbacks [reduce_lr, early_stop]factor0.5表示每次触发都将学习率乘以 0.5patience3是连续 3 个 epoch 验证 loss 没有改善就降学习率。restore_best_weightsTrue很关键它会让训练结束后模型权重回到验证 loss 最低的那一轮而不是最后一轮。如果你在 80 轮后看到准确率不升反降先确认这个参数是否打开。4.4 断点续训加载 pro1-10.ckpt 继续训练项目名称里的pro1-10更像是手动指定的 checkpoint 前缀而不是自动生成的文件名。如果要用 TensorFlow 的tf.train.CheckpointManager做自动化管理save()会产生ckpt-10这种带步数的文件名和当前项目的pro1-10命名风格不同。两种方式都可以区别在于自动管理器会额外保存一份checkpoint元文件方便tf.train.latest_checkpoint找到最新权重。ckpt tf.train.Checkpoint(modelmodel) manager tf.train.CheckpointManager( ckpt, directorycheckpoints, max_to_keep5) manager.save()续训时只用恢复模型还不够Adam 优化器的动量状态没有包含在model.load_weights里。如果希望学习率、梯度一阶矩、二阶矩也继续必须用完整 Checkpoint 恢复优化器。ckpt tf.train.Checkpoint(modelmodel, optimizeroptimizer) ckpt.restore(tf.train.latest_checkpoint(checkpoints))恢复后建议先用几个 batch 跑一下前向确认 loss 数值与中断前接近。如果 loss 突然变大大概率是characters.txt顺序变了或数据流水线的 shuffle 种子变了不要急着继续跑。5. 用 demo.py 做真实图片预测时的预处理和类名对齐技巧5.1 从 checkpoint 恢复模型并做好灰度反转demo.py负责加载一张真实图片并输出识别结果。真实图片和 EMNIST 训练样本的分布差异很大训练时图片是黑底白字用户用手机拍的白纸黑字则需要反转。from PIL import Image def preprocess_image(image_path): img Image.open(image_path).convert(L) img img.resize((28, 28), Image.LANCZOS) arr np.array(img, dtypenp.float32) arr 255.0 - arr # 白底黑字转成黑底白字 arr arr / 255.0 return arr.reshape(1, 28, 28, 1) model build_resnet_for_emnist(len(chars)) model.load_weights(checkpoints/pro1-10.ckpt) pred model.predict(preprocess_image(imgs/demo.png))255.0 - arr这一步很容易漏掉。如果用户图片已经是黑底白字再做一次反转会把背景变白、笔迹变黑预测结果完全错误。稳妥做法是先计算图片前景像素占比自动判断是否需要反转。5.2 输出 Top-3 而不是只取一个最大值只取最大索引会掩盖模型在相似类别上的犹豫。手写识别场景里2和Z、5和S本来就在人类眼中都容易混淆模型给出第二和第三概率才是有效调试信息。import numpy as np probs pred[0] top3_idx np.argsort(probs)[::-1][:3] for i in top3_idx: print(chars[i], probs[i])np.argsort默认升序[::-1]改成降序取前三个索引。结合characters.txt输出字符和置信度能快速判断是预处理问题还是类别本身难分。如果5和S的概率接近说明模型没有学到区分特征对话业务方合并类别往往比硬训练更有效。最后再强调一个实际工程点characters.txt一旦在训练后被修改现有 checkpoint 里的输出层权重就对不上了。任何一次类别数变化都要重新走训练流程。保持字符文件版本不变才能让pro1-10.ckpt反复使用。本文还有配套的精品资源点击获取
返回列表