JAX 模型导出实战:使用 jax2tf Examples 生成 TensorFlow SavedModel 并在 Keras、Serving 与 TF.js 中复用
2026/9/10 2:03:59 网站建设 项目流程

JAX 模型导出实战:使用 jax2tf Examples 生成 TensorFlow SavedModel 并在 Keras、Serving 与 TF.js 中复用

【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax

本指南以 JAX 仓库jax/experimental/jax2tf/examples目录下的示例为核心,系统讲解如何把训练好的 JAX(含 Flax、Haiku)MNIST 模型通过jax2tf.convert转换成标准 TensorFlow 函数并导出为 SavedModel,以及如何进一步将导出的模型接入 TensorFlow Hub/Keras、TensorFlow Serving 与 TensorFlow.js 生态。读完本文,你将掌握(predict_fn, params)二元组封装、参数以变量或常量入图两种控制方式、batch 多签名导出、Keras 特征复用微调与 Serving 端到端部署的完整实战流程。

一、示例目录全景:从训练到跨框架复用

jax2tf的核心目标是:把 JAX 函数转换成行为上“就像用 TensorFlow 写出来”的标准 TensorFlow 函数。基于这一原则,jax/experimental/jax2tf/examples/目录提供了从模型训练到多端部署的完整示例链:

文件作用
mnist_lib.py两个 MNIST 实现:纯 JAX 的PureJaxMNIST与 Flax CNN 的FlaxMNIST,含训练、评估、可视化代码
saved_model_lib.py核心工具函数convert_and_save_model,把(predict_fn, params)保存为 SavedModel
saved_model_main.py可执行入口:训练 → 转换 → 保存 → 重载验证全流程
keras_reuse_main.py演示把 jax2tf SavedModel 作为hub.KerasLayer嵌入更大的 Keras 模型并微调
serving/TensorFlow Serving 部署补充示例,含model_server_request.py客户端

这些示例围绕一个核心思想展开:jax2tf 与 SavedModel 解耦jax2tf.convert产出的只是标准 TensorFlow 函数,保存 SavedModel 用的是普通 TensorFlow 代码,因此用户对 SavedModel 中保存的元数据拥有完全控制权,示例函数只是参考起点,正式项目中可以按需复制扩展。

补充说明一点背景:自 JAX 0.4.14 起,JAX 与 TensorFlow 互操作默认走原生序列化(native serialization)模式,即把目标函数用标准 JAX API 降级为 StableHLO,再包一层薄的 TensorFlow op(XlaCallModule)供 TensorFlow 调用,其语义与性能都忠实于原生 JAX 执行(详见 jax2tf 主文档)。

二、核心模式:把模型整理成(predict_fn, params)二元组

无论用 Flax、Haiku 还是纯 JAX 训练,导入 SavedModel 前都要把模型整理成一对元素:

  • predict_fn:双参数函数,签名固定为(params: Parameters, inputs: Inputs) -> Outputs。两个参数都必须是numpy.ndarray或其(嵌套的)tuple/list/dict 组合。特别注意:多输入模型必须打包成只有两个参数的函数,例如把多个输入收进一个 tuple/list/dict。
  • params:类型为Parameters的模型参数,作为predict_fn的第一个输入,并被保存为 SavedModel 的变量(variables)。

这样做的意义在于控制“哪些参数以变量形式单独保存、哪些以内嵌常量形式固化在函数图中”,动机有两个:

  1. 规避 GraphDef 大小限制:参数可能非常大,超过 SavedModel 中 GraphDef 部分的 2GB 上限(变量区不受此限制);
  2. 支持微调:把参数保存为变量后,后续可以修改参数值,例如在 TensorFlow 侧继续 fine-tune。

2.1 Flax 模型的封装配方

class MyModel(nn.Module): ... model = MyModel(*config_args, **kwargs) # 构造模型 optimizer = ... # 训练模型 params = optimizer.target # 取出训练好的参数 predict_fn = lambda params, input: model.apply({"params": params}, input)

mnist_lib.py中的FlaxMNIST正是这一配方的实现:FlaxMNIST.predict调用model.apply({"params": params}, inputs, with_classifier=with_classifier),训练结束后用functools.partial固定with_classifier并返回(predict_fn, params)(见 mnist_lib.py)。

