
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_datasetsNNX 模块化模型定义训练/评估双阶段循环含 BatchNorm/Dropout 状态切换训练指标记录TensorBoard模型导出Orbax SavedModel从工程目录看examples/mnist 下文件职责划分非常清晰文件职责main.py程序入口解析命令行参数保持极简intentionally kept shorttrain.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, splittrain/test)的实现方式。运行前提本地需具备 Python 3 环境且 JAX 后端CPU/GPU/TPU可正常初始化。main.py会在启动时打印 JAX 进程与设备信息便于确认运行环境。三、启动训练一行命令跑通 MNISTREADME 给出的标准运行命令为python main.py --workdir/tmp/mnist --configconfigs/default.py拆解这两个必选参数main.py--workdir字符串类型模型数据TensorBoard 指标、导出模型的存储目录--configml_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_flagsMNIST 示例采用ml_collections 的 config_flags 机制定义并覆盖超参数。默认配置位于 configs/default.pydef 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 --configconfigs/default.py \ --config.learning_rate0.05 --config.num_epochs5上面这条命令会把学习率从默认 0.1 改成 0.05、训练轮数从 10 改成 5其余配置保持不变。这种默认配置文件 命令行增量覆盖的模式lock_configTrue可锁定配置防止意外修改非常适合做超参数扫描与实验复现。配置项作用与调参建议结合源码配置项默认值在训练循环中的实际作用learning_rate0.1传入optax.sgd(learning_rate, momentum)train.py直接决定步长momentum0.9SGD 动量系数加速收敛并抑制震荡batch_size128数据管道按此值切批drop_remainderTrue影响每轮迭代步数与显存占用num_epochs10外层训练循环的轮数train.pyREADME 基准输出即 10 轮的指标五、核心模型基于 NNX 的 CNN 实现本示例的最大看点在于它已经全面采用 Flax 新一代命令式 API——NNX而不是传统的 Linen。CNN 模型定义在 train.pyclass CNN(nnx.Module): 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): 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网络结构输入[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) ReLULinear(256→10)输出 10 类 logits。值得注意的 NNX 特性Dropout 需要显式传rngsself.dropout1(x, rngsrngs)随机性通过显式 RNG 传递可复现各层参数类nnx.Conv、nnx.Linear、nnx.BatchNorm分别见 flax/nnx/nn/linear.py、flax/nnx/nn/linear.py、flax/nnx/nn/normalization.py都由rngs初始化nnx.Dropoutflax/nnx/nn/stochastic.py与nnx.BatchNorm是有状态模块训练/评估模式切换通过model.train()/model.eval()完成见下文训练循环。六、训练与评估循环深度剖析train_and_evaluatetrain.py是全部逻辑的核心其流程如下1. 数据准备与模型实例化train_ds, test_ds get_datasets(config) model CNN(rngsnnx.Rngs(0)) optimizer nnx.Optimizer(model, optax.sgd(learning_rate, momentum), wrtnnx.Param) metrics nnx.MultiMetric( accuracynnx.metrics.Accuracy(), lossnnx.metrics.Average(loss), ) rngs nnx.Rngs(0)nnx.Rngs(0)以固定种子创建 RNG 流保证实验可复现nnx.Optimizerflax/nnx/training/optimizer.py绑定模型与 optax 优化器wrtnnx.Param表示只更新 Param 类型的变量如 BatchNorm 的均值/方差这类非 Param 变量不受影响nnx.MultiMetricflax/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_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.损失函数loss_fn使用optax.softmax_cross_entropy_with_integer_labels计算交叉熵并取均值train.pynnx.value_and_grad一次性同时得到损失值和梯度has_auxTrue携带 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.pyREADME 给出了一条参考输出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(...)写入 workdirTensorBoard 可读并在最后flush()。随后用Orbax export将模型导出为 SavedModeltrain.pyfrom 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), dtypetf.float32)] export_mgr ExportManager(jax_module, [ServingConfig(mnist_server, input_signaturesig)]) export_mgr.save(str(Path(workdir) / mnist_export))导出产物位于{workdir}/mnist_export输入签名固定为(1, 28, 28, 1)的 float32 张量服务名mnist_server——这意味着示例跑完后即可直接用于 TF Serving 之类的生产部署。七、数据加载与预处理细节get_datasetstrain.py展示了 tf.data 的标准流水线train_ds: tf.data.Dataset tfds.load(mnist, splittrain) test_ds: tf.data.Dataset tfds.load(mnist, splittest) # 像素归一化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_remainderTrue).prefetch(1)三个工程细节值得复制到自己的任务中归一化在数据管道内完成除以 255模型输入始终是[0,1]浮点shuffle(1024)用固定大小的缓冲池打乱避免全量洗牌的内存开销drop_remainderTrue丢弃尾批配合prefetch(1)隐藏 IO 延迟训练循环里as_numpy_iterator()直接取用。八、测试与基准如何验证你的示例单元测试train_test.pytest_cnn构造(1, 28, 28, 1)输入断言 CNN 输出形状为(1, 10)train_test.pytest_train_and_evaluate用tfds.testing.mock_data模拟 8 个样本、num_epochs1、batch_size8跑通完整训练评估流程验证代码路径可用性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 准确率default107.7m99.17%说明该数据来自官方示例的固定运行环境具体数值会因硬件CPU/GPU/TPU、JAX 版本与随机种子而略有波动参考价值在于数量级与收敛趋势约 99% 量级不宜当作硬性性能承诺。十、快速上手清单安装依赖见 requirements.txt确保 JAX 后端可用运行python main.py --workdir/tmp/mnist --configconfigs/default.py观察日志中每轮的train_loss / train_accuracy / test_loss / test_accuracy用--config.learning_rate0.05 --config.num_epochs5等参数做实验覆盖训练完成后用 TensorBoard 查看workdir下的标量曲线并在workdir/mnist_export拿到 SavedModel 用于部署需要复现官方指标可运行python -m mnist_benchmark类基准或参考 train_test.py 快速验证代码路径。十一、延伸阅读想理解 NNX 的模块与状态模型可深入 flax/nnx/module.py 与 flax/nnx/transforms/transforms.pyjit、value_and_grad的实现所在本示例使用的各层源码nnx.Conv/nnx.Linear见 flax/nnx/nn/linear.pynnx.BatchNorm见 flax/nnx/nn/normalization.pynnx.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),仅供参考