news 2026/9/10 1:34:01

JAX 模型接入 TensorFlow Serving 实战:基于 jax2tf 导出 SavedModel 并部署 gRPC/REST 推理服务

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
JAX 模型接入 TensorFlow Serving 实战:基于 jax2tf 导出 SavedModel 并部署 gRPC/REST 推理服务

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.functiontf.Moduletf.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_pb2prediction_service_pb2_grpc);
  • 使用tf_nightly是为了拿到足够新、能支持XlaCallModule的 TensorFlow 版本。

示例代码对第三方依赖的最小声明可参考 jax2tf/examples/requirements.txt(tensorflow_datasetstensorflow_hubflax),实际运行 serving 示例还需额外的grpciorequestsabsl-pymatplotlib等运行时依赖。

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_PATHSavedModel 的根目录默认/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_VERSIONServing 的模型版本号递增以保证模型服务器加载新版本

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 定义,各参数含义如下:

参数默认值说明
--modelmnist_flax可选mnist_flaxmnist_pure_jax
--model_path/tmp/jax2tf/saved_modelsSavedModel 保存根目录
--model_version1版本号,lower_bound=1,更大版本在 serving 时优先
--serving_batch_size1保存 serving 签名所用的 batch size;-1表示 batch 多态(校验器要求取值要么为-1,要么为正整数)
--num_epochs3训练轮数,lower_bound=1
--generate_modelTrue是否重新训练并保存;传--nogenerate_model可跳过训练、直接测试已有 SavedModel
--compile_modelTrue是否对 SavedModel 启用 TensorFlowjit_compile要用于 TensorFlow Serving 时必须开启
--show_modelTrue是否打印 SavedModel 详情;示例命令中用--noshow_model关闭
--test_savedmodelTrue加载 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=128

7.1 客户端脚本参数

对应 model_server_request.py 中的 flags:

参数默认值说明
--use_grpcTrue使用 gRPC API(默认);传--nouse_grpc切换为 HTTP REST API
--model_spec_name""导出时使用的模型名(如mnist_flax),对应模型服务器的模型名
--prediction_service_addrlocalhost:8500服务地址;本地 serving 时 gRPC 用localhost:8500,REST 用localhost:8501
--serving_batch_size1请求 batch size,lower_bound=1,必须与模型保存时的 batch 匹配,且需整除--count_images
--count_images16要测试的图片总数,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 * aa + b这类规格会直接报错)。更详细的规则、约束(polymorphic_constraints)与边界情况见 jax2tf 主文档的 shape-polymorphic 章节。

九、从源码理解全链路:JAX → SavedModel → Serving

把整条链路串起来看:

  1. 训练mnist_lib.pyPureJaxMNIST.train/FlaxMNIST.train产出(predict_fn, params)二元组;
  2. 转换saved_model_lib.convert_and_save_model调用jax2tf.convert,默认以原生序列化方式把 JAX 程序 lower 为 StableHLO,并封装进单个XlaCallModuleTF 算子;
  3. 保存tf.saved_model.saveXlaCallModule连同作为tf.Variable的参数一起写入 SavedModel,第一个input_signature成为默认 serving 签名;
  4. 服务:TensorFlow Serving 加载模型目录,在--xla_cpu_compilation_enabled=true(GPU 场景对应 GPU XLA 标志)开启 XLA 后,XlaCallModule反序列化、编译并执行内嵌的 StableHLO;
  5. 请求:客户端通过 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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/10 1:33:43

MATLAB实现计及碳交易与需求响应的微网虚拟电厂日前优化调度

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/10 1:33:13

用开源组件搭建微信用户画像系统:从数据采集到画像看板

上周有个做私域运营的朋友问我&#xff1a;"微信上几百个客户&#xff0c;聊天记录、朋友圈、小程序订单都散在各个地方&#xff0c;到底怎么整理出一份能用的用户画像&#xff1f;"这不是他一个人的问题。我见过太多人把"用户画像"理解成"把微信通讯…

作者头像 李华
网站建设 2026/9/10 1:32:44

深度学习入门学完,我用梯度检查点在8GB显存上跑通了7B模型

深度学习入门学完,我用梯度检查点在8GB显存上跑通了7B模型 周一例会后,Leader 突然丢给我一个任务:把开源的 7B 大模型部署到内部知识库做文本分类。我看了眼机器配置--一台 RTX 3060 12GB 的工作站,心里直接凉了半截。但话已经说出去了,只能硬着头皮上。 当天下午我就用 Hugg…

作者头像 李华
网站建设 2026/9/10 1:31:48

libmodbus在Windows平台Qt5 MinGW中的编译测试与上位机集成

简介&#xff1a;Windows 平台 Qt5 MinGW 环境下的 libmodbus 集成测试包&#xff0c;面向需要在 Qt 界面程序中集成 Modbus 通信的嵌入式与工业软件开发人员&#xff0c;重点解决 MinGW 工具链下 libmodbus 的编译链接、基础功能调用和界面联动问题。包内共 21 个文件&#xf…

作者头像 李华