多输入 Flax 模型只需把最后一行改为:

predict_fn = lambda params, input: model.apply({"params": params}, *input)

注意此时predict_fn只接收一个 tuple 输入,内部展开后作为多个输入传给model.apply

把所有参数内嵌到计算图中(变量区清空)的写法:

params = () predict_fn = lambda _, input: model.apply({"params": optimizer.target}, input)

原文给出了一组可复现的量级参考:mnist_lib.py中的 Flax MNIST 示例默认 GraphDef 约 150k、variables 区约 3MB;当把参数作为常量内嵌进 GraphDef 后,variables 区变为空,GraphDef 膨胀到约 13MB。内嵌方式可能让编译器通过常量折叠(constant-folding)生成更快的代码,代价是图变大且无法在 TensorFlow 侧改参数。

2.2 Haiku 模型的封装配方

Haiku(DeepMind 的神经网络库)模型整理方式类似:

model_fn = ... # 定义 Haiku 模型 net = hk.transform(model_fn) # 得到 (init, apply) 函数对 params = ... # 从 net.init() 出发训练你的模型 predict_fn = hk.without_apply_rng(net).apply

hk.transform把模型函数变换为显式参数的形式,hk.without_apply_rng(net).apply则去掉 apply 函数对 RNG 参数的要求,正好满足predict_fn双参数签名。原文档指出:同样的策略也适用于其他 JAX 神经网络库。

三、convert_and_save_model:把二元组写成 SavedModel

saved_model_lib.py提供了convert_and_save_model参考实现(saved_model_lib.py)。它几乎没有 jax2tf 特有逻辑——因为 jax2tf 的目标就是让你能用自己熟悉的 TensorFlow 代码保存模型。

3.1 函数签名与参数说明

