
Flax NNX MNIST 实战教程从数据加载、CNN 训练到 SavedModel 导出【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax本篇技术指南以 Flax 仓库中的 MNIST 教程及同名 notebook为骨架完整讲解如何使用 Flax NNX API 在 JAX 之上端到端构建、训练并部署一个手写数字分类卷积神经网络CNN。你将掌握 NNX 的模块化建模nnx.Module、随机数流管理nnx.Rngs、状态化变换nnx.jit/nnx.value_and_grad、视图切换nnx.view、训练指标nnx.MultiMetric、TensorBoard 监控以及通过 Orbax 将模型导出为 TensorFlow SavedModel 的完整链路。1. Flax NNX 与 MNIST 教程概览Flax 是一个构建在 JAX 之上的神经网络库主打灵活性。Flax NNX 是其新一代函数式与面向对象融合的 API模型被定义为普通的 Python 类nnx.Module参数、随机数流、优化器状态等都被视为可被 JAX 变换如jit、vmap、grad识别的对象同时保留原地in-place更新与引用语义代码更加简洁直观。本教程的完整可运行形式是仓库中的 mnist_tutorial.ipynb由 jupytext 维护Markdown 版本见 mnist_tutorial.md。如果你此前使用过 Flax Linen API可先阅读 Why Flax NNX 了解两者差异。本文默认你具备深度学习基本概念卷积、批归一化、Dropout、交叉熵、优化器。整个流程包含八个阶段与教程章节一一对应安装 Flaxpip install用 Hugging Facedatasets加载并批量化 MNIST用nnx.Module定义 CNN 模型创建nnx.Optimizer与nnx.MultiMetric定义训练/评估 step 函数nnx.jitnnx.value_and_grad定义测试集推理函数与可视化训练循环 TensorBoard 监控通过 Orbax 导出 SavedModel 供 LiteRT / TensorFlow Serving 使用2. 安装 Flax如果当前 Python 环境尚未安装flax从 PyPI 安装即可在 Google Colab / Jupyter 中取消注释对应代码单元# !pip install -U jax[cuda12] # !pip install -U flax说明jax[cuda12]提供了 CUDA 12 加速的 JAX 版本若使用 TPU 或纯 CPU 环境请按需选择对应的 JAX 安装方式。MNIST 这种规模的任务在 CPU 上同样可完成只是速度较慢。教程后续还用到datasets、matplotlib、optax、tensorboardX、orbax-export等第三方包可按需安装# !pip install -U datasets matplotlib optax tensorboardX orbax-export[all]3. 加载 MNIST 数据集使用 Hugging Facedatasets包加载 MNIST并对图像做归一化与批量化import numpy as np import matplotlib.pyplot as plt from datasets import load_dataset train_steps 1200 eval_every 200 batch_size 32 dataset load_dataset(mnist) train_ds dataset[train].shuffle(seed0) test_ds dataset[test] def make_batches(ds, batch_size): Yield batches of normalized (image, label) numpy arrays. for i in range(0, len(ds), batch_size): batch ds[i : i batch_size] if len(batch[label]) batch_size: # drop incomplete final batch break images np.stack([ np.array(img, dtypenp.float32)[..., None] / 255.0 for img in batch[image] ]) yield {image: images, label: np.array(batch[label])}关键细节归一化每个像素除以 255把灰度值映射到[0, 1]区间有助于梯度稳定。通道维度[..., None]为每张28×28的图像追加通道维得到形状(32, 28, 28, 1)的批次匹配 CNN 的NHWC数据布局。丢弃不完整批次最后一组样本数不足batch_size时直接跳过保证训练步中张量形状固定便于jit编译与指标统计。超参数train_steps1200为总训练步数eval_every200表示每隔 200 步做一次测试集评估batch_size32。4. 用 Flax NNX 定义 CNN 模型4.1 模型结构与源码对应通过继承nnx.Module定义模型层以实例属性的方式声明from flax import nnx # The Flax NNX API. from functools import partial from typing import Optional class CNN(nnx.Module): A simple CNN model. def __init__(self, *, rngs: nnx.Rngs): self.conv1 nnx.Conv(1, 32, kernel_size(3, 3), rngsrngs) self.batch_norm1 nnx.BatchNorm(32, rngsrngs) self.dropout1 nnx.Dropout(rate0.025) self.conv2 nnx.Conv(32, 64, kernel_size(3, 3), rngsrngs) self.batch_norm2 nnx.BatchNorm(64, rngsrngs) self.avg_pool partial(nnx.avg_pool, window_shape(2, 2), strides(2, 2)) self.linear1 nnx.Linear(3136, 256, rngsrngs) self.dropout2 nnx.Dropout(rate0.025) self.linear2 nnx.Linear(256, 10, rngsrngs) def __call__(self, x, rngs: nnx.Rngs | None None): x self.avg_pool(nnx.relu(self.batch_norm1(self.dropout1(self.conv1(x), rngsrngs)))) x self.avg_pool(nnx.relu(self.batch_norm2(self.conv2(x)))) x x.reshape(x.shape[0], -1) # flatten x nnx.relu(self.dropout2(self.linear1(x), rngsrngs)) x self.linear2(x) return x # Instantiate the model. model CNN(rngsnnx.Rngs(0)) # Visualize it. nnx.display(model)各组件与源码的对应关系组件用法源码位置nnx.Conv二维卷积封装lax.conv_general_dilated支持padding、strides、kernel_init等参数flax/nnx/nn/linear.pynnx.BatchNorm批归一化训练时更新mean/var运行统计评估时使用 running averageflax/nnx/nn/normalization.pynnx.Dropout随机丢弃deterministicFalse生效、deterministicTrue关闭flax/nnx/nn/stochastic.pynnx.Linear全连接层LinearGeneral的特例axis-1flax/nnx/nn/linear.pynnx.avg_pool平均池化window_shape(2,2)、strides(2,2)即 2×2 下采样flax/nnx/nn结构细节与维度推演conv1输入(28,28,1)3×3卷积输出 32 通道 →(28,28,32)默认paddingVALID或指定方式随后avg_pool下采样到(14,14,32)。conv2输出 64 通道 →(14,14,64)再经avg_pool得到(7,7,64)。展平7×7×64 3136正好对应linear1的输入维度3136linear1输出 256linear2输出 10对应 0~9 十个数字类别。nnx.relu为逐元素激活partial提前绑定avg_pool参数调用时无需重复传参。4.2 随机数流nnx.Rngs所有需要随机初始化的层Conv、Linear、BatchNorm都接收rngsnnx.Rngs(0)。从源码看Rngs 是一个围绕 JAX PRNG key 与计数器RngKeyRngCount见 rnglib.py的抽象每次请求 key 时计数器自增并通过jax.random.fold_in派生新 key从而保证同一模型内的不同层、不同调用各得到唯一且可复现的随机数。调用model(jnp.ones((1, 28, 28, 1)), nnx.Rngs(0))即用新的随机数流做一次前向传播。4.3 运行模型import jax.numpy as jnp # JAX NumPy y model(jnp.ones((1, 28, 28, 1)), nnx.Rngs(0)) y输出形状为(1, 10)的 logits 张量每一行对应 10 个类别的未归一化得分后续通过softmax_cross_entropy计算损失、通过argmax得到预测类别。nnx.display(model)则在笔记本中输出模型结构树便于直观检查各子模块与参数。5. 创建优化器与评估指标5.1nnx.Optimizernnx.Optimizer同时持有模型引用与 Optax 优化器从而在梯度计算后原地更新参数import optax learning_rate 0.005 momentum 0.9 optimizer nnx.Optimizer( model, optax.adamw(learning_rate, momentum), wrtnnx.Param ) metrics nnx.MultiMetric( accuracynnx.metrics.Accuracy(), lossnnx.metrics.Average(loss), ) nnx.display(optimizer)从 flax/nnx/training/optimizer.py 的源码可知wrtnnx.Param是一个 filter指定优化器只跟踪并更新模型中的Param类型变量即可训练参数这与后续nnx.value_and_grad的wrt需保持一致构造时调用tx.init(nnx.state(model, wrt))初始化 Optax 状态并以OptState/OptArray/OptVariable变量包装见 optimizer.py自 Flax 0.11.0 起nnx.Optimizer不再持有model属性update必须同时传入(model, grads)optimizer.py 中的参数校验即为此设计如需旧式ModelAndOptimizer请另行选用。选择optax.adamw带权重衰减的 Adam作为更新规则学习率0.005、动量0.9。5.2nnx.MultiMetricnnx.MultiMetric聚合多个指标nnx.metrics.Accuracy()分类准确率内部按logits与labels计算 argmax 后比较nnx.metrics.Average(loss)滚动平均值argnameloss表示通过update(loss...)传值源码中它以total与count两个MetricState变量累加求和取平均并支持reset()清零flax/nnx/training/metrics.py。6. 定义训练与评估 step 函数6.1 损失函数与训练步def loss_fn(model: CNN, batch, rngs: nnx.Rngs | None None): logits model(batch[image], rngs) loss optax.softmax_cross_entropy_with_integer_labels( logitslogits, labelsbatch[label] ).mean() return loss, logits nnx.jit def train_step(model: CNN, optimizer: nnx.Optimizer, metrics: nnx.MultiMetric, rngs: nnx.Rngs, batch): Train for a single step. grad_fn nnx.value_and_grad(loss_fn, has_auxTrue) (loss, logits), grads grad_fn(model, batch, rngs) metrics.update(lossloss, logitslogits, labelsbatch[label]) # In-place updates. optimizer.update(model, grads) # In-place updates. nnx.jit def eval_step(model: CNN, metrics: nnx.MultiMetric, batch): loss, logits loss_fn(model, batch) metrics.update(lossloss, logitslogits, labelsbatch[label]) # In-place updates.要点解析损失optax.softmax_cross_entropy_with_integer_labels直接以整数标签计算交叉熵无需 one-hot 编码has_auxTrue让value_and_grad同时返回非梯度辅助输出logits。nnx.jit是jax.jit的“状态化”版本函数输入输出可以是 NNX 对象模型、优化器、指标、Rngs内部经 XLA 编译加速特别适配 TPU/GPU。nnx.value_and_gradjax.value_and_grad的状态化版本自动微分求梯度。原地更新metrics.update(...)与optimizer.update(model, grads)都是原地修改无需显式返回状态。这是 Flax NNX 引用语义的关键特性——变换会尊重传入 NNX 对象的引用语义并自动传播状态更新详见 Why Flax NNX 与 transforms 指南代码因此更简洁。训练与评估的差异train_step额外接收rngs供 Dropout 使用eval_step不传rngs因为评估视图已设置deterministicTrueDropout 关闭、无需随机 key。7. 定义测试集推理函数训练前先借助nnx.view创建两个共享同一份权重但行为不同的视图train_model nnx.view(model, deterministicFalse, use_running_averageFalse) eval_model nnx.view(model, deterministicTrue, use_running_averageTrue) nnx.jit def pred_step(model: CNN, batch): logits model(batch[image], None) return logits.argmax(axis1) def plot_predictions(test_batch, pred): fig, axs plt.subplots(5, 5, figsize(6, 6)) for i, ax in enumerate(axs.flatten()): ax.imshow(test_batch[image][i, ..., 0], cmapbinary) # ax.set_title(flabel{pred[i]}) color green if test_batch[label][i] pred[i] else red ax.text(0.05, 0.05, str(pred[i]), transformax.transAxes, colorcolor) ax.axis(off) return fignnx.view机制从 flax/nnx/module.py 源码看nnx.view(node, **kwargs)返回一个“静态属性被 kwargs 更新”的新节点且新节点与原节点共享底层 JAX 数组引用——训练期间对train_model参数的更新会同步反映到eval_model。它还支持onlyFilter限定只修改特定模块如onlynnx.Dropout。两个视图的角色train_modeldeterministicFalseDropout 开启、use_running_averageFalseBatchNorm 使用批内统计并更新运行均值/方差eval_modeldeterministicTrueDropout 关闭、use_running_averageTrueBatchNorm 使用累积的运行统计。pred_step对每个测试批次做一次前向logits.argmax(axis1)得到每个样本的预测类别。plot_predictions在5×5网格上展示测试图像预测正确标绿色、错误标红色便于定性评估模型表现。8. 训练循环与 TensorBoard 监控8.1 初始化 TensorBoard 写入器教程使用tensorboardX无需安装完整 TensorFlow 即可获得兼容SummaryWriter接口from tensorboardX import SummaryWriter import tensorboard writer SummaryWriter()打开监控面板有两种方式在浏览器访问localhost:6006或在 Jupyter 中直接加载扩展# %load_ext tensorboard # %tensorboard --logdir runswriter.add_scalar(train_loss, value, step)把某个训练步的标量写入runs/目录TensorBoard 实时读取并绘制曲线writer.add_figure(inference, fig, step)记录用plot_predictions生成的推理可视化图。8.2 训练循环rngs nnx.Rngs(0) for step, batch in enumerate(make_batches(train_ds, batch_size)): if step train_steps: break # Run the optimization for one step and make a stateful update to the following: # - The train states model parameters # - The optimizer state # - The training loss and accuracy batch metrics train_step(train_model, optimizer, metrics, rngs, batch) if step 0 and (step % eval_every 0 or step train_steps - 1): # Evaluation period passed. # Log the training metrics. for metric, value in metrics.compute().items(): # Compute the metrics. writer.add_scalar(ftrain_{metric}, value, step) # Record the metrics. metrics.reset() # Reset the metrics for the test set. # Compute the metrics on the test set after each training epoch. for test_batch in make_batches(test_ds, batch_size): eval_step(eval_model, metrics, test_batch) # Show predicted labels on a single test batch pred pred_step(eval_model, test_batch) fig plot_predictions(test_batch, pred) writer.add_figure(inference, fig, step) # Log the test metrics. for metric, value in metrics.compute().items(): writer.add_scalar(ftest_{metric}, value, step) # Record the metrics. metrics.reset() # Reset the metrics for the next training epoch.训练循环的关键节奏每步调用train_step(train_model, optimizer, metrics, rngs, batch)状态化更新模型参数、优化器状态与训练指标每当step 0且step % 200 0或到达最后一步时执行评估先记录并reset训练指标再遍历整个测试集调用eval_step累计测试指标最后用pred_step生成一批预测并记录推理图metrics.compute()返回各指标当前值如train_accuracy、test_lossmetrics.reset()在训练/测试阶段切换之间清零。注意教程在评估时没有执行optimizer的更新或rngs的推进仅更新指标——这正是训练/评估解耦的设计。下图为本教程同目录下保存的 TensorBoard 实际运行截图tensorboard_screenshot.png可以看到test_loss随训练步数下降并收敛约 0.07、test_accuracy上升至约 0.976、train_loss/train_accuracy同步收敛右下角inference区块展示模型在测试图像上的预测结果绿色为正确、红色为错误整个 2×2 推理图布局即为本教程训练循环记录的指标形态9. 将模型导出为 SavedModelOrbaxFlax 模型适合研究与实验但生产推理服务如 LiteRT、TensorFlow Serving通常要求 SavedModel 格式。Orbax 的导出工具链可无缝完成这一转换。9.1 准备工作# !pip install -U orbax-export[all] from orbax.export import JaxModule, ExportManager, ServingConfig import tensorflow as tf9.2 包装 JaxModule用JaxModule把评估模型与其预测方法绑定def exported_predict(model, y): return model(y, None) jax_module JaxModule(eval_model, exported_predict)这里exported_predict的第二参数y是待预测的图像输入第二个参数None对应前向传播的rngs评估视图下 Dropout 已关闭无需随机数流。9.3 声明输入签名导出机制要求输入签名是tf.TensorSpec组成的 PyTreesig [tf.TensorSpec(shape(1, 28, 28, 1), dtypetf.float32)]这明确告知 TensorFlow Servingexported_predict接收形状(1, 28, 28, 1)、类型float32的单个张量即一张 28×28 单通道灰度图。9.4 导出并保存export_mgr ExportManager(jax_module, [ ServingConfig(mnist_server, input_signaturesig) ]) output_dir/tmp/mnist_export export_mgr.save(output_dir)ExportManager将输入签名与JaxModule打包并以mnist_server为服务配置名导出到/tmp/mnist_export。之后即可将该目录加载到 LiteRT 或 TensorFlow Serving 中提供推理服务。10. 小结与进阶路线至此你已经用 Flax NNX 完成了从数据加载、CNN 定义、优化器与指标配置、状态化训练/评估、TensorBoard 监控到 SavedModel 导出的全流程。教程在仓库中的实现与配套材料教程本体docs_nnx/mnist_tutorial.md、docs_nnx/mnist_tutorial.ipynb核心源码flax/nnx/module.py、flax/nnx/nn/linear.py、flax/nnx/nn/normalization.py、flax/nnx/nn/stochastic.py、flax/nnx/training/optimizer.py、flax/nnx/training/metrics.py、flax/nnx/rnglib.py进阶阅读Why Flax NNX、Flax NNX 概念、transforms 指南、view 指南更多可运行的 NNX 示例见 examples/nnx_toy_examples含函数式 API、lifted transforms、train state、数据并行、VAE、层扫描、参数手术、FSDP 优化器等NNX 层的完整 API 参考见 docs_nnx/api_reference/flax.nnx本教程所用模型结构简单、训练步数有限在真实项目中可根据硬件资源调整batch_size、learning_rate、train_steps与网络深度并将训练循环替换为 Flax GSPMD 数据并行 等大规模训练方案。【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考