辨识手写数字
这篇教学会使用 Keras 搭配 NumPy 训练手写数字模型,再搭配 OpenCV KNN 演算方法 ( cv2.ml.KNearest_load ),即时辨识出手写的阿拉伯数字。
快速导览:
因为程序中的 OpenCV 会需要使用镜头或 GPU,所以请使用本机环境 ( 参考:使用 Python 虚拟环境 ) 或使用 Anaconda Jupyter 进行实作 ( 参考:使用 Anaconda ) ,并安装 OpenCV 函数库 ( 参考:OpenCV 函数库 )。
什么是 KNN 算法?
KNN 算法的全名为 K Nearest Neighbor,也称为 K-近邻算法,是机器学习中的一种算法,意思是寻找 k 个最接近“某分类”的邻居,透过这些邻居来投票,以多数决定这个分类代表什么,举例来说,小明住的地方有 100 个邻居,小明生了一个小孩不知道要叫什么名字,于是一一询问邻居的意见,最后以多数邻居的意见作为小孩的命名。
参考:K-近邻算法
以下图为例,测试样本 ( 绿色圆形 ) 在 k=3 ( 实线圆圈 ) 的状态下,会被分配给红色三角形 ( 因为红色比较多 ),如果 k=5 ( 虚线圆圈 ) 的状态下,会被分配给蓝色正方形,何选择一个最佳的 K 值取决于数据内容。一般情况下,在分类时较大的 K 值能够减小噪声的影响,但会使类别之间的界限变得模糊。
安装 Keras
Keras 是一个开放源代码,可以进行神经网络学习的 Python 函数库,在 2017 年,Google TensorFlow 的核心库加入支援 Keras 的功能,尔后 Keras 常常会和 TensorFlow 互相搭配使用,许多功能也都必须建构在 TensorFlow 的基础上运作。
参考“Jupyter 安装 Tensorflow”文章,安装 TensorFlow 后,就会一并安装 Keras,简单步骤说明如下 ( 如果已经安装完成可略过此部分 ):
建立名为 tensorflow 的虚拟环境
conda create --name tensorflow python=3.9启动虚拟环境
conda activate tensorflow安装 Jupyter
conda install jupyter notebook安装 tensorflow
pip install tensorflow==2.5安装 OpenCV 和 OpenCV 进阶软件包
pip install opencv-python pip install opencv_contrib_python打开 Anaconda,进入 tensorflow 虚拟环境,启动 Jupyter
训练手写数字模型
训练手写数字模使用 keras 内建的“MNIST 手写字符数据集”进行训练,数据集内分成“训练集”和“测试集”,训练集有 60,000 张 28x28 像素灰度图像,作为深度学习与训练模型使用,测试集内有 10,000 同规格图像,作为测试训练模型使用,训练后可以辨识手写数字 0~9,下图为训练集的其中一部分手写数字图像。
下方的程序码执行后会进行模型训练,训练后会将模型储存为 mnist_knn.xml ( 文件大小约 200~250 MB ),储存后会使用测试集进行测试,测试过程需要耗费几分钟的时间 ( 范例的测试结果准确度为 96.88% ),如果不想测试也可移除该部分程序码,直接训练与储存模型。
import cv2
import numpy as np
from keras.datasets import mnist
from keras import utils
(x_train, y_train), (x_test, y_test) = mnist.load_data() # 載入訓練集
# 訓練集資料
x_train = x_train.reshape(x_train.shape[0],-1) # 轉換資料形狀
x_train = x_train.astype('float32')/255 # 轉換資料型別
y_train = y_train.astype(np.float32)
# 測試集資料
x_test = x_test.reshape(x_test.shape[0],-1) # 轉換資料形狀
x_test = x_test.astype('float32')/255 # 轉換資料型別
y_test = y_test.astype(np.float32)
knn=cv2.ml.KNearest_create() # 建立 KNN 訓練方法
knn.setDefaultK(5) # 參數設定
knn.setIsClassifier(True)
print('training...')
knn.train(x_train, cv2.ml.ROW_SAMPLE, y_train) # 開始訓練
knn.save('mnist_knn.xml') # 儲存訓練模型
print('ok')
print('testing...')
test_pre = knn.predict(x_test) # 讀取測試集並進行辨識
test_ret = test_pre[1]
test_ret = test_ret.reshape(-1,)
test_sum = (test_ret == y_test)
acc = test_sum.mean() # 得到準確率
print(acc)
根据模型,辨识手写数字
已经训练好 xml 模型档后,就可以开始进行辨识,下方的程序码会先取出一个正方形的区域,将这个区域的像素做二值化黑白的转换 ( 因为手写字通常是白底黑字,要转换成黑底白字 ),转换后将尺寸缩小到 28x28 进行辨识,就能得到手写字的辨识结果,同时也会将辨识的图像显示在原本图像的右上角。
import cv2
import numpy as np
cap = cv2.VideoCapture(0) # 啟用攝影鏡頭
print('loading...')
knn = cv2.ml.KNearest_load('mnist_knn.xml') # 載入模型
print('start...')
if not cap.isOpened():
print("Cannot open camera")
exit()
while True:
ret, img = cap.read()
if not ret:
print("Cannot receive frame")
break
img = cv2.resize(img,(540,300)) # 改變影像尺寸,加快處理效率
x, y, w, h = 400, 200, 60, 60 # 定義擷取數字的區域位置和大小
img_num = img.copy() # 複製一個影像作為辨識使用
img_num = img_num[y:y+h, x:x+w] # 擷取辨識的區域
img_num = cv2.cvtColor(img_num, cv2.COLOR_BGR2GRAY) # 顏色轉成灰階
# 針對白色文字,做二值化黑白轉換,轉成黑底白字
ret, img_num = cv2.threshold(img_num, 127, 255, cv2.THRESH_BINARY_INV)
output = cv2.cvtColor(img_num, cv2.COLOR_GRAY2BGR) # 顏色轉成彩色
img[0:60, 480:540] = output # 將轉換後的影像顯示在畫面右上角
img_num = cv2.resize(img_num,(28,28)) # 縮小成 28x28,和訓練模型對照
img_num = img_num.astype(np.float32) # 轉換格式
img_num = img_num.reshape(-1,) # 打散成一維陣列資料,轉換成辨識使用的格式
img_num = img_num.reshape(1,-1)
img_num = img_num/255
img_pre = knn.predict(img_num) # 進行辨識
num = str(int(img_pre[1][0][0])) # 取得辨識結果
text = num # 印出的文字內容
org = (x,y-20) # 印出的文字位置
fontFace = cv2.FONT_HERSHEY_SIMPLEX # 印出的文字字體
fontScale = 2 # 印出的文字大小
color = (0,0,255) # 印出的文字顏色
thickness = 2 # 印出的文字邊框粗細
lineType = cv2.LINE_AA # 印出的文字邊框樣式
cv2.putText(img, text, org, fontFace, fontScale, color, thickness, lineType) # 印出文字
cv2.rectangle(img,(x,y),(x+w,y+h),(0,0,255),3) # 標記辨識的區域
cv2.imshow('oxxostudio', img)
if cv2.waitKey(50) == ord('q'):
break # 按下 q 鍵停止
cap.release()
cv2.destroyAllWindows()
微信扫码关注
抖音扫码关注