finetune.py (3063B)
1 import torch 2 from datasets import load_dataset 3 from peft import LoraConfig, get_peft_model 4 from transformers import AutoModelForSeq2SeqLM, AutoTokenizer 5 from trl import SFTTrainer, SFTConfig 6 import os 7 8 #MODEL_ID = "google/flan-t5-large" 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" 12 DATASET_TRAIN_PATH = "train.jsonl" 13 DATASET_VAL_PATH = "validation.jsonl" 14 OUTPUT_DIR = "./t5gemma-s-s-ul2-it-sponsors" 15 LORA_R = 16 16 LORA_ALPHA = 32 17 LORA_DROPOUT = 0.05 18 LORA_TARGET_MODULES = [ 19 "q_proj", "k_proj", "v_proj", "o_proj", 20 "gate_proj", "up_proj", "down_proj", 21 ] 22 23 def main(): 24 tokenizer = AutoTokenizer.from_pretrained(MODEL_ID) 25 # A pad token is required for training 26 if tokenizer.pad_token is None: 27 tokenizer.pad_token = tokenizer.eos_token 28 29 model = AutoModelForSeq2SeqLM.from_pretrained( 30 MODEL_ID, 31 device_map="auto", 32 torch_dtype=torch.bfloat16, 33 attn_implementation='eager', 34 offload_folder="./offload", 35 offload_state_dict=True, 36 ) 37 model.config.use_cache = False # disable kv-cache during training 38 model.gradient_checkpointing_enable() # enable gradient checkpointing 39 40 train_dataset = load_dataset("json", data_files=DATASET_TRAIN_PATH, split="train") 41 validation_dataset = load_dataset("json", data_files=DATASET_VAL_PATH, split="train") 42 43 peft_config = LoraConfig( 44 r=LORA_R, 45 lora_alpha=LORA_ALPHA, 46 target_modules=LORA_TARGET_MODULES, 47 lora_dropout=LORA_DROPOUT, 48 bias="none", 49 task_type="SEQ_2_SEQ_LM" 50 ) 51 52 print("Applying LoRA config to the model...") 53 model = get_peft_model(model, peft_config) 54 model.print_trainable_parameters() 55 56 print("Model and LoRA config loaded. Setting up trainer...") 57 58 training_args = SFTConfig( 59 output_dir=OUTPUT_DIR, 60 61 # Training parameters 62 num_train_epochs=3, 63 per_device_eval_batch_size=1, 64 per_device_train_batch_size=1, 65 gradient_accumulation_steps=8, 66 learning_rate=1e-4, 67 lr_scheduler_type="linear", 68 optim="adamw_torch", 69 70 # Model and data parameters 71 max_length=5120, 72 dataset_text_field="text", 73 74 # Technical parameters 75 bf16=True, 76 max_grad_norm=1.0, 77 warmup_ratio=0.03, 78 79 # Logging, saving, and evaluation 80 logging_steps=25, 81 eval_strategy="epoch", 82 save_strategy="epoch", 83 load_best_model_at_end=True, 84 ) 85 training_args.eval_accumulation_steps = 4 86 87 trainer = SFTTrainer( 88 model=model, 89 processing_class=tokenizer, 90 args=training_args, 91 train_dataset=train_dataset, 92 eval_dataset=validation_dataset, 93 ) 94 95 print("Trainer initialized. Starting fine-tuning...") 96 trainer.train() 97 98 print("Training complete. Saving the best model...") 99 trainer.save_model(os.path.join(OUTPUT_DIR, "final_checkpoint")) 100 101 102 if __name__ == "__main__": 103 main()