- 深度学习
- 人工智能
- 机器学习
- 分布式训练
【免费下载链接】deeplearning4j
Suite of tools for deploying and training deep learning models using the JVM. Highlights include model import for keras, tensorflow, and onnx/pytorch, a modular and tiny c++ library for running math code and a java based math library on top of the core c++ library. Also includes samediff: a pytorch/tensorflow like library for running deep learn...
Omnihub 是 DeepLearning4J 仓库中contrib/omnihub目录下的一个 Python SDK,用于简化预训练模型的下载与格式转换:它屏蔽了 Keras、TensorFlow、ONNX、PyTorch、HuggingFace 五大模型 Zoo 的差异,将"下载权重 → 组织本地文件 → 冻结/导出为可部署模型"这些常见操作收敛为一套统一 API。读完本文,你将掌握 Omnihub 的安装方法、各框架 Hub 的调用方式、模型落地目录的存储规则,以及它背后的冻结(freezing)与微调(finetuning)设计动机。
背景:为什么需要一个跨框架的模型 Hub SDK
在 JVM 生态(DeepLearning4J、ND4J、SameDiff)中复用 TensorFlow 或 PyTorch 产出的模型文件时,通常会遇到两类高频工作流,而它们"都带有相当程度的摩擦,往往需要一次性教程和复制粘贴祈祷它能跑通"(contrib/omnihub/README.md):
- 微调(Finetuning):通常包括两步——① 解冻模型(把常量转换为变量,unfreezing);② 定制模型(在末尾添加新的目标函数和其他层)。不同框架做这两步的复杂度差异很大。
- 部署(Deployable):通常包括两步——① 冻结模型(把可训练参数转换为冻结常量,freezing);② 优化模型(量化、改变数据类型、删除多余算子以减小体积等)。
Omnihub 的目标正是封装每个框架的常见步骤(如冻结/解冻、模型下载),让这些工作流不必为每个模型、每个框架各写一遍胶水代码。它的实现思路是:每个框架都有自己的 "model hub",这个 hub 知道如何与该框架的模型 Zoo 交互,并对模型做预处理。
架构:一个基类 + 五个框架 Hub
Omnihub 的代码组织非常清晰,位于 contrib/omnihub/src/omnihub:
src/omnihub/ ├── model_hub.py # ModelHub 基类:下载、落盘、流式暂存 └── frameworks/ ├── keras.py # KerasModelHub:基于 tf.keras.applications ├── tensorflow.py # TensorflowModelHub:tfhub.dev + 冻结成 .pb ├── onnx.py # OnnxModelHub:onnx/model zoo 直连下载 ├── pytorch.py # PytorchModelHub:torchvision 导出 ONNX └── huggingface.py # HuggingFaceModelHub:transformers 多框架导出ModelHub 基类与存储目录规则
所有 Hub 都继承自 model_hub.py 中的ModelHub,它定义了三条核心约定:
- 存储根目录:优先读取环境变量
OMNIHUB_HOME,未设置时默认~/.omnihub(注意:README 示例中写的$HOME/.model_hub是早期文档描述,当前源码实际落盘到~/.omnihub或$OMNIHUB_HOME)。 - 按框架分子目录:每个 Hub 构造时传入
framework_name,模型会被组织到<root>/<framework_name>/下,例如~/.omnihub/keras/、~/.omnihub/onnx/。 - 两个核心方法:
download_model(model_path, **kwargs):按base_url/model_path拼接 URL 并流式下载,使用requests.get(..., stream=True)按 8192 字节分块写盘(见 model_hub.py)。stage_model(model_path, model_name):把下载好的模型复制到目标暂存目录(model_hub.py);另有stage_model_stream支持从文件流直接写盘。
基类还提供了默认的base_url拼接下载实现(f'{self.base_url}/{model_path}'),因此像 ONNX 这种纯直连下载的 Hub,几乎无需覆写任何逻辑(见 onnx.py)。
安装与快速上手
安装依赖并注册为 Python 包(requirements.txt、setup.py):
pip install -r requirements.txt python setup.py install依赖清单涵盖keras-applications、huggingface-hub、tensorflow-hub、onnx、requests、pytest、torch-model-archiver、tensorflow、transformers、scikit-learn等;setup.py通过package_dir={"": "src"}+find_packages(where="src")从src目录收集包,python_requires=">=3.6"。
一个最简的下载 + 暂存示例(对 README 中的片段做了 import 修正,KerasModelHub需从omnihub.frameworks.keras导入):
from omnihub.frameworks.keras import KerasModelHub keras_model_hub = KerasModelHub() model_path = keras_model_hub.download_model('vgg19/vgg19_weights_tf_dim_ordering_tf_kernels_notop.h5') keras_model_hub.stage_model(model_path, 'vgg19_weights_tf_dim_ordering_tf_kernels_notop.h5')执行后模型会被放到:
$HOME/.omnihub/keras/vgg19_weights_tf_dim_ordering_tf_kernels_notop.h5若设置了export OMNIHUB_HOME=/your/path,则目录变为$OMNIHUB_HOME/keras/...。
各框架 Hub 使用详解
KerasModelHub:基于 tf.keras.applications 的模型工厂
keras.py 是覆盖面最广的 Hub,BASE_URL指向 Google 的 keras-applications 存储。它的下载不是简单拉文件,而是解析路径后调用tf.keras.applications构造模型并保存权重:
- 路径即参数:
model_path被split('/')成两段,第一段是模型名(如vgg19),第二段是权重文件名。当文件名包含notop时,自动设置include_top=False(去掉顶部分类层,这是做迁移学习/微调的常见前提),否则include_top=True。 - 支持的模型族:
vgg16/vgg19、resnet50/101/152及其 v2 变体、densenet121/169/201、inceptionresnetv2、efficientnetb0~b7、mobilenet、mobilenetv2、inceptionv3、nasnet、nasnet_mobile、xception。 - 落盘位置:权重保存在
~/.keras/models/<weights_file>,随后通过ret.save(...)写出。代码中mobilenetv3分支被注释掉并标注了原因——MobileNetV3()缺少stack_fn和last_point_ch两个必填位置参数,这是从源码中可以观察到的已知限制(keras.py)。
结合 omnihub_bootstrap.py 中的批量清单,可以看到 Keras 的完整可用路径约定,例如resnet50/top、resnet101/notop、densenet201/top、efficientnetb0等。
TensorflowModelHub:tfhub.dev 下载 + 冻结为 .pb
tensorflow.py 演示了"下载 → 解压 → 冻结 → 写图"的完整链路:
- 以
https://tfhub.dev/<model_path>?tf-hub-format=compressed下载压缩包(tfhub 的 compressed 格式); - 校验
tarfile.is_tarfile,解压到临时目录; - 调用模块级函数
convert_saved_model(saved_model_dir):tf.saved_model.load加载 SavedModel,取signatures['serving_default'],再用convert_variables_to_constants_v2把变量冻结为常量,得到GraphDef; - 通过
tf.io.write_graph(..., as_text=False)以二进制.pb写入~/.omnihub/tensorflow/<name>.pb,最后删除临时 tar 包。
这正是 README 中"把可训练参数转换为冻结常量(freezing)"在代码层面的落地。测试用例中使用的路径是emilutz/vgg19-block4-conv2-unpooling-decoder/1(test_frameworks.py),实际可用模型以 tfhub.dev 上?tf-hub-format=compressed可下载的条目为准。
OnnxModelHub:最薄的直连下载
onnx.py 仅设置BASE_URL = 'https://media.githubusercontent.com/media/onnx/models/master',下载完全复用基类逻辑。典型路径如:
onnx_model_hub = OnnxModelHub() onnx_model_hub.download_model('vision/body_analysis/age_gender/models/age_googlenet.onnx')omnihub_bootstrap.py 中还罗列了更多 ONNX Model Zoo 路径:age_googlenet、gender_googlenet、arcfaceresnet100-8.onnx、emotion-ferplus-*.onnx、version-RFB-320/640.onnx、bvlcalexnet-12(-int8).onnx、caffenet-12(-int8).onnx、efficientnet-lite4-11.onnx等,覆盖年龄/性别、人脸、情感、检测、分类等任务。
PytorchModelHub:torchvision 权重一键导出 ONNX
pytorch.py 展示了"从训练框架产出可部署文件"的典型做法:
- 输入尺寸表:源码内置了两组默认尺寸——
MODEL_224_DEFAULTS(resnet18、vgg16、shufflenet_v2_x1_0、resnext50_32x4d、wide_resnet50_2、mnasnet1_0)为 224×224;MODEL_256_DEFAULTS(alexnet、squeezenet1_0、densenet161、googlenet、inception_v3以及动态生成的efficientnet_b0~b7、regnet_x/y_*)为 256×256;另有特例mobilenet_v2(32×32)、mobilenet_v3_large/small(320×320)、retinanet(512×512)。未在表中的模型名会导致KeyError。 - 导出流程:构造全 1 的
(1, 3, height, width)伪输入,检测模型(fasterrcnn、ssd、retinanet、maskrcnn、keypointrcnn)走models.detection[...],其余走models.__dict__[model_path],加载pretrained=True权重后用torch.onnx.export导出为~/.omnihub/pytorch/<name>.onnx,关键参数export_params=True、do_constant_folding=False、opset_version=13。
HuggingFaceModelHub:一个仓库多框架导出
huggingface.py 针对 HF 仓库"多框架共存"的特性,要求调用方必须通过framework_name指定目标框架,且download_model内部有assert 'framework_name' in kwargs强制校验:
- TensorFlow/Keras 路径:
TFAutoModel.from_pretrained加载后,用tf.function包装output_model.call,取concrete_function并convert_variables_to_constants_v2冻结,写出~/.omnihub/tensorflow/<name>.pb。 - PyTorch/ONNX 路径:
AutoModel.from_pretrained加载,从dummy_inputs中按main_input_name排序构造输入(主输入在前,其余辅助输入在后),再用torch.onnx.export(opset_version=13)导出到~/.omnihub/<framework_name>/<name>.onnx;也可通过download_function参数注入自定义下载/导出函数。 - URL 解析规则:类文档注释说明了 HF 使用 git LFS + 分支的下载约定——URL 公式为
https://huggingface.co + repo名 + resolve/<branch>/<file>,默认分支main(huggingface.py 的resolve_url即生成该路径)。
omnihub_bootstrap.py 中对gpt2、bert-base-uncased、t5-base、bert-base-chinese、google/electra-small-discriminator、facebook/wav2vec2-base-960h、facebook/bart-large-cnn分别以tensorflow和pytorch两个框架执行导出,是研究多框架导出的现成样例。
冻结与部署工作流在代码中的印证
README 描述的"冻结模型(freezing)"与"解冻模型(unfreezing)"并非空话,而是直接对应源码中的两处convert_variables_to_constants_v2调用:
| 场景 | 代码位置 | 输入 | 输出 |
|---|---|---|---|
| TF Hub SavedModel 冻结 | tensorflow.py | serving_default签名 | 二进制.pbGraphDef |
| HF TensorFlow 模型冻结 | huggingface.py | tf.function的 concrete function | 二进制.pbGraphDef |
PyTorch 一侧则通过torch.onnx.export(export_params=True)把训练参数固化进 ONNX 图,配合do_constant_folding=False保留原始计算结构,二者共同构成"把框架模型变成可部署独立文件"的能力。而 Keras 的include_top开关(notop后缀)正是"定制模型/迁移学习"的入口,呼应了 README 中微调工作流的第一步。
测试与批量引导脚本
- 单元测试:test_frameworks.py 覆盖五个 Hub 的真实下载,并用
assert os.path.exists(...)验证落盘位置(如~/.omnihub/keras/vgg19_weights_tf_dim_ordering_tf_kernels_notop.h5、~/.omnihub/onnx/age_googlenet.onnx、~/.omnihub/pytorch/resnet18.onnx)。运行方式:
cd contrib/omnihub && pytest- 批量引导:omnihub_bootstrap.py 是一份"一次性拉全常用模型"的脚本,内含 Keras 的 top/notop 双版本清单(resnet、densenet、inception、mobilenet、nasnet、xception、efficientnet 系列)、16 个 ONNX Model Zoo 模型、TF Hub 的 vgg19 解码器、PyTorch 的 resnet18,以及 HF 的 bart-large-cnn 双框架导出;其中部分 Keras 项(如 vgg16 top、mobilenetv3)因下载卡在最后几字节或构造参数缺失被注释,属于源码中记录的真实已知问题。
与 DeepLearning4J 生态的衔接
Omnihub 位于 contrib/omnihub(contrib 目录,非核心构建产物),其价值在于为 DeepLearning4J / ND4J / SameDiff 的模型导入管线提供统一、可复现的预训练模型获取途径:ONNX 与冻结后的.pb文件可直接对接 SameDiff 的模型导入(如 nd4j/samediff-import 下的 onnx/tensorflow 导入器),从而把"框架模型 → 独立文件 → JVM 图"的链路标准化。
小结
Omnihub 用约 200 行核心代码,把五大框架模型 Zoo 的下载、暂存、冻结与导出统一到了ModelHub基类的download_model/stage_model接口之下,并以"每个框架一个 Hub"的插件式设计隔离了各框架的差异。从 README 的设计动机(freezing/unfreezing、部署/微调两条工作流)到源码中的convert_variables_to_constants_v2与torch.onnx.export实现,它完整展示了一个轻量级跨框架模型仓库 SDK 应有的样子:简单、可扩展,且为上层 JVM 深度学习工具链的模型复用提供了坚实的文件基础。
- 深度学习
- 人工智能
- 机器学习
- 分布式训练
【免费下载链接】deeplearning4j
Suite of tools for deploying and training deep learning models using the JVM. Highlights include model import for keras, tensorflow, and onnx/pytorch, a modular and tiny c++ library for running math code and a java based math library on top of the core c++ library. Also includes samediff: a pytorch/tensorflow like library for running deep learn...
相关推荐
PaddleOCR ONNX转换:跨框架模型部署
PaddleOCR ONNX转换:跨框架模型部署 还在为OCR模型在不同框架间的部署兼容性而烦恼吗?PaddleOCR的ONNX转换功能让你轻松实现跨框架模型部
人工智能计算机视觉OCR深度学习大模型RAG【免费下载】 ONNXMLTools:一站式模型转换工具,助力AI模型跨平台部署
ONNXMLTools:一站式模型转换工具,助力AI模型跨平台部署 项目介绍 ONNXMLTools 是一个强大的开源工具,旨在将来自不同机器学习工具包的模型转
xcit_tiny_12_p16_384.fb_dist_in1k部署指南:轻量级Transformer模型的工业级应用
xcit_tiny_12_p16_384.fb_dist_in1k部署指南:轻量级Transformer模型的工业级应用 本文将为您提供一份完整的xcit_ti
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考