Flax 官方 MNIST 分类示例完全指南:从 CNN 训练到 SavedModel 导出
【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax
导读
本文围绕 Flax 仓库中的 examples/mnist 官方示例展开,带你完整走通一个基于 Flax NNX 的 MNIST 手写数字分类实战流程:从命令行启动训练、用 ml_collections 覆盖超参数、理解 CNN 网络结构与训练/评估循环,到最终将模型导出为 TensorFlow SavedModel。读完本文,你将掌握 Flax 示例工程的目录组织方式、config_flags配置覆盖机制,以及 NNX 模块化训练(含 BatchNorm/Dropout 状态切换、指标聚合、优化器原地更新)的完整套路,可直接迁移到自己的图像分类任务中。
一、示例概览:一个"麻雀虽小、五脏俱全"的 Flax 工程
examples/mnist/README.md 明确说明:该示例在 MNIST 数据集上训练一个简单的卷积网络(Trains a simple convolutional network on the MNIST dataset)。它不仅仅是"跑个准确率"的玩具代码,而是一个完整的工程示范,涵盖:
- 命令行入口(absl flags + ml_collections 配置)
- 数据加载与预处理(tensorflow_datasets)
- NNX 模块化模型定义
- 训练/评估双阶段循环(含 BatchNorm/Dropout 状态切换)
- 训练指标记录(TensorBoard)
- 模型导出(Orbax SavedModel)
从工程目录看,examples/mnist 下文件职责划分非常清晰:
| 文件 | 职责 |
|---|---|
| main.py | 程序入口,解析命令行参数,保持极简("intentionally kept short") |
| train.py | 训练库文件,包含 CNN 模型、数据加载、训练与评估循环 |
| configs/default.py | 默认超参数配置 |
| mnist_benchmark.py | 性能基准测试(CPU 全量训练) |
| train_test.py | 单元测试(模型形状 + 单步训练) |
| requirements.txt | 依赖清单 |
| mnist.ipynb | 交互式 Notebook 版本 |
其中 main.py 的文档字符串点明了这种分层设计的初衷:入口文件刻意保持简短,核心逻辑放在可以轻松被测试和导入的库文件(train.py)中。
二、运行环境与依赖安装
requirements.txt 给出了本示例的关键依赖(该清单基于 Flax 0.4.1 / JAX 0.3.4 时代,实际运行时建议使用当前仓库主分支对应的较新版本):
absl-py # 命令行 flags 与日志 clu # platform.work_unit() 工作单元管理 flax # 神经网络库本体 jax / jaxlib # 数值计算后端(jaxlib 需匹配 CUDA 版本) ml-collections # ConfigDict 超参数配置 optax # 优化器(SGD with momentum) tensorflow # 数据管道(tf.data) tensorflow-datasets # MNIST 数据源README 中的 Requirements 只有一条:TensorFlow datasetmnist会在需要时自动下载并准备,无需手动准备数据。这也是 train.py 中tfds.load('mnist', split='train'/'test')的实现方式。
运行前提:本地需具备 Python 3 环境,且 JAX 后端(CPU/GPU/TPU)可正常初始化。
main.py会在启动时打印 JAX 进程与设备信息,便于确认运行环境。
三、启动训练:一行命令跑通 MNIST
README 给出的标准运行命令为:
python main.py --workdir=/tmp/mnist --config=configs/default.py拆解这两个必选参数(main.py):
--workdir:字符串类型,模型数据(TensorBoard 指标、导出模型)的存储目录;--config:ml_collections 配置文件路径,加载训练超参数。注意 main.py 通过flags.mark_flags_as_required(['config', 'workdir'])将两者设为必填,漏传会直接报错。
main()入口函数(main.py)还做了几件容易被忽略但很重要的事:
- 禁用 TensorFlow 的 GPU 可见性(
tf.config.experimental.set_visible_devices([], 'GPU')):防止 TF 提前占用显存,把 GPU 完全留给 JAX; - 记录 JAX 进程/设备信息(
jax.process_index()、jax.local_devices()); - 通过
clu.platform.work_unit()设置任务状态并创建 workdir 工件; - 调用
train.train_and_evaluate(FLAGS.config, FLAGS.workdir)进入训练主流程。
四、超参数配置与命令行覆盖(config_flags)
MNIST 示例采用ml_collections 的 config_flags 机制定义并覆盖超参数。默认配置位于 configs/default.py:
def get_config(): config = ml_collections.ConfigDict() config.learning_rate = 0.1 # SGD 学习率 config.momentum = 0.9 # SGD 动量 config.batch_size = 128 # 批大小 config.num_epochs = 10 # 训练轮数 return configREADME 特别强调 config_flags 允许在命令行直接覆盖配置字段,语法为--config.字段名=新值:
python main.py \ --workdir=/tmp/mnist --config=configs/default.py \ --config.learning_rate=0.05 --config.num_epochs=5上面这条命令会把学习率从默认 0.1 改成 0.05、训练轮数从 10 改成 5,其余配置保持不变。这种"默认配置文件 + 命令行增量覆盖"的模式(lock_config=True可锁定配置防止意外修改)非常适合做超参数扫描与实验复现。
配置项作用与调参建议(结合源码)
| 配置项 | 默认值 | 在训练循环中的实际作用 |
|---|---|---|
learning_rate | 0.1 | 传入optax.sgd(learning_rate, momentum)(train.py),直接决定步长 |
momentum | 0.9 | SGD 动量系数,加速收敛并抑制震荡 |
batch_size | 128 | 数据管道按此值切批(drop_remainder=True),影响每轮迭代步数与显存占用 |
num_epochs | 10 | 外层训练循环的轮数(train.py),README 基准输出即 10 轮的指标 |
五、核心模型:基于 NNX 的 CNN 实现
本示例的最大看点在于:它已经全面采用 Flax 新一代命令式 API——NNX,而不是传统的 Linen。CNN 模型定义在 train.py:
class CNN(nnx.Module): def __init__(self, rngs: nnx.Rngs): self.conv1 = nnx.Conv(1, 32, kernel_size=(3, 3), rngs=rngs) self.batch_norm1 = nnx.BatchNorm(32, rngs=rngs) self.dropout1 = nnx.Dropout(rate=0.025) self.conv2 = nnx.Conv(32, 64, kernel_size=(3, 3), rngs=rngs) self.batch_norm2 = nnx.BatchNorm(64, rngs=rngs) self.avg_pool = partial(nnx.avg_pool, window_shape=(2, 2), strides=(2, 2)) self.linear1 = nnx.Linear(3136, 256, rngs=rngs) self.dropout2 = nnx.Dropout(rate=0.025) self.linear2 = nnx.Linear(256, 10, rngs=rngs) def __call__(self, x, rngs: nnx.Rngs): x = self.avg_pool(nnx.relu(self.batch_norm1(self.dropout1(self.conv1(x), rngs=rngs)))) 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), rngs=rngs)) x = self.linear2(x) return x网络结构(输入[B, 28, 28, 1]灰度图):
Conv(1→32, 3×3)+BatchNorm(32)+Dropout(0.025)+ ReLU + 2×2 平均池化;Conv(32→64, 3×3)+BatchNorm(64)+ ReLU + 2×2 平均池化;- 展平为 3136 维 →
Linear(3136→256)+ Dropout(0.025) + ReLU; Linear(256→10)输出 10 类 logits。
值得注意的 NNX 特性:
- Dropout 需要显式传
rngs(self.dropout1(x, rngs=rngs)),随机性通过显式 RNG 传递,可复现; - 各层参数类(
nnx.Conv、nnx.Linear、nnx.BatchNorm分别见 flax/nnx/nn/linear.py、flax/nnx/nn/linear.py、flax/nnx/nn/normalization.py)都由rngs初始化; nnx.Dropout(flax/nnx/nn/stochastic.py)与nnx.BatchNorm是有状态模块:训练/评估模式切换通过model.train()/model.eval()完成(见下文训练循环)。
六、训练与评估循环深度剖析
train_and_evaluate(train.py)是全部逻辑的核心,其流程如下:
1. 数据准备与模型实例化
train_ds, test_ds = get_datasets(config) model = CNN(rngs=nnx.Rngs(0)) optimizer = nnx.Optimizer(model, optax.sgd(learning_rate, momentum), wrt=nnx.Param) metrics = nnx.MultiMetric( accuracy=nnx.metrics.Accuracy(), loss=nnx.metrics.Average('loss'), ) rngs = nnx.Rngs(0)nnx.Rngs(0)以固定种子创建 RNG 流,保证实验可复现;nnx.Optimizer(flax/nnx/training/optimizer.py)绑定模型与 optax 优化器,wrt=nnx.Param表示只更新 Param 类型的变量(如 BatchNorm 的均值/方差这类非 Param 变量不受影响);nnx.MultiMetric(flax/nnx/training/metrics.py)同时聚合准确率与平均损失。
2. 训练步骤:JIT 编译 + 原地更新
@nnx.jit def train_step(model, optimizer, metrics, batch, rngs): grad_fn = nnx.value_and_grad(loss_fn, has_aux=True) (loss, logits), grads = grad_fn(model, batch, rngs) metrics.update(loss=loss, logits=logits, labels=batch['label']) # In-place updates. optimizer.update(model, grads) # In-place updates.- 损失函数
loss_fn使用optax.softmax_cross_entropy_with_integer_labels计算交叉熵并取均值(train.py); nnx.value_and_grad一次性同时得到损失值和梯度,has_aux=True携带 logits 供指标更新使用;@nnx.jit对训练步骤做 JIT 编译加速;- 关键点:metrics 和 optimizer 的更新都是 in-place 的(原地修改状态),这是 NNX 区别于纯函数式 Linen 的标志性设计。
3. 每轮循环:状态切换与指标计算
for epoch in range(1, config.num_epochs + 1): model.train() # 切换到训练模式(启用 dropout,更新 BN 统计量) for batch in train_ds.as_numpy_iterator(): train_step(model, optimizer, metrics, batch, rngs) train_metrics = metrics.compute() metrics.reset() model.eval() # 切换到评估模式(关闭 dropout,使用 BN 运行均值/方差) for batch in test_ds.as_numpy_iterator(): eval_step(model, metrics, batch) eval_metrics = metrics.compute() metrics.reset()这里体现的 NNX 状态语义非常实用:model.train()/model.eval()在模块内部递归切换所有子模块模式——训练时 BatchNorm 更新 running statistics、Dropout 生效;评估时 Dropout 关闭、BatchNorm 使用累计统计量。测试集评估在每轮训练后执行一次,因此日志能看到每轮的 train/test 两组指标。
4. 日志输出格式
训练日志由 absl logging 输出(train.py),README 给出了一条参考输出(100% 为百分比化后的准确率,实际示例中模型输出的是 0~1 的小数,日志中乘以 100 展示):
I1009 17:56:42.674334 3280981 train.py:175] epoch: 10, train_loss: 0.0073, train_accuracy: 99.75, test_loss: 0.0294, test_accuracy: 99.255. 指标落盘与模型导出
每轮结束后,训练/测试损失与准确率通过summary_writer.scalar(...)写入 workdir(TensorBoard 可读),并在最后flush()。随后用Orbax export将模型导出为 SavedModel(train.py):
from orbax.export import JaxModule, ExportManager, ServingConfig def exported_predict(model, y): return model(y, None) model.eval() jax_module = JaxModule(model, exported_predict) sig = [tf.TensorSpec(shape=(1, 28, 28, 1), dtype=tf.float32)] export_mgr = ExportManager(jax_module, [ServingConfig('mnist_server', input_signature=sig)]) export_mgr.save(str(Path(workdir) / 'mnist_export'))导出产物位于{workdir}/mnist_export,输入签名固定为(1, 28, 28, 1)的 float32 张量,服务名mnist_server——这意味着示例跑完后即可直接用于 TF Serving 之类的生产部署。
七、数据加载与预处理细节
get_datasets(train.py)展示了 tf.data 的标准流水线:
train_ds: tf.data.Dataset = tfds.load('mnist', split='train') test_ds: tf.data.Dataset = tfds.load('mnist', split='test') # 像素归一化:uint8 → float32,除以 255 缩放到 [0, 1] train_ds = train_ds.map(lambda sample: { 'image': tf.cast(sample['image'], tf.float32) / 255, 'label': sample['label'], }) # 训练集 shuffle 缓冲 1024 个样本 train_ds = train_ds.shuffle(1024) # 按 batch_size 切批、丢弃不完整批次、prefetch(1) 预取加速 train_ds = train_ds.batch(batch_size, drop_remainder=True).prefetch(1)三个工程细节值得复制到自己的任务中:
- 归一化在数据管道内完成(除以 255),模型输入始终是
[0,1]浮点; shuffle(1024)用固定大小的缓冲池打乱,避免全量洗牌的内存开销;drop_remainder=True丢弃尾批,配合prefetch(1)隐藏 IO 延迟,训练循环里as_numpy_iterator()直接取用。
八、测试与基准:如何验证你的示例
单元测试(train_test.py)
test_cnn:构造(1, 28, 28, 1)输入,断言 CNN 输出形状为(1, 10)(train_test.py);test_train_and_evaluate:用tfds.testing.mock_data模拟 8 个样本、num_epochs=1、batch_size=8跑通完整训练评估流程,验证代码路径可用性(train_test.py);- 测试文件还硬编码了
CNN_PARAMS = 825_034作为参数量参照——你可以自行核对模型的 82.5 万参数。
性能基准(mnist_benchmark.py)
该文件把整个训练流程封装进flax.testing.Benchmark:
- 执行 CPU 全量训练(
main.main([])),统计总墙钟时间; - 从 TensorBoard summaries 读取每轮 eval_accuracy,计算
sec_per_epoch与最终准确率; - 断言最终准确率落在
[0.98, 1.0]区间(mnist_benchmark.py),并上报三项指标:wall_time、sec_per_epoch、accuracy。
这也解释了 README 基准表格的来历——它是可复现的自动化基准,而非一次性手工记录。
九、官方参考指标
README 记录了 default 配置(10 epochs)下的官方基准输出,可作为你自己运行的对照基线:
| 名称 | Epochs | 墙钟时间 | Top-1 准确率 |
|---|---|---|---|
| default | 10 | 7.7m | 99.17% |
说明:该数据来自官方示例的固定运行环境,具体数值会因硬件(CPU/GPU/TPU)、JAX 版本与随机种子而略有波动;参考价值在于数量级与收敛趋势(约 99% 量级),不宜当作硬性性能承诺。
十、快速上手清单
- 安装依赖(见 requirements.txt),确保 JAX 后端可用;
- 运行
python main.py --workdir=/tmp/mnist --config=configs/default.py; - 观察日志中每轮的
train_loss / train_accuracy / test_loss / test_accuracy; - 用
--config.learning_rate=0.05 --config.num_epochs=5等参数做实验覆盖; - 训练完成后用 TensorBoard 查看
workdir下的标量曲线,并在workdir/mnist_export拿到 SavedModel 用于部署; - 需要复现官方指标可运行
python -m mnist_benchmark类基准,或参考 train_test.py 快速验证代码路径。
十一、延伸阅读
- 想理解 NNX 的模块与状态模型,可深入 flax/nnx/module.py 与 flax/nnx/transforms/transforms.py(
jit、value_and_grad的实现所在); - 本示例使用的各层源码:
nnx.Conv/nnx.Linear见 flax/nnx/nn/linear.py,nnx.BatchNorm见 flax/nnx/nn/normalization.py,nnx.Dropout见 flax/nnx/nn/stochastic.py; - 指标与优化器 API 见 flax/nnx/training/metrics.py 与 flax/nnx/training/optimizer.py;
- 仓库中其他官方示例(如 examples/imagenet、examples/sst2)采用同样的
main.py + configs/工程骨架,可对比学习更复杂的模型与训练策略。
【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考