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 任務,前處理器負責:

  1. 添加起始與結束文字標記。
  2. 填充標記序列並產生遮罩。
  3. 產生訓練用的「預期輸出」(即將輸入字串向右平移一個位置)。

訓練與 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)特別有用。

使用者可透過定義 DeviceMeshLayoutMap,將模型切分至多個 GPU 或 TPU 上。雖然大多數模型可透過 get_layout_map(device_mesh) 提供合理的預設值,使用者亦可自行定義自訂布局圖以最佳化效能。例如,在 TPU v5e 上的自訂布局圖可將 epoch 時間從 62 秒縮減至 54 秒。

Sources