ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

使用 jax2tf 将 Flax MNIST 模型导出为 TensorFlow Lite:端侧图像分类完整实战指南

使用 jax2tf 将 Flax MNIST 模型导出为 TensorFlow Lite:端侧图像分类完整实战指南 使用 jax2tf 将 Flax MNIST 模型导出为 TensorFlow Lite端侧图像分类完整实战指南【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/gh_mirrors/jax/jax导读本文基于 JAX 仓库中 jax/experimental/jax2tf/examples/tflite/mnist 目录下的完整示例讲解如何用 JAX Flax 训练一个 MNIST 手写数字卷积神经网络再借助jax2tf将其转换为 TensorFlow 函数、导出为 TF Lite 模型最终部署到 Android 端侧应用做推理。读完本文你将掌握jax2tf.convert的核心用法、enable_xlaFalse的必要性、TF Lite 转换与后训练量化的完整流程以及端侧运行 Select TF Ops 模型的依赖配置方法。示例整体流程从 JAX 训练到端侧推理该示例完整覆盖了一条训练 → 图导出 → 端侧格式转换 → 移动端部署的流水线对应目录结构如下jax/experimental/jax2tf/examples/tflite/mnist/ ├── README.md # 示例说明文档 └── mnist.py # 训练 转换 量化 评估的入口脚本示例还依赖同目录上一级的两个模块位于 jax/experimental/jax2tf/examples/ 下mnist_lib.py定义了两套 MNIST 模型与训练代码——纯 JAX 实现PureJaxMNIST和 Flax Linen 实现FlaxMNIST以及数据集加载函数load_mnistrequirements.txt声明示例所需的第三方依赖。整个流程可以拆解为四个阶段数据准备通过 TensorFlow Datasets 下载 MNIST 数据集并做归一化、one-hot 编码、分批次等预处理模型训练用 Flax Linen 训练一个简单的 CNN共 10 个 epoch演示目的转换导出用jax2tf.convert把预测函数转换为 TensorFlow 函数再用tf.lite.TFLiteConverter转成 TF Lite 浮点模型并应用后训练量化生成量化模型部署验证用tf.lite.Interpreter加载模型在测试集上评估精度将最终量化模型写入文件供 Android 应用使用。示例的灵感来源有两处模型训练部分参考了 Flax 官方的 MNIST 分类示例端侧部署部分参考了 TensorFlow Lite 官方的 Android 手写数字分类器 Codelab——将两者结合就形成了JAX 训练、端侧推理的完整闭环。环境与依赖准备运行示例需要以下依赖见 requirements.txtFlax提供 Linen API 以定义和训练神经网络TensorFlow用于tf.function包装转换结果、TF Lite 转换器以及最终的解释器评估TensorFlow Datasets负责下载和预处理 MNIST 数据集NumPy用于精度统计等后处理逻辑。示例代码从 TensorFlow Datasets 下载 MNIST 数据集并在喂入神经网络前完成预处理。load_mnist的实现见 mnist_lib.py将图像像素值除以 255 归一化到[0, 1]标签转换为 10 维 one-hot 向量并依次执行cache()、shuffle(1000)、batch(batch_size, drop_remainderTrue)。其中训练批大小train_batch_size 128评估批大小test_batch_size 16特意让训练与评估使用不同的 batch size以验证转换后的模型对输入形状的适配能力。运行训练与转换一条命令完成全流程在满足依赖的前提下运行入口脚本即可一次性完成训练、SavedModel/函数转换、TF Lite 导出与精度评估python mnist.py假设数据集已下载整个训练过程大约耗时 1 分钟。数据集会被直接加载到/tmp/jax2tf/mnist下的data/目录中这是 TensorFlow Datasets 的默认缓存行为。mnist.py 通过absl.flags暴露了三个可调参数方便针对不同场景定制Flag默认值作用--tflite_file_path/tmp/mnist.tflite最终 TF Lite 模型文件的保存路径--serving_batch_size4转换时 serving signature 使用的批大小即输入签名中的 batch 维度--num_epochs10训练轮数epoch注意--serving_batch_size会同时用于两处——既作为tf.TensorSpec输入签名中的 batch 维度也作为测试集评估时的批大小。这意味着转换出的模型对输入形状是固定批大小的monomorphic后续部署时需要按该批大小喂入数据或改用 shape polymorphism 支持动态形状详见下文 jax2tf 能力扩展部分。脚本执行完毕后会在终端打印如下关键信息对应 mnist.py浮点模型大小KB量化模型大小KB及其占浮点模型大小的百分比浮点 TF Lite 模型在测试集上的精度——通常与原 Flax 模型精度一致因为两者本质上是同一模型的不同存储格式量化模型的精度以及相对浮点模型的精度下降accuracy drop数值。脚本最后把量化后的模型写入--tflite_file_path指定的路径默认/tmp/mnist.tflite。核心环节一用 jax2tf 将 Flax 模型转换为 TensorFlow 函数训练完成后需要从模型参数构建一个纯预测函数再交给jax2tf.convert转换。示例中的做法是见 mnist.pydef predict(image): return flax_predict(flax_params, image) # Convert your Flax model to TF function. tf_predict tf.function( jax2tf.convert(predict, enable_xlaFalse), input_signature[ tf.TensorSpec( shape[_SERVING_BATCH_SIZE, 28, 28, 1], dtypetf.float32, nameinput) ], autographFalse)这里有两个关键细节jax2tf.convert(predict, enable_xlaFalse)这是整个转换的核心调用。jax2tf.convert接受一个 JAX 函数其参数与返回值应为 JAX 数组或由 tuple/list/dict 组成的 pytree返回一个只使用 TensorFlow ops 实现、可在 TensorFlow 程序中直接调用的版本。enable_xlaFalse的含义与作用详见下一节。tf.functioninput_signature将转换结果包装为带输入签名的 TF 函数。签名声明输入形状为[serving_batch_size, 28, 28, 1]、类型tf.float32、名称input——这正是后续get_concrete_function()和 TF Lite 转换器所依赖的入口。设置autographFalse可避免 TensorFlow autograph 对 JAX 转换代码做额外的改写。从实现层面看jax2tf.convert定义于 jax/experimental/jax2tf/jax2tf.py是一个被api_util.api_hook标记的公开 API其完整签名还包括polymorphic_shapes、with_gradient、native_serialization等高级参数。对于本示例而言with_gradient默认为True会通过转换jax.vjp(fun)为输出函数附加tf.custom_gradient从而支持 TensorFlow 反向模式自动微分纯推理场景可将其保留默认值不影响 TF Lite 导出enable_xlaTrue时默认转换器会尽量使用最简化的 XLA TF ops 来降低某些 JAX 原语而这些 op 正是 TFLite/TF.js 转换器无法解析的enable_xlaFalse时转换器会更努力地用非 XLA 的 TensorFlow ops 完成 lowering若做不到则直接报错中止。该模式与native_serialization互斥native_serialization要求enable_xlaTrue。核心环节二理解enable_xlaFalse的动机与限制文档中特别强调enable_xlaFalse这个参数它指示转换器避免使用一批只有 XLA 编译器才支持的特殊 TensorFlow ops因为接下来的 TFLite 转换器还无法理解这些 op。要理解这一点需要回溯 jax2tf 的 lowering 机制。对于大多数 JAX 原语都能找到语义完全匹配的原生 TF op例如jax.lax.abs等价于tf.abs。但对于没有对应 TF op 的 JAX 原语例如jax.lax.conv_general_dilated转换器在enable_xlaTrue模式下会使用一层薄封装在 HLO op 之上的特殊 TF ops。这类 op只有链接了 XLA 的运行时才能执行而 TF.js 和 TFLite 转换器都不具备这一条件。因此 jax2tf 在 impl_no_xla.py 中为这些 op 提供了无 XLA 降级实现用 TensorFlow 原生支持的 ops 重新组合出等价语义。enable_xlaFalse即强制走这条路径。详细的支持矩阵记录在 no_xla_limitations.md 中核心要点如下XLA op对应 JAX 原语无 XLA 支持程度XlaDotlax.dot_general完整支持XlaDynamicSlicelax.dynamic_slice完整支持XlaDynamicUpdateSlicelax.dynamic_update_slice完整支持XlaPadlax.pad完整支持XlaConvlax.conv_general_dilated部分支持XlaGatherlax.gather部分支持XlaReduceWindowlax.reduce_window部分支持XlaScatterlax.scatter系列部分支持XlaSelectAndScatterlax._select_and_scatter_add不支持XlaReducelax.reduce/lax.argmin/lax.argmax不支持XlaVariadicSortlax.sort不支持对于与本示例直接相关的XlaConvCNN 卷积的 lowering无 XLA 支持的具体范围是仅支持 1D 和 2D 卷积lhs.ndim 3 or 4普通卷积和空洞atrous/dilated卷积通过tf.nn.conv2d实现转置卷积lhs_dilation ! 1通过tf.nn.conv2d_transpose实现仅支持SAME/VALID两种 padding深度可分离卷积in_channels feature_group_count 1通过tf.nn.depthwise_conv2d实现不支持 batch groups 与一般意义上的 feature groups多个大卷积叠加时可能存在相对较高的数值误差。本示例使用的 Flax CNN 只包含两个nn.Conv32 通道和 64 通道均为普通 3×3 卷积与两个nn.avg_pool池化正好落在上述受支持范围内因此可以顺利无 XLA 转换。avg_pool对应的lax.reduce_window通过tf.nn.avg_pool实现文档也提示这种转换可能产生先取均值再乘回窗口大小的冗余计算属正常现象可自行优化。核心环节三转换为 TF Lite 格式得到 TF 函数后即可交给 TF Lite 转换器。示例使用from_concrete_functions这个 API见 mnist.py# Convert your TF function to the TF Lite format. converter tf.lite.TFLiteConverter.from_concrete_functions( [tf_predict.get_concrete_function()], tf_predict) converter.target_spec.supported_ops [ tf.lite.OpsSet.TFLITE_BUILTINS, # enable TensorFlow Lite ops. tf.lite.OpsSet.SELECT_TF_OPS # enable TensorFlow ops. ] tflite_float_model converter.convert()两点需要特别说明get_concrete_function()从带输入签名的tf.function中取出具体的函数实现作为转换输入。因为前面已经通过tf.TensorSpec固定了输入形状这里可以得到唯一的 concrete function。SELECT_TF_OPS是必需的由于enable_xlaFalse的 lowering 结果中仍会包含原生 TensorFlow ops而非纯 TF Lite ops必须在supported_ops中同时启用TFLITE_BUILTINS和SELECT_TF_OPSTF Lite 运行时才能正确执行这些 TensorFlow ops。跳过这一项会导致转换或端侧运行失败。转换器本身还支持多种入口例如从 SavedModel 目录转换from_concrete_functions适合模型还只存在于内存的场景。核心环节四后训练量化因为转换出的 TF Lite 模型与从普通 TensorFlow SavedModel 转换出的模型在格式上没有差别所以可以无缝套用 TF Lite 标准的后训练量化流程# Re-convert the model to TF Lite using quantization. converter.optimizations [tf.lite.Optimize.DEFAULT] tflite_quantized_model converter.convert()tf.lite.Optimize.DEFAULT启用默认优化策略典型表现为权重/激活量化显著缩小模型体积。示例脚本随后会同时评估浮点模型与量化模型在测试集上的精度并打印两者之差accuracy drop。文档与代码注释都提醒量化模型的精度偶尔会高于原浮点模型这是量化过程中的正常现象不必惊讶。模型最终以量化形式保存体积相对浮点版本大幅缩小更适合移动端部署与低延迟推理。部署到 Android 应用TF Lite 模型导出后可以参照构建手写数字分类器 Android 应用的 Codelab 教程将其集成进 Android 应用。关键点是凡是通过SELECT_TF_OPS转换得到的 TF Lite 模型客户端必须使用包含 TensorFlow op 支持库的 TF Lite 运行时否则会遇到未知 op 的解析错误。具体做法是在应用的 Gradle 依赖中加入 TF Lite 的 nightly 依赖以及 Select TF Ops 支持库dependencies { implementation org.tensorflow:tensorflow-lite:0.0.0-nightly // This dependency adds the necessary TF op support. implementation org.tensorflow:tensorflow-lite-select-tf-ops:0.0.0-nightly }其中tensorflow-lite-select-tf-ops负责在端侧提供 TensorFlow op 的运行时支持。此外还可以通过abiFilters限制 ABI 种类减小 TensorFlow op 相关依赖的包体积android { defaultConfig { ndk { abiFilters armeabi-v7a, arm64-v8a } } }注意输入图片的预处理例如本示例中的除以 255 归一化需要与应用侧的图像获取逻辑对齐同时 serving 批大小已在转换时固定默认 4端侧喂入数据时应与之匹配。部署前的本地验证用 TFLite Interpreter 评估模型在把模型发布到 Android 之前示例脚本已经内置了本地验证环节。evaluate_tflite_model函数见 mnist.py演示了标准的 TF Lite 推理 API 用法用tf.lite.Interpreter(model_contenttflite_model)基于内存中的模型字节构建解释器调用allocate_tensors()分配张量通过get_input_details()[0][index]与get_output_details()[0][index]获取输入输出张量索引对测试集逐批执行set_tensor(...)→invoke()→ 读取输出对输出做np.argmax(axis1)取概率最大的类别与 ground-truth 标签比较并统计准确率。这套流程与 Android 端Interpreter的用法一致可作为端侧集成的预演。jax2tf 能力的进一步扩展本示例是 jax2tf 的入门级用法固定形状、无梯度、无原生序列化。如果希望把固定批大小扩展为动态形状可以关注jax2tf.convert的polymorphic_shapes参数见 jax2tf.py 的 docstring例如传入batch, ...将 batch 维度声明为符号变量从而生成可接受任意批大小的 TF 函数——但这属于实验性功能部分 JAX 程序可能被拒绝。更多已知问题与限制可查阅 jax2tf 主 README 与 primitives_with_limited_support.md。总结本示例以 MNIST 手写数字识别为载题完整演示了 JAX 生态与 TensorFlow 移动端生态的衔接方式Flax 训练 →jax2tf.convert(..., enable_xlaFalse)→tf.function签名固化 →TFLiteConverter转换 →Optimize.DEFAULT量化 → Android Gradle 依赖集成。其中enable_xlaFalse与SELECT_TF_OPS是保证 TFLite 兼容性的两个关键开关理解它们背后的 XLA op 限制矩阵见 no_xla_limitations.md是迁移更大、更复杂模型到端侧时排查转换失败问题的出发点。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/gh_mirrors/jax/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进