summaryrefslogtreecommitdiff
path: root/finetune.py
diff options
context:
space:
mode:
Diffstat (limited to 'finetune.py')
-rw-r--r--finetune.py22
1 files changed, 13 insertions, 9 deletions
diff --git a/finetune.py b/finetune.py
index bca7ac1..05b558f 100644
--- a/finetune.py
+++ b/finetune.py
@@ -1,15 +1,17 @@
1import torch 1import torch
2from datasets import load_dataset 2from datasets import load_dataset
3from peft import LoraConfig 3from peft import LoraConfig, get_peft_model
4from transformers import AutoModelForCausalLM, AutoModelForSeq2SeqLM, AutoTokenizer 4from transformers import AutoModelForSeq2SeqLM, AutoTokenizer
5from trl import SFTTrainer, SFTConfig 5from trl import SFTTrainer, SFTConfig
6import os 6import os
7 7
8#MODEL_ID = "google/gemma-3-270m-it" 8#MODEL_ID = "google/flan-t5-large"
9MODEL_ID = "google/gemma-3-1b-it" 9#MODEL_ID = "google/t5gemma-b-b-prefixlm-it"
10#MODEL_ID = "google/long-t5-tglobal-base"
11MODEL_ID = "google/t5gemma-s-s-ul2-it"
10DATASET_TRAIN_PATH = "train.jsonl" 12DATASET_TRAIN_PATH = "train.jsonl"
11DATASET_VAL_PATH = "validation.jsonl" 13DATASET_VAL_PATH = "validation.jsonl"
12OUTPUT_DIR = "./gemma3-1b-it-sponsors" 14OUTPUT_DIR = "./t5gemma-s-s-ul2-it-sponsors"
13LORA_R = 16 15LORA_R = 16
14LORA_ALPHA = 32 16LORA_ALPHA = 32
15LORA_DROPOUT = 0.05 17LORA_DROPOUT = 0.05
@@ -24,7 +26,7 @@ def main():
24 if tokenizer.pad_token is None: 26 if tokenizer.pad_token is None:
25 tokenizer.pad_token = tokenizer.eos_token 27 tokenizer.pad_token = tokenizer.eos_token
26 28
27 model = AutoModelForCausalLM.from_pretrained( 29 model = AutoModelForSeq2SeqLM.from_pretrained(
28 MODEL_ID, 30 MODEL_ID,
29 device_map="auto", 31 device_map="auto",
30 torch_dtype=torch.bfloat16, 32 torch_dtype=torch.bfloat16,
@@ -44,9 +46,13 @@ def main():
44 target_modules=LORA_TARGET_MODULES, 46 target_modules=LORA_TARGET_MODULES,
45 lora_dropout=LORA_DROPOUT, 47 lora_dropout=LORA_DROPOUT,
46 bias="none", 48 bias="none",
47 task_type="CAUSAL_LM" 49 task_type="SEQ_2_SEQ_LM"
48 ) 50 )
49 51
52 print("Applying LoRA config to the model...")
53 model = get_peft_model(model, peft_config)
54 model.print_trainable_parameters()
55
50 print("Model and LoRA config loaded. Setting up trainer...") 56 print("Model and LoRA config loaded. Setting up trainer...")
51 57
52 training_args = SFTConfig( 58 training_args = SFTConfig(
@@ -57,7 +63,6 @@ def main():
57 per_device_eval_batch_size=1, 63 per_device_eval_batch_size=1,
58 per_device_train_batch_size=1, 64 per_device_train_batch_size=1,
59 gradient_accumulation_steps=8, 65 gradient_accumulation_steps=8,
60 gradient_checkpointing=True,
61 learning_rate=1e-4, 66 learning_rate=1e-4,
62 lr_scheduler_type="linear", 67 lr_scheduler_type="linear",
63 optim="adamw_torch", 68 optim="adamw_torch",
@@ -83,7 +88,6 @@ def main():
83 model=model, 88 model=model,
84 processing_class=tokenizer, 89 processing_class=tokenizer,
85 args=training_args, 90 args=training_args,
86 peft_config=peft_config,
87 train_dataset=train_dataset, 91 train_dataset=train_dataset,
88 eval_dataset=validation_dataset, 92 eval_dataset=validation_dataset,
89 ) 93 )