from typing import Dict, Optional, List from dataclasses import dataclass, field import transformers from supported_models import MODEL_HF_PATH, MODEL_FAMILIES @dataclass class ModelArguments: model_id: str = field(default="llava-1.5-7b") model_local_path: Optional[str] = field(default=None) def __post_init__(self): assert self.model_id in MODEL_HF_PATH, f"Unknown model_id: {self.model_id}" self.model_hf_path: str = MODEL_HF_PATH[self.model_id] assert self.model_id in MODEL_FAMILIES, f"Unknown model_id: {self.model_id}" self.model_family_id: str = MODEL_FAMILIES[self.model_id] if not self.model_local_path: self.model_local_path = self.model_hf_path @dataclass class DataArguments: data_path: str = field( default=None, metadata={"help": "Path to the training data json file."} ) eval_data_path: Optional[str] = field( default=None, metadata={"help": "Path to the evaluation data json file."} ) image_folder: Optional[str] = field(default=None) video_folder: Optional[str] = field(default=None) num_frames: Optional[int] = field(default=8) user_key: Optional[str] = field(default="human") assistant_key: Optional[str] = field(default="gpt") @dataclass class TrainingArguments(transformers.TrainingArguments): model_max_length: int = field( default=1024, metadata={ "help": "Maximum sequence length. Sequences will be right padded (and possibly truncated)." }, ) use_flash_attn: bool = field(default=False) train_vision_encoder: bool = field(default=False) train_vision_projector: bool = field(default=False) mask_question_tokens: bool = field(default=True) def __post_init__(self): super().__post_init__() self.remove_unused_columns = False @dataclass class LoraArguments: use_lora: bool = field(default=True) use_vision_lora: bool = field(default=True) q_lora: bool = field(default=False) lora_r: int = field(default=8) lora_alpha: int = field(default=16) lora_dropout: float = field(default=0.05) lora_weight_path: str = "" lora_bias: str = "none"