JAX 神经网络实战:用 PyTorch 数据加载器训练 MNIST 全连接网络
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
导读
本文以 JAX 官方 Notebook(docs/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 don't 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] # 输入 784(28x28 展平)、两个 512 隐层、输出 10(MNIST 类别数) 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, scale=1e-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.py:
random.key(seed)(core.py L231-L254):以整数种子创建标量 PRNG key。文档明确指出"接受标量种子",若传入数组会抛出TypeError(见 L225-L228),提示改用vmap实现批量打 key;key 的 dtype 由jax_default_prng_impl配置决定。random.split(key, num)(core.py L318):将 key 拆分为num个相互独立的新 key,默认num=2。这正是上面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 取 None),image 沿第 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)的含义是:第一个参数params(PyTree 结构)在所有批次间共享,第二个参数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, dtype=jnp.float32): """创建 x 的 k 类 one-hot 编码。""" return jnp.array(x[:, None] == jnp.arange(k), dtype) def accuracy(params, images, targets): target_class = jnp.argmax(targets, axis=1) predicted_class = jnp.argmax(batched_predict(params, images), axis=1) 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 用grad+jit定义单步更新
@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, dtype=jnp.float32))default_collate:PyTorch 内置的批处理函数,把样本列表堆叠为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 数据集定义训练数据,download=True 时首次自动下载到本地 mnist_dataset = MNIST('/tmp/mnist/', download=True, transform=flatten_and_cast) # 用自定义 collate 函数创建 DataLoader,产出 numpy 数组批次 training_generator = DataLoader(mnist_dataset, batch_size=batch_size, collate_fn=numpy_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/', download=True, train=False) test_images = jnp.array(mnist_dataset_test.test_data.numpy().reshape(len(mnist_dataset_test.test_data), -1), dtype=jnp.float32) test_labels = one_hot(np.array(mnist_dataset_test.test_labels), n_targets)注意两点:
- 评估时直接使用完整数据集(不经 DataLoader 分批),
accuracy内部用batched_predict一次前向即可算完,因为vmap的批处理维度是任意的。 - 训练集用
np.array(...)得到 NumPy 数组,测试集用jnp.array(..., dtype=jnp.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同构的梯度 PyTree |
jit | 即时编译加速 | 装饰update,让前向+反向+更新整体编译为 XLA |
vmap | 自动向量化 / 批处理 | vmap(predict, in_axes=(None, 0))一行升级批量预测 |
random | 显式 key 的可复现随机数 | random.key/split/normal初始化全部参数 |
tree_util.tree_map | PyTree 递归映射 | 把 PyTorch 张量批量转为 NumPy 数组 |
正如文档结尾所总结的:"We've 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列表化配置,扩展为任意深度网络。 - 加入正则化与优化器:本示例使用朴素 SGD(
step_size=0.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 Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考