File size: 4,120 Bytes
c85c557
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
#!/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()