Googleは新たに、20億件(2B)および70億件(7B)のパラメータを持つオープンソース大規模言語モデル「Gemma」をリリースした。このモデル群は、消費者向けハードウェア(GPU、TPU、CPU、モバイルデバイス)での実行を念頭に設計されており、8,192トークンまでの長文処理が可能である。
Gemmaは以下の4種類が提供される:
gemma-2b:2Bパラメータのベースモデルgemma-2b-it:2Bパラメータの指示微調整版(instruction-tuned)gemma-7b:7Bパラメータのベースモデルgemma-7b-it:7Bパラメータの指示微調整版
性能比較
Open LLM Leaderboardにおけるベースモデルのスコア(数値が高いほど高性能)は以下の通り:
| モデル | ライセンス | トレーニングトークン数 | スコア |
|---|---|---|---|
| Llama 2 70B Chat | Llama 2 License | 2T | 67.87 |
| Gemma-7B | Gemma License | 6T | 63.75 |
| DeciLM-7B | Apache 2.0 | 不明 | 61.55 |
| Phi-2 (2.7B) | MIT | 1.4T | 61.33 |
| Mistral-7B-v0.1 | Apache 2.0 | 不明 | 60.97 |
| Llama 2 7B | Llama 2 License | 2T | 54.32 |
| Gemma-2B | Gemma License | 2T | 46.51 |
7BクラスではMistral-7Bと同等の性能を発揮し、2Bクラスでも競争力のある結果を示している。ただし、チャット用途にはMT-BenchやLMSYS Arenaなどの専用ベンチマークでの評価が推奨される。
プロンプト形式
ベースモデルは自由なテキスト生成に対応するが、指示微調整版(-it)は特定の対話フォーマットを要求する:
<start_of_turn>user
Hello<end_of_turn>
<start_of_turn>model
Hi there!<end_of_turn>
Hugging Face Transformersのapply_chat_templateを利用することで、このフォーマットを自動的に適用できる。
Transformersによる推論例
以下はgemma-7b-itをbfloat16精度でGPU上で実行するコード例(約18GB VRAM必要):
from transformers import AutoTokenizer, pipeline
import torch
model_id = "google/gemma-7b-it"
tokenizer = AutoTokenizer.from_pretrained(model_id)
pipe = pipeline(
"text-generation",
model=model_id,
tokenizer=tokenizer,
torch_dtype=torch.bfloat16,
device_map="auto"
)
messages = [{"role": "user", "content": "Explain quantum computing in simple terms."}]
prompt = pipe.tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
outputs = pipe(
prompt,
max_new_tokens=256,
do_sample=True,
temperature=0.7,
top_p=0.95
)
print(outputs[0]["generated_text"][len(prompt):])
4ビット量子化により、VRAM使用量を約9GBに抑えることも可能:
from transformers import BitsAndBytesConfig
quantization_config = BitsAndBytesConfig(load_in_4bit=True)
pipe = pipeline(
"text-generation",
model=model_id,
tokenizer=tokenizer,
quantization_config=quantization_config,
device_map="auto"
)
JAX/Flaxサポート
GemmaはPyTorchだけでなく、JAX/Flaxでも利用可能。TPUやマルチGPU環境での高速推論に適している:
from transformers import FlaxGemmaForCausalLM, AutoTokenizer
import jax.numpy as jnp
model_id = "google/gemma-2b"
tokenizer = AutoTokenizer.from_pretrained(model_id, padding_side="left")
model, params = FlaxGemmaForCausalLM.from_pretrained(
model_id,
dtype=jnp.bfloat16,
revision="flax"
)
inputs = tokenizer("Tokyo is", return_tensors="jax")
output = model.generate(**inputs, params=params, max_new_tokens=20)
print(tokenizer.decode(output.sequences[0], skip_special_tokens=True))
Google Cloudとの統合
GemmaはVertex AIおよびGoogle Kubernetes Engine(GKE)を通じてクラウド上でのデプロイ・ファインチューニングが可能。Hugging Face Hubのモデルページから「Deploy → Google Cloud」を選択することで、ワンクリックでデプロイ設定が開始される。
推論エンドポイント
Hugging FaceのInference Endpointsは、Text Generation Inference(TGI)をバックエンドとしてGemmaを本番環境向けに最適化された状態で提供する。OpenAI互換APIもサポートしており、既存アプリケーションへの統合が容易:
from openai import OpenAI
client = OpenAI(
base_url="YOUR_ENDPOINT_URL/v1/",
api_key="YOUR_HF_TOKEN"
)
stream = client.chat.completions.create(
model="tgi",
messages=[{"role": "user", "content": "What is open source?"}],
stream=True
)
for chunk in stream:
if chunk.choices[0].delta.content:
print(chunk.choices[0].delta.content, end="")
TRLによる効率的ファインチューニング
QLoRAと4ビット量子化を組み合わせることで、単一のA10G GPU(24GB VRAM)上でGemma-7Bのファインチューニングが可能。以下はOpenAssistantデータセットを使用したサンプルコマンド:
accelerate launch \
examples/scripts/sft.py \
--model_name google/gemma-7b \
--dataset_name OpenAssistant/oasst_top1_2023-08-25 \
--per_device_train_batch_size 2 \
--learning_rate 2e-4 \
--use_peft \
--peft_lora_r 16 \
--peft_lora_alpha 32 \
--target_modules q_proj k_proj v_proj o_proj \
--load_in_4bit
この設定では、アテンション層の線形プロジェクションのみをLoRAで微調整し、MLP層は固定することでメモリ効率を最大化している。