基于MATLAB的卷积神经网络(CNN)实现手写数字识别:
1. 准备工作空间
clc;
clear all;
close all;
2. 导入数据
假设你已经下载了MNIST手写数字数据集,并将其解压到当前工作目录下的HandWrittenDataset文件夹中。
digitDatasetPath = fullfile('./', 'HandWrittenDataset');
imds = imageDatastore(digitDatasetPath, 'IncludeSubfolders', true, 'LabelSource', 'foldernames');
% 数据集图片个数
countEachLabel(imds);
numTrainFiles = 17; % 每个数字有22个样本,取17个样本作为训练数据
[imdsTrain, imdsValidation] = splitEachLabel(imds, numTrainFiles, 'randomize');
% 查看图片的大小
img = readimage(imds, 1);
size(img);
3. 定义卷积神经网络的结构
layers = [
imageInputLayer([28 28 1]) % 输入层
convolution2dLayer(5, 6, 'Padding', 2) % 卷积层
batchNormalizationLayer
reluLayer
maxPooling2dLayer(2, 'Stride', 2) % 池化层
convolution2dLayer(5, 16)
batchNormalizationLayer
reluLayer
maxPooling2dLayer(2, 'Stride', 2)
convolution2dLayer(5, 120)
batchNormalizationLayer
reluLayer
fullyConnectedLayer(10) % 全连接层
softmaxLayer
classificationLayer
];
4. 训练神经网络
% 设置训练参数
options = trainingOptions('sgdm', ...
'MaxEpochs', 50, ...
'ValidationData', imdsValidation, ...
'ValidationFrequency', 5, ...
'Verbose', false, ...
'Plots', 'training-progress'); % 显示训练进度
% 训练神经网络,保存网络
net = trainNetwork(imdsTrain, layers, options);
save('CSNet.mat', 'net');
5. 使用网络进行分类并计算准确性
% 手写数据
YPred = classify(net, imdsValidation);
YValidation = imdsValidation.Labels;
% 计算正确率
accuracy = sum(YPred == YValidation) / numel(YValidation);
% 绘制预测结果
figure;
nSample = 10;
ind = randperm(size(YPred, 1), nSample);
for i = 1:nSample
subplot(2, fix((nSample + 1) / 2), i);
imshow(char(imdsValidation.Files(ind(i))));
title(['预测:' char(YPred(ind(i)))]);
if char(YPred(ind(i))) == char(YValidation(ind(i)))
xlabel(['真实:' char(YValidation(ind(i)))]);
else
xlabel(['真实:' char(YValidation(ind(i)))], 'Color', 'r');
end
end
项目 :在MATLAB中利用卷积神经网络实现手写数字的识别 www.youwenfan.com/contentcnd/95846.html
6. 事项
- 如果你的MATLAB版本较旧,可能需要手动下载MNIST数据集并进行预处理。
- 在训练过程中,可以根据训练进度图调整训练参数,如学习率、批大小等。
- 为了提高模型的泛化能力,可以尝试使用数据增强技术。