gradio-forge / scripts /finetune.py
Sokheng's picture
init
3f47273
Raw
History Blame Contribute Delete
2.98 kB
import torch
from datasets import load_dataset
from trl import SFTConfig, SFTTrainer
from unsloth import FastLanguageModel
MAX_SEQ_LENGTH = 512
MODEL_ID = "Qwen/Qwen3-8B"
OUTPUT_DIR = "outputs/gradio-forge-7b"
HF_REPO = "SokhengDin/gradio-forge-7b"
DATASET_PATH = "data/finetune_dataset.jsonl"
SYSTEM_PROMPT = open("prompts/system.txt", encoding="utf-8").read().strip()
def load_base_model() -> tuple:
"""Load base model with Unsloth 4-bit quantization."""
model, tokenizer = FastLanguageModel.from_pretrained(
model_name = MODEL_ID,
max_seq_length = MAX_SEQ_LENGTH,
load_in_4bit = True,
)
return model, tokenizer
def add_lora(model) -> object:
"""Attach LoRA adapters to the model."""
return FastLanguageModel.get_peft_model(
model,
r = 16,
target_modules = [
"q_proj", "k_proj", "v_proj", "o_proj",
"gate_proj", "up_proj", "down_proj",
],
lora_alpha = 16,
lora_dropout = 0,
bias = "none",
use_gradient_checkpointing = "unsloth",
)
def format_example(example: dict, tokenizer) -> dict:
"""Format a prompt/completion pair as a full chat-template string."""
text = tokenizer.apply_chat_template(
[
{"role": "system", "content": SYSTEM_PROMPT},
{"role": "user", "content": example["prompt"]},
{"role": "assistant", "content": example["completion"]},
],
tokenize = False,
add_generation_prompt = False,
)
return {"text": text}
def main() -> None:
model, tokenizer = load_base_model()
model = add_lora(model)
dataset = load_dataset("json", data_files=DATASET_PATH, split="train")
dataset = dataset.map(lambda ex: format_example(ex, tokenizer))
trainer = SFTTrainer(
model = model,
tokenizer = tokenizer,
train_dataset = dataset,
args = SFTConfig(
dataset_text_field = "text",
max_seq_length = MAX_SEQ_LENGTH,
output_dir = OUTPUT_DIR,
num_train_epochs = 3,
per_device_train_batch_size = 4,
gradient_accumulation_steps = 4,
warmup_steps = 10,
learning_rate = 2e-4,
logging_steps = 10,
save_strategy = "epoch",
fp16 = not torch.cuda.is_bf16_supported(),
bf16 = torch.cuda.is_bf16_supported(),
report_to = "none",
),
)
trainer.train()
model.push_to_hub(HF_REPO)
tokenizer.push_to_hub(HF_REPO)
print(f"Published to {HF_REPO}")
if __name__ == "__main__":
main()