资讯动态

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

发布时间:2026/9/10 1:34:01 来源:尧图企业网站定制
JAX 模型接入 TensorFlow Serving 实战基于 jax2tf 导出 SavedModel 并部署 gRPC/REST 推理服务【免费下载链接】jaxComposable transformations of PythonNumPy 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 Servingtensorflow/servingDocker 镜像Google 内部版本的模型服务器说明位于 serving 示例目录的internal子目录中当前仓库不包含该部分代码。整个 serving 示例由两部分组成saved_model_main.py负责训练 MNIST 模型并导出 SavedModelmodel_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_datasetsTFDS用于下载 MNIST 数据集示例通过tfds.load(mnist)读取见 mnist_lib.pytensorflow_serving_api提供 gRPC 请求所需的 protobuf 类型predict_pb2、prediction_service_pb2_grpc使用tf_nightly是为了拿到足够新、能支持XlaCallModule的 TensorFlow 版本。示例代码对第三方依赖的最小声明可参考 jax2tf/examples/requirements.txttensorflow_datasets、tensorflow_hub、flax实际运行 serving 示例还需额外的grpcio、requests、absl-py与matplotlib等运行时依赖。2.2 安装 TensorFlow Serving Docker 镜像DOCKER_IMAGEtensorflow/serving:nightly docker pull ${DOCKER_IMAGE}这里同样选用 nightly 版本以获得对最新 XLA 算子的支持。三、设置变量在导出模型前先定义一批贯穿全流程的快捷变量# 快捷变量 # SavedModel 的保存位置 MODEL_PATH/tmp/jax2tf/saved_models # 示例模型可选 mnist_flax 与 mnist_pure_jax MODELmnist_flax # SavedModel 的 batch size。用 -1 表示 batch 多态任意批量 # 或用严格正数表示固定 batch size SERVING_BATCH_SIZE_SAVE-1 # 发送给模型的 batch size。若 SERVING_BATCH_SIZE_SAVE 不是 -1则二者必须相等 SERVING_BATCH_SIZE16 # 每次修改模型参数并重新导出后将该值加 1版本号递增 MODEL_VERSION$(( 1 ${MODEL_VERSION:-0} ))各变量含义如下变量作用取值说明MODEL_PATHSavedModel 的根目录默认/tmp/jax2tf/saved_modelsMODEL选择示例模型mnist_flaxFlax CNN或mnist_pure_jax纯 JAX MLPSERVING_BATCH_SIZE_SAVE导出模型时的 batch size-1表示 batch 多态正整数表示固定 batchSERVING_BATCH_SIZE客户端请求的 batch size固定 batch 导出时必须与上面相等MODEL_VERSIONServing 的模型版本号递增以保证模型服务器加载新版本MODEL_VERSION的递增机制与 TensorFlow Serving 的版本管理约定一致模型以模型名/版本号/的目录结构存放模型服务器会自动发现并加载更大版本号的新模型因此修改参数后只需重新导出并递增版本号无需重启服务器。四、训练并导出 SavedModel使用saved_model_main.py完成训练与导出该脚本的完整说明见 examples/README.mdpython ${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_flax或mnist_pure_jax--model_path/tmp/jax2tf/saved_modelsSavedModel 保存根目录--model_version1版本号lower_bound1更大版本在 serving 时优先--serving_batch_size1保存 serving 签名所用的 batch size-1表示 batch 多态校验器要求取值要么为-1要么为正整数--num_epochs3训练轮数lower_bound1--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 sizeserving 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.pyPureJaxMNIST纯 JAX 实现隐藏层尺寸[784, 512, 512, 10]使用jnp.dot tanh前向、jax.grad更新参数、jax.jit加速训练逻辑简单直观适合快速理解FlaxMNISTFlaxnn.Module实现的 CNNConv(32) → relu → avg_pool → Conv(64) → relu → avg_pool → Dense(256) → Dense(10) log_softmax用optax.sgd(learning_rate0.001, momentum0.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_gradientwith_gradient, polymorphic_shapes[None, polymorphic_shapes]) # 将参数包装为 tf.Variable使保存器把参数作为变量单独存储 param_vars tf.nest.map_structure( lambda param: tf.Variable(param, trainablewith_gradient), params) tf_graph tf.function(lambda inputs: tf_fn(param_vars, inputs), autographFalse, jit_compilecompile_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_gradientTrue时jax2tf 会用tf.custom_gradient注解降低后的函数在 TensorFlow 求导时惰性调用 JAX 的jax.vjp来计算梯度从而保证 TF 侧微分与 JAX 微分结果一致同时保存时需配合tf.saved_model.SaveOptions(experimental_custom_gradientsTrue)源码中已自动处理若 JAX 函数本身不可反向微分如使用lax.while_loop导出会报ValueError: Error when tracing gradients for SavedModel此时应传with_gradientFalse。五、检查导出的 SavedModelsaved_model_cli show --all --dir ${MODEL_PATH}/${MODEL}/${MODEL_VERSION}输出中如果某个签名的shape首维是-1说明这是batch 多态模型对任意批量都可用。saved_model_cli也是确认输入名示例中为inputs与输出名如output_0的依据——客户端代码正是从 dump 结果中获知这些名字的见 model_server_request.py 的注释。六、启动本地模型服务器开启 XLAdocker run -p 8500:8500 -p 8501:8501 \ --mount typebind,source${MODEL_PATH}/${MODEL}/,target/models/${MODEL} \ -e MODEL_NAME${MODEL} -t --rm --nameserving ${DOCKER_IMAGE} \ --xla_cpu_compilation_enabledtrue 各要素说明-p 8500:8500gRPC 服务端口-p 8501:8501HTTP REST 服务端口--mount typebind,...把模型目录挂载到容器内/models/${MODEL}-e MODEL_NAME${MODEL}告诉模型服务器要加载的模型名--xla_cpu_compilation_enabledtrue关键标志开启 CPU 侧的 XLA 编译jax2tf 生成的XlaCallModule算子依赖它才能在模型服务器中执行-t --rm --nameserving分配伪终端、退出即删除容器、指定容器名。模型服务器的版本管理特性意味着只要递增${MODEL_VERSION}并重新导出无需重启服务器运行中的模型服务器会自动加载更新的版本。七、发送推理请求gRPC 与 REST 双通道python ${JAX2TF_EXAMPLES}/serving/model_server_request.py --model_spec_name${MODEL} \ --use_grpc --prediction_service_addrlocalhost:8500 \ --serving_batch_size${SERVING_BATCH_SIZE} \ --count_images1287.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:8500REST 用localhost:8501--serving_batch_size1请求 batch sizelower_bound1必须与模型保存时的 batch 匹配且需整除--count_images--count_images16要测试的图片总数lower_bound17.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, dtypeimages.dtype, shapeimages.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 为 1612544 16 × 784而模型服务器加载的模型 batch size 为 1784 1 × 784。请检查导出时的--serving_batch_size与发请求时的--serving_batch_size是否一致或者改用 batch 多态方式导出-1。八、进阶实验切换模型与批量策略8.1 使用纯 JAX 模型MODELmnist_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_SAVE16 SERVING_BATCH_SIZE16然后重做导出步骤与发送请求步骤无需重启模型服务器。注意此时--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_enabledtrueGPU 场景对应 GPU XLA 标志开启 XLA 后XlaCallModule反序列化、编译并执行内嵌的 StableHLO请求客户端通过 gRPCPredictionService或 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_enabledtrue启动 Serving → gRPC/REST 双通道发请求并统计准确率同时覆盖了最常见的 batch 不匹配排错与多种批量策略实验。所有命令与参数均可直接对照仓库源码saved_model_main.py、saved_model_lib.py、model_server_request.py、mnist_lib.py逐行验证可作为把 JAX/Flax 模型投入生产化推理服务的最小可运行模板。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价