Guest User

Train.py

a guest
May 6th, 2025
148
0
Never
Not a member of Pastebin yet? Sign Up, it unlocks many cool features!
Python 4.79 KB | None | 0 0
  1. import wandb
  2. import os
  3. import json
  4. import torch
  5. from accelerate import Accelerator
  6. from data.dataset import create_splits
  7. from model.model_loader import load_gen_model_and_processor, load_mini_gen_model_and_processor
  8. from configs.config import TrainingConfig, get_sft_configs
  9. from trl import SFTTrainer
  10. from qwen_vl_utils import process_vision_info
  11. from transformers import Qwen2_5_VLProcessor, AutoProcessor
  12. from trl import SFTConfig
  13. from peft import get_peft_model, LoraConfig
  14.  
  15. config = TrainingConfig()
  16.  
  17. def main():
  18.     accelerator = Accelerator()
  19.  
  20.     training_args = SFTConfig(output_dir=config.output_dir,
  21.                                run_name=config.wandb_run_name,
  22.                                num_train_epochs=config.num_train_epochs,
  23.                                per_device_train_batch_size=1,  
  24.                                per_device_eval_batch_size=1,  
  25.                                gradient_accumulation_steps=8,
  26.                                gradient_checkpointing=True,
  27.                                learning_rate=config.lr,
  28.                                lr_scheduler_type="constant",
  29.                                logging_steps=10,
  30.                                eval_steps=10,
  31.                                eval_strategy="steps",
  32.                                save_strategy="steps",
  33.                                save_steps=20,
  34.                                metric_for_best_model="eval_loss",
  35.                                greater_is_better=False,
  36.                                load_best_model_at_end=True,
  37.                                fp16=True,
  38.                                bf16 = False,                      
  39.                                max_grad_norm=config.max_grad_norm,
  40.                                warmup_ratio=config.warmup_ratio,
  41.                                push_to_hub=False,
  42.                                report_to="wandb",
  43.                                gradient_checkpointing_kwargs={"use_reentrant": False},
  44.                                dataset_kwargs={"skip_prepare_dataset": True},
  45.                                deepspeed="configs/ds_config.json")  
  46.  
  47.     wandb.init(
  48.         project=config.wandb_project,
  49.         name=config.wandb_run_name,
  50.         config=config
  51.     )
  52.     model, processor = load_gen_model_and_processor(config)
  53.     model.config.use_cache = False
  54.  
  55.     # collects data from the dataset and prepares labels (predictors) for the model to
  56.     # compute loss over the assistant's response only
  57.  
  58.     def collate_fn(samples):
  59.         """each example is a dictionary of system, user, labels and image inputs like
  60.        [
  61.            {'role': 'system', 'content': [...]},
  62.            {'role': 'user',   'content': [...]},
  63.            {'role': 'assistant','content': [...]}
  64.        ]"""
  65.  
  66.         prompts = [processor.apply_chat_template(sample, tokenize=False) for sample in samples]
  67.  
  68.         # process vision inputs (returns tuple, so get the image tensor)
  69.         image_inputs = [process_vision_info(sample)[0] for sample in samples]
  70.  
  71.         batch = processor(
  72.             text=prompts,
  73.             images=image_inputs,
  74.             return_tensors="pt",
  75.             padding=True
  76.         )
  77.  
  78.         labels = batch["input_ids"].clone()
  79.         labels[labels == processor.tokenizer.pad_token_id] = -100
  80.  
  81.         # qwen-specific image tokens
  82.         if isinstance(processor, AutoProcessor):
  83.             image_tokens = [151652, 151653, 151655]  
  84.         else:
  85.             image_tokens = [processor.tokenizer.convert_tokens_to_ids(processor.image_token)]
  86.  
  87.         for image_token_id in image_tokens:
  88.             labels[labels == image_token_id] = -100
  89.  
  90.         batch["labels"] = labels
  91.         return batch
  92.  
  93.     train_dataset, eval_dataset, test_dataset = create_splits(config.json_path, config.image_dir, config.train, config.val, config.test)
  94.  
  95.     output_test_dir = os.path.join(config.output_dir, "test")
  96.     os.makedirs(output_test_dir, exist_ok=True)
  97.  
  98.     test_data = list(test_dataset)
  99.  
  100.     test_file_path = os.path.join(output_test_dir, "test_data.json")
  101.     with open(test_file_path, "w") as f:
  102.         json.dump(test_data, f, indent=2)
  103.  
  104.     peft_config = LoraConfig(
  105.             lora_alpha=config.lora_alpha,
  106.             lora_dropout=config.lora_dropout,
  107.             r=config.lora_r,
  108.             bias="none",
  109.             target_modules=["q_proj", "v_proj"],
  110.             task_type="CAUSAL_LM",
  111.         )
  112.  
  113.     trainer = SFTTrainer(
  114.         model=model,
  115.         args=training_args,
  116.         train_dataset=train_dataset,
  117.         eval_dataset=eval_dataset,
  118.         data_collator=collate_fn,  
  119.         peft_config=peft_config,
  120.         processing_class=processor.tokenizer,
  121.     )
  122.  
  123.     trainer.train()
  124.  
  125.     trainer.save_model(config.output_dir)
  126.  
  127.  
  128.  
  129.  
  130. if __name__ == "__main__":
  131.     main()
  132.  
Advertisement
Add Comment
Please, Sign In to add comment