模型推理服务人工智能后端大模型MLOpsLLMOps【免费下载链接】BentoMLThe easiest way to serve AI apps and models - Build Model Inference APIs, Job queues, LLM apps, Multi-model pipelines, and more!项目地址https://gitcode.com/gh_mirrors/be/BentoML点击查看免费下载本篇技术指南围绕 BentoML 官方 API 参考文档docs/source/reference/bentoml/frameworks/flax.rst中定义的bentoml.flax.save_model、bentoml.flax.load_model与bentoml.flax.get三个核心接口展开结合仓库源码 flax.py 与集成测试 flax.py 逐一解析其参数语义、序列化机制与 Runner 推理链路。读者读完本文后将能够在自己的 JAX / Flax 训练脚本中直接保存flax.linen.Module模型、从模型仓库加载权重并接入 BentoML 推理服务同时理解其背后基于 msgpack 的 state dict 持久化与多设备CPU / GPU / TPU放置策略。一、接口总览BentoML 为 Flax 提供的最小 API 面bentoml.flax模块是 BentoML 针对 Flax基于 JAX 的神经网络框架提供的官方适配层。与 PyTorch、TensorFlow 等其他框架适配一样它遵循「保存save_model→ 加载load_model→ 查询get→ 运行get_runnable」的统一模型管理模型让 Flax 模型也能被 BentoML 的模型仓库、Bento 打包和推理服务统一管理。从源码 flax.py 顶部可以看到模块的核心常量MODULE_NAME bentoml.flax MODEL_FILENAME saved_model.msgpack API_VERSION v1 __all__ [load_model, save_model, get_runnable, get, JaxArrayContainer]其中MODEL_FILENAME为saved_model.msgpack表明 Flax 模型的 state dict参数与可变状态通过msgpack二进制格式落盘JaxArrayContainer则是 JAX 数组的 Runner 数据容器定义于 common/jax.py负责在服务请求与批量推理之间对jax.Array/jnp.ndarray进行序列化、拆包与合并。此外模块在导入时会强制校验两项依赖见 flax.pyflax至少需要flax.linen与flax.serializationtensorflow源码注释明确说明JAX 依赖 XLA而 XLA 是 TensorFlow 的一部分因此必须安装 TensorFlow 才能使用bentoml.flax。缺少任一依赖都会抛出MissingDependencyException。因此在使用本模块前请先执行pip install flax tensorflow jax jaxlib。二、保存 Flax 模型bentoml.flax.save_model 参数全解save_model负责将一个flax.linen.Module实例连同其训练好的 state dict 一起保存到 BentoML 本地模型仓库。其完整签名如下def save_model( name: Tag | str, module: nn.Module, state: dict[str, t.Any] | FrozenDict[str, t.Any] | struct.PyTreeNode, *, signatures: ModelSignaturesType | None None, labels: dict[str, str] | None None, custom_objects: dict[str, t.Any] | None None, external_modules: t.List[ModuleType] | None None, metadata: dict[str, t.Any] | None None, ) - bentoml.Model2.1 核心参数说明参数类型说明namestr或Tag模型在仓库中的名称必须是合法的bentoml.Tag名称例如mnist或mnist:v1moduleflax.linen.Module要保存的模块实例。源码会校验其类型若不是nn.Module将抛出BentoMLExceptionstatedict/FrozenDict/PyTreeNode训练得到的模型状态通常是net.init(...)返回的参数字典含params键也可以是 FlaxFrozenDict或任意struct.PyTreeNodesignaturesModelSignaturesType预测方法的签名定义默认值为{__call__: {batchable: False}}详见下文 2.2labelsdict[str, str]管理标签例如{training-set: data-1}便于在模型管理界面中筛选custom_objectsdict[str, t.Any]随模型一起保存的自定义对象例如归一化器、预处理函数当前使用 cloudpickle 序列化该实现后续可能变更external_moduleslist[ModuleType]用户自定义的附加 Python 模块随模型或自定义对象一并保存例如 tokenizer 模块、预处理模块、模型配置模块metadatadict[str, t.Any]与模型关联的元数据例如{bias: 4}要求为默认 Python 类型str/int等主要用于模型管理 UI 展示返回值为bentoml.Model对象可用于后续load_model、to_runnable或打包进 Bento。2.2 签名signatures的默认行为从源码可见当未显式传入signatures时模块会采用默认签名并打印日志提示if signatures is None: signatures {__call__: {batchable: False}} logger.info(Using the default model signature for Flax (%s) for model %s., signatures, name)也就是说Flax 模型默认把__call__作为推理入口且默认不开启批处理batchable: False。若希望启用 Runner 的自动批处理需要在保存时显式指定例如signatures{__call__: {batchable: True, batch_dim: 0}}集成测试 models/flax.py 中正是这样配置的。2.3 保存流程的内部实现从源码可以看到save_model的实际工作步骤类型校验module必须是nn.Module否则抛错记录框架上下文构造ModelContext记录flax、jax、jaxlib三个包的版本号通过get_pkg_version获取这些版本信息会写入模型元数据供环境还原时参考注入自定义对象将module本身存入custom_objects[_module]——这是后续load_model反序列化时重建模型结构的依据创建模型仓库条目调用bentoml.models._create(...)创建模型目录写入moduleMODULE_NAME、api_versionv1、签名、标签、选项PartialKwargsModelOptions支持partial_kwargs、自定义对象、外部模块与元数据序列化 state使用flax.serialization.to_bytes(state)将参数状态编码为 msgpack 字节流写入模型目录下的saved_model.msgpack文件。官方文档给出的保存示例训练 MNIST 分类器后import jax rng, init_rng jax.random.split(rng) state create_train_state(init_rng, config) for epoch in range(1, config.num_epochs 1): rng, input_rng jax.random.split(rng) state, train_loss, train_accuracy train_epoch( state, train_ds, config.batch_size, input_rng ) _, test_loss, test_accuracy apply_model( state, test_ds[image], test_ds[label] ) # 训练完成后保存模型 tag bentoml.flax.save_model(mnist, CNN(), state)三、加载 Flax 模型bentoml.flax.load_model 参数与设备语义load_model从本地模型仓库加载flax.linen.Module实例与其 state dict返回(module, state_dict)二元组。完整签名def load_model( bento_model: str | Tag | bentoml.Model, init: bool True, device: str | XlaBackend cpu, ) - tuple[nn.Module, dict[str, t.Any]]3.1 参数语义参数类型默认值说明bento_modelstr/Tag/bentoml.Model必填模型仓库中的标签或已获取的bentoml.Model实例initboolTrue是否初始化 state dict。为True时所有权重与值被转换为jnp.ndarrayjax.tree_util.tree_map(jnp.array, state_dict)为False时参数会通过jax.device_put直接放到指定加速设备上devicestr或XlaBackendcpustate dict 放置的目标设备仅在initFalse时生效源码中的设备放置逻辑if init: state_dict jax.tree_util.tree_map(jnp.array, state_dict) else: state_dict jax.tree_util.tree_map( lambda s: jax.device_put(s, jax.devices(device)[0]), state_dict )这里initTrue默认把参数统一转换为 JAX 数组便于直接参与计算而initFalse则保留参数在 GPU / TPU 显存中避免不必要的设备间拷贝。注意源码注释指出加载时会将 TensorFlow 的可见 GPU 设备清空tf.config.experimental.set_visible_devices([], GPU)以防止 TensorFlow 抢占 GPU 显存导致 JAX 无法分配。3.2 内部实现要点若传入的是标签而非bentoml.Model实例内部先调用get()完成解析校验模型的info.module是否为bentoml.flax或__name__否则抛出NotFound检查custom_objects中是否存在_module键若缺失说明模型损坏或不是由save_model保存的抛出BentoMLException读取saved_model.msgpack通过flax.serialization.from_bytes(module, f.read())反序列化得到 state dict若发生UnpicklingError、msgpack.exceptions.ExtraData或UnicodeDecodeError则统一包装为BentoMLException抛出按init参数完成数组类型转换或设备放置后返回(module, state_dict)。官方文档的加载示例import bentoml import jax net, state_dict bentoml.flax.load_model(mnist:latest) predict_fn jax.jit(lambda s: net.apply({params: state_dict[params]}, x)) results predict_fn(jnp.ones((1, 28, 28, 1)))该示例展示了典型用法加载后用jax.jit对前向计算进行编译加速并以net.apply({params: state_dict[params]}, x)的方式执行推理。四、查询模型bentoml.flax.getget用于按标签从模型仓库中获取bentoml.Model对象def get(tag_like: str | Tag) - bentoml.Model实现上它委托给bentoml.models.get(tag_like)并额外校验模型是否由本框架保存model bentoml.models.get(tag_like) if model.info.module not in (MODULE_NAME, __name__): raise NotFound( fModel {model.tag} was saved with module {model.info.module}, failed to load with {MODULE_NAME}. ) return model即get只能返回由bentoml.flax.save_model保存的模型若用其他框架如bentoml.pytorch保存的同名模型将抛出NotFound。官方示例import bentoml model bentoml.flax.get(mnist:latest)get返回的bentoml.Model对象可直接传给load_model也可通过to_runnable接入 Runner 推理服务。五、从保存到推理get_runnable 与 Runner 集成虽然get_runnable在源码中被标记为 Private API官方建议使用bentoml.Model.to_runnable但它揭示了 Flax 模型在 BentoML 服务端的完整推理链路值得深入理解。5.1 Runnable 的资源声明与初始化class FlaxRunnable(bentoml.legacy.Runnable): SUPPORTED_RESOURCES (tpu, nvidia.com/gpu, cpu) SUPPORTS_CPU_MULTI_THREADING True def __init__(self): super().__init__() self.device backend.get_backend().platform self.model, self.state_dict load_model(bento_model, deviceself.device) self.params self.state_dict[params] self.methods_cache: t.Dict[str, t.Callable[..., t.Any]] {}关键点声明支持TPU、NVIDIA GPU、CPU三类资源并支持 CPU 多线程初始化时通过backend.get_backend().platform探测当前 JAX 后端平台并将该平台作为device传入load_model此时init保持默认True参数在 CPU 上以jnp.ndarray形式存在调用时由 JAX 自行调度到后端设备params self.state_dict[params]提取可训练参数供model.apply使用用methods_cache缓存编译后的推理方法避免重复生成。5.2 推理方法的动态生成gen_run_method会根据模型上的方法名如__call__或自定义方法动态生成运行函数def mapping(item: jnp.ndarray | ext.NpNDArray | ext.PdDataFrame) - jnp.ndarray: if LazyTypeext.NpNDArray.isinstance(item): return jnp.asarray(item) if LazyTypeext.PdDataFrame.isinstance(item): return jnp.asarray(item.to_numpy()) return item def run_method(self, *args): params Paramsjnp.ndarray.map(mapping) arg params.args[0] if len(params.args) 1 else params.args return self.model.apply({params: self.params}, arg, methodmethod)这里体现了两个重要设计输入自动转换numpy.ndarray与pandas.DataFrame都会被自动转换为jnp.ndarray保证与 JAX 计算兼容不做 jit源码注释明确说明不能在此处jax.jit因为多线程场景下不应干预 JAX 的 tracing参考 JAX 官方并发文档从而避免线程安全问题。此外PartialKwargsModelOptions即partial_kwargs选项允许为特定推理方法预设部分参数通过functools.partial注入这在多签名模型或固定超参数推理时非常实用。5.3 与签名联动add_runnable_method遍历bento_model.info.signatures将每个签名中的batchable、batch_dim、input_spec、output_spec应用到FlaxRunnable.add_method(...)从而把保存时定义的签名自动注册为 Runner 的批量推理方法——这正解释了为什么save_model中signatures的batchable/batch_dim设置会直接决定服务端的批处理行为。六、JAX 数组的数据容器JaxArrayContainer当 Runner 开启批处理时JaxArrayContainer定义于 common/jax.py负责 JAX 数组的批量组装与拆解batches_to_batch沿batch_dim轴用jnp.concatenate合并多个子批次并返回切分索引itertools.accumulate累计各子批大小batch_to_batches用jnp.split按索引将合并后的批次切回多个子批次to_payload/from_payload将jax.Array转为numpy.ndarray后pickle.dumps生成 Payload反方向则用pickle.loads后jnp.asarray还原。模块通过DataContainerRegistry.register_container分别注册了jax.numpy.ndarray与jax.Array两种类型从而让 BentoML 服务端能透明处理 JAX 数组的请求序列化与响应反序列化。七、端到端实战从保存到 Runner 推理附仓库测试佐证仓库的集成测试 models/flax.py 提供了一个可直接借鉴的完整范例包含一个纯 Linen 实现的 MLP 模型import chex import jax.numpy as jnp from flax import linen as nn from jax import random from jax.nn import initializers import bentoml class Dense(nn.Module): features: int nn.compact def __call__(self, x: jnp.ndarray) - jnp.ndarray: kernel self.param( kernel, initializers.lecun_normal(), (x.shape[-1], self.features) ) y: jnp.ndarray jnp.dot(x, kernel) return y class MLP(nn.Module): nn.compact def __call__(self, x: jnp.ndarray) - jnp.ndarray: x Dense(3)(x) x Dense(3)(x) return x def init_mlp_state(): net MLP() params net.init({params: random.PRNGKey(42)}, jnp.ones((10,))) return params保存与签名配置mlp FrameworkTestModel( namemlp, modelMLP(), save_kwargs{ state: init_mlp_state(), signatures: {__call__: {batchable: True, batch_dim: 0}}, }, ... )注意这里保存时显式启用了__call__的批处理batchableTrue, batch_dim0对应真实场景可写为bentoml.flax.save_model(mlp, MLP(), init_mlp_state(), signatures{__call__: {batchable: True, batch_dim: 0}})加载后即可校验前向输出形状一致测试中通过chex.assert_equal_shape断言net, state_dict bentoml.flax.load_model(mlp:latest) logits net.apply({params: state_dict[params]}, jnp.ones((10,)))在tests/integration/frameworks/test_frameworks.py中test_get、test_load、test_runnable、test_runner_batching、test_runner_nvidia_gpu、test_service等用例会依次验证get查询、load_model加载、to_runnable构建 Runner、GPU 上运行、以及最终组装成推理服务bentoml serve的完整流程。若需在本地复现可参考 test_frameworks.py 中的测试组织方式用pytest tests/integration/frameworks -k flax定向运行 Flax 相关用例。八、使用建议与注意事项依赖完整性bentoml.flax同时依赖flax、jax、jaxlib与tensorflowXLA 运行时打包为 Bento 或部署到云端时务必在requirements.txt/bentofile.yaml中声明全部依赖默认签名不可批处理若不显式指定signatures服务端默认batchableFalse需要吞吐优化时请像测试示例那样显式声明batchableTrue与batch_dim0设备放置策略默认load_model(initTrue)将参数转为jnp.ndarray并留在 CPU由 JAX 运行时调度若追求极致性能且设备内存充足可用initFalse配合device参数把参数直接放入 GPU / TPU模型一致性校验get与load_model都会校验模型的info.module是否为bentoml.flax跨框架复用同名模型会得到NotFound这是有意为之的类型安全设计状态文件格式模型落盘文件为saved_model.msgpackmsgpack 序列化custom_objects中的_module记录了模块结构二者共同构成可完整还原的模型工件。九、关联文档与源码索引API 参考文档本文依据docs/source/reference/bentoml/frameworks/flax.rst核心实现src/bentoml/_internal/frameworks/flax.pyJAX 数据容器src/bentoml/_internal/frameworks/common/jax.py集成测试模型定义tests/integration/frameworks/models/flax.py通用框架集成测试tests/integration/frameworks/test_frameworks.py框架相关参数选项定义src/bentoml/_internal/models/model.py赞分享模型推理服务人工智能后端大模型MLOpsLLMOps【免费下载链接】BentoMLThe easiest way to serve AI apps and models - Build Model Inference APIs, Job queues, LLM apps, Multi-model pipelines, and more!项目地址https://gitcode.com/gh_mirrors/be/BentoML点击查看免费下载相关推荐视频修复神器untrunc5分钟拯救损坏的MP4文件终极指南视频修复神器untrunc5分钟拯救损坏的MP4文件终极指南 你是否曾因视频文件突然损坏而痛心疾首当珍贵的家庭录像、重要的工作记录或专业的拍摄素材因传输中断模型推理服务人工智能后端大模型MLOpsLLMOpsBentoML Diffusers API 完整指南import_model / save_model / load_model / get 实战详解BentoML Diffusers API 完整指南import_model / save_model / load_model / get 实战详解 导读模型推理服务人工智能后端大模型MLOpsLLMOpsBentoML ONNX 框架 API 参考save_model / load_model / get 完整指南BentoML ONNX 框架 API 参考save_model / load_model / get 完整指南 本篇以 BentoML 官方 API 参考文模型推理服务人工智能后端大模型MLOpsLLMOps上一篇MolmoAct2-LIBERO-LeRobot在仿真环境中的评估lerobot-eval命令参数配置与结果解读下一篇Ethermint治理机制终极指南如何通过提案管理EVM参数创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
企业数字化 ERP 产品动态
相关推荐
Fission 仓库开发与维护实战指南:构建、测试、代码生成与架构速查 云原生后端 【免费下载链接】fission Fast and Simple Serverless Functions for Kubernetes 项目地址: https://gitcode.com/gh_mirrors/fi/fission 点击查看 免费下载 本篇技术指南以 Fission(Kubernetes 原生 Serverless 框架,Go 编写&am… · 2026/9/25 2:14:29
PaddleNLP 文本信息抽取应用实战:基于 UIE 微调的数据标注、模型训练与封闭域蒸馏全流程指南 人工智能大模型预训练微调LoRARLHF强化学习分布式训练 【免费下载链接】PaddleNLP Easy-to-use and powerful LLM and SLM library with awesome model zoo. 项目地址: https://gitcode.com/gh_mirrors/pa/PaddleNLP 点击查看 免费下载 本文以 PaddleNLP 信息抽取应… · 2026/9/25 2:14:28
xberg 批处理提取 API 实战:基于 C FFI 的 extract_batch 字节批量抽取 后端AI 应用NLP 【免费下载链接】xberg Polyglot document intelligence with a Rust core: extract text, metadata, images, tables, and structured data from 106 formats across 140 file extensions, plus code intelligence for 371 languages. Fifteen bindings, with … · 2026/9/25 2:44:01
Dart SDK 中 Observatory 开发者工具实战指南:激活、Web 服务与 DDC 调试开发 编程语言编译器语言运行时标准库开发工具 【免费下载链接】sdk The Dart SDK, including the VM, JS and Wasm compilers, analysis, core libraries, and more. 项目地址: https://gitcode.com/gh_mirrors/sdk1/sdk 点击查看 免费下载 Observatory 是 Dart VM 团队… · 2026/9/25 2:44:01
ESP32-C5-WROOM-1U 双频 Wi-Fi 6 模组硬件设计与固件配置实战 /* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/25 2:44:01
jc date 解析器实战指南:将 date 命令输出转换为结构化 JSON 时间数据 开发工具 【免费下载链接】jc CLI tool and python library that converts the output of popular command-line tools, file-types, and common strings to JSON, YAML, or Dictionaries. This allows piping of output to tools like jq and simplifying automation scripts.… · 2026/9/25 2:43:55
创维E900V22D刷机全攻略:S905L3SB芯片兼容性解析与救砖实战 /* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/25 1:00:31
MQTT协议原理与Broker服务器搭建实战:从Mosquitto到EMQX /* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/25 1:00:37