Chuyển đến nội dung chính

Lesson 11: Supervised Fine-Tuning (SFT) — Instruction Tuning

Explore the Supervised Fine-Tuning (SFT) technique to turn pre-trained LLM into an instruction-following model. Learn how to prepare data, configure SFTTrainer, apply chat templates and optimize hyperparameters to fine-tune Mistral-7B effectively.

🧠 AI & ML — Lesson 10 Lesson 11: Supervised Fine-Tuning (SFT) — Instruction Tuning

AI & LLM: From Basics to Advanced

Part 3: Training & Fine-tuning LLMs

xdev.asia

Lesson 11: Supervised Fine-Tuning (SFT) — Instruction Tuning

1. Why is Fine-Tuning needed?

Pre-trained Model vs Instruction-Following Model

A pre-trained LLM (like GPT-2, LLaMA base, Mistral base) is trained with a single goal: predict the next token in the text. As a result, the model learned grammar, world knowledge, and the ability to reason — but it didn't know how to answer questions or follow instructions.

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..."

This is pure text completion — not useful to the end user.

Instruction-following model (ChatGPT, Claude, Mistral-Instruct) has been fine-tuned to:

  • Understand and follow user instructions
  • Reply in desired format (JSON, bullet points, etc.)
  • Maintain roles (assistant, expert, etc.) defined in system prompt
  • Reject harmful or inappropriate requests

The process of turning a base model into an instruction model is Supervised Fine-Tuning (SFT).

When should you fine-tune?

SituationSolution
Model needs to understand domain-specific knowledgeFine-tune on domain data
Need output in fixed formatSFT with format example
Want the model to "speak" in its own styleSFT with style data
Need to optimize inference costsFine-tune smaller model
Prompt engineering is not enoughSFT is often more effective

2. Instruction Dataset Format

Basic structure: (Instruction, Input, Output)

SFT datasets usually have the form triple:

{
  "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."
}
  • instruction: Specific requirements (mandatory)
  • input: Context or additional data (can be empty)
  • output: Desired answer (ground truth)

Conversation Format (Multi-turn)

To train the multi-turn conversation model:

{
  "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. Popular Datasets

Alpaca (Stanford, 2023)

  • 52,000 instruction-following examples generated by GPT-3.5
  • Format: instruction + input + output
  • Pave the way for the self-instruct movement
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

  • Real conversations from ChatGPT users (shared voluntarily)
  • Multi-turn, diverse topics
  • Suitable for chat models

Dolly (Databricks)

  • 15,000 examples written by Databricks staff
  • Completely open-source, no legal issues
  • Higher quality than Alpaca (written by human)

FLAN Collection

  • Compiled from hundreds of different NLP tasks
  • Very good for general instruction following
  • Used in FLAN-T5, FLAN-UL2

4. SFT with Hugging Face TRL

Install

pip install trl transformers datasets accelerate peft

SFTTrainer

SFTTrainer from the TRL (Transformer Reinforcement Learning) library which is the standard tool for SFT:

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()

Data Collators

DataCollatorForCompletionOnlyLM Only calculate loss on assistant response, not loss on instruction/input:

from trl import DataCollatorForCompletionOnlyLM

response_template = "[/INST]"  # Llama-2 format
collator = DataCollatorForCompletionOnlyLM(
    response_template=response_template,
    tokenizer=tokenizer
)

This is very important: if you calculate loss on prompts, the model will learn to generate... prompts instead of responses.


5. Chat Templates

ChatML Format (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 Format (Meta)

<|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|>

Mistral/Llama-2 Format

<s>[INST] <<SYS>>
You are a helpful assistant.
<</SYS>>

Hello! [/INST] Hi! How can I help you? </s>

Apply Chat Template automatically

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. Training Hyperparameters

Learning Rate

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ụ

Epochs and Overfitting

num_train_epochs = 3
# 1-3 epochs thường đủ cho SFT
# Nhiều hơn: nguy cơ overfitting, model chỉ nhớ training data

Batch Size and Gradient Accumulation

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

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. Mixed Precision, Gradient Checkpointing

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
)

Gradient Checkpointing

# 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. Full Code Example: Fine-tune Mistral-7B with TRL SFTTrainer

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. Evaluation: Loss Curve & Qualitative Testing

Track Training Loss

# 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

Qualitative Testing

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)

Simple benchmark

# 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

Summary

In this lesson we learned:

  • Pre-trained model only does text completion; SFT turns it into an instruction follower
  • SFT data needs at least instruction and output; many popular datasets such as Alpaca, ShareGPT, Dolly
  • TRL SFTTrainer helps simplify the fine-tuning process, especially with DataCollatorForCompletionOnlyLM
  • Chat templates (ChatML, Llama-3, Mistral) define the conversation structure — using the wrong template will give bad results
  • Combining LoRA + 4-bit quantization allows fine-tuning 7B models on consumer GPUs
  • Always evaluate qualitatively (realistic testing) in parallel with the loss curve

The next article will delve into PEFT, LoRA and QLoRA — techniques that help fine-tune large models with minimal VRAM.