podcast-sponsor-remove

Attempt at identify sponsored segments in audio transcripts and removing them.
Log | Files | Refs

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