convert_and_save_model( jax_fn, # (params, inputs) -> outputs 的 JAX 函数 params, # 模型参数,将保存为 SavedModel 的变量 model_dir, # 模型保存目录 *, input_signatures, # 输入签名序列(tf.TensorSpec 或其嵌套结构) polymorphic_shapes=None, # 形状多态(batch 维度用 None 表示) with_gradient=False, # 是否保存梯度(传给 jax2tf.convert) compile_model=True, # 是否启用 TF jit_compile(Serving 必需) saved_model_options=None) # 传给 tf.saved_model.save 的选项

各参数的关键约束:

  • input_signatures:必须至少给出一个,与jax_fn的第二个参数(输入)结构匹配。第一个 signature 会被保存为默认 serving signature;其余 signature 只用于按对应形状对jax_fn进行额外 tracing 和转换(即多 batch 预热)。
  • polymorphic_shapes:非空时传给jax2tf.convert作用于jax_fn第二个参数,此时只支持单个input_signature,且多态维度应写None
  • with_gradient:决定是否保存自定义梯度。开启后函数内部会要求tf.saved_model.SaveOptions(experimental_custom_gradients=True)(见 3.3)。
  • compile_model:为 Serving 场景必须开启(对应tf.function(jit_compile=True))。

3.2 内部实现三步走

从源码看,该函数的实现分为三步:

  1. 转换tf_fn = jax2tf.convert(jax_fn, with_gradient=with_gradient, polymorphic_shapes=[None, polymorphic_shapes])—— 第一个参数(参数)不做多态,第二个参数(输入)按需多态。
  2. 参数变量化param_vars = tf.nest.map_structure(lambda param: tf.Variable(param, trainable=with_gradient), params)—— 用tf.nest保持与params完全一致的嵌套结构,每个参数变成tf.Variable。若想要更有意义的变量名,文档建议用dm-tree包的tree.map_structure_with_path
  3. 打包保存tf_graph = tf.function(lambda inputs: tf_fn(param_vars, inputs), autograph=False, jit_compile=compile_model)把参数闭包进图;随后为第一个 signature 生成默认 serving signature,其余 signature 仅触发 tracing 缓存,最后用_ReusableSavedModelWrapper包装后调用tf.saved_model.save

_ReusableSavedModelWrapper继承tf.train.Checkpoint,实现了 TensorFlow Hub reusable saved models 接口 所需的属性:variables(展平后的参数)、trainable_variablesregularization_losses(默认为空列表,如需为正则项预留可追加无输入tf.function),并把tf_graph绑定为__call__。正是这个包装类让导出的 SavedModel 天然兼容 TensorFlow Hub。

3.3 梯度保存的两个易错点

jax2tf.convert默认会为降级后的 primal 函数标注tf.custom_gradient:TensorFlow 侧求导时惰性调用 JAX 的jax.vjp计算梯度,从而保证与 JAX 求导结果一致(尊重自定义梯度规则)。围绕梯度有两个坑:

  1. 若用with_gradient=True却忘记给tf.saved_model.saveSaveOptions(experimental_custom_gradients=True),加载时会收到警告Importing a function (...) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.,真正求导时可能抛TypeError: An op outside of the function building code is being passed a "Graph" tensor...convert_and_save_model内部在with_gradient=True时会自动补上该选项。
  2. 若 JAX 函数本身不可反向微分(例如使用了lax.while_loop),保存会失败并报ValueError: Error when tracing gradients for SavedModel。此时要么传with_gradient=False(内部改用tf.raw_ops.PreventGradient包裹,求导时报错),要么显式设experimental_custom_gradients=False,代价都是加载后的函数无法求梯度(详见 jax2tf 主文档“Saved model and differentiation”一节)。

四、saved_model_main.py:端到端命令行演练

saved_model_main.py是可执行文件,演示完整序列:训练 MNIST 模型 → 得到 (推断函数, 参数) → 用 jax2tf 按一个或多个 batch 尺寸转换 → 保存 SavedModel 并 dump 内容 → 用 TensorFlow 重载并验证与 JAX 推断结果一致 → 可选绘制训练数字与推断结果图像。

4.1 命令行参数速查

Flag默认值说明
--modelmnist_flax二选一:mnist_flax/mnist_pure_jax
--model_classifier_layerTrue是否包含分类层;False时只输出 logits/特征,供复用场景使用
--model_path/tmp/jax2tf/saved_modelsSavedModel 保存根路径
--model_version1版本号(Serving 需要,更大版本优先加载),下界为 1
--serving_batch_size1serving signature 的 batch 尺寸;-1表示 batch 多态转换保存;校验器要求其大于 0 或等于 -1
--num_epochs3训练轮数,下界 1
--generate_modelTrue是否训练并保存新模型;False时只测试已加载的 SavedModel(即--nogenerate_model
--compile_modelTrue是否启用 TensorFlow jit_compiler(Serving 必须)
--show_modelTrue是否用saved_model_cli show --all --dir ...打印 SavedModel 细节
--show_imagesFalse是否绘制训练样本与推断结果图
--test_savedmodelTrue是否用 TensorFlow 重载 SavedModel 并与 JAX 推断结果对比验证

4.2 默认三签名导出与 batch 多态

默认情况下,示例会针对三个 batch 尺寸(1、16、128)分别转换推断函数:三个签名分别是serving_batch_size(默认 1)、train_batch_size(128)、test_batch_size(16)对应的tf.TensorSpec,其中第一个作为默认 serving signature,后两个仅用于触发对应形状的 tracing。在 dump 出的 SavedModel 中可以看到这三份函数。

--serving_batch_size=-1时,改用 batch 多态路径:

input_signatures = [tf.TensorSpec((None,) + mnist_lib.input_shape, tf.float32)] polymorphic_shapes = "(batch, ...)"

此时只保存一个签名,batch 维度以None表示,polymorphic_shapes中的...表示其余维度沿用签名中的静态形状。加载后在 Serving/客户端侧可使用任意 batch 尺寸。

4.3 数值一致性验证与容差

train_and_save()在保存后(--test_savedmodel默认开启)会重载 SavedModel 并做一致性校验:

pure_restored_model = tf.saved_model.load(model_dir) test_input = np.ones((test_batch_size,) + mnist_lib.input_shape, dtype=np.float32) np.testing.assert_allclose( pure_restored_model(tf.convert_to_tensor(test_input)), predict_fn(predict_params, test_input), **tolerances)

容差按执行设备选取(见 saved_model_main.py 的tf_accelerator_and_tolerances):TPU 为atol=1e-6, rtol=1e-6,GPU 为atol=1e-6, rtol=1e-4,CPU 为atol=1e-5, rtol=1e-5

仓库自带的测试 saved_model_main_test.py 以参数化方式覆盖mnist_pure_jax/mnist_flax×serving_batch_size ∈ {1, -1}model_classifier_layer=False的特征提取模式,使用--mock_data假数据与 1 个 epoch 做轻量回归验证,可直接作为 CI 集成参考。

4.4 运行示例

依赖见 requirements.txt:tensorflow_datasetstensorflow_hubflax(另有 JAX 与 TensorFlow 本体)。典型运行:

python jax/experimental/jax2tf/examples/saved_model_main.py \ --model=mnist_flax \ --serving_batch_size=1 \ --num_epochs=3

SavedModel 默认落在/tmp/jax2tf/saved_models/mnist_flax/1/(路径拼接规则见savedmodel_dir()model_path/model_name[_features]/model_version)。

五、TensorFlow Hub 与 Keras:把特征提取器嵌入大模型微调

5.1 原理:去掉分类层导出特征

saved_model_main.py导出的 SavedModel 已实现 reusable saved models 接口,可直接被 Hub 使用。为了演示复用,MNIST 模型都带with_classifier关键字:置False时模型去掉最后一层分类层,SavedModel 只计算 logits(特征)。keras_reuse_main.py正是靠FLAGS.model_classifier_layer = False先导出“特征提取器”版模型,再将其嵌入更大的 Keras 模型。

5.2 Keras 复用完整流程

# 1) 训练并导出特征提取器(复用 saved_model_main 的全部逻辑) FLAGS.model_classifier_layer = False saved_model_main.train_and_save() feature_model_dir = saved_model_main.savedmodel_dir() # 2) 用 tf.distribute.OneDeviceStrategy 建立 Keras 模型 strategy = tf.distribute.OneDeviceStrategy(tf_accelerator) with strategy.scope(): images = tf.keras.layers.Input( mnist_lib.input_shape, batch_size=mnist_lib.train_batch_size) keras_feature_extractor = hub.KerasLayer(feature_model_dir, trainable=True) features = keras_feature_extractor(images) predictor = tf.keras.layers.Dense(10, activation="softmax") predictions = predictor(features) keras_model = tf.keras.Model(images, predictions) # 3) 在 TensorFlow 中编译并训练 keras_model.compile( loss=tf.keras.losses.categorical_crossentropy, optimizer=tf.keras.optimizers.SGD(learning_rate=0.01), metrics=["accuracy"]) keras_model.fit(train_ds, epochs=FLAGS.num_epochs, validation_data=test_ds)

关键点:

  • hub.KerasLayer(feature_model_dir, trainable=True)把 jax2tf SavedModel 变成 Keras 层,trainable=True使导出时保存的参数变量在微调中可更新(因为参数是以变量而非常量保存的,这正是二元组封装的价值所在);
  • tf.distribute.OneDeviceStrategy是高阶 API 下对tf.device(...)的等价物,CPU/GPU/TPU 均可用;高性能训练则应换用适当复制的 TF Distribution Strategy;
  • 示例用高阶 Keras API 展示,但同样可以在底层 TensorFlow 中把 restored features 放在tf.GradientTape下手动做梯度更新;
  • saved_model_main.py的所有 flag 对keras_reuse_main.py同样适用,例如可用--model=mnist_flax选择 Flax 模型。

5.3 数据集注意事项

mnist_lib.load_mnist中数据集 pipeline 的关键细节是drop_remainder=True,注释明确说明这对 Keras 使用很重要(保证每个 batch 形状一致)。训练与评估刻意使用了不同 batch 大小:train_batch_size = 128test_batch_size = 16,这也与 4.2 节三签名导出一一对应。

六、TensorFlow Serving 部署

jax2tf 生成的 SavedModel 与普通 TensorFlow SavedModel 的唯一区别是:函数图中可能包含需要启用 XLA 才能执行的 TF op,在模型服务器中通过命令行 flag 开启即可。完整步骤见 serving/README.md:

6.1 环境准备

pip install -e jax pip install flax jaxlib tensorflow_datasets tensorflow_serving_api tf_nightly DOCKER_IMAGE=tensorflow/serving:nightly docker pull ${DOCKER_IMAGE}

(原文档还提到 Google 内部版本的模型服务器支持 XLA,本文以开源 TensorFlow model server 为例,可通过--xla_cpu_compilation_enabled=true启用 CPU XLA。)

6.2 导出与启动服务器

MODEL_PATH=/tmp/jax2tf/saved_models MODEL=mnist_flax SERVING_BATCH_SIZE_SAVE=-1 # -1 表示 batch 多态;正数表示固定 batch SERVING_BATCH_SIZE=16 # 请求侧 batch,须与 SAVE 一致(除非 SAVE=-1) MODEL_VERSION=$(( 1 + ${MODEL_VERSION:-0 } )) # 训练并导出(SavedModel 落在 ${MODEL_PATH}/${MODEL}/${MODEL_VERSION}) python ${JAX2TF_EXAMPLES}/saved_model_main.py --model=${MODEL} \ --model_path=${MODEL_PATH} --model_version=${MODEL_VERSION} \ --serving_batch_size=${SERVING_BATCH_SIZE_SAVE} \ --compile_model --noshow_model # 检查导出结果(shape 首维为 -1 即 batch 多态模型) saved_model_cli show --all --dir ${MODEL_PATH}/${MODEL}/${MODEL_VERSION} # 以 XLA 编译启用状态启动模型服务器(8500 gRPC / 8501 HTTP REST) docker run -p 8500:8500 -p 8501:8501 \ --mount type=bind,source=${MODEL_PATH}/${MODEL}/,target=/models/${MODEL} \ -e MODEL_NAME=${MODEL} -t --rm --name=serving ${DOCKER_IMAGE} \ --xla_cpu_compilation_enabled=true &

要点:模型服务器会按版本号自动加载更新的模型,因此修改参数重新导出时只需递增MODEL_VERSION,无需重启服务器

6.3 发送推理请求

python ${JAX2TF_EXAMPLES}/serving/model_server_request.py --model_spec_name=${MODEL} \ --use_grpc --prediction_service_addr=localhost:8500 \ --serving_batch_size=${SERVING_BATCH_SIZE} --count_images=128

常见错误排查:若报Input to reshape is a tensor with 12544 values, but the requested shape has 784,说明请求 batch(12544/784=16)与服务器中加载模型的 batch(784/784=1)不一致——请保证导出与请求使用相同的--serving_batch_size,或改用 batch 多态导出(-1)后自由调整请求侧 batch;--count_images应为所选 batch 的整数倍。

6.4 实验变体

  • MODEL=mnist_pure_jax可切换到纯 JAX 实现的更简单模型,其余步骤不变;
  • batch 多态模型(SAVE=-1)下可任意改变SERVING_BATCH_SIZE重发请求;
  • 固定 batch 场景:SERVING_BATCH_SIZE_SAVE=16SERVING_BATCH_SIZE=16,只需重做导出与请求两步,无需重启服务器。

七、TensorFlow JavaScript 转换

jax2tf 生成的 SavedModel 也可以借助 SavedModel 的转换器转成 TensorFlow.js 可用格式。需要注意:当前这些转换器可能因部分 op 尚未实现而拒绝某些 jax2tf 生成的 SavedModel。一个部分可行的变通方案是jax2tf.convertenable_xla=False,让 jax2tf 避开有问题的 op,从而显著提高转换覆盖率——大多数(并非全部)Flax 示例都能以此方式成功转换。官方另附 Quickdraw 的 TF.js 示例(examples/tf_js/quickdraw/README.md),本仓库对应目录未包含该示例文件,但jax2tf.py中仍保留enable_xla参数可供使用。

八、小结与延伸阅读

jax/experimental/jax2tf/examples/为蓝本,一套完整的 JAX → TensorFlow 生产链路可以概括为:整理(predict_fn, params)二元组 →convert_and_save_model保存(可选 batch 多态)→ 用saved_model_cli校验 → 接入 Hub/Keras 微调、TensorFlow Serving 在线推理或 TF.js 端上部署。整个过程除jax2tf.convert一步外全部是标准 TensorFlow 代码,保证了保存元数据的完全可控与生态兼容。

若需深入底层机制,建议继续阅读 jax2tf 主文档(覆盖原生序列化、形状多态、call_tf反向调用、分片/分区支持等),并参考 saved_model_test.py 与 jax2tf_test.py 中的系统级测试来验证自定义场景。

【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询