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"

この柔軟性により、最適化されたパフォーマンスのために XLA コンパイルを伴う JAX を利用できます。

Keras-Hub とモデル統合

keras-hub(旧称 KerasNLP と KerasCV)は、Keras 用の事前学習済みモデルのコレクションです。Llama 3、Gemma、StableDiffusion、Segment Anything など、人気モデルの標準的な Keras 実装を提供します。

Llama 3.2 は keras_hubLlama3CausalLM クラスを使ってロードできます:

from keras_hub import models.Llama3CausalLM
model = Llama3CausalLM.from_preset("hf://meta-llama/Llama-3.2-1B-Instruct", dtype="bfloat16")

「バッテリー同梱」LLM 機能

Keras の LLM は、トークナイザーをモデルオブジェクトに直接統合することで使いやすさを実現しています。これにより、生の文字列に対して高レベルの操作が可能になります:

  • 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 でアクセスできます。

Preprocessor の概念

Keras の Preprocessor はデータ変換のための包括的ツールです。CausalLM タスクに対して、Preprocessor は以下を処理します:

  1. 開始および終了テキストトークンの追加。
  2. トークンシーケンスのパディングとマスクの生成。
  3. トレーニング用の「期待出力」の生成(入力文字列を 1 つシフトしたもの)。

トレーニングと Hub 統合

Keras には model.fit(ds) で利用できる組み込みトレーナーが含まれています。このトレーナーは、分散トレーニング、混合精度、量子化、LoRA や QLoRA といったパラメータ効率の高いファインチューニング手法など、Keras の機能と互換性があります。

ファインチューニング済みモデルは、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 上のカスタムレイアウトマップはエポック時間を 62 秒から 54 秒に短縮できます。

Sources