news 2026/8/24 9:36:39

在浏览器中训练 MNIST:用 jax-js 搭建完整神经网络训练循环实战

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
在浏览器中训练 MNIST:用 jax-js 搭建完整神经网络训练循环实战

在浏览器中训练 MNIST:用 jax-js 搭建完整神经网络训练循环实战

【免费下载链接】jax-jsJAX in JavaScript – ML library for the web, running on WebGPU & Wasm项目地址: https://gitcode.com/gh_mirrors/ja/jax-js

jax-js是一个纯 JavaScript 编写的机器学习库,把 JAX 风格的高性能计算内核带到网页端——它能把数组运算自动翻译成WebGPUWebAssembly (Wasm)内核,让你不需要服务器、不安装任何重型依赖,就能在自己的浏览器里完成一次完整的MNIST 手写数字识别神经网络训练:加载数据、前向传播、反向传播、Adam 优化、测试集评估,一条龙跑通。

1️⃣ 为什么能在浏览器里跑深度学习?

传统训练需要 Python + CUDA GPU,而 jax-js 的巧妙之处在于:

  • 零外部依赖:库从零手写,gzip 后仅约 80 KB;
  • 多后端切换webgpu(GPU 加速,性能最佳)、wasm(多线程 CPU,兼容性最好)、webgl(旧设备兜底);
  • JAX 式 APInumpy数组、grad自动微分、jit算子融合、vmap向量化,与 Python 的 JAX 高度同构;
  • optax 优化器:配套的@jax-js/optax提供 Adam、SGD 等主流优化算法。

完整可运行的 MNIST 训练 Demo 源码位于 website/src/routes/mnist/+page.svelte,官方站点上可直接点 "Run" 观看实时训练曲线。

2️⃣ 一键准备环境:初始化 WebGPU 后端

浏览器端训练的第一步是探测并启动可用的计算后端。推荐优先使用webgpu

import { init, defaultDevice } from "@jax-js/jax"; const devices = await init(); // 启动所有可用后端 if (devices.includes("webgpu")) { defaultDevice("webgpu"); // 优先 GPU }

💡 Chrome / Edge 上 WebGPU 支持最完整,训练速度比 Wasm 后端快一个数量级。后端能力对照表见 FEATURES.md。

3️⃣ 在浏览器加载 MNIST 数据集

MNIST 有 6 万张训练图 + 1 万张测试图,每张是 28×28 灰度图。jax-js 的 Demo 直接用浏览器原生的DecompressionStream解压 gzip 文件,并解析 IDX 二进制格式,无需服务端支持,数据加载还带缓存:

  • 数据抓取与解析:website/src/lib/dataset/mnist.ts
  • 像素值归一化到[0, 1]后 reshape 成[batch, 28, 28]float32数组,即可送入网络。
const X = np.array(buf).mul(1 / 255).reshape([-1, 28, 28]);

4️⃣ 搭建模型:三层 MLP 的前向传播

官方 Demo 提供两个可选模型,这里以最经典的784 → 256 → 128 → 10三层 MLP 为例。权重用random.uniform按 Xavier 风格初始化,前向传播就三组"矩阵乘 + ReLU",最后用logSoftmax输出对数概率:

const z1 = np.dot(x, w1).add(b1); const a1 = nn.relu(z1); // ……同理 z2/a2 → z3 return nn.logSoftmax(z3);
  • 激活函数库:src/library/nn.ts
  • 卷积模型(ConvNet:两层卷积 + 池化 + 全连接)也写在同一文件里,准确率更高。

5️⃣ 损失函数:负对数似然

分类任务用交叉熵损失。技巧是logSoftmax输出直接乘以oneHot标签再取负均值,比"先 exp 再 log"数值更稳定:

const loss = (params, x, y) => predict(params, x).mul(nn.oneHot(y, 10)).sum().mul(-1 / batchSize);

6️⃣ 训练循环核心:valueAndGrad + Adam

这是整篇文章最精华的一步。JAX 风格的valueAndGrad一次调用同时返回损失值和梯度,再配合@jax-js/optax的 Adam 更新参数,就构成了完整的训练循环:

