使用 Transformers.js 制作基于机器学习的网页游戏

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 格式。这可以通过使用 🤗 Optimum 库并在 Transformers.js 仓库中提供的转换脚本来实现:

python -m scripts.convert --model_id <model_id>

技术实现

为了防止计算密集型推理过程阻塞主 UI 线程(该线程负责渲染和用户输入),游戏实现了 Web Workers API。推理逻辑被隔离在一个单独的工作线程(worker.js)中,该线程初始化 image-classification 管道并处理灰度图像以返回预测结果。

游戏设计与优化

实时推理循环

与原始的 "Quick, Draw!" 游戏不同,后者每隔几秒进行一次预测,Doodle Dash 利用浏览器内推理的高频性能来创建更快的游戏循环:

  • 目标:玩家尝试在 60 秒内绘制尽可能多的涂鸦。
  • 机制:在正确预测后,画布会立即清除,提示一个新单词。
  • 惩罚:跳过一个单词会使玩家剩余时间减少 3 秒。
  • 得分调整:为了防止模型仅仅划掉标签,游戏会减少前 n 个错误标签的得分,其中 n 随时间增加。

数据集精炼

尽管模型尺寸较小(~20MB),为了保持游戏质量,开发者对原始的 345 个类别进行了过滤。如果标签符合以下情况,则会被移除:

  • 与其他标签过于相似(例如,“谷仓” vs. “房屋”)。
  • 理解或绘制细节过于困难(例如,“动物迁徙” 或 “大脑”)。
  • 存在歧义(例如,“蝙蝠”)。

此过滤过程得到的最终集合包含超过 300 个不同的类别。

Sources