时间: 2020-09-03 00:08:26 人气: 2269 评论: 0
PyTorch 是一个开源的深度学习框架,对深度学习算法进行训练和优化,它被广泛应用在人工智能领域。更多详细介绍可参考 PyTorch 官网。
本案例使用的场景是计算机视觉领域基本任务之一:手写数字识别,输入一张手写数字的图像,然后识别图像中手写的是哪个数字(0 - 9)。
本文通过智能钛机器学习平台提供的 PyTorch 框架搭建一个简单的神经网络模型实现 MNIST 手写数字识别。通过本文的学习,您可以掌握如下操作:
本案例使用的 MNIST 数据集可参考 MNIST官网。该数据集由来自 250 个不同人手写的数字构成,共包含 60,000 个训练数据,10,000 个测试数据,每个数据都是一张 28 像素 * 28 像素大小的灰度图像。
部分手写数字图像示例如下:
利用智能钛机器学习平台完成手写数字识别任务,我们需要完成以下几个步骤:
整体工作流示例如下:
1. 数据集准备 为方便用户操作,我们将本案例所需数据集 MNIST 上传到公共访问路径下,用户在代码中可直接通过公共路径访问该数据集。
注意: 上传到公共存储桶中文件的访问路径格式为
/cos_public/XXXX,上传到个人存储桶中文件的访问路径为cos_person/XXXX。
2. 代码准备
本案例使用 PyTorch 框架搭建简单的卷积神经网络来完成手写数字图像识别的任务。
为方便用户直接进行后续工作流的搭建,本文提供案例源代码 mnist.py 供用户直接下载体验。
在 mnist.py 源代码中,访问公共 COS 下 MNIST 数据集的示例代码片段如下,访问路径为:/cos_public/mnist。
1.在智能钛控制台的左侧导航栏中,选择【框架】>【深度学习】>【PyTorch】,并拖入画布中。
2.右键【PyTorch】,选择【重命名】,输入新名称:手写数字识别,单击【确定】。
3.单击【手写数字识别】,在右侧弹出的配置栏中配置框架参数。
--batch-size 64 --test-batch-size 1000 --epochs 10 --lr 0.01 --momentum 0.5
以上列表中的参数对应源代码 mnist.py 中以下部分:(用户可参考此处格式配置自定义代码中的参数)
注意: 若自定义代码中,用户给未给参数命名,则可在代码中可通过默认参数 args[0] 读取用户填写的第一个取值,args[1] 读取第二个取值,以此类推。
4.配置资源参数,用户可直接选择平台提供的默认值,也可根据自身代码调整资源分配。
5.运行工作流 单击画布左上角的【运行】,即可开始运行工作流,待运行成功(运行大概需要 4 min)。
1.右键【手写数字识别】,单击【PyTorch 控制台】可查看该工作流运行相关日志。
2.在弹框中, 选择单击 【stdout.log】即可在日志中查看手写数字识别任务的训练过程和测试结果。
本案例实验结果展示了:手写数字识别任务一共训练了10个 Epoch,并详细输出了每个 Epoch 过程中各 batch_size 数据下的损失值 Loss 变化过程。在第10个 Epoch 训练后,模型在测试数据集上取得最佳准确率 98%。
腾讯云一站式机器学习平台智能钛TI-ONE试运营阶段限时0折,欢迎大家积极试用。
https://cloud.tencent.com/product/tio
更多优质技术文章请关注官方知乎机构号:
https://www.zhihu.com/org/teng-xun-zhi-neng-tai-ji-qi-xue-xi-ping-tai/activities
更多优质技术文章请关注官方微信公众号: