JAX 模型接入 TensorFlow Serving 实战:基于 jax2tf 导出 SavedModel 并部署 gRPC/REST 推理服务
【免费下载链接】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/serving/README.md 为核心,完整讲解如何把训练好的 JAX(含 Flax)模型通过jax2tf转换成标准 TensorFlow 函数、保存为 SavedModel,再部署到开源的 TensorFlow Serving 模型服务器上对外提供 gRPC 与 HTTP REST 推理服务。读完本文,你将掌握从模型训练、批量/形状多态导出、Docker 启动服务到客户端发请求、排查批次不匹配问题的完整闭环,并理解 jax2tf 生成 SavedModel 的底层机制(XlaCallModule、参数变量化与形状多态)。
一、为什么 jax2tf 的 SavedModel 可以给 TensorFlow Serving 用
jax2tf的目标是把 JAX 函数转换成"行为上如同直接用 TensorFlow 编写"的 Python 函数。这意味着转换得到的函数可以走标准的 TensorFlow 代码路径完成 tracing 与保存,例如tf.function、tf.Module与tf.saved_model.save。正如 jax2tf 主文档 所强调的:"jax2tf 调用之后的一切都是标准 TensorFlow 代码,SavedModel 的保存并不属于 jax2tf API 的一部分,用户对 SavedModel 中保存哪些元数据拥有完全的控制权"。
唯一与原生 TensorFlow 模型不同的地方在于:jax2tf 生成的函数图中可能包含XLA TF 算子(默认原生序列化模式下,整个 JAX 编译单元被包裹在一个名为XlaCallModule的薄 TensorFlow 算子中,内部承载序列化后的 StableHLO 程序),模型服务器需要开启 CPU/GPU 的 XLA 才能执行。这一要求通过模型服务器的一个命令行标志即可满足,除此之外与 TensorFlow 直接产出的 SavedModel 没有任何差别。
注意:本文针对的是开源的 TensorFlow Serving(
tensorflow/servingDocker 镜像);Google 内部版本的模型服务器说明位于 serving 示例目录的internal子目录中,当前仓库不包含该部分代码。
整个 serving 示例由两部分组成:
- saved_model_main.py:负责训练 MNIST 模型并导出 SavedModel;
- model_server_request.py:负责向模型服务器发送推理请求并统计准确率。
二、环境准备:安装 JAX、TensorFlow 与 Serving Docker 镜像
如果本地已装好 JAX 和 TensorFlow Serving,可以跳过大部分安装步骤,但必须设置下面两个环境变量:
JAX2TF_EXAMPLES DOCKER_IMAGE从零开始的完整准备步骤如下。
2.1 克隆 JAX 源码并安装 Python 依赖
git clone https://github.com/jax-ml/jax JAX2TF_EXAMPLES=$(pwd)/jax/jax/experimental/jax2tf/examples pip install -e jax pip install flax jaxlib tensorflow_datasets tensorflow_serving_api tf_nightly其中:
pip install -e jax以可编辑模式安装 JAX 源码本身(示例脚本位于仓库内,运行时会导入jax包);flax用于 Flax 版本的 MNIST 模型;tensorflow_datasets(TFDS)用于下载 MNIST 数据集,示例通过tfds.load("mnist")读取,见 mnist_lib.py;tensorflow_serving_api提供 gRPC 请求所需的 protobuf 类型(predict_pb2、prediction_service_pb2_grpc);- 使用
tf_nightly是为了拿到足够新、能支持XlaCallModule的 TensorFlow 版本。
示例代码对第三方依赖的最小声明可参考 jax2tf/examples/requirements.txt(tensorflow_datasets、tensorflow_hub、flax),实际运行 serving 示例还需额外的grpcio、requests、absl-py与matplotlib等运行时依赖。
2.2 安装 TensorFlow Serving Docker 镜像
DOCKER_IMAGE=tensorflow/serving:nightly docker pull ${DOCKER_IMAGE}这里同样选用 nightly 版本以获得对最新 XLA 算子的支持。
三、设置变量
在导出模型前,先定义一批贯穿全流程的快捷变量:
# 快捷变量 # SavedModel 的保存位置 MODEL_PATH=/tmp/jax2tf/saved_models # 示例模型,可选 "mnist_flax" 与 "mnist_pure_jax" MODEL=mnist_flax # SavedModel 的 batch size。用 -1 表示 batch 多态(任意批量), # 或用严格正数表示固定 batch size SERVING_BATCH_SIZE_SAVE=-1 # 发送给模型的 batch size。若 SERVING_BATCH_SIZE_SAVE 不是 -1,则二者必须相等 SERVING_BATCH_SIZE=16 # 每次修改模型参数并重新导出后,将该值加 1(版本号递增) MODEL_VERSION=$(( 1 + ${MODEL_VERSION:-0} ))各变量含义如下:
| 变量 | 作用 | 取值说明 |
|---|---|---|
MODEL_PATH | SavedModel 的根目录 | 默认/tmp/jax2tf/saved_models |
MODEL | 选择示例模型 | mnist_flax(Flax CNN)或mnist_pure_jax(纯 JAX MLP) |
SERVING_BATCH_SIZE_SAVE | 导出模型时的 batch size | -1表示 batch 多态;正整数表示固定 batch |
SERVING_BATCH_SIZE | 客户端请求的 batch size | 固定 batch 导出时必须与上面相等 |
MODEL_VERSION | Serving 的模型版本号 | 递增以保证模型服务器加载新版本 |
MODEL_VERSION的递增机制与 TensorFlow Serving 的版本管理约定一致:模型以模型名/版本号/的目录结构存放,模型服务器会自动发现并加载更大版本号的新模型,因此修改参数后只需重新导出并递增版本号,无需重启服务器。
四、训练并导出 SavedModel
使用saved_model_main.py完成训练与导出(该脚本的完整说明见 examples/README.md):
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命令执行后,SavedModel 会落在${MODEL_PATH}/${MODEL}/${MODEL_VERSION}目录下。
4.1 关键命令行参数(与源码对应)
结合 saved_model_main.py 中的 absl flags 定义,各参数含义如下:
| 参数 | 默认值 | 说明 |
|---|---|---|
--model | mnist_flax | 可选mnist_flax或mnist_pure_jax |
--model_path | /tmp/jax2tf/saved_models | SavedModel 保存根目录 |
--model_version | 1 | 版本号,lower_bound=1,更大版本在 serving 时优先 |
--serving_batch_size | 1 | 保存 serving 签名所用的 batch size;-1表示 batch 多态(校验器要求取值要么为-1,要么为正整数) |
--num_epochs | 3 | 训练轮数,lower_bound=1 |
--generate_model | True | 是否重新训练并保存;传--nogenerate_model可跳过训练、直接测试已有 SavedModel |
--compile_model | True | 是否对 SavedModel 启用 TensorFlowjit_compile,要用于 TensorFlow Serving 时必须开启 |
--show_model | True | 是否打印 SavedModel 详情;示例命令中用--noshow_model关闭 |
--test_savedmodel | True | 加载 SavedModel 用 TensorFlow 跑推理,并与 JAX 模型结果做数值比对 |
当--serving_batch_size=-1时,脚本构造的是 batch 多态的输入签名与形状多态描述:
input_signatures = [tf.TensorSpec((None,) + mnist_lib.input_shape, tf.float32)] polymorphic_shapes = "(batch, ...)"而当指定固定 batch 时,脚本会为 3 个 batch size(serving batch、训练 batch 128、测试 batch 16)各 trace 一份具体化签名,其中第一个签名会作为默认的 serving 签名保存(详见 saved_model_main.py)。mnist_lib.input_shape = (28, 28, 1)(不含 batch 维),训练与测试 batch 大小分别为 128 与 16(见 mnist_lib.py)。
4.2 模型内部结构:纯 JAX 与 Flax 两种实现
仓库为这个示例提供了两个 MNIST 实现(见 mnist_lib.py):
PureJaxMNIST:纯 JAX 实现,隐藏层尺寸[784, 512, 512, 10],使用jnp.dot + tanh前向、jax.grad更新参数、jax.jit加速,训练逻辑简单直观,适合快速理解;FlaxMNIST:Flaxnn.Module实现的 CNN(Conv(32) → relu → avg_pool → Conv(64) → relu → avg_pool → Dense(256) → Dense(10) + log_softmax),用optax.sgd(learning_rate=0.001, momentum=0.9)优化,使用model.apply({"params": params}, inputs)做前向。
两者训练完成后都返回一个二元组:(predict_fn, params),其中predict_fn是签名为(params, inputs) -> outputs的双参数函数,这正是后续convert_and_save_model需要的形态。
4.3 导出背后的关键函数convert_and_save_model
真正执行转换与保存的是 saved_model_lib.py 中的convert_and_save_model,它演示了"把 JAX 模型参数保存为 SavedModel 变量(而非常量)"的标准做法:
tf_fn = jax2tf.convert( jax_fn, with_gradient=with_gradient, polymorphic_shapes=[None, polymorphic_shapes]) # 将参数包装为 tf.Variable,使保存器把参数作为变量单独存储 param_vars = tf.nest.map_structure( lambda param: tf.Variable(param, trainable=with_gradient), params) tf_graph = tf.function(lambda inputs: tf_fn(param_vars, inputs), autograph=False, jit_compile=compile_model) # 第一个 input_signature 保存为默认 serving 签名 signatures[tf.saved_model.DEFAULT_SERVING_SIGNATURE_DEF_KEY] = \ tf_graph.get_concrete_function(input_signatures[0])这里的要点(也是 jax2tf 主文档 中反复强调的最佳实践):
- 参数必须作为函数入参传入,并用
tf.Variable包装。如果直接闭包捕获参数常量,参数会被嵌入计算图(GraphDef),既可能突破 GraphDef 的 2GB 上限,也无法支持后续微调;包装成tf.Variable后参数存放在 SavedModel 的variables区域,不受 2GB 限制; with_gradient=True时,jax2tf 会用tf.custom_gradient注解降低后的函数,在 TensorFlow 求导时惰性调用 JAX 的jax.vjp来计算梯度,从而保证 TF 侧微分与 JAX 微分结果一致;同时保存时需配合tf.saved_model.SaveOptions(experimental_custom_gradients=True)(源码中已自动处理);- 若 JAX 函数本身不可反向微分(如使用
lax.while_loop),导出会报ValueError: Error when tracing gradients for SavedModel,此时应传with_gradient=False。
五、检查导出的 SavedModel
saved_model_cli show --all --dir ${MODEL_PATH}/${MODEL}/${MODEL_VERSION}输出中如果某个签名的shape首维是-1,说明这是batch 多态模型(对任意批量都可用)。saved_model_cli也是确认输入名(示例中为"inputs")与输出名(如"output_0")的依据——客户端代码正是从 dump 结果中获知这些名字的(见 model_server_request.py 的注释)。
六、启动本地模型服务器(开启 XLA)
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 &各要素说明:
-p 8500:8500:gRPC 服务端口;-p 8501:8501:HTTP REST 服务端口;--mount type=bind,...:把模型目录挂载到容器内/models/${MODEL};-e MODEL_NAME=${MODEL}:告诉模型服务器要加载的模型名;--xla_cpu_compilation_enabled=true:关键标志,开启 CPU 侧的 XLA 编译,jax2tf 生成的XlaCallModule算子依赖它才能在模型服务器中执行;-t --rm --name=serving:分配伪终端、退出即删除容器、指定容器名。
模型服务器的版本管理特性意味着:只要递增${MODEL_VERSION}并重新导出,无需重启服务器,运行中的模型服务器会自动加载更新的版本。
七、发送推理请求(gRPC 与 REST 双通道)
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=1287.1 客户端脚本参数
对应 model_server_request.py 中的 flags:
| 参数 | 默认值 | 说明 |
|---|---|---|
--use_grpc | True | 使用 gRPC API(默认);传--nouse_grpc切换为 HTTP REST API |
--model_spec_name | "" | 导出时使用的模型名(如mnist_flax),对应模型服务器的模型名 |
--prediction_service_addr | localhost:8500 | 服务地址;本地 serving 时 gRPC 用localhost:8500,REST 用localhost:8501 |
--serving_batch_size | 1 | 请求 batch size,lower_bound=1,必须与模型保存时的 batch 匹配,且需整除--count_images |
--count_images | 16 | 要测试的图片总数,lower_bound=1 |
7.2 请求流程与准确率统计
脚本主流程(model_server_request.py)先校验count_images % serving_batch_size == 0,然后从 TFDS 加载 MNIST 测试集,按指定 batch 分批调用模型服务器,逐批计算预测数字与标签数字的一致率并打印运行准确率。
gRPC 路径的核心调用如下:
channel = grpc.insecure_channel(_PREDICTION_SERVICE_ADDR.value) stub = prediction_service_pb2_grpc.PredictionServiceStub(channel) request = predict_pb2.PredictRequest() request.model_spec.name = _MODEL_SPEC_NAME.value request.model_spec.signature_name = tf.saved_model.DEFAULT_SERVING_SIGNATURE_DEF_KEY # 输入名 "inputs" 可在 SavedModel dump 中查到 request.inputs["inputs"].CopyFrom( tf.make_tensor_proto(images, dtype=images.dtype, shape=images.shape)) response = stub.Predict(request) # 输出名 "output_0" 同样可在 SavedModel dump 中查到; # 也可以直接取第一个输出 outputs, = response.outputs.values() return tf.make_ndarray(outputs)REST 路径则向http://<addr>/v1/models/<model_name>:predict发送{"inputs": <json 数组>}的 POST 请求,并校验 HTTP 状态码(model_server_request.py)。
7.3 常见错误:batch size 不匹配
如果看到如下报错:
Input to reshape is a tensor with 12544 values, but the requested shape has 784含义是:请求的 batch size 为 16(12544 = 16 × 784),而模型服务器加载的模型 batch size 为 1(784 = 1 × 784)。请检查导出时的--serving_batch_size与发请求时的--serving_batch_size是否一致,或者改用 batch 多态方式导出(-1)。
八、进阶实验:切换模型与批量策略
8.1 使用纯 JAX 模型
MODEL=mnist_pure_jax然后从导出步骤(第四节)重新执行即可。mnist_pure_jax不依赖 Flax,前向计算就是显式的矩阵乘法加激活函数,便于把注意力集中在 jax2tf 与 Serving 的集成上。
8.2 任意 batch 发送(batch 多态模型)
如果导出时用了SERVING_BATCH_SIZE_SAVE=-1,可以随意改变SERVING_BATCH_SIZE的值,直接从发送请求步骤(第七节)重试即可——这正是形状多态(shape polymorphism)带来的便利:一份 SavedModel 服务任意批量。
8.3 改成固定 batch 导出
SERVING_BATCH_SIZE_SAVE=16 SERVING_BATCH_SIZE=16然后重做导出步骤与发送请求步骤(无需重启模型服务器)。注意此时--count_images必须是所选 batch size 的整数倍(脚本会显式校验)。
8.4 关于形状多态的底层机制
batch 多态之所以能用,是因为jax2tf.convert支持polymorphic_shapes参数:以"(batch, ...)"这样的形状描述符声明哪些维度是"维度变量",JAX 在 tracing 时对这些维度做符号化处理(如_占位符从tf.TensorSpec对应维度取值、...展开为一系列_)。其正确性契约是:tf.function(jax2tf.convert(f, polymorphic_shapes)).get_concrete_function(sig)(x)的结果与f(x)一致。需要留意的是,形状多态只适用于中间形状能表达为维度变量简单表达式(线性多项式)的程序,且维度变量必须能从输入形状唯一解出(如a * a、a + b这类规格会直接报错)。更详细的规则、约束(polymorphic_constraints)与边界情况见 jax2tf 主文档的 shape-polymorphic 章节。
九、从源码理解全链路:JAX → SavedModel → Serving
把整条链路串起来看:
- 训练:
mnist_lib.py中PureJaxMNIST.train/FlaxMNIST.train产出(predict_fn, params)二元组; - 转换:
saved_model_lib.convert_and_save_model调用jax2tf.convert,默认以原生序列化方式把 JAX 程序 lower 为 StableHLO,并封装进单个XlaCallModuleTF 算子; - 保存:
tf.saved_model.save将XlaCallModule连同作为tf.Variable的参数一起写入 SavedModel,第一个input_signature成为默认 serving 签名; - 服务:TensorFlow Serving 加载模型目录,在
--xla_cpu_compilation_enabled=true(GPU 场景对应 GPU XLA 标志)开启 XLA 后,XlaCallModule反序列化、编译并执行内嵌的 StableHLO; - 请求:客户端通过 gRPC(
PredictionService)或 REST(/v1/models/<name>:predict)调用默认 serving 签名,获得与 JAX 原生推理一致的结果。
需要记住的限制(详见 jax2tf 主文档的 Known issues 章节):
- 原生序列化的模块是平台相关的,在非序列化平台上执行会报
The current platform CPU is not among the platforms required by the module [CUDA]; - 默认序列化仅接受 StableHLO 等有稳定性保证的 dialect 与受允许的自定义调用(如 GPU 上的 PRNG 自定义调用);
- SavedModel 只保存一阶梯度;
- 恢复后的模型运行不需要 JAX,只需要带 XLA 的 TensorFlow。
十、小结
本文给出了在 JAX 仓库内把 jax2tf 与 TensorFlow Serving 结合起来的完整方案:设置环境变量 → 训练并导出(固定 batch 或 batch 多态)→saved_model_cli检查 → 带--xla_cpu_compilation_enabled=true启动 Serving → gRPC/REST 双通道发请求并统计准确率,同时覆盖了最常见的 batch 不匹配排错与多种批量策略实验。所有命令与参数均可直接对照仓库源码(saved_model_main.py、saved_model_lib.py、model_server_request.py、mnist_lib.py)逐行验证,可作为把 JAX/Flax 模型投入生产化推理服务的最小可运行模板。
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考