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