
在浏览器中训练 MNIST用 jax-js 搭建完整神经网络训练循环实战【免费下载链接】jax-jsJAX in JavaScript – ML library for the web, running on WebGPU Wasm项目地址: https://gitcode.com/gh_mirrors/ja/jax-jsjax-js是一个纯 JavaScript 编写的机器学习库把 JAX 风格的高性能计算内核带到网页端——它能把数组运算自动翻译成WebGPU与WebAssembly (Wasm)内核让你不需要服务器、不安装任何重型依赖就能在自己的浏览器里完成一次完整的MNIST 手写数字识别神经网络训练加载数据、前向传播、反向传播、Adam 优化、测试集评估一条龙跑通。1️⃣ 为什么能在浏览器里跑深度学习传统训练需要 Python CUDA GPU而 jax-js 的巧妙之处在于零外部依赖库从零手写gzip 后仅约 80 KB多后端切换webgpuGPU 加速性能最佳、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 后端浏览器端训练的第一步是探测并启动可用的计算后端。推荐优先使用webgpuimport { 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.tsvalueAndGrad等核心变换从主包导出src/index.tsDemo 默认配置10 个 epoch、batchSize 1000MLP/ 250ConvNet、学习率 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的adamGPU 加速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),仅供参考