レッスン 11: 教師ありファインチューニング (SFT) — 命令チューニング
1. なぜ微調整が必要なのでしょうか?
事前トレーニング済みモデルと指示に従うモデル
事前トレーニングされた LLM (GPT-2、LLaMA ベース、Mistral ベースなど) は、テキスト内の 次のトークンを予測するという 1 つの目標を持ってトレーニングされます。その結果、モデルは文法、世界の知識、推論能力を学習しましたが、質問に答えたり、指示に従う方法を知りませんでした。
Input: "Thủ đô của Việt Nam là"
Output: "thành phố lớn nhất ở miền Bắc, nằm bên bờ sông Hồng..."
これは純粋なテキスト補完であり、エンド ユーザーにとっては役に立ちません。
命令追従モデル (ChatGPT、Claude、Mistral-Instruct) は次のように微調整されました。
- ユーザーの指示を理解し、従う
- 希望の形式 (JSON、箇条書きなど) で返信します。
- システム プロンプトで定義された 役割 (アシスタント、エキスパートなど) を維持します
- 有害または不適切なリクエストを拒否します
ベース モデルを命令モデルに変換するプロセスは 教師あり微調整 (SFT) です。
いつ微調整する必要がありますか?
| 状況 | ソリューション |
|---|---|
| モデルはドメイン固有の知識を理解する必要があります | ドメイン データの微調整 |
| 固定フォーマットでの出力が必要 | SFT とフォーマットの例 |
| モデルが独自のスタイルで「話す」ようにしたい | スタイル データを含む SFT |
| 推論コストを最適化する必要がある | 小型モデルを微調整する |
| 迅速なエンジニアリングだけでは不十分 | SFT は多くの場合、より効果的です。 |
2. 命令データセットの形式
基本構造:(命令、入力、出力)
SFT データセットは通常、トリプル の形式になります。
{
"instruction": "Dịch đoạn văn sau sang tiếng Anh.",
"input": "Hôm nay trời đẹp, tôi muốn đi dạo.",
"output": "The weather is nice today, I want to go for a walk."
}
- 説明: 特定の要件 (必須)
- 入力: コンテキストまたは追加データ (空でも可)
- 出力: 望ましい答え (グラウンド トゥルース)
会話形式(マルチターン)
マルチターン会話モデルをトレーニングするには:
{
"conversations": [
{"role": "system", "content": "Bạn là trợ lý AI hữu ích."},
{"role": "user", "content": "Python là gì?"},
{"role": "assistant", "content": "Python là ngôn ngữ lập trình..."},
{"role": "user", "content": "Nó có ưu điểm gì?"},
{"role": "assistant", "content": "Python có nhiều ưu điểm..."}
]
}
3. 人気のあるデータセット
アルパカ (スタンフォード、2023)
- GPT-3.5 によって生成された 52,000 の命令に従うサンプル
- 形式: 命令 + 入力 + 出力
- 自己指導運動への道を切り開く
from datasets import load_dataset
alpaca = load_dataset("tatsu-lab/alpaca")
print(alpaca["train"][0])
# {'instruction': 'Give three tips for staying healthy.',
# 'input': '',
# 'output': '1. Eat a balanced diet...'}
ShareGPT
- ChatGPT ユーザーからの実際の会話 (自主的に共有)
- マルチターン、多様なトピック
- チャットモデルに適しています
ドリー (Databricks)
- 15,000 例 Databricks スタッフによって作成されました
- 完全にオープンソースであり、法的問題はありません ・アルパカより高品質(人間が書いたもの)
FLAN コレクション
- 何百もの異なる NLP タスクからコンパイル
- 以下の一般的な指導に非常に適しています
- FLAN-T5、FLAN-UL2で使用
4. ハグフェイスTRLを使用したSFT
インストール
pip install trl transformers datasets accelerate peft
SFTトレーナー
SFTTrainer SFT の標準ツールである TRL (Transformer Reinforcement Learning) ライブラリから:
from trl import SFTTrainer, SFTConfig
from transformers import AutoModelForCausalLM, AutoTokenizer
model = AutoModelForCausalLM.from_pretrained("mistralai/Mistral-7B-v0.1")
tokenizer = AutoTokenizer.from_pretrained("mistralai/Mistral-7B-v0.1")
trainer = SFTTrainer(
model=model,
train_dataset=dataset,
args=SFTConfig(
output_dir="./sft-output",
num_train_epochs=3,
per_device_train_batch_size=4,
learning_rate=2e-5,
),
)
trainer.train()
データ照合者
DataCollatorForCompletionOnlyLM 命令/入力での損失ではなく、アシスタント応答での損失のみを計算します。
from trl import DataCollatorForCompletionOnlyLM
response_template = "[/INST]" # Llama-2 format
collator = DataCollatorForCompletionOnlyLM(
response_template=response_template,
tokenizer=tokenizer
)
これは非常に重要です。プロンプトの損失を計算すると、モデルは応答の代わりにプロンプトを生成することを学習します。
5. チャット テンプレート
ChatML 形式 (OpenAI)
<|im_start|>system
You are a helpful assistant.<|im_end|>
<|im_start|>user
Hello!<|im_end|>
<|im_start|>assistant
Hi! How can I help you?<|im_end|>
Llama-3 形式 (メタ)
<|begin_of_text|><|start_header_id|>system<|end_header_id|>
You are a helpful assistant.<|eot_id|>
<|start_header_id|>user<|end_header_id|>
Hello!<|eot_id|>
<|start_header_id|>assistant<|end_header_id|>
Hi! How can I help?<|eot_id|>
ミストラル/ラマ-2 フォーマット
<s>[INST] <<SYS>>
You are a helpful assistant.
<</SYS>>
Hello! [/INST] Hi! How can I help you? </s>
チャット テンプレートを自動的に適用する
messages = [
{"role": "system", "content": "Bạn là trợ lý AI hữu ích."},
{"role": "user", "content": "Python là gì?"}
]
# Tokenizer sẽ tự áp dụng đúng template
formatted = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True
)
print(formatted)
6. ハイパーパラメータのトレーニング
学習率
learning_rate = 2e-5 # Điểm khởi đầu tốt cho full fine-tuning
# Với LoRA: 1e-4 đến 3e-4
# Quá cao: model "quên" kiến thức cũ (catastrophic forgetting)
# Quá thấp: học rất chậm, không hội tụ
エポックと過学習
num_train_epochs = 3
# 1-3 epochs thường đủ cho SFT
# Nhiều hơn: nguy cơ overfitting, model chỉ nhớ training data
バッチサイズとグラジエントの蓄積
per_device_train_batch_size = 4
gradient_accumulation_steps = 4
# Effective batch size = 4 * 4 = 16
# Dùng gradient accumulation khi VRAM không đủ để batch lớn
ウォームアップ
warmup_ratio = 0.03
# 3% đầu của training: learning rate tăng dần từ 0
# Giúp tránh gradient explosion ở đầu training
7. 混合精度、勾配チェックポイント
BF16/FP16
from transformers import TrainingArguments
args = TrainingArguments(
bf16=True, # Dùng BF16 nếu có Ampere GPU (A100, 3090)
fp16=False, # FP16 cho GPU cũ hơn (V100, T4)
# BF16 ổn định hơn FP16, ít bị overflow hơn
)
グラデーションチェックポイント
# Tiết kiệm VRAM bằng cách không lưu tất cả activations
# Đánh đổi: chậm hơn ~30% do phải tính lại activations
model.gradient_checkpointing_enable()
args = TrainingArguments(
gradient_checkpointing=True,
gradient_checkpointing_kwargs={"use_reentrant": False}
)
8. 完全なコード例: TRL SFTTrainer を使用した Mistral-7B の微調整
import torch
from datasets import load_dataset
from transformers import (
AutoModelForCausalLM,
AutoTokenizer,
BitsAndBytesConfig,
)
from peft import LoraConfig, get_peft_model
from trl import SFTTrainer, SFTConfig
# --- 1. Load model (4-bit để tiết kiệm VRAM) ---
model_id = "mistralai/Mistral-7B-v0.3"
bnb_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.bfloat16,
bnb_4bit_use_double_quant=True,
)
model = AutoModelForCausalLM.from_pretrained(
model_id,
quantization_config=bnb_config,
device_map="auto",
trust_remote_code=True,
)
model.config.use_cache = False
tokenizer = AutoTokenizer.from_pretrained(model_id)
tokenizer.pad_token = tokenizer.eos_token
tokenizer.padding_side = "right"
# --- 2. LoRA Config ---
lora_config = LoraConfig(
r=16,
lora_alpha=32,
target_modules=["q_proj", "v_proj", "k_proj", "o_proj"],
lora_dropout=0.05,
bias="none",
task_type="CAUSAL_LM",
)
model = get_peft_model(model, lora_config)
model.print_trainable_parameters()
# trainable params: 83,886,080 || all params: 3,836,829,696 || ~2.19%
# --- 3. Dataset ---
dataset = load_dataset("tatsu-lab/alpaca", split="train[:5000]")
def format_prompt(example):
"""Chuyển Alpaca format sang Mistral instruction format."""
if example["input"]:
prompt = (
f"[INST] {example['instruction']}\n\n"
f"Input: {example['input']} [/INST] "
f"{example['output']}"
)
else:
prompt = (
f"[INST] {example['instruction']} [/INST] "
f"{example['output']}"
)
return {"text": prompt}
dataset = dataset.map(format_prompt)
# --- 4. Training Config ---
sft_config = SFTConfig(
output_dir="./mistral-sft-alpaca",
num_train_epochs=3,
per_device_train_batch_size=4,
gradient_accumulation_steps=4,
learning_rate=2e-4,
lr_scheduler_type="cosine",
warmup_ratio=0.03,
bf16=True,
gradient_checkpointing=True,
logging_steps=25,
save_steps=500,
evaluation_strategy="no",
max_seq_length=2048,
dataset_text_field="text",
report_to="wandb", # hoặc "tensorboard"
)
# --- 5. Trainer ---
trainer = SFTTrainer(
model=model,
args=sft_config,
train_dataset=dataset,
tokenizer=tokenizer,
)
trainer.train()
# --- 6. Lưu model ---
trainer.save_model("./mistral-sft-alpaca/final")
tokenizer.save_pretrained("./mistral-sft-alpaca/final")
print("Training complete!")
9. 評価: 損失曲線と定性テスト
トレーニングの損失を追跡する
# Training loss nên giảm dần và ổn định
# - Giảm quá nhanh rồi tăng lại: overfitting
# - Không giảm: learning rate quá thấp hoặc data có vấn đề
# - Dao động lớn: learning rate quá cao
# Với wandb: loss curve được vẽ tự động
# Mục tiêu: train loss < 1.0 sau vài epoch
定性的テスト
from transformers import pipeline
pipe = pipeline(
"text-generation",
model="./mistral-sft-alpaca/final",
tokenizer=tokenizer,
device_map="auto",
)
test_prompts = [
"[INST] Giải thích khái niệm recursion trong lập trình. [/INST]",
"[INST] Viết hàm Python tính số Fibonacci. [/INST]",
"[INST] Tóm tắt bài báo sau: ... [/INST]",
]
for prompt in test_prompts:
output = pipe(
prompt,
max_new_tokens=256,
temperature=0.7,
do_sample=True,
)
print(output[0]["generated_text"])
print("-" * 50)
簡単なベンチマーク
# So sánh model trước và sau fine-tune trên test set
# Đo lường:
# - ROUGE score cho summarization
# - Exact match cho extraction tasks
# - Human evaluation cho general quality
概要
このレッスンでは次のことを学びました。
- 事前トレーニングされたモデルはテキスト補完のみを行います。 SFT はそれを命令フォロワーに変えます
- SFT データには少なくとも 命令 と 出力 が必要です。 Alpaca、ShareGPT、Dolly などの多くの人気のあるデータセット
- TRL SFTTrainer は、特に次のような微調整プロセスを簡素化するのに役立ちます。
DataCollatorForCompletionOnlyLM - チャット テンプレート (ChatML、Llama-3、Mistral) は会話の構造を定義します - 間違ったテンプレートを使用すると悪い結果が生じます
- LoRA + 4 ビット量子化 を組み合わせることで、コンシューマ GPU で 7B モデルを微調整できます
- 損失曲線と並行して常に 定性的評価 (現実的なテスト) を行う
次の記事では、最小限の VRAM で大規模なモデルを微調整するのに役立つテクニックである PEFT、LoRA、および QLoRA について詳しく説明します。