Hugging Face 与 KerasHub 集成

Hugging Face 与 KerasHub 引入了一种共享的模型保存格式,使 KerasHub 用户能够直接加载来自 Hugging Face Hub、使用 Transformers 库创建的模型。此集成消除了之前 KerasHub 用户只能访问专为 KerasHub 构建的模型的限制,为他们提供了超过 30 万个预训练模型的库。

直接访问 Transformers 模型

KerasHub 现在可以使用 from_preset 方法直接加载 Transformers 库的检查点。这使用户能够使用大量原本未使用 Keras 创建的微调模型。

最初,此集成支持以下架构:

  • Gemma(版本 1 和 2)
  • Llama 3
  • PaliGemma

多框架部署

由于 KerasHub 模型可以使用 TensorFlowJAXPyTorch 后端运行,此集成允许用户仅用一行代码将 Hugging Face 检查点加载到任意这些框架中。这一功能简化了模型迁移的过程,以满足特定需求,例如部署到 TFLite 进行服务或使用 JAX 进行研究。

技术实现

该集成通过在两个库之间映射配置变量、权重名称和分词器词汇表来实现。由于 Transformers 模型以 JSON 配置文件、分词器文件和 safetensors 权重的形式存储,只要两个库都具备相应架构的建模代码,KerasHub 就能创建兼容的检查点。此转换过程由库内部处理,用户无需手动转换。

使用与配置

要使用此集成,用户必须升级到 keras-hub 并使用 keras>=3.3.3

文本生成

用户可以加载 Transformers 模型并使用 .generate 方法生成文本。例如,从 Hub 加载 Llama 3 模型:

from keras_hub.models import Llama3CausalLM

causal_lm = Llama3CausalLM.from_preset(
    "hf://NousResearch/Hermes-2-Pro-Llama-3-8B"
)

prompts = ["Your prompt here"]
causal_lm.generate(prompts, max_length=200)

精度与后端控制

KerasHub 允许轻松调整模型精度和底层计算后端:

  • 更改精度: 可以在加载模型之前通过 keras.config.set_dtype_policy("bfloat16") 设置精度。
  • 切换后端: 通过设置环境变量 os.environ["KERAS_BACKEND"] = "jax",用户可以使用 JAX 后端运行已加载的 Transformers 检查点。

支持的模型

除了 Llama 3,集成还明确支持以下模型:

  • Gemma 2: 用户可以直接加载 Gemma 2 模型(例如 google/gemma-2-9b)。
  • PaliGemma: 任何 PaliGemma safetensor 检查点(包括微调版本)都可以集成到 KerasHub 流水线中。

Sources