在浏览器中训练 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 风格的高性能计算内核带到网页端——它能把数组运算自动翻译成WebGPU与WebAssembly (Wasm)内核,让你不需要服务器、不安装任何重型依赖,就能在自己的浏览器里完成一次完整的MNIST 手写数字识别神经网络训练:加载数据、前向传播、反向传播、Adam 优化、测试集评估,一条龙跑通。
1️⃣ 为什么能在浏览器里跑深度学习?
传统训练需要 Python + CUDA GPU,而 jax-js 的巧妙之处在于:
- 零外部依赖:库从零手写,gzip 后仅约 80 KB;
- 多后端切换:
webgpu(GPU 加速,性能最佳)、wasm(多线程 CPU,兼容性最好)、webgl(旧设备兜底); - JAX 式 API:
numpy数组、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/optax的adam |
| 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),仅供参考