import os
import torch
from unsloth import FastLanguageModel
from trl import SFTTrainer
from transformers import TrainingArguments
from datasets import load_dataset
# --- Configuration ---
MODEL_NAME = "unsloth/llama-3-8b-instruct-bnb-4bit"
DATASET_PATH = os.path.expanduser("~/projects/web-dev-swarm/datasets/frontend/raw/frontend_gold_standard.jsonl")
OUTPUT_DIR = os.path.expanduser("~/projects/web-dev-swarm/outputs/frontend_specialist_lora")
MAX_SEQ_LENGTH = 2048
LORA_RANK = 16
LORA_ALPHA = 32
LORA_DROPOUT = 0
# --- Prompt Template ---
prompt_style = """Below is an instruction that describes a task, paired with an input that provides further context. Write a response that appropriately completes the request.
### Instruction:
{}
### Input:
{}
### Response:
{}"""
def train():
print(f"🚀 Loading model: {MODEL_NAME}")
model, tokenizer = FastLanguageModel.from_pretrained(
model_name = MODEL_NAME,
max_seq_length = MAX_SEQ_LENGTH,
load_in_4bit = True,
)
print("🛠 Applying LoRA adapters...")
model = FastLanguageModel.get_peft_model(
model,
r = LORA_RANK,
target_modules = ["q_proj", "k_proj", "v_proj", "o_proj",
"gate_proj", "up_proj", "down_proj",],
lora_alpha = LORA_ALPHA,
lora_dropout = LORA_DROPOUT,
bias = "none",
use_gradient_checkpointing = "unsloth",
random_state = 3407,
)
EOS_TOKEN = tokenizer.eos_token
def formatting_prompts_func(examples):
instructions = examples["instruction"]
inputs = examples["input"]
outputs = examples["output"]
texts = []
for instruction, input, output in zip(instructions, inputs, outputs):
text = prompt_style.format(instruction, input, output) + EOS_TOKEN
texts.append(text)
return { "text" : texts, }
print(f"📂 Loading dataset: {DATASET_PATH}")
dataset = load_dataset("json", data_files=DATASET_PATH, split="train")
dataset = dataset.map(formatting_prompts_func, batched = True)
print("⚙️ Configuring Training Arguments...")
trainer = SFTTrainer(
model = model,
tokenizer = tokenizer,
train_dataset = dataset,
dataset_text_field = "text",
max_seq_length = MAX_SEQ_LENGTH,
args = TrainingArguments(
per_device_train_batch_size = 2,
gradient_accumulation_steps = 4,
warmup_steps = 5,
max_steps = 60,
learning_rate = 2e-4,
fp16 = not torch.cuda.is_bf16_supported(),
bf16 = torch.cuda.is_bf16_supported(),
logging_steps = 1,
optim = "adamw_8bit",
weight_decay = 0.01,
lr_scheduler_type = "linear",
seed = 3407,
output_dir = OUTPUT_DIR,
save_strategy = "no",
),
)
print("🔥 Starting training...")
trainer_stats = trainer.train()
print(f"💾 Saving LoRA adapters to {OUTPUT_DIR}...")
model.save_pretrained(OUTPUT_DIR)
tokenizer.save_pretrained(OUTPUT_DIR)
print("✅ Training complete!")
print(f"Stats: {trainer_stats}")
if __name__ == "__main__":
train()