lora_finetune.py

script

← Back to skill

Content hash: 2ec7c4cd53440b0b7391a59656a8751064ae5ac7799f68c3786f08b6b8bbea66
#!/usr/bin/env python3
"""LoRA / QLoRA fine-tuning configuration and training script template.

Demonstrates: PEFT config setup, data formatting with chat templates,
training loop with SFTTrainer, and adapter merge/export.

RUN WITH: accelerate launch lora_finetune.py
REQUIRES: pip install transformers peft trl bitsandbytes datasets accelerate
"""

from __future__ import annotations

import json
from dataclasses import dataclass
from typing import Any


# ── Configuration ───────────────────────────────────────────────────────

@dataclass
class LoRAConfig:
    model_name: str = "meta-llama/Llama-3.2-3B-Instruct"
    r: int = 16
    lora_alpha: int = 32
    lora_dropout: float = 0.05
    target_modules: tuple[str, ...] = ("q_proj", "k_proj", "v_proj", "o_proj")
    bias: str = "none"
    task_type: str = "CAUSAL_LM"

    # Training
    learning_rate: float = 2e-4
    num_epochs: int = 2
    per_device_batch_size: int = 2
    gradient_accumulation_steps: int = 4
    warmup_ratio: float = 0.03
    max_seq_length: int = 2048
    use_4bit: bool = True

    # Paths
    output_dir: str = "./lora-adapter"
    dataset_path: str = "data/train.jsonl"


# ── Data formatting ─────────────────────────────────────────────────────

def format_chat_example(example: dict, tokenizer) -> dict:
    """Format a single example using the model's chat template."""
    messages = [
        {"role": "system", "content": example.get("system", "You are a helpful assistant.")},
        {"role": "user", "content": example["instruction"]},
        {"role": "assistant", "content": example["response"]},
    ]
    text = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=False)
    return {"text": text}


def load_and_format_dataset(path: str, tokenizer) -> Any:
    """Load JSONL dataset and apply chat template formatting."""
    from datasets import Dataset

    examples = []
    with open(path) as f:
        for line in f:
            examples.append(json.loads(line))

    dataset = Dataset.from_list(examples)
    # Format function depends on tokenizer — kept as template
    return dataset


# ── PEFT model setup (template) ─────────────────────────────────────────

def setup_peft_model(config: LoRAConfig):
    """Setup model with LoRA/QLoRA adapters. Template — requires real model loading."""
    import torch
    from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig
    from peft import LoraConfig as PeftLoraConfig, get_peft_model, TaskType

    # Load tokenizer
    tokenizer = AutoTokenizer.from_pretrained(config.model_name)
    if tokenizer.pad_token is None:
        tokenizer.pad_token = tokenizer.eos_token

    # QLoRA quantization config
    bnb_config = None
    if config.use_4bit:
        bnb_config = BitsAndBytesConfig(
            load_in_4bit=True,
            bnb_4bit_quant_type="nf4",
            bnb_4bit_compute_dtype=torch.bfloat16,
            bnb_4bit_use_double_quant=True,
        )

    # Load model
    model_kwargs = {
        "torch_dtype": torch.bfloat16,
        "device_map": "auto",
    }
    if bnb_config:
        model_kwargs["quantization_config"] = bnb_config

    model = AutoModelForCausalLM.from_pretrained(config.model_name, **model_kwargs)

    # LoRA config
    peft_config = PeftLoraConfig(
        task_type=TaskType.CAUSAL_LM,
        r=config.r,
        lora_alpha=config.lora_alpha,
        lora_dropout=config.lora_dropout,
        target_modules=list(config.target_modules),
        bias=config.bias,
    )

    model = get_peft_model(model, peft_config)
    model.print_trainable_parameters()

    return model, tokenizer


# ── Training loop (template) ────────────────────────────────────────────

def run_training(config: LoRAConfig, model, tokenizer, dataset):
    """Run SFTTrainer training loop."""
    from transformers import TrainingArguments
    from trl import SFTTrainer

    training_args = TrainingArguments(
        output_dir=config.output_dir,
        per_device_train_batch_size=config.per_device_batch_size,
        gradient_accumulation_steps=config.gradient_accumulation_steps,
        learning_rate=config.learning_rate,
        num_train_epochs=config.num_epochs,
        warmup_ratio=config.warmup_ratio,
        logging_steps=10,
        save_strategy="epoch",
        bf16=True,
        gradient_checkpointing=True,
        report_to="none",
    )

    trainer = SFTTrainer(
        model=model,
        args=training_args,
        train_dataset=dataset,
        tokenizer=tokenizer,
        max_seq_length=config.max_seq_length,
    )

    trainer.train()

    # Save adapter
    model.save_pretrained(config.output_dir)
    tokenizer.save_pretrained(config.output_dir)

    return trainer


# ── Merge & Export ──────────────────────────────────────────────────────

def merge_and_export(config: LoRAConfig, output_path: str):
    """Merge LoRA adapter into base model and export."""
    import torch
    from transformers import AutoModelForCausalLM, AutoTokenizer
    from peft import PeftModel

    base_model = AutoModelForCausalLM.from_pretrained(
        config.model_name, torch_dtype=torch.bfloat16, device_map="auto"
    )
    model = PeftModel.from_pretrained(base_model, config.output_dir)
    model = model.merge_and_unload()
    model.save_pretrained(output_path)
    tokenizer = AutoTokenizer.from_pretrained(config.model_name)
    tokenizer.save_pretrained(output_path)
    print(f"Merged model saved to {output_path}")


# ── Demo (config only, no actual training) ──────────────────────────────

if __name__ == "__main__":
    config = LoRAConfig()
    print("LoRA Training Configuration:")
    print(f"  Model: {config.model_name}")
    print(f"  Rank (r): {config.r}, Alpha: {config.lora_alpha}")
    print(f"  Target modules: {config.target_modules}")
    print(f"  Learning rate: {config.learning_rate}")
    print(f"  Epochs: {config.num_epochs}")
    print(f"  4-bit (QLoRA): {config.use_4bit}")
    print(f"  Effective batch: {config.per_device_batch_size * config.gradient_accumulation_steps}")
    print(f"  Output: {config.output_dir}")
    print("\nTo train, run: accelerate launch lora_finetune.py")
    print("Ensure model, tokenizer, and dataset are available.")
    print("\nāœ“ LoRA config template ready.")