diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index b3116c5c04..8a164193ea 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -110,17 +110,21 @@ class UnslothTrainer: model_name: str, max_seq_length: int = 2048, load_in_4bit: bool = True, - hf_token: Optional[str] = None) -> bool: + hf_token: Optional[str] = None, + is_dataset_multimodal: bool = False) -> bool: """Load model for training (supports both text and vision models)""" try: print("\nClearing GPU memory before training...") clear_gpu_cache() - # Detect if this is a vision model first - self.is_vlm = is_vision_model(model_name) + # Detect if this is a vision model AND dataset is multimodal + # A vision-capable model with a text-only dataset should use FastLanguageModel + self.is_vlm = is_vision_model(model_name) and is_dataset_multimodal self.model_name = model_name - logger.info(f"Model type detected: {'Vision' if self.is_vlm else 'Text'}") + logger.info(f"Model architecture is vision: {is_vision_model(model_name)}") + logger.info(f"Dataset is multimodal: {is_dataset_multimodal}") + logger.info(f"Using VLM path: {self.is_vlm}") # Reset training state for new run self._update_progress( diff --git a/studio/backend/core/training/training.py b/studio/backend/core/training/training.py index d9aaa8ca0e..251b246f35 100644 --- a/studio/backend/core/training/training.py +++ b/studio/backend/core/training/training.py @@ -96,7 +96,8 @@ class TrainingBackend: # Optional parameters custom_format_mapping: dict = None, subset: str = None, - split: str = "train") -> bool: + split: str = "train", + is_dataset_multimodal: bool = False) -> bool: """ Start training. @@ -150,7 +151,8 @@ class TrainingBackend: model_name=model_name, max_seq_length=max_seq_length, load_in_4bit=load_in_4bit if use_lora_actual else False, # Only 4bit for LoRA - hf_token=hf_token if hf_token.strip() else None + hf_token=hf_token if hf_token.strip() else None, + is_dataset_multimodal=is_dataset_multimodal, ) if not success or self.trainer.should_stop: diff --git a/studio/backend/models/training.py b/studio/backend/models/training.py index e0839ae485..dcc6aff0ae 100644 --- a/studio/backend/models/training.py +++ b/studio/backend/models/training.py @@ -55,6 +55,7 @@ class TrainingStartRequest(BaseModel): finetune_language_layers: bool = Field(False, description="Finetune language layers") finetune_attention_modules: bool = Field(False, description="Finetune attention modules") finetune_mlp_modules: bool = Field(False, description="Finetune MLP modules") + is_dataset_multimodal: bool = Field(False, description="Whether the dataset contains multimodal (image) data") # Logging parameters enable_wandb: bool = Field(False, description="Enable Weights & Biases logging") diff --git a/studio/backend/routes/training.py b/studio/backend/routes/training.py index 5bf59d85d2..13f7135ba4 100644 --- a/studio/backend/routes/training.py +++ b/studio/backend/routes/training.py @@ -176,6 +176,7 @@ async def start_training( "finetune_language_layers": request.finetune_language_layers, "finetune_attention_modules": request.finetune_attention_modules, "finetune_mlp_modules": request.finetune_mlp_modules, + "is_dataset_multimodal": request.is_dataset_multimodal, "enable_wandb": request.enable_wandb, "wandb_token": request.wandb_token or "", "wandb_project": request.wandb_project or "", diff --git a/studio/frontend/src/features/studio/sections/params-section.tsx b/studio/frontend/src/features/studio/sections/params-section.tsx index 0b85cafe7e..9a3495120d 100644 --- a/studio/frontend/src/features/studio/sections/params-section.tsx +++ b/studio/frontend/src/features/studio/sections/params-section.tsx @@ -109,7 +109,7 @@ function SliderRow({ export function ParamsSection(): ReactElement { const store = useTrainingConfigStore(); const isLora = store.trainingMethod !== "full"; - const isVision = store.isVisionModel; + const showVisionLora = store.isVisionModel && store.isDatasetMultimodal === true; const [loraOpen, setLoraOpen] = useState(false); const [hyperOpen, setHyperOpen] = useState(false); @@ -350,7 +350,7 @@ export function ParamsSection(): ReactElement { /> {/* Vision checkboxes */} - {isVision && ( + {showVisionLora && (
{startError}
)} + {isIncompatible && ( ++ Text model is not compatible with a multimodal dataset. Switch to a vision model or choose a text-only dataset. +
+ )} {/* Save / Clear */}