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