ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

TensorFlow Hub 入门指南:在 TensorFlow 中一键复用预训练模型

TensorFlow Hub 入门指南:在 TensorFlow 中一键复用预训练模型 文档开发工具教程【免费下载链接】docsTensorFlow documentation项目地址https://gitcode.com/gh_mirrors/doc/docs点击查看免费下载TensorFlow Hub 是 TensorFlow 生态中面向可复用机器学习的开放模型仓库与配套 Python 库它汇聚了文本嵌入、图像分类、TF.js 与 TFLite 等大量预训练模型并允许社区贡献者发布自己的模型。本文将围绕本仓库 overview.md 的核心内容带你从安装tensorflow_hub库开始掌握hub.KerasLayer、hub.load()的用法理解 SavedModel 的加载、微调、缓存与托管机制最终能够在一行代码内把成熟模型接入自己的 TensorFlow 程序并具备向 TensorFlow Hub 发布模型的完整知识。什么是 TensorFlow HubTensorFlow Hub 由两部分构成开放模型仓库托管了大量可直接使用的预训练模型覆盖文本嵌入如 NNLM、Universal Sentence Encoder、图像分类、TF.js 与 TFLite 模型等多种形态且面向社区贡献者开放任何人都可以发布经过审核的模型。tensorflow_hubPython 库提供下载与复用模型的 API让在程序中接入一个预训练模型这件事的代码量降到最低。其核心设计理念是同一 URL 既可以写在代码里加载模型又可以在浏览器中查看模型的文档页面模型句柄handle跨系统可移植、便于查阅。该理念在 hosting.md 中被总结为托管协议的关键特性。最小的加载示例来自 overview.mdimport tensorflow_hub as hub model hub.KerasLayer(https://tfhub.dev/google/nnlm-en-dim128/2) embeddings model([The rain in Spain., falls, mainly, In the plain!]) print(embeddings.shape) # (4,128)这里hub.KerasLayer会先下载并缓存模型再返回一个可以直接调用的 Keras Layer输入 4 条英文文本输出形状为(4, 128)的嵌入向量。这就是 TensorFlow Hub以最少代码复用预训练模型的直观体现。安装 tensorflow_hub安装说明详见 installation.md。tensorflow_hub可同时兼容 TensorFlow 1 与 TensorFlow 2官方建议新用户直接使用 TensorFlow 2。配合 TensorFlow 2 使用推荐$ pip install tensorflow2.0.0 $ pip install --upgrade tensorflow-hub注意tensorflow-hub版本必须是 0.5.0 或更新。TensorFlow Hub 的 TF1 风格 API 仍可在 TensorFlow 2 的 v1 兼容模式下工作。配合 TensorFlow 1旧版截至tensorflow_hub0.11.0 版本TensorFlow 1.15 是 1.x 中唯一仍受支持的版本它默认表现为 TF1 兼容行为但内部包含大量 TF2 特性因此可以部分使用 TensorFlow Hub 的 TF2 风格 API$ pip install tensorflow1.15,2.0 $ pip install --upgrade tensorflow-hub使用预发布版本tf-nightly与tf-hub-nightly由源码自动构建、未经发布测试适合希望抢先体验最新代码的开发者无需从源码自行构建$ pip install tf-nightly $ pip install --upgrade tf-hub-nightly理解模型句柄Model Handle在 tf2_saved_model.md 中模型句柄被定义为 SavedModel 的加载来源可以是文件系统路径本地已解压的 SavedModel 目录合法的 TensorFlow Hub 模型 URL如https://tfhub.dev/...Kaggle Models URL——它镜像 TensorFlow Hub 句柄语义上与对应的 TensorFlow Hub 句柄等价。在极少数需要拿到模型实际文件系统位置下载并解压之后或把句柄解析为路径之后的场景可以调用hub.resolve(handle)获取。默认情况下代码中应始终使用规范的 TensorFlow Hub URL因为它跨系统可移植且便于文档导航。在 TensorFlow 2 中复用 SavedModelTensorFlow 2 的 SavedModel 格式是官方推荐的预训练模型共享方式它取代了旧的 TF1 Hub 格式并带来了一套新 API底层函数hub.load()与它的 Keras 包装类hub.KerasLayer。相关完整指南见 tf2_saved_model.md。在 Keras 中使用hub.KerasLayerhub.KerasLayer以模型 URL或文件系统路径初始化随后提供 SavedModel 的计算逻辑与预训练权重。它通常与其它tf.keras.layers组合来构建 Keras 模型import tensorflow as tf import tensorflow_hub as hub hub_url https://tfhub.dev/google/nnlm-en-dim128/2 embed hub.KerasLayer(hub_url) embeddings embed([A long sentence., single-word, http://example.com]) print(embeddings.shape, embeddings.dtype)在此基础上可以按 Keras 常规方式搭建一个文本分类器model tf.keras.Sequential([ embed, tf.keras.layers.Dense(16, activationrelu), tf.keras.layers.Dense(1, activationsigmoid), ])完整训练与评估示例可参考 tf2_text_classification.ipynb。需要注意hub.KerasLayer中的模型权重默认是不可训练的非 trainable同一层对象在 Keras 中多次应用时权重共享要启用微调需显式设置trainableTrue。在 Estimator 中使用使用 TensorFlow Estimator API 做分布式训练的用户可以在自己的model_fn中把hub.KerasLayer与其它tf.keras.layers组合使用从而在分布式训练含参数服务器中加载 TensorFlow Hub 的 SavedModel。底层 APIhub.load()hub.load(handle)会下载并解压 SavedModel若句柄不是本地路径然后调用 TensorFlow 内置的tf.saved_model.load()完成加载因此它可以加载任意合法的 SavedModel这是它相对 TF1 时代hub.Module的关键差异。加载后得到的obj有多种调用方式详见 TensorFlow SavedModel 指南服务签名serving signatures以签名名 - 具体函数的字典形式存在可调用tensors_out obj.signaturesserving_default输入输出均为按名字索引的张量字典且受签名对形状和 dtype 的约束tf.function装饰的方法恢复为可调用的 tf.function 对象支持保存前已追踪过的各种张量/非张量参数组合若存在合适的obj.__call__追踪obj本身可像 Python 函数一样调用如output_tensor obj(input_tensor, trainingFalse)。加载后SavedModel 中可训练的变量仍以可训练状态恢复tf.GradientTape默认会追踪它们微调前可查看obj.trainable_variables是否建议只重训其中一部分。在低层代码中加载 TF1 Hub 格式模型迁移指南 migration_tf2.md 给出了新旧 API 的对应关系。旧写法# 已弃用TensorFlow 1 m hub.Module(handle, tags{foo, bar}) tensors_out_dict m(dict(x1..., x2...), signaturesig, as_dictTrue)推荐改写为# TensorFlow 2 m hub.load(path, tags{foo, bar}) tensors_out_dict m.signaturessig这里m.signatures是按键签名名索引的 TensorFlow 具体函数concrete functions字典调用时会计算其全部输出与 TF1 图模式的惰性求值不同。从tensorflow_hub0.7 起TF1 Hub 格式的旧模型也可配合hub.KerasLayer使用KerasLayer还额外暴露tags、signature、output_key、signature_outputs_as_dict等参数以满足旧模型与旧 SavedModel 的特定用法。创建并发布可复用的 SavedModel从 Keras 模型导出从 TensorFlow 2 开始tf.keras.Model.save()与tf.keras.models.save_model()默认导出 SavedModel 格式而非 HDF5导出的结果可直接被hub.load()、hub.KerasLayer等适配器使用。分享完整模型时用include_optimizerFalse保存即可分享模型中的一部分时应先把这部分做成独立的tf.keras.Model再保存。可以从一开始就把代码组织成待分享片段 完整模型piece_to_share tf.keras.Model(...) full_model tf.keras.Sequential([piece_to_share, ...]) full_model.fit(...) piece_to_share.save(...)也可以在事后从完整模型中切出片段要求切分点与模型分层对齐full_model tf.keras.Model(...) sharing_input full_model.get_layer(...).get_output_at(0) sharing_output full_model.get_layer(...).get_output_at(0) piece_to_share tf.keras.Model(sharing_input, sharing_output) piece_to_share.save(..., include_optimizerFalse)从低层 TensorFlow 导出Reusable SavedModel 接口如果你的模型不止提供服务签名而希望被嵌入更大的模型甚至被微调官方强烈建议实现 reusable_saved_models.md 中定义的Reusable SavedModel 接口。它要求加载后的obj具备以下属性__call__必选实现模型前向计算的 tf.functionvariables列出__call__所有可能调用中用到的全部 tf.Variable含可训练与不可训练为空时可省略trainable_variablesvariables中可训练变量的子集即微调时要更新的变量模型创建者可刻意省略某些原本可训练的变量以声明它们不应在微调中被修改为空尤其不支持微调时可省略regularization_losses零输入、返回单个标量 float 张量的 tf.function 列表用于表达权重正则化项微调时建议原样并入总损失典型如权重正则器由于无输入无法表达活动正则器。__call__的约定接受一个位置参数可以是单个张量、张量列表或按名字索引的张量字典形状与 dtype 由创建者定义批大小维度通常应留空可选关键字参数trainingPython 布尔值默认False用于区分如 dropout、批归一化等训练/推理行为还可定义更多带默认值的 kwargs张量型用于定制数值超参如 dropout 率Python 值型用于在追踪函数内做离散选择。输出同样可以是单个张量、张量列表或张量字典。该接口只使用 TensorFlow 2 原生原语不依赖 Keras、Sonnet 等任何模型构建库因此可以跨模型库复用。一个最小示例来自 tf2_saved_model.mdclass MyMulModel(tf.train.Checkpoint): def __init__(self, v_init): super().__init__() self.v tf.Variable(v_init) self.variables [self.v] self.trainable_variables [self.v] self.regularization_losses [ tf.function(input_signature[])(lambda: 0.001 * self.v**2), ] tf.function(input_signature[tf.TensorSpec(shapeNone, dtypetf.float32)]) def __call__(self, inputs): return tf.multiply(inputs, self.v) tf.saved_model.save(MyMulModel(2.0), /tmp/my_mul) layer hub.KerasLayer(/tmp/my_mul) print(layer([10., 20.])) # [20., 40.] layer.trainable True print(layer.trainable_weights) # [2.] print(layer.losses) # 0.004任务相关的更具体约定如文本任务、图像任务的输入输出规范由 common_saved_model_apis 定义以让同类模型易于互换。发布流程发布模型的入口见 publish.mdTensorFlow Hub 的模型发布目前通过 Kaggle Models 的早期接入计划EAP进行需要向官方提交 Kaggle 用户名、期望的组织 slug 与方形头像 URL 以开通权限之后按官方文档创建并发布模型。发布后请注意模型版本文档元数据可以更新但版本的资产模型文件不可变需要修改模型内容时必须发布新版本并建议在文档中附上版本变更日志见 common_issues.md。微调Fine-Tuning微调指对导入的 SavedModel 中已训练过的变量与外围模型一起继续训练。它可能提升模型质量但通常会让训练更苛刻耗时更长、更依赖优化器及其超参、过拟合风险上升CNN 场景往往还要求数据增强。因此官方建议消费者先建立良好的训练机制且仅在模型发布者推荐微调时才进行详见 tf2_saved_model.md。微调改变的是连续的模型参数不会改变硬编码的变换如文本分词、token 到嵌入矩阵条目的映射。消费者视角创建可微调的层只需传入trainableTruelayer hub.KerasLayer(..., trainableTrue)这会将该 SavedModel 声明的可训练权重与权重正则器加入 Keras 模型并让 SavedModel 的计算以训练模式运行考虑 dropout 等。端到端示例见 tf2_image_retraining.ipynb可选微调。高级用户可以把微调结果重新导出为 SavedModel 以替代原模型loaded_obj hub.load(https://tfhub.dev/...) hub_layer hub.KerasLayer(loaded_obj, trainableTrue, ...) model keras.Sequential([..., hub_layer, ...]) model.compile(...) model.fit(...) export_module_dir os.path.join(os.getcwd(), finetuned_model_export) tf.saved_model.save(loaded_obj, export_module_dir)创建者视角模型创建者应提前想好消费者如何微调并在文档中给出指引从 Keras 保存通常能自动满足微调的全部机制保存权重正则化损失、声明可训练变量、为trainingTrue/False分别追踪__call__等选择利于梯度流动的模型接口例如输出 logits 而非 softmax 概率或 top-k 预测若模型使用 dropout、批归一化等涉及超参的技术应设置为在多种目标任务和批大小下都合理的取值单层上的权重正则器会随模型保存含正则强度系数但优化器内部的正则如tf.keras.optimizers.Ftrl.l1_regularization_strength...不会被保存需在文档中告知消费者。下载缓存机制tensorflow_hub支持两种下载模式详见 caching.md压缩下载并本地缓存默认把模型作为压缩归档下载后缓存在磁盘上适合大多数环境除非磁盘空间紧张而网络带宽与延迟极佳从远程存储直接读取把模型直接从远程存储GCS读入 TensorFlow不需要缓存目录适合磁盘小、网络快的环境。无论哪种方式Python 代码中仍应使用规范的模型 URL。缓存位置与 TFHUB_CACHE_DIR默认缓存位置是本地临时目录/tmp/tfhub_modules即os.path.join(tempfile.gettempdir(), tfhub_modules)的计算结果。可通过环境变量TFHUB_CACHE_DIR推荐或命令行标志--tfhub_cache_dir自定义。若希望跨重启持久缓存可在~/.bashrc中添加export TFHUB_CACHE_DIR$HOME/.cache/tfhub_modules注意使用持久化位置时没有自动清理机制需要自行管理磁盘空间。从远程存储直接读取os.environ[TFHUB_MODEL_LOAD_FORMAT] UNCOMPRESSED或设置命令行标志--tfhub_model_load_formatUNCOMPRESSED。此模式下无需缓存目录对磁盘空间小、网络快的环境尤其有用。Colab TPU 场景的两种方案在 Colab 中运行 TPU 时压缩下载会与 TPU 运行时冲突——计算被委派给另一台机器它默认无法访问缓存目录。两种解决办法使用 TPU worker 可访问的 GCS bucket启用上述直接读取模式或把缓存指到自己的 bucketimport os os.environ[TFHUB_CACHE_DIR] gs://my-bucket/tfhub-modules-cache让所有读取都经由 Colab 主机load_options tf.saved_model.LoadOptions(experimental_io_device/job:localhost) reloaded_model hub.load(https://tfhub.dev/..., optionsload_options)另外模型下载过程本身可通过hub.resolve(handle)在需要时获得实际文件系统位置。模型托管协议tfhub.dev 平台与tensorflow_hub库共同遵循一套 HTTP(S) 托管协议其关键特征是代码中加载模型与浏览器中查看模型文档使用同一个 URL。URL 约定发布者https://tfhub.dev/publisher集合https://tfhub.dev/publisher/collection/collection_name模型带版本https://tfhub.dev/publisher/model_name/version不带版本https://tfhub.dev/publisher/model_name会解析到最新版本。不同模型类型通过 URL 参数指定下载格式类型模型 URL 示例下载类型URL 参数下载 URLTensorFlowSavedModel / TF1 Hub 格式https://tfhub.dev/google/spice/2.tar.gz?tf-hub-formatcompressedhttps://tfhub.dev/google/spice/2?tf-hub-formatcompressedTF Litehttps://tfhub.dev/google/lite-model/spice/1.tflite?lite-formattflitehttps://tfhub.dev/google/lite-model/spice/1?lite-formattfliteTF.jshttps://tfhub.dev/google/tfjs-model/spice/2/default/1.tar.gz?tfjs-formatcompressedhttps://tfhub.dev/google/tfjs-model/spice/2/default/1?tfjs-formatcompressed部分模型还支持直接远程读取TensorFlow 模型追加?tf-hub-formatuncompressed会返回 GCS 上未压缩模型所在目录的路径TF.js 模型追加?tfjs-formatfile则直接返回.json如.../model.json?tfjs-formatfile。注意远程直读可能增加延迟。压缩托管与加载流程模型以 tar.gz 压缩包形式存储库默认自动下载并解压缓存。手动下载可模拟协议wget https://tfhub.dev/tensorflow/albert_en_xxlarge/1?tf-hub-formatcompressed归档根目录即模型目录根应包含 SavedModel./saved_model.pb、./variables/、./assets/等TF1 Hub 格式的 tar 包还含./tfhub_module.pb。可以用tar -cz -f model.tar.gz --owner0 --group0 -C /tmp/export-model/ .从 SavedModel 目录打包。调用hub.KerasLayer、hub.load等 API 时库下载、解压并本地缓存模型协议要求模型 URL 带版本且同一版本内容不可变以便无限期缓存。非压缩托管当设置TFHUB_MODEL_LOAD_FORMATUNCOMPRESSED或命令行标志--tfhub_model_load_format时库会在模型 URL 后追加?tf-hub-formatuncompressed该请求在 303 响应体中返回 GCS 上未压缩模型目录的路径库随后直接从该位置读取模型。从 TF1 迁移到 TF2迁移要点详见 migration_tf2.md 与 model_compatibility.mdTF2 下应使用hub.KerasLayer与其它 Keras 层一起构建tf.keras.Model及底层hub.load()hub.ModuleAPI 仅保留用于 TF1 及 TF2 的 v1 兼容模式且只能加载 TF1 Hub 格式的模型新 API 在 TensorFlow 1.15eager 与图模式和 TensorFlow 2 中均可工作能加载 TF2 SavedModel并能在受限条件下加载 TF1 Hub 格式旧模型在 Estimator 中使用 TF2 SavedModel 做参数服务器训练或变量位于远程设备的 TF1 Session时需要在tf.Session的 ConfigProto 中设置experimental.share_cluster_devices_in_sessionTrue否则会报 Assigned device ... does not match any device.从 TF2.2 起该选项不再是实验性的可去掉.experimental前缀。两种模型格式的能力矩阵TF1 Hub 格式 / TF2 SavedModelTF1 Hub 格式在 TF1或 TF2 v1 兼容模式中加载/推理、微调、创建均完全支持通过hub.Module及trainableTrue、tags[train]等在纯 TF2 中加载/推理需改用hub.loadm.signaturessig或hub.KerasLayersignaturesig微调与创建不受支持。TF2 SavedModelTF1.15 之前的版本不支持在 TF1.15/TF2 中均可通过hub.load或hub.KerasLayer加载推理微调在 TF2 中完全支持hub.KerasLayer(handle, trainableTrue)或hub.load后调用m(inputs, trainingis_training)在 TF1 兼容模式下仅当hub.KerasLayer用于tf.keras.Model.fit()或包装该 Model 的 Estimator 时受支持创建在两种环境下均可tf.saved_model.save()也可在兼容模式内调用。常见问题排查common_issues.md 整理了高频故障与对策AutoTrackable object is not callable常见于用hub.load()在 TF2 中加载 TF1 Hub 格式模型。修复方式是通过签名调用embed hub.load(https://tfhub.dev/google/nnlm-en-dim128/1) embed.signaturesdefault无法下载模块多为网络栈问题而非库缺陷常见两类——EOF occurred in violation of protocolPython 版本不支持服务器 TLS 要求如 python 2.7.5需升级 Pythoncannot verify tfhub.devs certificate网络中有软件拦截 .dev 域名的解析需重新配置相关软件。另外还有写入缓存目录/tmp/tfhub_modules失败的情况可参考缓存章节更换位置。若仍不行可手动模拟协议下载并解压后改用本地路径$ mkdir /tmp/moduleA $ curl -L https://tfhub.dev/google/universal-sentence-encoder/2?tf-hub-formatcompressed | tar -zxvC /tmp/moduleA $ python import tensorflow_hub as hub hub.Module(/tmp/moduleA)对预初始化模块反复做推理TF2 SavedModel 可在初始化阶段hub.load(...)一次请求阶段直接调用tf.function 调用已针对性能优化TF1 Hub 模块则需在初始化阶段构建图、占位符并初始化 Session请求阶段通过session.run(embedded_text, feed_dict{...})喂数据。生产服务建议考虑 TensorFlow Serving 等可扩展、免 Python 的方案。无法更改模型 dtype如 float32 转 bfloat16SavedModel 内的操作绑定固定数据类型加载后无法事后更改发布者可选择发布不同 dtype 的模型。更新模型版本版本的文档元数据可更新但版本资产不可变改动模型内容需发布新版本并建议在文档中写清版本变更日志。下一步在 lib_overview.md 中了解tensorflow_hub库的整体 API 与公平性、安全性注意事项模型可视为任意 TensorFlow 图的程序引用不可信来源的模型需关注安全影响复用在大数据集上训练的模型时也应留意其训练数据与潜在偏差对自身场景的影响通过 tf2_saved_model.md 掌握 TF2 SavedModel 的加载、创建与微调细节通过 reusable_saved_models.md 与 common_saved_model_apis 学习可复用模型的接口规范参考教程文本分类 tf2_text_classification.ipynb、图像分类 tf2_image_retraining.ipynb以及 tutorials 目录下的更多实战 notebook阅读 caching.md、hosting.md 与 common_issues.md解决缓存、托管与故障排查问题参考 publish.md 了解模型发布流程并可通过 contribute.md 参与 TensorFlow Hub 社区建设。赞分享文档开发工具教程【免费下载链接】docsTensorFlow documentation项目地址https://gitcode.com/gh_mirrors/doc/docs点击查看免费下载相关推荐TensorFlow Hub 库tensorflow_hub使用指南加载、缓存与复用预训练模型TensorFlow Hub 库tensorflow_hub使用指南加载、缓存与复用预训练模型 tensorflow_hub 是 TensorFlow 官文档开发工具教程DeepSeek-V2.5模型文件解析55个safetensors分片的组成与加载DeepSeek V2.5模型文件解析55个safetensors分片的组成与加载 DeepSeek V2.5是DeepSeek AI推出的升级版语言模型融人工智能深度学习大模型dumpDex与易开发集成如何将脱壳功能嵌入你的开发工具dumpDex与易开发集成如何将脱壳功能嵌入你的开发工具 dumpDex是一款功能强大的Android脱壳工具需要Xposed框架支持并且已与易开发平台无逆向工程应用安全移动开发上一篇KMS_VL_ALL_AIO 深度使用指南一个脚本搞定 Windows 与 Office 的自动续期激活下一篇KMS_VL_ALL_AIO 激活工具实战教程装一次就让 Windows 和 Office 自动续期不再每 180 天返工一次创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进