diff options
Diffstat (limited to 'finetune.py')
| -rw-r--r-- | finetune.py | 22 |
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 @@ | |||
| 1 | import torch | 1 | import torch |
| 2 | from datasets import load_dataset | 2 | from datasets import load_dataset |
| 3 | from peft import LoraConfig | 3 | from peft import LoraConfig, get_peft_model |
| 4 | from transformers import AutoModelForCausalLM, AutoModelForSeq2SeqLM, AutoTokenizer | 4 | from transformers import AutoModelForSeq2SeqLM, AutoTokenizer |
| 5 | from trl import SFTTrainer, SFTConfig | 5 | from trl import SFTTrainer, SFTConfig |
| 6 | import os | 6 | import os |
| 7 | 7 | ||
| 8 | #MODEL_ID = "google/gemma-3-270m-it" | 8 | #MODEL_ID = "google/flan-t5-large" |
| 9 | MODEL_ID = "google/gemma-3-1b-it" | 9 | #MODEL_ID = "google/t5gemma-b-b-prefixlm-it" |
| 10 | #MODEL_ID = "google/long-t5-tglobal-base" | ||
| 11 | MODEL_ID = "google/t5gemma-s-s-ul2-it" | ||
| 10 | DATASET_TRAIN_PATH = "train.jsonl" | 12 | DATASET_TRAIN_PATH = "train.jsonl" |
| 11 | DATASET_VAL_PATH = "validation.jsonl" | 13 | DATASET_VAL_PATH = "validation.jsonl" |
| 12 | OUTPUT_DIR = "./gemma3-1b-it-sponsors" | 14 | OUTPUT_DIR = "./t5gemma-s-s-ul2-it-sponsors" |
| 13 | LORA_R = 16 | 15 | LORA_R = 16 |
| 14 | LORA_ALPHA = 32 | 16 | LORA_ALPHA = 32 |
| 15 | LORA_DROPOUT = 0.05 | 17 | LORA_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 | ) |
