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()
|