Keras 中的 Llama 3.2
Llama 3.2 現已可在 Keras 中直接使用,支援載入標準的 Hugging Face 檢查點(包括 safetensors),如有需要會即時轉換。此整合讓開發者能在 Keras 生態系統中利用 Llama 3.2,受惠於多後端彈性與整合的訓練工具。
支援 JAX、PyTorch 與 TensorFlow 的多後端
Keras 作為多後端建模函式庫,讓相同模型能在 JAX、PyTorch 或 TensorFlow 上執行。後端需在匯入 Keras 前透過環境變數指定:
import os
os.environ["KERAS_BACKEND"] = "jax" # Options: "jax", "torch", or "tensorflow"
此彈性允許使用者利用 JAX 搭配 XLA 編譯,以獲得最佳效能。
Keras-Hub 與模型整合
keras-hub(前身為 KerasNLP 與 KerasCV)是 Keras 的預訓練模型集合。它提供流行模型的標準 Keras 實作,包括 Llama 3、Gemma、StableDiffusion 與 Segment Anything。
可使用 keras_hub 中的 Llama3CausalLM 類別載入 Llama 3.2:
from keras_hub import models.Llama3CausalLM
model = Llama3CausalLM.from_preset("hf://meta-llama/Llama-3.2-1B-Instruct", dtype="bfloat16")
「開箱即用」的 LLM 功能
Keras LLM 旨在易於使用,將 tokenizer 直接整合至模型物件中。這使得可對原始字串執行高階操作:
- Generation:
model.generate("Hi there!")直接從字串輸入產生文字輸出。 - Training:
model.fit(strings)允許直接在字串列表或資料集上進行訓練。
聊天與指令微調
指令微調的變體,例如 Llama-3.2-1B-Instruct,支援使用特定標籤格式的逐回合對話。Llama 3.2 所需的格式包含 <|start_header_id|>system<|end_header_id|>、<|start_header_id|>user<|end_header_id|> 與 <|eot_id|> 等標籤。格式化後,這些字串即可直接傳入 model.generate()。
低階模型存取
對於需要更多控制的使用者,Keras 提供對底層元件的存取:
- Tokenizer: 可透過
model.preprocessor.tokenizer存取。此元件將文字轉換為整數向量。 - Backbone: 核心模型架構可透過
model.backbone存取。
前處理器概念
Keras 中的 Preprocessor 是一個完整的資料轉換工具。對於 CausalLM 任務,前處理器負責:
- 添加起始與結束文字標記。
- 填充標記序列並產生遮罩。
- 產生訓練用的「預期輸出」(即將輸入字串向右平移一個位置)。
訓練與 Hub 整合
Keras 內建訓練器,可透過 model.fit(ds) 存取。此訓練器相容於 Keras 的功能,包括分散式訓練、混合精度、量化,以及 LoRA 與 QLoRA 等參數效率微調方法。
微調後的模型可在本機使用 model.save_to_preset() 儲存後,透過 keras_hub.upload_preset() 直接上傳至 Hugging Face Hub。
分散式模型平行化
Keras 透過 JAX 與 XLA 編譯器提供簡化的進階模型平行化管道,對於過大無法在單一加速器上執行的模型(例如 Llama 3.1 8B)特別有用。
使用者可透過定義 DeviceMesh 與 LayoutMap,將模型切分至多個 GPU 或 TPU 上。雖然大多數模型可透過 get_layout_map(device_mesh) 提供合理的預設值,使用者亦可自行定義自訂布局圖以最佳化效能。例如,在 TPU v5e 上的自訂布局圖可將 epoch 時間從 62 秒縮減至 54 秒。
Sources
- Original“Llama 3.2 in Keras”