summaryrefslogtreecommitdiff
path: root/finetune.py
blob: 05b558fd6b271826672d5b0f4b03ffa3a2ebb199 (plain)
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
import torch
from datasets import load_dataset
from peft import LoraConfig, get_peft_model
from transformers import AutoModelForSeq2SeqLM, AutoTokenizer
from trl import SFTTrainer, SFTConfig
import os

#MODEL_ID = "google/flan-t5-large"
#MODEL_ID = "google/t5gemma-b-b-prefixlm-it"
#MODEL_ID = "google/long-t5-tglobal-base"
MODEL_ID = "google/t5gemma-s-s-ul2-it"
DATASET_TRAIN_PATH = "train.jsonl"
DATASET_VAL_PATH = "validation.jsonl"
OUTPUT_DIR = "./t5gemma-s-s-ul2-it-sponsors"
LORA_R = 16
LORA_ALPHA = 32
LORA_DROPOUT = 0.05
LORA_TARGET_MODULES = [
    "q_proj", "k_proj", "v_proj", "o_proj",
    "gate_proj", "up_proj", "down_proj",
]

def main():
    tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
    # A pad token is required for training
    if tokenizer.pad_token is None:
        tokenizer.pad_token = tokenizer.eos_token

    model = AutoModelForSeq2SeqLM.from_pretrained(
        MODEL_ID,
        device_map="auto",
        torch_dtype=torch.bfloat16,
        attn_implementation='eager',
        offload_folder="./offload",
        offload_state_dict=True,
    )
    model.config.use_cache = False         # disable kv-cache during training
    model.gradient_checkpointing_enable()  # enable gradient checkpointing

    train_dataset = load_dataset("json", data_files=DATASET_TRAIN_PATH, split="train")
    validation_dataset = load_dataset("json", data_files=DATASET_VAL_PATH, split="train")

    peft_config = LoraConfig(
        r=LORA_R,
        lora_alpha=LORA_ALPHA,
        target_modules=LORA_TARGET_MODULES,
        lora_dropout=LORA_DROPOUT,
        bias="none",
        task_type="SEQ_2_SEQ_LM"
    )

    print("Applying LoRA config to the model...")
    model = get_peft_model(model, peft_config)
    model.print_trainable_parameters()

    print("Model and LoRA config loaded. Setting up trainer...")

    training_args = SFTConfig(
        output_dir=OUTPUT_DIR,
        
        # Training parameters
        num_train_epochs=3,
        per_device_eval_batch_size=1,
        per_device_train_batch_size=1,
        gradient_accumulation_steps=8,
        learning_rate=1e-4,
        lr_scheduler_type="linear",
        optim="adamw_torch",
        
        # Model and data parameters
        max_length=5120,
        dataset_text_field="text",
        
        # Technical parameters
        bf16=True,
        max_grad_norm=1.0,
        warmup_ratio=0.03,
        
        # Logging, saving, and evaluation
        logging_steps=25,
        eval_strategy="epoch",
        save_strategy="epoch",
        load_best_model_at_end=True,
    )
    training_args.eval_accumulation_steps = 4

    trainer = SFTTrainer(
        model=model,
        processing_class=tokenizer,
        args=training_args,
        train_dataset=train_dataset,
        eval_dataset=validation_dataset,
    )

    print("Trainer initialized. Starting fine-tuning...")
    trainer.train()

    print("Training complete. Saving the best model...")
    trainer.save_model(os.path.join(OUTPUT_DIR, "final_checkpoint"))


if __name__ == "__main__":
    main()