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 模型可以使用 TensorFlow、JAX 或 PyTorch 后端运行,此集成允许用户仅用一行代码将 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 流水线中。