使用 Teachable Machine
Teachable Machine 是 Google 所推出的无程序码机器学习平台,只需要简单的步骤,就能够在浏览器上训练模型,透过训练的模型辨识图片、声音或是姿势,这篇教学将会介绍如何使用 Teachable Machine。
快速导览:
什么是 Teachable Machine?
Teachable Machine 是 Google 所推出的无程序码机器学习平台,更简单来说,Teachable Machine 是一个网页工具,只需要打开浏览器,就能在不需要专业知识和撰写程序码的情况下,轻松的为网站和应用程序训练机器学习模型。
前往 Teachable Machine:https://teachablemachine.withgoogle.com/
Teachable Machine 目前提供了“图片、声音和姿势”共三种训练模型,只要经过“搜集和训练”的步骤,就能够建立自己的模型,由于 Teachable Machine 背后应用了开源的机器学习函数库 Tensorflow.js,因此可以将训练好的模型以 Tensorflow.js、Keras、或 Tensorflow Lite 格式输出,在任何的网页或是应用程序中呼叫使用。
注意,Python 只能使用图片专案所训练的模型。
建立分类,训练模型
打开 Teachable Machine 网站后,点击“开始使用”开始训练模型的新专案 ( 最下方可以切换语系为繁体中文 )。
点击 “图片专案”,选择“标准图片模型”,就可以进入图片模型训练流程。
训练流程主要有三个步骤,“添加分类内容”、“训练模型”和“预览训练结果”。
添加分类内容:
每个分类的内容可以透过摄影镜头 Webcam 或上传 Upload 的方式增加图片,最少需要有两个分类,点击下方 Add a class 的按钮可以增加分类,下图的范例的分类使用 oxxo ( 画面中有人 ) 以及维他命 ( 画面中有维他命的罐子 ) 两种。
训练模型:
分类建立完成后,点击“训练模型”,就会开始进行图片模型的训练,出现“模型已训练完成”的文字表示训练完成。
预览训练结果:
训练完成后,就能从预览模型的视窗里,测试自己训练的模型准确度。
在 Python 中使用模型
点击右上方的“导出模型”,选择 Tensorflow,勾选 Keras,就能下载 Keras.h5 模型供 Python 使用。
下载模型并压缩,将 keras_model.h5 放到和 Python 程序同样的路径下,就可以开始编辑 Python 程序。
在 Anaconda Jupyter 安装好 Tensorflow 和 OpenCV 后,执行下方的程序码,就可以看到透过 OpenCV 播放摄影镜头的影片,并判断现在出现的图像是什么分类。
import tensorflow as tf
import cv2
import numpy as np
model = tf.keras.models.load_model('keras_model.h5', compile=False) # 載入 model
data = np.ndarray(shape=(1, 224, 224, 3), dtype=np.float32) # 設定資料陣列
cap = cv2.VideoCapture(0) # 設定攝影機鏡頭
if not cap.isOpened():
print("Cannot open camera")
exit()
while True:
ret, frame = cap.read() # 讀取攝影機影像
if not ret:
print("Cannot receive frame")
break
img = cv2.resize(frame , (398, 224)) # 改變尺寸
img = img[0:224, 80:304] # 裁切為正方形,符合 model 大小
image_array = np.asarray(img) # 去除換行符號和結尾空白,產生文字陣列
normalized_image_array = (image_array.astype(np.float32) / 127.0) - 1 # 轉換成預測陣列
data[0] = normalized_image_array
prediction = model.predict(data) # 預測結果
a,b= prediction[0] # 取得預測結果
if a>0.9:
print('oxxo')
if b>0.9:
print('維他命')
cv2.imshow('oxxostudio', img)
if cv2.waitKey(500) == ord('q'):
break # 按下 q 鍵停止
cap.release()
cv2.destroyAllWindows()
微信扫码关注
抖音扫码关注