使用 Transformers.js 製作具備 ML 能力的網頁遊戲
Hugging Face 已詳細說明 Doodle Dash 的開發,這是一個完全在使用者瀏覽器內運行的即時機器學習 (ML) 驅動網頁遊戲。透過使用 Transformers.js,該遊戲消除了伺服器延遲,使模型能夠在使用者繪圖時每秒進行超過 60 次預測。
模型訓練與架構
Doodle Dash 的核心是一個在 Google 的「Quick, Draw!」資料集子集上訓練的素描偵測模型,該資料集包含超過 500 萬張圖畫,涵蓋 345 個類別。
模型選擇
開發者對 apple/mobilevit-small 進行了微調,這是一個在 ImageNet-1k 上預訓練的輕量級 Vision Transformer (ViT)。由於其適合行動裝置的架構和小巧的體積(僅 560 萬個參數,檔案大小約為 20 MB),此模型非常適合在瀏覽器中執行。
微調過程
微調工作流程包含以下技術步驟:
- 資料載入:匯入「Quick, Draw!」資料集子集。
- 預處理:使用
MobileViTImageProcessor轉換資料。 - 配置:定義 collate 函數與評估指標。
- 模型初始化:載入預訓練的
MobileViTForImageClassification模型。 - 訓練:使用
Trainer和TrainingArguments輔助類別。 - 評估:使用 🤗 Evaluate 函式庫驗證效能。
使用 Transformers.js 進行瀏覽器部署
Transformers.js 是一個 JavaScript 函式庫,可讓 🤗 Transformers 模型在無需後端伺服器的情況下直接在瀏覽器中執行。它提供了一個與 Python 函式庫功能等價的 API。
轉換模型至 ONNX
因為 Transformers.js 使用 ONNX Runtime 進行執行,PyTorch 模型必須轉換為 ONNX 格式。這可以透過在 Transformers.js 儲存庫中提供的轉換腳本,使用 🤗 Optimum 函式庫來完成:
python -m scripts.convert --model_id <model_id>
技術實作
為防止計算密集型的推理程序阻塞主要 UI 執行緒(該執行緒負責渲染和使用者輸入),遊戲實作了 Web Workers API。推理邏輯被隔離在一個獨立的工作執行緒 (worker.js) 中,該執行緒會初始化 image-classification 管線,並處理灰階圖像以返回預測結果。
遊戲設計與最佳化
即時推理迴圈
與原始的「Quick, Draw!」遊戲每隔數秒進行一次預測不同,Doodle Dash 利用瀏覽器內推理的高頻效能來建立更快速的遊戲迴圈:
- 目標:玩家在 60 秒內嘗試繪製盡可能多的塗鴉。
- 機制:在正確預測後,畫布會立即清除,並提示新的單字。
- 懲罰:跳過單字將使玩家剩餘時間減少 3 秒。
- 分數調整:為防止模型僅僅劃掉標籤,遊戲會對前
n個錯誤標籤減分,其中n會隨時間增加。
資料集精煉
儘管模型大小較小(約 20 MB),開發者仍然篩選了原始的 345 個類別,以維持遊戲品質。如果標籤符合以下條件,則會被移除:
- 與其他標籤過於相似(例如:「barn」vs.「house」)。
- 難以理解或無法以足夠細節繪製(例如:「animal migration」或「brain」)。
- 模糊不清(例如:「bat」)。
此篩選過程最終得到超過 300 個不同的類別。