基于深度学习的数字识别GUI的设计
·
基于深度学习的数字识别GUI的设计
用matlab的deeplearning工具箱搭建了CNN来识别手写数字的GUI。
一.训练CNN
采用的是matlab自带的数字训练集和验证集,搭建的CNN的代码如下:
clc,clear;
digitalDatasetPath = fullfile(matlabroot,'toolbox','nnet','nndemos','nndatasets','DigitDataset');
imds = imageDatastore(digitalDatasetPath,'IncludeSubfolders',true,'LabelSource','foldernames');
figure;
perm = randperm(10000,20);
for i = 1:20
subplot(4,5,i);
imshow(imds.Files{perm(i)});
end
labelCount = countEachLabel(imds);
img = readimage(imds,1);
[m,n] = size(img);
numTrainFiles = 750;
[imdsTrain,imdsValidation] = splitEachLabel(imds,numTrainFiles,'randomize');
%定义CNN架构
layers = [
imageInputLayer([m n 1])
convolution2dLayer(3,8,'Padding','same')
batchNormalizationLayer
reluLayer
maxPooling2dLayer(2,'Stride',2)
convolution2dLayer(3,32,'Padding','same')
batchNormalizationLayer
reluLayer
fullyConnectedLayer(10)%输出类别有10个类
softmaxLayer %对全连接层的输出归一化
classificationLayer]; %分类层
options = trainingOptions('sgdm',... %求解器参数设置
'InitialLearnRate',0.01,...
'MaxEpochs',4,...
'Shuffle','every-epoch',...%每一轮都需要验证
'ValidationData',imdsValidation,...%验证的数据集
'ValidationFrequency',30,...
'Verbose',false,...
'Plots','training-progress');
net = trainNetwork(imdsTrain,layers,options);
save('CNNstructure','net');
%验证的图像分类并且计算准确度
YPred = classify(net,imdsValidation);
YValidation = imdsValidation.Labels;
accuracy = sum(YPred == YValidation)/numel(YValidation)
训练的过程如下图所示:
![[外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传(img-dwzENLlT-1612683572929)(C:\Users\pc\Desktop\深度学习数字识别\训练结果.png)]](https://i-blog.csdnimg.cn/blog_migrate/3fdf95128a562cda361b67f85d8064eb.png#pic_center)
二.结果测试
将训练好的神经网络存储于当前文件夹:
![[外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传(img-MpOfl7Gn-1612683572935)(C:\Users\pc\Desktop\深度学习数字识别\CNNSTRUCTURE.png)]](https://i-blog.csdnimg.cn/blog_migrate/3ae08ff4c6bd80d1b117d63456e1e8a3.png#pic_center)
并且用以下代码去测试网络的准确性,当然输入的图像必须是28×2828\times2828×28像素的图(因为训练用的图是28×2828 \times 2828×28像素的图)。
clear;
load CNNstructure.mat;
[filename pathname filterindex] = uigetfile(...
{'*.png','图像文件(*.png)';...
'*.jpg','图像文件(*.jpg)';...
'*.*','所有文件(*.*)'},...
'选择图像文件','MultiSelect','off',...
pwd);
filePath = 0;
if isequal(filename,0)||isequal(pathname,0) %只要这里面表示返回或没有合适的文件时候
return;
end
filePath = fullfile(pathname,filename);
im = imread(filePath);
figure(1)
imshow(im(:,:,1));
figure(2)
prediction = classify(net,imresize(im(:,:,1),[28 28]));
imshow(im(:,:,1));
title(char(prediction));
三.搭建GUI
搭建的框架如下图所示:
![[外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传(img-ZBEhPhBq-1612683572937)(C:\Users\pc\Desktop\深度学习数字识别\guiFig.png)]](https://i-blog.csdnimg.cn/blog_migrate/ec4992b746e51d78d0dc00dc0981d07a.png#pic_center)
导入图像按钮的回调函数:
% --- Executes on button press in import.
function import_Callback(hObject, eventdata, handles)
% hObject handle to import (see GCBO)
% eventdata reserved - to be defined in a future version of MATLAB
% handles structure with handles and user data (see GUIDATA)
[filename pathname filterindex] = uigetfile(...
{'*.png','图像文件(*.png)';...
'*.jpg','图像文件(*.jpg)';...
'*.*','所有文件(*.*)'},...
'选择图像文件','MultiSelect','off',...
pwd);
if isequal(filename,0)||isequal(pathname,0) %只要这里面表示返回或没有合适的文件时候
return;
end
filePath = fullfile(pathname,filename);
im = imread(filePath);
handles.figure = imresize(im(:,:,1));
guidata(hObject, handles);
imshow(im);
开始识别的回调函数
% --- Executes on button press in test.
function test_Callback(hObject, eventdata, handles)
% hObject handle to test (see GCBO)
% eventdata reserved - to be defined in a future version of MATLAB
% handles structure with handles and user data (see GUIDATA)
load CNNstructure.mat;
prediction = classify(net,handles.figure);
imshow(handles.figure);
title(char(prediction));
set(handles.result,'String',char(prediction));
最后的结果显示:
![[外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传(img-uRgjODHw-1612683572947)(C:\Users\pc\Desktop\深度学习数字识别\测试结果.png)]](https://i-blog.csdnimg.cn/blog_migrate/1e4c25b1a73a51d64fa86d1c5abab0df.png#pic_center)
更多推荐
所有评论(0)