使用 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),此模型非常適合在瀏覽器中執行。

微調過程

微調工作流程包含以下技術步驟:

  1. 資料載入:匯入「Quick, Draw!」資料集子集。
  2. 預處理:使用 MobileViTImageProcessor 轉換資料。
  3. 配置:定義 collate 函數與評估指標。
  4. 模型初始化:載入預訓練的 MobileViTForImageClassification 模型。
  5. 訓練:使用 TrainerTrainingArguments 輔助類別。
  6. 評估:使用 🤗 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 個不同的類別。

Sources