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 流程中。