使用 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,非常适合在浏览器中执行。
微调过程
微调工作流程涉及以下技术步骤:
- 数据加载:导入 "Quick, Draw!" 数据集子集。
- 预处理:使用
MobileViTImageProcessor转换数据。 - 配置:定义 collate 函数和评估指标。
- 模型初始化:加载预训练的
MobileViTForImageClassification模型。 - 训练:利用
Trainer和TrainingArguments辅助类。 - 评估:使用 🤗 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 个不同的类别。