使用OPENCV訓練手寫數字識別分類器

來源:互聯網
上載者:User

使用OPENCV訓練手寫數字識別分類器 

1,下載訓練資料和測試資料檔案,這裡用的是MNIST手寫數字圖片庫,其中訓練資料庫中為60000個,測試資料庫中為10000個
2,建立訓練資料和測試資料檔案讀取函數,注意位元組順序為大端
3,確定字元特徵方式為最簡單的8×8網格內的字元點數


4,建立SVM,訓練並讀取,結果如下
 1000個訓練樣本,測試資料正確率80.21%(並沒有體現SVM小樣本高準確率的特性啊)
  10000個訓練樣本,測試資料正確率95.45%
  60000個訓練樣本,測試資料正確率97.67%

5,編寫手寫輸入的GUI程式,並進行驗證,效果還可以接受。

 

以下為主要代碼,以供參考

(類似的也實現了隨機樹分類器,比較發現在相同的樣本數情況下,SVM準確率略高)

 #include "stdafx.h"</p><p>#include <fstream><br />#include "opencv2/opencv.hpp"<br />#include <vector></p><p>using namespace std;<br />using namespace cv;</p><p>#define SHOW_PROCESS 0<br />#define ON_STUDY 0</p><p>class NumTrainData<br />{<br />public:<br />NumTrainData()<br />{<br />memset(data, 0, sizeof(data));<br />result = -1;<br />}<br />public:<br />float data[64];<br />int result;<br />};</p><p>vector<NumTrainData> buffer;<br />int featureLen = 64;</p><p>void swapBuffer(char* buf)<br />{<br />char temp;<br />temp = *(buf);<br />*buf = *(buf+3);<br />*(buf+3) = temp;</p><p>temp = *(buf+1);<br />*(buf+1) = *(buf+2);<br />*(buf+2) = temp;<br />}</p><p>void GetROI(Mat& src, Mat& dst)<br />{<br />int left, right, top, bottom;<br />left = src.cols;<br />right = 0;<br />top = src.rows;<br />bottom = 0;</p><p>//Get valid area<br />for(int i=0; i<src.rows; i++)<br />{<br />for(int j=0; j<src.cols; j++)<br />{<br />if(src.at<uchar>(i, j) > 0)<br />{<br />if(j<left) left = j;<br />if(j>right) right = j;<br />if(i<top) top = i;<br />if(i>bottom) bottom = i;<br />}<br />}<br />}</p><p>//Point center;<br />//center.x = (left + right) / 2;<br />//center.y = (top + bottom) / 2;</p><p>int width = right - left;<br />int height = bottom - top;<br />int len = (width < height) ? height : width;</p><p>//Create a squre<br />dst = Mat::zeros(len, len, CV_8UC1);</p><p>//Copy valid data to squre center<br />Rect dstRect((len - width)/2, (len - height)/2, width, height);<br />Rect srcRect(left, top, width, height);<br />Mat dstROI = dst(dstRect);<br />Mat srcROI = src(srcRect);<br />srcROI.copyTo(dstROI);<br />}</p><p>int ReadTrainData(int maxCount)<br />{<br />//Open image and label file<br />const char fileName[] = "../res/train-images.idx3-ubyte";<br />const char labelFileName[] = "../res/train-labels.idx1-ubyte";</p><p>ifstream lab_ifs(labelFileName, ios_base::binary);<br />ifstream ifs(fileName, ios_base::binary);</p><p>if( ifs.fail() == true )<br />return -1;</p><p>if( lab_ifs.fail() == true )<br />return -1;</p><p>//Read train data number and image rows / cols<br />char magicNum[4], ccount[4], crows[4], ccols[4];<br />ifs.read(magicNum, sizeof(magicNum));<br />ifs.read(ccount, sizeof(ccount));<br />ifs.read(crows, sizeof(crows));<br />ifs.read(ccols, sizeof(ccols));</p><p>int count, rows, cols;<br />swapBuffer(ccount);<br />swapBuffer(crows);<br />swapBuffer(ccols);</p><p>memcpy(&count, ccount, sizeof(count));<br />memcpy(&rows, crows, sizeof(rows));<br />memcpy(&cols, ccols, sizeof(cols));</p><p>//Just skip label header<br />lab_ifs.read(magicNum, sizeof(magicNum));<br />lab_ifs.read(ccount, sizeof(ccount));</p><p>//Create source and show image matrix<br />Mat src = Mat::zeros(rows, cols, CV_8UC1);<br />Mat temp = Mat::zeros(8, 8, CV_8UC1);<br />Mat img, dst;</p><p>char label = 0;<br />Scalar templateColor(255, 0, 255 );</p><p>NumTrainData rtd;</p><p>//int loop = 1000;<br />int total = 0;</p><p>while(!ifs.eof())<br />{<br />if(total >= count)<br />break;</p><p>total++;<br />cout << total << endl;</p><p>//Read label<br />lab_ifs.read(&label, 1);<br />label = label + '0';</p><p>//Read source data<br />ifs.read((char*)src.data, rows * cols);<br />GetROI(src, dst);</p><p>#if(SHOW_PROCESS)<br />//Too small to watch<br />img = Mat::zeros(dst.rows*10, dst.cols*10, CV_8UC1);<br />resize(dst, img, img.size());</p><p>stringstream ss;<br />ss << "Number " << label;<br />string text = ss.str();<br />putText(img, text, Point(10, 50), FONT_HERSHEY_SIMPLEX, 1.0, templateColor);</p><p>//imshow("img", img);<br />#endif</p><p>rtd.result = label;<br />resize(dst, temp, temp.size());<br />//threshold(temp, temp, 10, 1, CV_THRESH_BINARY);</p><p>for(int i = 0; i<8; i++)<br />{<br />for(int j = 0; j<8; j++)<br />{<br />rtd.data[ i*8 + j] = temp.at<uchar>(i, j);<br />}<br />}</p><p>buffer.push_back(rtd);</p><p>//if(waitKey(0)==27) //ESC to quit<br />//break;</p><p>maxCount--;</p><p>if(maxCount == 0)<br />break;<br />}</p><p>ifs.close();<br />lab_ifs.close();</p><p>return 0;<br />}</p><p>void newRtStudy(vector<NumTrainData>& trainData)<br />{<br />int testCount = trainData.size();</p><p>Mat data = Mat::zeros(testCount, featureLen, CV_32FC1);<br />Mat res = Mat::zeros(testCount, 1, CV_32SC1);</p><p>for (int i= 0; i< testCount; i++)<br />{ </p><p>NumTrainData td = trainData.at(i);<br />memcpy(data.data + i*featureLen*sizeof(float), td.data, featureLen*sizeof(float));</p><p>res.at<unsigned int>(i, 0) = td.result;<br />}</p><p>/////////////START RT TRAINNING//////////////////<br /> CvRTrees forest;<br /> CvMat* var_importance = 0;</p><p> forest.train( data, CV_ROW_SAMPLE, res, Mat(), Mat(), Mat(), Mat(),<br /> CvRTParams(10,10,0,false,15,0,true,4,100,0.01f,CV_TERMCRIT_ITER));<br />forest.save( "new_rtrees.xml" );<br />}</p><p>int newRtPredict()<br />{<br /> CvRTrees forest;<br />forest.load( "new_rtrees.xml" );</p><p>const char fileName[] = "../res/t10k-images.idx3-ubyte";<br />const char labelFileName[] = "../res/t10k-labels.idx1-ubyte";</p><p>ifstream lab_ifs(labelFileName, ios_base::binary);<br />ifstream ifs(fileName, ios_base::binary);</p><p>if( ifs.fail() == true )<br />return -1;</p><p>if( lab_ifs.fail() == true )<br />return -1;</p><p>char magicNum[4], ccount[4], crows[4], ccols[4];<br />ifs.read(magicNum, sizeof(magicNum));<br />ifs.read(ccount, sizeof(ccount));<br />ifs.read(crows, sizeof(crows));<br />ifs.read(ccols, sizeof(ccols));</p><p>int count, rows, cols;<br />swapBuffer(ccount);<br />swapBuffer(crows);<br />swapBuffer(ccols);</p><p>memcpy(&count, ccount, sizeof(count));<br />memcpy(&rows, crows, sizeof(rows));<br />memcpy(&cols, ccols, sizeof(cols));</p><p>Mat src = Mat::zeros(rows, cols, CV_8UC1);<br />Mat temp = Mat::zeros(8, 8, CV_8UC1);<br />Mat m = Mat::zeros(1, featureLen, CV_32FC1);<br />Mat img, dst;</p><p>//Just skip label header<br />lab_ifs.read(magicNum, sizeof(magicNum));<br />lab_ifs.read(ccount, sizeof(ccount));</p><p>char label = 0;<br />Scalar templateColor(255, 0, 0);</p><p>NumTrainData rtd;</p><p>int right = 0, error = 0, total = 0;<br />int right_1 = 0, error_1 = 0, right_2 = 0, error_2 = 0;<br />while(ifs.good())<br />{<br />//Read label<br />lab_ifs.read(&label, 1);<br />label = label + '0';</p><p>//Read data<br />ifs.read((char*)src.data, rows * cols);<br />GetROI(src, dst);</p><p>//Too small to watch<br />img = Mat::zeros(dst.rows*30, dst.cols*30, CV_8UC3);<br />resize(dst, img, img.size());</p><p>rtd.result = label;<br />resize(dst, temp, temp.size());<br />//threshold(temp, temp, 10, 1, CV_THRESH_BINARY);<br />for(int i = 0; i<8; i++)<br />{<br />for(int j = 0; j<8; j++)<br />{<br />m.at<float>(0,j + i*8) = temp.at<uchar>(i, j);<br />}<br />}</p><p>if(total >= count)<br />break;</p><p>char ret = (char)forest.predict(m); </p><p>if(ret == label)<br />{<br />right++;<br />if(total <= 5000)<br />right_1++;<br />else<br />right_2++;<br />}<br />else<br />{<br />error++;<br />if(total <= 5000)<br />error_1++;<br />else<br />error_2++;<br />}</p><p>total++;</p><p>#if(SHOW_PROCESS)<br />stringstream ss;<br />ss << "Number " << label << ", predict " << ret;<br />string text = ss.str();<br />putText(img, text, Point(10, 50), FONT_HERSHEY_SIMPLEX, 1.0, templateColor);</p><p>imshow("img", img);<br />if(waitKey(0)==27) //ESC to quit<br />break;<br />#endif</p><p>}</p><p>ifs.close();<br />lab_ifs.close();</p><p>stringstream ss;<br />ss << "Total " << total << ", right " << right <<", error " << error;<br />string text = ss.str();<br />putText(img, text, Point(50, 50), FONT_HERSHEY_SIMPLEX, 1.0, templateColor);<br />imshow("img", img);<br />waitKey(0);</p><p>return 0;<br />}</p><p>void newSvmStudy(vector<NumTrainData>& trainData)<br />{<br />int testCount = trainData.size();</p><p>Mat m = Mat::zeros(1, featureLen, CV_32FC1);<br />Mat data = Mat::zeros(testCount, featureLen, CV_32FC1);<br />Mat res = Mat::zeros(testCount, 1, CV_32SC1);</p><p>for (int i= 0; i< testCount; i++)<br />{ </p><p>NumTrainData td = trainData.at(i);<br />memcpy(m.data, td.data, featureLen*sizeof(float));<br />normalize(m, m);<br />memcpy(data.data + i*featureLen*sizeof(float), m.data, featureLen*sizeof(float));</p><p>res.at<unsigned int>(i, 0) = td.result;<br />}</p><p>/////////////START SVM TRAINNING//////////////////<br />CvSVM svm = CvSVM();<br />CvSVMParams param;<br />CvTermCriteria criteria;</p><p>criteria= cvTermCriteria(CV_TERMCRIT_EPS, 1000, FLT_EPSILON);<br />param= CvSVMParams(CvSVM::C_SVC, CvSVM::RBF, 10.0, 8.0, 1.0, 10.0, 0.5, 0.1, NULL, criteria); </p><p>svm.train(data, res, Mat(), Mat(), param);<br />svm.save( "SVM_DATA.xml" );<br />}</p><p>int newSvmPredict()<br />{<br />CvSVM svm = CvSVM();<br />svm.load( "SVM_DATA.xml" );</p><p>const char fileName[] = "../res/t10k-images.idx3-ubyte";<br />const char labelFileName[] = "../res/t10k-labels.idx1-ubyte";</p><p>ifstream lab_ifs(labelFileName, ios_base::binary);<br />ifstream ifs(fileName, ios_base::binary);</p><p>if( ifs.fail() == true )<br />return -1;</p><p>if( lab_ifs.fail() == true )<br />return -1;</p><p>char magicNum[4], ccount[4], crows[4], ccols[4];<br />ifs.read(magicNum, sizeof(magicNum));<br />ifs.read(ccount, sizeof(ccount));<br />ifs.read(crows, sizeof(crows));<br />ifs.read(ccols, sizeof(ccols));</p><p>int count, rows, cols;<br />swapBuffer(ccount);<br />swapBuffer(crows);<br />swapBuffer(ccols);</p><p>memcpy(&count, ccount, sizeof(count));<br />memcpy(&rows, crows, sizeof(rows));<br />memcpy(&cols, ccols, sizeof(cols));</p><p>Mat src = Mat::zeros(rows, cols, CV_8UC1);<br />Mat temp = Mat::zeros(8, 8, CV_8UC1);<br />Mat m = Mat::zeros(1, featureLen, CV_32FC1);<br />Mat img, dst;</p><p>//Just skip label header<br />lab_ifs.read(magicNum, sizeof(magicNum));<br />lab_ifs.read(ccount, sizeof(ccount));</p><p>char label = 0;<br />Scalar templateColor(255, 0, 0);</p><p>NumTrainData rtd;</p><p>int right = 0, error = 0, total = 0;<br />int right_1 = 0, error_1 = 0, right_2 = 0, error_2 = 0;<br />while(ifs.good())<br />{<br />//Read label<br />lab_ifs.read(&label, 1);<br />label = label + '0';</p><p>//Read data<br />ifs.read((char*)src.data, rows * cols);<br />GetROI(src, dst);</p><p>//Too small to watch<br />img = Mat::zeros(dst.rows*30, dst.cols*30, CV_8UC3);<br />resize(dst, img, img.size());</p><p>rtd.result = label;<br />resize(dst, temp, temp.size());<br />//threshold(temp, temp, 10, 1, CV_THRESH_BINARY);<br />for(int i = 0; i<8; i++)<br />{<br />for(int j = 0; j<8; j++)<br />{<br />m.at<float>(0,j + i*8) = temp.at<uchar>(i, j);<br />}<br />}</p><p>if(total >= count)<br />break;</p><p>normalize(m, m);<br />char ret = (char)svm.predict(m); </p><p>if(ret == label)<br />{<br />right++;<br />if(total <= 5000)<br />right_1++;<br />else<br />right_2++;<br />}<br />else<br />{<br />error++;<br />if(total <= 5000)<br />error_1++;<br />else<br />error_2++;<br />}</p><p>total++;</p><p>#if(SHOW_PROCESS)<br />stringstream ss;<br />ss << "Number " << label << ", predict " << ret;<br />string text = ss.str();<br />putText(img, text, Point(10, 50), FONT_HERSHEY_SIMPLEX, 1.0, templateColor);</p><p>imshow("img", img);<br />if(waitKey(0)==27) //ESC to quit<br />break;<br />#endif</p><p>}</p><p>ifs.close();<br />lab_ifs.close();</p><p>stringstream ss;<br />ss << "Total " << total << ", right " << right <<", error " << error;<br />string text = ss.str();<br />putText(img, text, Point(50, 50), FONT_HERSHEY_SIMPLEX, 1.0, templateColor);<br />imshow("img", img);<br />waitKey(0);</p><p>return 0;<br />}</p><p>int main( int argc, char *argv[] )<br />{<br />#if(ON_STUDY)<br />int maxCount = 60000;<br />ReadTrainData(maxCount);</p><p>//newRtStudy(buffer);<br />newSvmStudy(buffer);<br />#else<br />//newRtPredict();<br />newSvmPredict();<br />#endif<br />return 0;<br />}

 

 

聯繫我們

該頁面正文內容均來源於網絡整理,並不代表阿里雲官方的觀點,該頁面所提到的產品和服務也與阿里云無關,如果該頁面內容對您造成了困擾,歡迎寫郵件給我們,收到郵件我們將在5個工作日內處理。

如果您發現本社區中有涉嫌抄襲的內容,歡迎發送郵件至: info-contact@alibabacloud.com 進行舉報並提供相關證據,工作人員會在 5 個工作天內聯絡您,一經查實,本站將立刻刪除涉嫌侵權內容。

A Free Trial That Lets You Build Big!

Start building with 50+ products and up to 12 months usage for Elastic Compute Service

  • Sales Support

    1 on 1 presale consultation

  • After-Sales Support

    24/7 Technical Support 6 Free Tickets per Quarter Faster Response

  • Alibaba Cloud offers highly flexible support services tailored to meet your exact needs.