JAX 神经网络实战:用 PyTorch 数据加载器训练 MNIST 全连接网络
JAX 神经网络实战用 PyTorch 数据加载器训练 MNIST 全连接网络【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax导读本文以 JAX 官方 Notebookdocs/notebooks/Neural_Network_and_Data_Loading.md为核心完整演示如何在 JAX 中零第三方神经网络框架地手写并训练一个多层感知机MLP用vmap自动向量化单样本预测、用grad自动求导、用jit编译加速同时借助 PyTorch 的DataLoader完成数据加载打通PyTorch 管数据、JAX 管计算的经典协作链路。读完本文你将掌握 JAX 三大核心变换的组合用法、PRNG 随机数与 PyTree 的实战操作以及一套可复制到任意自定义数据集上的训练循环模板。一、为什么选择JAX 计算 PyTorch 数据加载的组合JAX 的设计哲学是专注于程序变换与加速器上的 NumPy 语义grad自动微分、jit即时编译、vmap自动向量化等变换是它的核心价值而数据加载、清洗这类工程问题并不在 JAX 库的职责范围内。正如文档中所述JAX is laser-focused on program transformations and accelerator-backed NumPy, so we dont include data loading or munging in the JAX library——JAX 刻意不做数据加载而是鼓励复用生态中已有的优秀方案。PyTorch 的torch.utils.data.DataLoader提供了成熟的多进程预取、采样器、批处理collate机制是现成的优质选择。两者组合时唯一的适配点是PyTorch 默认产出torch.Tensor而 JAX 需要 NumPy 数组。解决方式是在DataLoader的collate_fn参数中注入一个自定义 collate 函数把张量批量转换为 NumPy 数组细节见第五节。由于 NumPy 数组可以直接作为 JAX 数组使用这一层shim垫片足够轻量且不引入任何额外的数据加载库。仓库中的 examples/mnist_classifier.py、examples/mnist_classifier_fromscratch.py 也提供了类似的纯 JAX 手写 MNIST 训练示例可作为本文之外的补充对照实现。二、环境与依赖准备本文代码基于 JAX、NumPy 以及 PyTorch含 torchvision运行。首次在 Notebook 环境中执行时需要安装!pip install torch torchvisionJAX 本体与 NumPy 一般已随环境安装若需独立安装可参考仓库根目录的 pyproject.toml 与 setup.py 中的依赖声明。运行平台支持 CPU / GPU / TPU本文的jit编译与vmap向量化在任意后端上均可工作。导入所需的全部模块import jax.numpy as jnp from jax import grad, jit, vmap from jax import random这里jax.numpy即jnp是 NumPy 风格的数组 API所有计算都以它为基础确保后续可以被grad/jit/vmap变换。三、超参数与网络参数初始化含 PRNG 正确用法3.1 超参数一览layer_sizes [784, 512, 512, 10] # 输入 78428x28 展平、两个 512 隐层、输出 10MNIST 类别数 step_size 0.01 # SGD 学习率 num_epochs 8 # 训练轮数 batch_size 128 # 每批样本数 n_targets 10 # 分类数3.2 用 JAX PRNG 初始化权重JAX 的随机数体系与 NumPy 完全不同它采用显式可拆分的 PRNG key 体系Threefry/Philox 等实现随机函数不依赖全局状态而是接收一个 key 数组作为输入。初始化代码# 为单个全连接层随机初始化权重和偏置 def random_layer_params(m, n, key, scale1e-2): w_key, b_key random.split(key) # 将一个 key 拆成两个独立 key return scale * random.normal(w_key, (n, m)), scale * random.normal(b_key, (n,)) # 初始化 sizes 指定的全连接网络各层参数 def init_network_params(sizes, key): keys random.split(key, len(sizes)) # 按层数批量拆分 key return [random_layer_params(m, n, k) for m, n, k in zip(sizes[:-1], sizes[1:], keys)] params init_network_params(layer_sizes, random.key(0))关键 API 的底层实现都在 jax/_src/random/core.pyrandom.key(seed)core.py L231-L254以整数种子创建标量 PRNG key。文档明确指出接受标量种子若传入数组会抛出TypeError见 L225-L228提示改用vmap实现批量打 keykey 的 dtype 由jax_default_prng_impl配置决定。random.split(key, num)core.py L318将 key 拆分为num个相互独立的新 key默认num2。这正是上面random_layer_params拆出w_key/b_key、init_network_params按层数批量拆分的依据——每个 key 只被使用一次保证了随机序列的可复现性与独立性。random.normal(key, shape)core.py L911-L939以给定 key 采样标准正态分布dtype默认在jax_enable_x64开启时为 float64、否则为 float32。权重形状为(n, m)输出维度在前这是为了与预测函数中的jnp.dot(w, activations)直接对齐。参数集合params是一个 Python list其中每个元素是(weight, bias)二元组——这正是 JAX 中的一种PyTree嵌套容器结构后续grad会以同样的结构返回梯度vmap可以通过in_axes(None, 0)声明其不变性。补充说明random.key是当前推荐 API老代码中可能见到random.PRNGKey后者是key的兼容别名新代码应统一使用random.key。四、定义预测函数并用vmap自动批处理4.1 面向单样本的预测函数首先以单个样本为单位定义前向计算不做任何批处理假设from jax.scipy.special import logsumexp def relu(x): return jnp.maximum(0, x) def predict(params, image): # 前向传播除最后一层外每层先线性变换再过 ReLU activations image for w, b in params[:-1]: outputs jnp.dot(w, activations) b activations relu(outputs) # 最后一层输出 logits并减去 logsumexp 得到对数概率log-softmax final_w, final_b params[-1] logits jnp.dot(final_w, activations) final_b return logits - logsumexp(logits)其中logsumexp来自 jax/_src/scipy/special.py与 SciPy 语义一致用于数值稳定的 log-softmax 归一化。验证它对单样本工作正常random_flattened_image random.normal(random.key(1), (28 * 28,)) preds predict(params, random_flattened_image) print(preds.shape) # (10,)而对批量输入形状(10, 28*28)直接调用则会因维度不匹配抛错——这正是设计成单样本函数的原因random_flattened_images random.normal(random.key(1), (10, 28 * 28)) try: preds predict(params, random_flattened_images) except TypeError: print(Invalid shapes!)4.2 用vmap一行升级为批量版本JAX 的vmap将函数沿指定轴自动向量化语义上等价于手写循环但无 Python 层开销、无性能损失# 生成批处理版本params 不批处理in_axes 取 Noneimage 沿第 0 轴批处理 batched_predict vmap(predict, in_axes(None, 0)) # batched_predict 与 predict 的调用签名完全一致 batched_preds batched_predict(params, random_flattened_images) print(batched_preds.shape) # (10, 10)in_axes(None, 0)的含义是第一个参数paramsPyTree 结构在所有批次间共享第二个参数image的第 0 维即批量维。vmap会自动把predict内部的jnp.dot、relu、logsumexp全部向量化返回形状(batch, n_targets)。这一模式是 JAX 中先写标量/单样本逻辑再自动批处理的惯用法与 docs/notebooks/automatic-vectorization.md即automatic-vectorization教程所讲的核心思想一致。至此训练所需的全部零件已经齐备batched_predict负责批量前向grad负责对参数求导jit负责编译提速。五、损失函数、准确率与单步参数更新5.1 工具函数def one_hot(x, k, dtypejnp.float32): 创建 x 的 k 类 one-hot 编码。 return jnp.array(x[:, None] jnp.arange(k), dtype) def accuracy(params, images, targets): target_class jnp.argmax(targets, axis1) predicted_class jnp.argmax(batched_predict(params, images), axis1) return jnp.mean(predicted_class target_class) def loss(params, images, targets): preds batched_predict(params, images) return -jnp.mean(preds * targets) # 平均负对数似然交叉熵等价形式one_hot利用广播比较x[:, None] jnp.arange(k)生成(batch, k)的布尔矩阵再转成 float32。loss直接对 log-softmax 输出与 one-hot 标签做点积求平均即标准的多分类交叉熵。5.2 用gradjit定义单步更新jit def update(params, x, y): grads grad(loss)(params, x, y) # 对第一个参数params自动求梯度 return [(w - step_size * dw, b - step_size * db) for (w, b), (dw, db) in zip(params, grads)]这里体现了两点 JAX 精髓grad(loss)默认对第一个参数求导即网络参数params由于params是 PyTree返回的grads与params结构完全同构每层各有一组(dw, db)可直接逐层做 SGD 更新。jit装饰器把整个前向 反向 参数更新过程编译为 XLA 计算图跨调用复用编译结果大幅降低 Python 解释与调度开销——这是训练循环得以高效运行的关键。JAX 的编译机制详见 docs/201/jit.md 与 docs/jit-compilation.md。六、用 PyTorch DataLoader 加载数据6.1 安装并导入import numpy as np from jax.tree_util import tree_map from torch.utils.data import DataLoader, default_collate from torchvision.datasets import MNIST6.2 两个关键的适配函数def numpy_collate(batch): collate 函数负责把一批样本组合成 batch。 default_collate 先产出 PyTorch 张量tree_map 再将其整体转为 numpy 数组。 return tree_map(np.asarray, default_collate(batch)) def flatten_and_cast(pic): 将 PIL 图像转换为展平的一维 numpy 数组。 return np.ravel(np.array(pic, dtypejnp.float32))default_collatePyTorch 内置的批处理函数把样本列表堆叠为torch.Tensor。tree_map来自 jax/_src/tree_util.py是jax.tree.map的别名。其实现本质是先tree_flatten拍平叶子 → 对每个叶子应用f→ 再unflatten恢复结构L394-L400。此处的作用是递归遍历 collate 产出的数据结构列表/元组/字典等任意 PyTree把每一片torch.Tensor都替换为np.asarray(...)转换后的 NumPy 数组——无论数据是图像、标签还是自定义结构一行tree_map即可全部转换这正是 PyTree 工具在生态互操作上的典型应用。flatten_and_cast作为MNIST数据集的transform把 PIL 图像转成float32的展平一维数组形状(784,)与网络输入维度对齐。6.3 构建数据集与 DataLoader# 用 torchvision 数据集定义训练数据downloadTrue 时首次自动下载到本地 mnist_dataset MNIST(/tmp/mnist/, downloadTrue, transformflatten_and_cast) # 用自定义 collate 函数创建 DataLoader产出 numpy 数组批次 training_generator DataLoader(mnist_dataset, batch_sizebatch_size, collate_fnnumpy_collate)DataLoader内部的多进程加载、shuffle、批切分等机制全部由 PyTorch 负责JAX 侧只消费 NumPy 数组。6.4 加载完整训练集与测试集用于评估# 完整训练集用于训练过程中检查准确率 train_images np.array(mnist_dataset.train_data).reshape(len(mnist_dataset.train_data), -1) train_labels one_hot(np.array(mnist_dataset.train_labels), n_targets) # 完整测试集 mnist_dataset_test MNIST(/tmp/mnist/, downloadTrue, trainFalse) test_images jnp.array(mnist_dataset_test.test_data.numpy().reshape(len(mnist_dataset_test.test_data), -1), dtypejnp.float32) test_labels one_hot(np.array(mnist_dataset_test.test_labels), n_targets)注意两点评估时直接使用完整数据集不经 DataLoader 分批accuracy内部用batched_predict一次前向即可算完因为vmap的批处理维度是任意的。训练集用np.array(...)得到 NumPy 数组测试集用jnp.array(..., dtypejnp.float32)得到 JAX 数组——两者对 JAX 均可直接使用展示了jnp与np的无缝互操作。版本提示mnist_dataset.train_data/train_labels是 torchvision 旧版属性较新版本推荐改用mnist_dataset.data/mnist_dataset.targets语义相同。七、训练循环import time for epoch in range(num_epochs): start_time time.time() for x, y in training_generator: y one_hot(y, n_targets) # 标签转 one-hot params update(params, x, y) # jit 编译的单步更新 epoch_time time.time() - start_time train_acc accuracy(params, train_images, train_labels) test_acc accuracy(params, test_images, test_labels) print(Epoch {} in {:0.2f} sec.format(epoch, epoch_time)) print(Training set accuracy {}.format(train_acc)) print(Test set accuracy {}.format(test_acc))循环结构非常简洁外层按epoch遍历内层从training_generator逐批取出(x, y)其中x形状为(batch_size, 784)、y为(batch_size,)的整数标签。标签先经one_hot转为(batch_size, 10)再交给update完成求梯度 SGD 更新。每轮结束后分别在全量训练集与测试集上评估准确率打印耗时。由于update已被jit编译训练中每个 batch 的更新都以接近编译后原生的速度执行time统计的是每轮的墙钟耗时含数据迭代与评估开销。八、回顾一次训练走遍 JAX 三大核心变换训练结束时本示例已经完整使用了 JAX 的核心 API变换/API作用在本示例中的用法grad自动微分对第一个参数求导grad(loss)(params, x, y)得到与params同构的梯度 PyTreejit即时编译加速装饰update让前向反向更新整体编译为 XLAvmap自动向量化 / 批处理vmap(predict, in_axes(None, 0))一行升级批量预测random显式 key 的可复现随机数random.key/split/normal初始化全部参数tree_util.tree_mapPyTree 递归映射把 PyTorch 张量批量转为 NumPy 数组正如文档结尾所总结的Weve now used the whole of the JAX API:gradfor derivatives,jitfor speedups andvmapfor auto-vectorization.整个计算过程全部以 NumPy 风格书写jnp模型构建不依赖任何神经网络库数据加载借力 PyTorch 生态训练则可在 CPU/GPU/TPU 上运行。九、扩展方向与实战建议改用更规范的 API 与结构init_network_params的random.split(key, len(sizes))拆分出的 key 在每层内又被random_layer_params二次拆分多个 key 依次使用保证了每层权重与偏置的独立采样。若层数更多可把layer_sizes列表化配置扩展为任意深度网络。加入正则化与优化器本示例使用朴素 SGDstep_size0.01。仓库中的 examples/mnist_classifier_fromscratch.py 展示了更完整的可复现训练脚本如需动量/Adam 等优化器JAX 生态中已有成熟的实现可参考。其他数据源本文用collate_fn适配 PyTorch同样的思路适用于 TensorFlow 的tf.data等任何能产出批数据的 API只需把张量转成 NumPy 数组。JAX 官方还提供了更原生的 docs/501/data-loading.md 数据加载指南供进阶阅读。注意版本兼容性本文代码以 notebook 编写时的 API 为准random.key、train_data等接口在新版本中可能有更名或行为调整如train_data→data运行前请对照所安装的 jax/torchvision 版本确认。通过本文你已经掌握了一条可复用的通用链路手写单样本前向 →vmap自动批处理 →grad求梯度 →jit编译更新 → 外部生态数据加载器供数这套模式可以直接迁移到卷积网络、Transformer 以及任何自定义数据集上。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考