gemma4-dev-agent / scripts /train_lora.py
EzioDevio's picture
Upload folder using huggingface_hub
c85c557 verified
Raw
History Blame Contribute Delete
4.12 kB
#!/usr/bin/env python3
"""
LoRA Fine-Tuning Script for Gemma Developer Agent
Optimized for code generation, tool calling, and multi-step reasoning.
"""
import os
import torch
from datasets import load_dataset
from transformers import (
AutoModelForCausalLM,
AutoTokenizer,
BitsAndBytesConfig,
TrainingArguments,
)
from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training
from trl import SFTTrainer
# Configuration Constants
MODEL_ID = os.getenv("MODEL_ID", "google/gemma-2-2b-it") # Adjust to your Gemma target model
DATASET_PATH = os.getenv("DATASET_PATH", "data/agent_instructions.jsonl")
OUTPUT_DIR = os.getenv("OUTPUT_DIR", "./outputs/gemma-lora-agent")
def setup_model_and_tokenizer(model_id: str):
"""Loads tokenizer and model with QLoRA 4-bit quantization for efficient training."""
print(f"Loading model and tokenizer for {model_id}...")
tokenizer = AutoTokenizer.from_pretrained(model_id, trust_remote_code=True)
tokenizer.pad_token = tokenizer.eos_token
tokenizer.padding_side = "right" # Necessary for training
# 4-bit quantization config (QLoRA) to fit on standard GPUs
bnb_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16,
bnb_4bit_use_double_quant=True,
)
model = AutoModelForCausalLM.from_pretrained(
model_id,
quantization_config=bnb_config,
device_map="auto",
trust_remote_code=True,
)
# Prepare model for k-bit training
model = prepare_model_for_kbit_training(model)
return model, tokenizer
def get_peft_config():
"""Configures LoRA target modules for comprehensive linear layer adaptation."""
return LoraConfig(
r=32, # Rank dimension
lora_alpha=64, # Scaling parameter
target_modules=[
"q_proj", "k_proj", "v_proj", "o_proj",
"gate_proj", "up_proj", "down_proj"
],
lora_dropout=0.05,
bias="none",
task_type="CAUSAL_LM",
)
def main():
# 1. Load Model & Tokenizer
model, tokenizer = setup_model_and_tokenizer(MODEL_ID)
peft_config = get_peft_config()
# 2. Load Dataset
print(f"Loading training dataset from {DATASET_PATH}...")
if os.path.exists(DATASET_PATH):
dataset = load_dataset("json", data_files=DATASET_PATH, split="train")
else:
print(f"Warning: {DATASET_PATH} not found. Loading dummy dataset for demonstration.")
from datasets import Dataset
dataset = Dataset.from_dict({
"text": [
"<bos><start_of_turn>user\nRefactor file.py to add type hints.<end_of_turn>\n<start_of_turn>model\n```python\n# refactored code\n```<end_of_turn><eos>"
] * 10
})
# 3. Training Arguments
training_args = TrainingArguments(
output_dir=OUTPUT_DIR,
per_device_train_batch_size=2,
gradient_accumulation_steps=4,
learning_rate=2e-4,
logging_steps=10,
num_train_epochs=3,
max_grad_norm=0.3,
warmup_ratio=0.03,
fp16=not torch.cuda.is_bf16_supported(),
bf16=torch.cuda.is_bf16_supported(),
optim="paged_adamw_8bit",
save_strategy="epoch",
evaluation_strategy="no",
report_to="none",
)
# 4. Supervised Fine-Tuning Trainer (TRL)
trainer = SFTTrainer(
model=model,
train_dataset=dataset,
peft_config=peft_config,
dataset_text_field="text",
max_seq_length=2048,
tokenizer=tokenizer,
args=training_args,
)
print("Starting LoRA fine-tuning...")
trainer.train()
# 5. Save Adapter Weights
print(f"Saving LoRA adapter weights to {OUTPUT_DIR}/final_adapter")
trainer.model.save_pretrained(os.path.join(OUTPUT_DIR, "final_adapter"))
tokenizer.save_pretrained(os.path.join(OUTPUT_DIR, "final_adapter"))
print("Training complete!")
if __name__ == "__main__":
main()