const solver = adam(learningRate); let optState = solver.init(tree.ref(params)); for (const [X, y] of batches) { const [lossVal, lossGrad] = valueAndGrad(loss)(tree.ref(params), X, y); [updates, optState] = solver.update(lossGrad, optState); params = applyUpdates(params, updates); await blockUntilReady(params); // 等待 GPU 完成 }
  • Adam 实现:packages/optax/src/alias.ts
  • valueAndGrad等核心变换从主包导出:src/index.ts

Demo 默认配置:10 个 epoch、batchSize 1000(MLP)/ 250(ConvNet)、学习率 0.005,每轮结束在测试集上评估准确率并绘制 Train Loss / Test Accuracy 实时曲线。

7️⃣ 性能关键:jit 算子融合

在 GPU 上,瓶颈常常是内存带宽而非算力。用jit包裹前向函数,可把"矩阵乘 → 加法 → ReLU"等多个算子融合成单个内核,减少内核调度与显存往返开销:

const predict = jit((params, x) => { /* 前向传播 */ });

这就是 jax-js 相比手写内核库在神经网络场景下的独特优势。

8️⃣ 训练完成后:手绘数字实时推理

Demo 还内置了一个彩蛋画布——训练结束后,你可以直接在鼠标/触屏上画一个数字,图像经过居中归一化后送入模型,实时显示 0~9 十个类别的概率条。从"训练"到"交互推理"完全发生在同一页面,这正是浏览器端 ML 最迷人的地方。

📌 小结

环节jax-js 对应能力
数据加载原生 fetch + gzip 解压 + 缓存
张量运算numpy模块(兼容 NumPy API)
自动微分valueAndGrad/grad
优化器@jax-js/optaxadam
GPU 加速webgpu后端 +jit融合

无需服务器、无需 Python 环境,一份 TypeScript 代码就能在浏览器里完成MNIST 完整训练循环——这就是 jax-js 给 Web 端深度学习带来的改变。动手的下一步:把 Demo 里的 MLP 换成 ConvNet,或把学习率滑到 0.01 看看收敛速度的变化。

【免费下载链接】jax-jsJAX in JavaScript – ML library for the web, running on WebGPU & Wasm项目地址: https://gitcode.com/gh_mirrors/ja/jax-js

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/8/24 9:29:12

经验模型与插值方法实战指南:从原理到建模应用

1. 从“拍脑袋”到“有章法”:经验模型与插值的实战价值在数学建模竞赛或者实际工程问题里,我们常常会遇到一种尴尬的局面:题目给的数据要么少得可怜,要么分布得七零八落,根本不够支撑一个漂亮的理论模型。比如&#x…

作者头像 李华
网站建设 2026/8/24 9:29:09

Backtrader-Bench:基于LLM自我生成MCQ的量化交易智能体评估框架

1. 项目概述:当LLM智能体遇上量化交易,如何科学评估?最近,关于“LLM驱动的自主智能体”的讨论热度不减,尤其是在金融量化交易这个对决策精度和逻辑严谨性要求极高的领域。大家可能都看过Lilian Weng那篇关于智能体架构…

作者头像 李华
网站建设 2026/8/24 9:28:29

wgpu-py 云端部署指南:Headless GPU 服务器与 Lavapipe 软件渲染实践

wgpu-py 云端部署指南:Headless GPU 服务器与 Lavapipe 软件渲染实践 【免费下载链接】wgpu-py WebGPU for Python 项目地址: https://gitcode.com/gh_mirrors/wg/wgpu-py wgpu-py 是将 WebGPU 图形 API 引入 Python 的开源库,为 Python 提供强大…

作者头像 李华
网站建设 2026/8/24 9:27:22

SmileToUnlock完全使用指南:6个自定义属性让App启动页更有趣

SmileToUnlock完全使用指南:6个自定义属性让App启动页更有趣 【免费下载链接】SmileToUnlock This library uses ARKit Face Tracking in order to catch users smile. 项目地址: https://gitcode.com/gh_mirrors/smi/SmileToUnlock SmileToUnlock 是一个基于…

作者头像 李华