拓冰建站拓冰建站
首页 / 资讯中心 / 正文

从零构建手写数字识别系统:基于MNIST的机器学习全流程实践

1. 项目概述从零构建一个手写数字识别系统最近几年无论是学生做课程设计、毕业设计还是开发者入门机器学习手写数字识别Handwritten Digit Recognition几乎成了一个“必刷”的经典项目。原因很简单它麻雀虽小五脏俱全。数据集MNIST干净规整、问题定义10分类清晰明确、模型从简单的逻辑回归到复杂的卷积神经网络都能上阵演练非常适合用来理解机器学习从数据准备、模型构建、训练调优到评估部署的完整流程。这个项目标题“手写数字识别系统”听起来可能有些“老生常谈”但真正能把它做透、做明白并理解其中每一个技术决策背后原因的人并不多。很多人止步于跑通代码却忽略了为什么用这个模型、为什么这样预处理、参数为什么这么设等核心问题。今天我就结合自己带学生和实际项目中的经验抛开那些框架自带的“一键运行”示例从头拆解如何构建一个稳健、可解释且具备一定扩展性的手写数字识别系统。我们会深入每个环节的“为什么”并分享那些教科书和官方文档里通常不会写的“踩坑”实录。2. 核心需求解析与方案选型2.1 问题定义与核心目标手写数字识别本质上是一个**多分类Multi-class Classification**问题具体来说是10分类数字0-9。我们的目标是构建一个模型或系统输入一张手写数字的灰度图像通常已预处理为固定大小如28x28像素输出一个0-9之间的整数作为识别结果。这个项目的核心目标可以分解为几个层次基础功能实现能够正确识别MNIST测试集中绝大部分样本达到较高的准确率如95%。流程完整性完整覆盖数据加载、预处理、模型定义、训练、评估、预测的机器学习全流程。技术深度探索不满足于调用现成API要理解不同模型如全连接网络、卷积网络在此任务上的表现差异及其原因。系统化思维将模型封装成一个可以接收新输入如图片文件并返回识别结果的“系统”考虑输入输出的接口设计。2.2 技术栈选型与考量为什么是Python为什么是这些库这是项目开始前必须想清楚的问题。编程语言Python。这是机器学习领域的事实标准拥有最庞大、最成熟的生态库NumPy, Pandas, Scikit-learn, TensorFlow, PyTorch。社区资源丰富任何问题几乎都能找到解决方案。核心计算库NumPy。所有数值计算的基础处理多维数组我们的图像数据效率极高。TensorFlow/PyTorch的张量操作也与之深度兼容。机器学习框架TensorFlow/Keras 或 PyTorch。这是主要选择。TensorFlow/Keras尤其是其高阶APItf.keras对于新手和快速原型开发非常友好。它提供了清晰的模块化接口Sequential或Functional API让开发者能像搭积木一样构建网络同时隐藏了许多底层细节让初学者更专注于模型结构本身。其生态系统如TensorBoard用于可视化也非常完善。PyTorch以动态计算图和更“Pythonic”的编程风格著称在研究领域和需要高度定制化的场景中更受欢迎。它的设计让调试和理解模型内部状态更加直观。选择建议对于课程大作业或快速入门我强烈推荐从tf.keras开始。它的学习曲线更平缓代码更简洁能让你更快地看到成果建立信心。本项目的后续示例也将主要基于tf.keras。辅助工具库Matplotlib/Seaborn用于数据可视化查看图像样本、绘制损失/准确率曲线、绘制混淆矩阵。Scikit-learn虽然我们可能用深度学习框架构建模型但它的工具函数如train_test_split,classification_report,confusion_matrix在数据划分和模型评估阶段依然非常有用。OpenCV/PIL如果你的系统需要处理来自摄像头、扫描仪或用户上传的原始图片文件这些图像处理库将用于前期的图像读取、缩放、灰度化等预处理。注意不要陷入“框架之争”。对于这个项目无论选TensorFlow还是PyTorch都能很好地完成。关键是理解其背后的机器学习原理框架只是工具。2.3 数据集MNIST的深入理解MNIST数据集是此项目的基石它包含60000张训练图像和10000张测试图像每张都是28x28像素的灰度图像素值在0-255之间标签是0-9的数字。关于MNIST有几个关键点常被忽略“过于简单”的争议MNIST因其干净、规整而常被诟病“太玩具”无法代表真实世界的复杂图像。这没错但它的价值恰恰在于此——作为一个基准测试Benchmark和教学工具它让你能排除数据质量的干扰专注于理解和比较不同模型架构、优化算法、正则化技术的效果。在学车时你会在空旷的停车场练习而不是直接上晚高峰的市区道路。数据分布MNIST的各类别0-9样本数量基本均衡这避免了类别不平衡问题。但在构建真实系统时你必须检查并处理可能的数据不平衡。内置与下载tf.keras.datasets.mnist.load_data()可以自动下载并加载数据非常方便。但第一次使用需要联网。3. 数据预处理不止是归一化数据预处理是模型成功的基石这里有很多细节决定成败。3.1 标准化与归一化这是最关键的一步。原始图像的像素值是0-255的整数。直接输入网络会带来两个问题1) 数值范围大导致梯度不稳定训练缓慢2) 不同特征像素点的尺度不一致但我们的模型通常假设所有输入特征在相似范围内。两种常见方法归一化Normalization将像素值缩放到 [0, 1] 区间。x_train x_train.astype(float32) / 255.0标准化Standardization将像素值调整为均值为0标准差为1的分布。x_train (x_train - mean) / std其中mean和std通常在训练集上计算。为什么这里通常用归一化对于图像数据尤其是MNIST这种像素值有明确物理意义亮度且范围固定0-255的数据归一化到[0,1]是最直观、最常用的方法。标准化在特征尺度差异大时更有效但MNIST各像素点本质是同质的亮度值归一化足矣且计算更简单。3.2 标签编码One-Hot Encoding我们的标签y是整数0-9。对于多分类问题网络的输出层通常使用Softmax激活函数输出一个10维的概率向量每个维度代表属于该类别的概率。因此我们需要将整数标签转换为独热编码One-Hot Encoding。例如标签“3”转换为[0, 0, 0, 1, 0, 0, 0, 0, 0, 0]。在tf.keras中可以使用tf.keras.utils.to_categorical(y, num_classes10)轻松完成。为什么必须这么做因为我们的损失函数如分类交叉熵categorical_crossentropy要求目标标签与预测输出概率的格式一致都是概率分布的形式。使用整数标签配合sparse_categorical_crossentropy损失函数是另一种选择它内部处理了编码问题但理解One-Hot编码对于掌握分类问题的本质很重要。3.3 数据增强Data Augmentation的考量对于MNIST标准流程通常不包含数据增强因为其测试集已经能很好评估模型泛化能力。但在实际手写识别系统中用户书写可能存在旋转、轻微平移、笔画粗细不一等情况。为了提升模型鲁棒性可以引入数据增强。对于图像分类常见增强操作包括随机旋转小角度如±10度随机平移小幅度的上下左右移动随机缩放轻微放大或缩小添加噪声在tf.keras中可以使用ImageDataGenerator或tf.keras.layers.experimental.preprocessing.RandomRotation等层在模型内部实现实时增强。实操心得对于课程大作业可以先在不使用数据增强的情况下在MNIST测试集上达到一个很高的基准分数如99%。然后再尝试加入数据增强观察其在你自己收集的、更具挑战性的手写样本上的效果提升。这能清晰展示数据增强的价值。4. 模型构建从全连接网络到卷积网络这是项目的核心。我们将构建两个经典模型进行对比理解其演进逻辑。4.1 模型一多层感知机MLP / 全连接网络这是最基础的神经网络模型也称为深度前馈网络。import tensorflow as tf from tensorflow import keras from tensorflow.keras import layers model_mlp keras.Sequential([ layers.Flatten(input_shape(28, 28)), # 将28*28的二维图像展平为784维向量 layers.Dense(128, activationrelu), # 第一个隐藏层128个神经元 layers.Dropout(0.2), # 丢弃层防止过拟合 layers.Dense(64, activationrelu), # 第二个隐藏层64个神经元 layers.Dense(10, activationsoftmax) # 输出层10个神经元使用softmax ])关键点解析Flatten层卷积网络处理二维空间信息但全连接层要求输入是一维向量。Flatten层将(28,28)的形状转换为(784,)。Dense层全连接层每个神经元都与上一层的所有神经元相连。units参数定义该层神经元的数量。activation指定激活函数ReLU是目前最常用的因为它能有效缓解梯度消失问题计算也快。Dropout层在训练过程中随机“丢弃”暂时禁用一部分神经元这里是20%。这是一种非常有效的正则化技术强迫网络不过度依赖某些特定的神经元从而提升泛化能力防止过拟合。输出层10个神经元对应10个类别。softmax激活函数将10个神经元的原始输出logits转换为一个概率分布所有输出值之和为1。为什么选择这样的结构128 - 64这是一种常见的维度递减设计高层神经元捕捉更抽象、更全局的特征。神经元数量是超参数需要通过实验调整。Dropout放在哪通常放在激活函数层之后。这里的Dropout(0.2)意味着前一层Dense(128)输出的每个神经元在每次训练迭代中都有20%的概率被置零。4.2 模型二卷积神经网络CNNCNN是处理图像数据的“王者”它通过卷积核自动学习图像的空间层次特征。model_cnn keras.Sequential([ # 第一个卷积块 layers.Conv2D(32, (3, 3), activationrelu, input_shape(28, 28, 1)), # 32个3x3卷积核 layers.MaxPooling2D((2, 2)), # 2x2最大池化 # 第二个卷积块 layers.Conv2D(64, (3, 3), activationrelu), layers.MaxPooling2D((2, 2)), # 分类头 layers.Flatten(), layers.Dense(64, activationrelu), layers.Dropout(0.5), # CNN模型容量大Dropout率可以设高一些 layers.Dense(10, activationsoftmax) ])关键点解析Conv2D层这是卷积层。filters32表示使用32个不同的卷积核滤波器每个核会学习提取一种特征如边缘、角点。kernel_size(3,3)是卷积核大小。input_shape(28,28,1)需要明确输入图像的通道数灰度图为1RGB图为3。MaxPooling2D层池化层用于下采样。(2,2)池化窗口将2x2区域内的最大值作为输出其作用是降低维度减少参数和计算量。引入平移不变性即使目标在图像中轻微移动池化后得到的特征可能不变。扩大感受野让后续层能看到更广范围的输入信息。从卷积到全连接的过渡经过几轮“卷积-池化”后特征图仍然是三维的高度宽度通道数。在输入全连接层之前必须用Flatten层将其展平为一维向量。Dropout(0.5)CNN模型学习能力很强更容易过拟合因此通常使用更高的Dropout率如0.5来加强正则化。为什么CNN比MLP更适合图像MLP的Flatten操作破坏了图像固有的空间局部相关性。一个像素和它上下左右的像素关系最密切但Flatten后这个像素可能与屏幕上任意远的像素相连。CNN的卷积操作通过局部连接和权值共享显式地利用了这种空间局部性和平移不变性用更少的参数学到了更有效的特征。5. 模型训练与超参数调优5.1 编译模型配置学习过程model_cnn.compile(optimizeradam, losscategorical_crossentropy, # 如果标签未one-hot则用sparse_categorical_crossentropy metrics[accuracy])优化器OptimizerAdam是目前最流行、默认效果往往就不错的优化算法。它结合了动量Momentum和自适应学习率RMSProp的优点。对于初学者无脑选Adam通常是个好起点。损失函数Losscategorical_crossentropy分类交叉熵是多分类问题的标准损失函数。它衡量模型输出的概率分布与真实标签的One-Hot编码之间的差异。评估指标Metricsaccuracy准确率是最直观的指标即预测正确的样本比例。5.2 训练模型喂入数据并迭代history model_cnn.fit(x_train, y_train_onehot, batch_size64, epochs10, validation_split0.2) # 从训练集中划分20%作为验证集批次大小Batch Size一次迭代一个Step输入网络的样本数量。较小的Batch Size如32, 64带来更多的权重更新次数和可能更好的泛化但训练更慢、更震荡。较大的Batch Size如256, 512训练更稳定、更快但可能泛化能力稍差且需要更多内存。64或128是常见的起始选择。轮数Epochs整个训练集完整通过网络一次称为一个Epoch。需要训练多少轮这没有固定答案需要通过观察验证集损失/准确率曲线来决定。通常训练到验证集指标不再提升甚至开始下降过拟合为止。验证集Validation Set使用validation_split从训练集中自动划分一部分这里是20%作为验证集。验证集用于在训练过程中监控模型在未见过的数据上的表现是判断过拟合和决定早停Early Stopping的关键。切记测试集Test Set在最终评估前绝对不能用于训练或调参。5.3 超参数调优实战超参数是训练前设定的不是模型学到的。调优是一个实验过程。学习率Learning Rate这是最重要的超参数之一。Adam有默认学习率通常为0.001但有时需要调整。如果训练损失下降很慢可以尝试增大如果损失剧烈震荡或变成NaN则需要减小。可以使用tf.keras.optimizers.Adam(learning_rate0.0001)进行设置。网络结构与容量隐藏层的层数和每层的神经元/卷积核数量。原则是从简单开始例如先只用1个隐藏层如果欠拟合训练集准确率也低再增加层数或神经元数。对于CNN可以尝试增加卷积块如3个Conv2DPooling层或增加卷积核数量如从32-64-128。正则化强度主要是Dropout率。如果模型在训练集上表现很好但在验证集上差很多过拟合可以尝试增大Dropout率如从0.2提高到0.5或者在更多的层后加入Dropout。批量归一化Batch Normalization在激活函数前加入layers.BatchNormalization()层可以加速训练、允许使用更高的学习率并有一定的正则化效果。它通过规范化每一层的输入分布来实现。在CNN中通常放在Conv2D层之后、激活函数之前。实操心得调优策略不要同时调整多个超参数采用网格搜索Grid Search或随机搜索Random Search时也应有一个基准模型。我的习惯是先用一组保守参数中等学习率、简单网络、加入Dropout训练一个基准模型。如果欠拟合先增加网络容量更多层/神经元。如果过拟合先增强正则化加大Dropout或加入BN层。最后再微调学习率。务必使用验证集来指导调优并用最终的测试集仅做一次最终评估。6. 模型评估与可视化不止看准确率训练完成后不能只看测试集上一个准确率数字就结束。6.1 绘制学习曲线import matplotlib.pyplot as plt def plot_history(history): fig, (ax1, ax2) plt.subplots(1, 2, figsize(12, 4)) # 绘制损失曲线 ax1.plot(history.history[loss], labelTraining Loss) ax1.plot(history.history[val_loss], labelValidation Loss) ax1.set_xlabel(Epoch) ax1.set_ylabel(Loss) ax1.legend() ax1.set_title(Loss over Epochs) # 绘制准确率曲线 ax2.plot(history.history[accuracy], labelTraining Accuracy) ax2.plot(history.history[val_accuracy], labelValidation Accuracy) ax2.set_xlabel(Epoch) ax2.set_ylabel(Accuracy) ax2.legend() ax2.set_title(Accuracy over Epochs) plt.show() plot_history(history)如何解读理想情况训练和验证损失同步下降准确率同步上升最终趋于平稳且两者差距很小。过拟合训练损失持续下降但验证损失在某个点后开始上升或不再下降训练准确率远高于验证准确率。这说明模型记住了训练数据的噪声而非一般规律。欠拟合训练损失和验证损失都很高且准确率都较低。这说明模型复杂度不够无法捕捉数据中的模式。6.2 混淆矩阵与分类报告准确率可能会掩盖模型在特定类别上的弱点。混淆矩阵能清晰展示错误分类的细节。from sklearn.metrics import confusion_matrix, classification_report import seaborn as sns import numpy as np # 获取测试集预测结果概率 y_pred_proba model_cnn.predict(x_test) # 将概率转换为类别标签 y_pred np.argmax(y_pred_proba, axis1) # 如果测试标签是one-hot格式需要转换回来 y_true np.argmax(y_test_onehot, axis1) # 计算混淆矩阵 cm confusion_matrix(y_true, y_pred) # 使用Seaborn绘制热力图 plt.figure(figsize(10,8)) sns.heatmap(cm, annotTrue, fmtd, cmapBlues, xticklabelsrange(10), yticklabelsrange(10)) plt.xlabel(Predicted Label) plt.ylabel(True Label) plt.title(Confusion Matrix) plt.show() # 打印详细的分类报告 print(classification_report(y_true, y_pred, target_names[str(i) for i in range(10)]))从混淆矩阵中你能发现什么例如你可能会发现模型经常把“4”和“9”、“5”和“6”、“3”和“8”混淆。这是因为这些数字在手写体上本身形状就相似。这为你指明了模型改进的方向也许需要收集更多这些易混淆数字的样本或者设计针对性的数据增强如模拟连笔、旋转。6.3 可视化错误样本分析被模型错误分类的样本是调试和理解模型局限性的黄金方法。# 找出预测错误的索引 error_indices np.where(y_pred ! y_true)[0] # 随机查看几个错误样本 num_samples_to_show 5 for i in range(num_samples_to_show): idx error_indices[i] img x_test[idx].reshape(28, 28) true_label y_true[idx] pred_label y_pred[idx] plt.imshow(img, cmapgray) plt.title(fTrue: {true_label}, Pred: {pred_label}) plt.axis(off) plt.show()看看这些被分错的图片是书写极其潦草还是图片本身模糊这能让你直观感受到当前模型的边界在哪里。7. 构建识别系统从模型到应用训练出一个高准确率的模型只是第一步。一个完整的“系统”需要能接收新数据并返回预测结果。7.1 模型保存与加载训练好的模型需要保存下来供后续使用。# 保存整个模型架构权重优化器状态 model_cnn.save(my_digit_recognizer.h5) # 在另一个脚本或应用中加载模型 from tensorflow.keras.models import load_model loaded_model load_model(my_digit_recognizer.h5)7.2 设计预测接口我们需要一个函数能够处理用户可能提供的各种输入格式例如保存为.png或.jpg的图片文件。从画板程序传来的图像数组。甚至是通过摄像头捕获的帧。下面是一个处理图片文件的示例函数import numpy as np from PIL import Image import tensorflow as tf def predict_digit_from_image_file(image_path, model): 从图片文件路径预测手写数字 Args: image_path: 图片文件路径 model: 加载好的Keras模型 Returns: digit (int): 预测的数字 confidence (float): 预测置信度最高概率 # 1. 读取并预处理图像 img Image.open(image_path).convert(L) # 转换为灰度图 img img.resize((28, 28)) # 调整大小为28x28 img_array np.array(img) # 2. 预处理归一化并适配模型输入形状 # 注意MNIST是黑底白字如果用户输入是白底黑字可能需要反转 # 这里假设输入是黑底白字与MNIST一致。 img_array img_array.astype(float32) / 255.0 # 如果图片背景是白色数字是黑色则需要反转颜色 if np.mean(img_array) 0.5: # 简单判断平均像素值大于0.5可能是白底 img_array 1 - img_array img_array img_array.reshape(1, 28, 28, 1) # 添加批次和通道维度 - (1,28,28,1) # 3. 预测 predictions model.predict(img_array, verbose0) # verbose0不显示预测进度条 predicted_digit np.argmax(predictions[0]) confidence np.max(predictions[0]) return predicted_digit, confidence # 使用示例 # digit, conf predict_digit_from_image_file(user_written_4.png, loaded_model) # print(fPredicted digit: {digit} with confidence {conf:.2%})这个函数里的几个关键细节颜色空间转换convert(L)确保图像是灰度图。尺寸调整模型输入是28x28必须调整。颜色反转这是一个极易被忽略但至关重要的坑MNIST数据集是黑底像素值0白字像素值1。但用户用画图工具保存的图片很可能是白底黑字。如果直接输入模型会认不出来。代码中通过计算图像平均像素值做了一个简单判断和反转。更稳健的做法是让用户固定一种格式或在系统说明中明确要求。形状适配模型预测需要批次维度所以要用reshape(1, 28, 28, 1)。7.3 构建简单的图形用户界面可选但推荐对于课程大作业一个简单的GUI能极大提升项目的完整度和展示效果。可以使用tkinterPython标准库或Gradio更简单快速搭建。使用Gradio的极简示例import gradio as gr # 复用上面的预测函数 def predict_image(image): # image 是 gradio 传入的 PIL Image 对象 img image.convert(L).resize((28, 28)) img_array np.array(img).astype(float32) / 255.0 if np.mean(img_array) 0.5: img_array 1 - img_array img_array img_array.reshape(1, 28, 28, 1) predictions loaded_model.predict(img_array, verbose0) return {str(i): float(predictions[0][i]) for i in range(10)} # 返回所有类别的概率字典 # 创建界面 iface gr.Interface( fnpredict_image, inputsgr.Image(typepil, image_modeL, sourcecanvas), # 启用画板 outputsgr.Label(num_top_classes3), # 显示概率最高的3个结果 title手写数字识别系统, description在下方画板中写一个数字0-9点击提交进行识别。 ) iface.launch()运行这段代码会自动在浏览器中打开一个交互页面用户可以直接用鼠标写字并实时看到识别结果和置信度体验非常好。8. 项目进阶与扩展思考完成基础系统后你可以从以下方向进行深化这会让你的项目从“合格”变为“优秀”。8.1 模型优化与对比实验更先进的CNN架构尝试经典的网络结构如LeNet-5正是为MNIST设计的、VGG的简化版甚至微调Fine-tune一个在ImageNet上预训练的小型模型如MobileNetV2的输入层观察效果。集成学习训练多个不同的模型例如不同初始化的CNN、MLP、SVM等然后对它们的预测结果进行投票硬投票或平均概率软投票这通常能提升最终性能。超参数自动化调优使用Keras Tuner或Optuna库自动搜索最佳的超参数组合。8.2 处理更真实的数据自制数据集自己用笔在纸上写数字拍照然后进行预处理去背景、二值化、居中、缩放创建一个小的“真实世界”测试集。你会发现模型准确率很可能下降这就是现实与理想实验室数据的差距。探索其他数据集尝试更复杂的数据集如Fashion-MNIST衣物分类、CIFAR-10小物体彩色图像分类挑战会更大。8.3 部署考量如果想让别人真正用起来需要考虑模型轻量化使用TensorFlow Lite将模型转换为.tflite格式以便部署到移动设备或嵌入式设备上。API服务化使用Flask或FastAPI将模型封装成RESTful API这样任何能发送HTTP请求的客户端网页、手机App都可以调用你的识别服务。持续学习设计一个反馈机制当用户指出预测错误时系统能否将这张图片和正确标签加入训练集进行增量学习这是一个更高级的话题。9. 常见问题与排查技巧实录这里记录一些我和学生们在实际操作中踩过的坑和解决方案。Q1: 训练一开始损失Loss就是NaN非数字。可能原因1学习率过高。这是最常见的原因。尝试将学习率降低一个数量级例如从0.001降到0.0001。可能原因2数据未归一化。确保输入数据已经除以255.0缩放到[0,1]或进行了标准化。可能原因3网络结构太深或不稳定。尝试加入Batch Normalization层或者使用更小的网络初始化。Q2: 模型在训练集上准确率很高99%但在验证集/测试集上很低80%过拟合严重。解决方案1增加正则化。提高Dropout率或在更多层后加入Dropout。对于CNN可以尝试加入L2权重正则化 (kernel_regularizerkeras.regularizers.l2(0.001))。解决方案2使用数据增强。对训练图像进行随机旋转、平移、缩放可以显著提升模型泛化能力。解决方案3简化模型。减少网络层数或神经元数量降低模型容量。解决方案4早停Early Stopping。使用tf.keras.callbacks.EarlyStopping回调函数当验证集损失连续几个Epoch不再下降时自动停止训练防止过度拟合。Q3: 训练速度非常慢。检查1是否使用了GPU在代码开头运行print(tf.config.list_physical_devices(GPU))确认TensorFlow是否识别到了GPU。确保安装了对应版本的CUDA和cuDNN。检查2批次大小Batch Size是否过小在GPU内存允许的范围内适当增大Batch Size如从32增到128可以更充分利用GPU并行计算能力加快训练。检查3数据加载是否成为瓶颈对于大型数据集可以使用tf.data.DatasetAPI进行高效的数据流水线处理和预取。Q4: 自己画的数字模型总是预测错。排查1颜色反转问题。这是头号杀手务必确认你预处理时输入模型的图像格式黑底白字与MNIST训练数据一致。使用上文预测函数中的颜色反转逻辑。排查2数字是否居中MNIST中的数字大致位于图像中心。你自己写的数字如果太偏角落模型可能不认识。可以在预处理中加入一个简单的“重心居中”算法。排查3笔画粗细。你的笔迹可能比MNIST标准数据粗或细很多。可以尝试对图像进行形态学操作腐蚀、膨胀来调整笔画粗细或者使用数据增强时模拟不同笔画粗细。Q5: 如何进一步提升准确率例如从99%到99.5%技巧1模型集成。训练5-10个结构相同但随机初始化不同的模型对它们的预测结果取平均。技巧2测试时增强Test Time Augmentation, TTA。预测时对同一张输入图像进行几种不同的增强如轻微旋转、平移得到多个预测结果然后取平均。这能稳定预测。技巧3更精细的超参数搜索。对学习率、Dropout率、网络深度/宽度进行系统性的网格搜索或随机搜索。心态调整在MNIST上从99%提升到99.5%需要付出巨大的努力但收益可能并不显著。这更像一个学术练习。在实际应用中你需要权衡精度提升带来的价值与所付出的计算和工程成本。
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门