From f04c684d8ac010ca39ee79e83e2848cfb801892a Mon Sep 17 00:00:00 2001 From: Manan17 Date: Tue, 3 Mar 2026 09:34:32 +0000 Subject: [PATCH] variable changes and some cleanup --- studio/backend/core/training/trainer.py | 56 +++++++++---------- studio/backend/core/training/training.py | 7 +-- studio/backend/models/datasets.py | 2 +- studio/backend/models/training.py | 2 +- studio/backend/routes/datasets.py | 4 +- studio/backend/routes/training.py | 2 +- .../backend/utils/datasets/dataset_utils.py | 56 +++++++++---------- .../utils/datasets/format_detection.py | 4 +- studio/backend/utils/hardware/hardware.py | 15 +---- .../sections/dataset-preview-dialog.tsx | 6 +- .../studio/sections/params-section.tsx | 2 +- .../studio/sections/training-section.tsx | 2 +- .../src/features/studio/studio-page.tsx | 2 +- .../src/features/training/api/mappers.ts | 2 +- .../training/hooks/use-training-actions.ts | 12 ++-- .../training/stores/training-config-store.ts | 24 ++++---- .../src/features/training/types/api.ts | 2 +- .../src/features/training/types/config.ts | 2 +- .../src/features/training/types/datasets.ts | 2 +- 19 files changed, 93 insertions(+), 111 deletions(-) diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index 8b760c7599..8208f71694 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -30,11 +30,6 @@ from trl import SFTTrainer, SFTConfig logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) -# Process-level flag: set True after CUDA-heavy audio preprocessing (Whisper/DAC/BiCodec). -# Once CUDA has been used for audio processing, fork-based multiprocessing (num_proc>1) -# deadlocks because forked children inherit CUDA's internal thread locks. -# This flag is never reset — once contaminated, the process stays contaminated. -_CUDA_AUDIO_PREPROCESSING_DONE = False @dataclass class TrainingProgress: @@ -75,6 +70,7 @@ class UnslothTrainer: self.is_audio = False self.is_audio_vlm = False # Multimodal model (e.g. Gemma 3N) trained on audio data self._audio_type = None # 'csm', 'whisper', 'snac', 'xcodec2', 'bicodec', 'dac' + self._cuda_audio_used = False # Set once after audio CUDA preprocessing; never cleared self._spark_tts_repo_dir = None # Path to downloaded Spark-TTS repo (for BiCodecTokenizer) self.model_name = None @@ -335,7 +331,7 @@ class UnslothTrainer: max_seq_length: int = 2048, load_in_4bit: bool = True, hf_token: Optional[str] = None, - is_dataset_multimodal: bool = False, + is_dataset_image: bool = False, is_dataset_audio: bool = False) -> bool: """Load model for training (supports both text and vision models)""" try: @@ -381,14 +377,14 @@ class UnslothTrainer: # Uses FastModel + SFTTrainer with audio collator self.is_audio_vlm = not self.is_audio and is_vision_model(model_name) and is_dataset_audio # VLM: vision model with image dataset (mutually exclusive with audio VLM) - self.is_vlm = not self.is_audio and not self.is_audio_vlm and is_vision_model(model_name) and is_dataset_multimodal + self.is_vlm = not self.is_audio and not self.is_audio_vlm and is_vision_model(model_name) and is_dataset_image self.model_name = model_name self.max_seq_length = max_seq_length logger.info(f"Audio type: {self._audio_type}") if not self.is_audio: logger.info(f"Model architecture is vision: {is_vision_model(model_name)}") - logger.info(f"Dataset is multimodal: {is_dataset_multimodal}, audio: {is_dataset_audio}") + logger.info(f"Dataset has images: {is_dataset_image}, audio: {is_dataset_audio}") logger.info(f"Using VLM path: {self.is_vlm}, Audio VLM: {self.is_audio_vlm}") # Reset training state for new run @@ -555,10 +551,26 @@ class UnslothTrainer: print("Model loaded successfully") return True + except OSError as e: + if "could not get source code" in str(e) and not getattr(self, '_source_code_retried', False): + # Unsloth's patching can leave stale state that makes + # inspect.getsource() fail when switching model families + # (e.g. gemma3 → gemma3n). The load always succeeds on a + # second attempt because the failed first call's partial + # imports clean up the stale state as a side effect. + self._source_code_retried = True + print(f"\n'could not get source code' — retrying once...\n") + return self.load_model(model_name, max_seq_length, load_in_4bit, hf_token, + is_dataset_image, is_dataset_audio) + logger.error(f"Error loading model: {e}") + self._update_progress(error=str(e), is_training=False) + return False except Exception as e: logger.error(f"Error loading model: {e}") self._update_progress(error=str(e), is_training=False) return False + finally: + self._source_code_retried = False def prepare_model_for_training(self, use_lora: bool = True, @@ -1204,9 +1216,7 @@ class UnslothTrainer: import gc gc.collect() torch.cuda.empty_cache() - - global _CUDA_AUDIO_PREPROCESSING_DONE - _CUDA_AUDIO_PREPROCESSING_DONE = True + self._cuda_audio_used = True if not processed_examples: raise ValueError( @@ -1394,9 +1404,7 @@ class UnslothTrainer: import gc gc.collect() torch.cuda.empty_cache() - - global _CUDA_AUDIO_PREPROCESSING_DONE - _CUDA_AUDIO_PREPROCESSING_DONE = True + self._cuda_audio_used = True if not processed_examples: raise ValueError( @@ -1579,13 +1587,7 @@ class UnslothTrainer: import gc gc.collect() torch.cuda.empty_cache() - - # Mark process as CUDA-contaminated from audio preprocessing. - # Fork-based multiprocessing (num_proc>1) will deadlock after this - # because forked children inherit CUDA's internal thread locks from - # Whisper/DAC processing that can't be released. - global _CUDA_AUDIO_PREPROCESSING_DONE - _CUDA_AUDIO_PREPROCESSING_DONE = True + self._cuda_audio_used = True if not processed_examples: raise ValueError( @@ -2236,7 +2238,7 @@ class UnslothTrainer: "output_dir": output_dir, "report_to": ["wandb"] if training_args.get('enable_wandb', False) else "none", "include_num_input_tokens_seen": True, # Enable token counting - "dataset_num_proc": safe_num_proc(max(1, os.cpu_count() // 4)), + "dataset_num_proc": 1 if (self.is_audio or self.is_audio_vlm or self._cuda_audio_used) else safe_num_proc(max(1, os.cpu_count() // 4)), "max_seq_length": training_args.get('max_seq_length', 2048), } @@ -2424,19 +2426,11 @@ class UnslothTrainer: try: from unsloth.chat_templates import train_on_responses_only - # After CUDA-heavy audio preprocessing (Whisper/DAC/BiCodec/SNAC), - # fork-based multiprocessing deadlocks because children inherit - # CUDA's internal thread locks. Use single-process mode instead. - toro_num_proc = config_args.get("dataset_num_proc", safe_num_proc(max(1, os.cpu_count() // 4))) - if _CUDA_AUDIO_PREPROCESSING_DONE: - toro_num_proc = 1 - print("Using single-process train_on_responses (CUDA audio preprocessing detected)\n") - self.trainer = train_on_responses_only( self.trainer, instruction_part=instruction_part, response_part=response_part, - num_proc=toro_num_proc, + num_proc=config_args["dataset_num_proc"], ) print("Train on responses only configured successfully\n") diff --git a/studio/backend/core/training/training.py b/studio/backend/core/training/training.py index a2020cf68c..4bff972447 100644 --- a/studio/backend/core/training/training.py +++ b/studio/backend/core/training/training.py @@ -116,7 +116,7 @@ class TrainingBackend: train_split: str = "train", eval_split: str = None, eval_steps: float = 0.00, - is_dataset_multimodal: bool = False, + is_dataset_image: bool = False, is_dataset_audio: bool = False) -> bool: """ Start training. @@ -146,9 +146,6 @@ class TrainingBackend: import torch as _torch if _torch.cuda.is_available(): _torch.cuda.synchronize() - # Reset torch dynamo/compiler caches — Unsloth's compiled SFTTrainer - # and model.for_training() set class-level state that persists between - # runs (e.g. BiCodec text trainer pollutes subsequent VLM runs). _torch._dynamo.reset() _torch.compiler.reset() import gc @@ -182,7 +179,7 @@ class TrainingBackend: 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, - is_dataset_multimodal=is_dataset_multimodal, + is_dataset_image=is_dataset_image, is_dataset_audio=is_dataset_audio, ) diff --git a/studio/backend/models/datasets.py b/studio/backend/models/datasets.py index a144cf65cb..b545c4be76 100644 --- a/studio/backend/models/datasets.py +++ b/studio/backend/models/datasets.py @@ -27,7 +27,7 @@ class CheckFormatResponse(BaseModel): requires_manual_mapping: bool detected_format: str columns: List[str] - is_multimodal: bool = False + is_image: bool = False is_audio: bool = False multimodal_columns: Optional[List[str]] = None suggested_mapping: Optional[Dict[str, str]] = None diff --git a/studio/backend/models/training.py b/studio/backend/models/training.py index 72d8b256ac..afa04f491c 100644 --- a/studio/backend/models/training.py +++ b/studio/backend/models/training.py @@ -65,7 +65,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") + is_dataset_image: bool = Field(False, description="Whether the dataset contains image data") is_dataset_audio: bool = Field(False, description="Whether the dataset contains audio data") # Logging parameters diff --git a/studio/backend/routes/datasets.py b/studio/backend/routes/datasets.py index af50b2b72e..822f6c972b 100644 --- a/studio/backend/routes/datasets.py +++ b/studio/backend/routes/datasets.py @@ -183,7 +183,7 @@ def check_format(request: CheckFormatRequest): # Run lightweight format check on the preview slice result = check_dataset_format(preview_slice, is_vlm=request.is_vlm) - logger.info(f"Format check result: requires_mapping={result['requires_manual_mapping']}, format={result['detected_format']}, is_multimodal={result.get('is_multimodal', False)}") + logger.info(f"Format check result: requires_mapping={result['requires_manual_mapping']}, format={result['detected_format']}, is_image={result.get('is_image', False)}") # Generate preview samples preview_samples = None @@ -211,7 +211,7 @@ def check_format(request: CheckFormatRequest): requires_manual_mapping=result["requires_manual_mapping"], detected_format=result["detected_format"], columns=result["columns"], - is_multimodal=result.get("is_multimodal", False), + is_image=result.get("is_image", False), is_audio=result.get("is_audio", False), multimodal_columns=result.get("multimodal_columns"), suggested_mapping=result.get("suggested_mapping"), diff --git a/studio/backend/routes/training.py b/studio/backend/routes/training.py index 1fb49f493f..27c7fad91a 100644 --- a/studio/backend/routes/training.py +++ b/studio/backend/routes/training.py @@ -178,7 +178,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, + "is_dataset_image": request.is_dataset_image, "is_dataset_audio": request.is_dataset_audio, "enable_wandb": request.enable_wandb, "wandb_token": request.wandb_token or "", diff --git a/studio/backend/utils/datasets/dataset_utils.py b/studio/backend/utils/datasets/dataset_utils.py index 6ddd2c753c..6b3cc70105 100644 --- a/studio/backend/utils/datasets/dataset_utils.py +++ b/studio/backend/utils/datasets/dataset_utils.py @@ -67,8 +67,8 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict: multimodal_info = detect_multimodal_dataset(dataset) is_audio = multimodal_info.get("is_audio", False) - if multimodal_info["is_multimodal"] and not is_audio: - is_vlm = True # Route to VLM detection for image datasets only + if multimodal_info["is_image"]: + is_vlm = True # Route to VLM detection for image datasets # Common audio fields for all return paths audio_fields = { @@ -88,7 +88,7 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict: "suggested_mapping": None, "detected_image_column": vlm_structure.get("image_column"), "detected_text_column": vlm_structure.get("text_column"), - "is_multimodal": multimodal_info["is_multimodal"], + "is_image": multimodal_info["is_image"], "multimodal_columns": multimodal_info.get("multimodal_columns"), **audio_fields, } @@ -105,7 +105,7 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict: "suggested_mapping": None, "detected_image_column": None, "detected_text_column": multimodal_info.get("detected_text_column"), - "is_multimodal": True, + "is_image": False, "multimodal_columns": multimodal_info.get("audio_columns"), **audio_fields, } @@ -124,7 +124,7 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict: "suggested_mapping": heuristic_mapping, "detected_image_column": None, "detected_text_column": None, - "is_multimodal": False, + "is_image": False, "multimodal_columns": None, **audio_fields, } @@ -136,7 +136,7 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict: "suggested_mapping": None, "detected_image_column": None, "detected_text_column": None, - "is_multimodal": False, + "is_image": False, "multimodal_columns": None, **audio_fields, } @@ -149,7 +149,7 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict: "suggested_mapping": None, "detected_image_column": None, "detected_text_column": None, - "is_multimodal": False, + "is_image": False, "multimodal_columns": None, **audio_fields, } @@ -278,7 +278,7 @@ def format_dataset( "chat_column": chat_column, "is_standardized": True, "requires_manual_mapping": False, - "is_multimodal": multimodal_info["is_multimodal"], + "is_image": multimodal_info["is_image"], "multimodal_info": multimodal_info, "warnings": [f"Applied user-provided column mapping ({format_type}): {custom_format_mapping}"] } @@ -290,7 +290,7 @@ def format_dataset( "chat_column": None, "is_standardized": False, "requires_manual_mapping": True, - "is_multimodal": multimodal_info["is_multimodal"], + "is_image": multimodal_info["is_image"], "multimodal_info": multimodal_info, "warnings": [f"Failed to apply user mapping: {e}"] } @@ -301,7 +301,7 @@ def format_dataset( warnings = [] # Add multimodal warning if detected - if multimodal_info["is_multimodal"]: + if multimodal_info["is_image"]: warnings.append( f"Multimodal dataset detected. Found columns: {multimodal_info['multimodal_columns']}" ) @@ -318,7 +318,7 @@ def format_dataset( "chat_column": None, "is_standardized": True, "requires_manual_mapping": False, - "is_multimodal": multimodal_info["is_multimodal"], + "is_image": multimodal_info["is_image"], "multimodal_info": multimodal_info, "warnings": [] } @@ -338,7 +338,7 @@ def format_dataset( "chat_column": detected["chat_column"], "is_standardized": True, "requires_manual_mapping": False, - "is_multimodal": multimodal_info["is_multimodal"], + "is_image": multimodal_info["is_image"], "multimodal_info": multimodal_info, "warnings": [] } @@ -351,7 +351,7 @@ def format_dataset( "chat_column": detected["chat_column"], "is_standardized": False, "requires_manual_mapping": True, - "is_multimodal": multimodal_info["is_multimodal"], + "is_image": multimodal_info["is_image"], "multimodal_info": multimodal_info, "warnings": warnings } @@ -364,7 +364,7 @@ def format_dataset( "chat_column": detected["chat_column"], "is_standardized": True, "requires_manual_mapping": False, - "is_multimodal": multimodal_info["is_multimodal"], + "is_image": multimodal_info["is_image"], "multimodal_info": multimodal_info, "warnings": warnings } @@ -415,7 +415,7 @@ def format_dataset( "chat_column": "conversations", "is_standardized": True, "requires_manual_mapping": False, - "is_multimodal": multimodal_info["is_multimodal"], + "is_image": multimodal_info["is_image"], "multimodal_info": multimodal_info, "warnings": warnings } @@ -438,7 +438,7 @@ def format_dataset( "chat_column": detected["chat_column"], "is_standardized": True, "requires_manual_mapping": False, - "is_multimodal": multimodal_info["is_multimodal"], + "is_image": multimodal_info["is_image"], "multimodal_info": multimodal_info, "warnings": warnings } @@ -453,7 +453,7 @@ def format_dataset( "chat_column": detected["chat_column"], "is_standardized": False, "requires_manual_mapping": True, - "is_multimodal": multimodal_info["is_multimodal"], + "is_image": multimodal_info["is_image"], "multimodal_info": multimodal_info, "warnings": warnings } @@ -469,7 +469,7 @@ def format_dataset( "chat_column": None, "is_standardized": True, "requires_manual_mapping": False, - "is_multimodal": multimodal_info["is_multimodal"], + "is_image": multimodal_info["is_image"], "multimodal_info": multimodal_info, "warnings": [] } @@ -492,7 +492,7 @@ def format_dataset( "chat_column": None, "is_standardized": True, "requires_manual_mapping": False, - "is_multimodal": multimodal_info["is_multimodal"], + "is_image": multimodal_info["is_image"], "multimodal_info": multimodal_info, "warnings": [] } @@ -506,7 +506,7 @@ def format_dataset( "chat_column": detected["chat_column"], "is_standardized": False, "requires_manual_mapping": True, - "is_multimodal": multimodal_info["is_multimodal"], + "is_image": multimodal_info["is_image"], "multimodal_info": multimodal_info, "warnings": warnings } @@ -523,7 +523,7 @@ def format_dataset( "chat_column": "conversations", "is_standardized": True, "requires_manual_mapping": False, - "is_multimodal": multimodal_info["is_multimodal"], + "is_image": multimodal_info["is_image"], "multimodal_info": multimodal_info, "warnings": [] } @@ -541,7 +541,7 @@ def format_dataset( "chat_column": detected["chat_column"], "is_standardized": True, "requires_manual_mapping": False, - "is_multimodal": multimodal_info["is_multimodal"], + "is_image": multimodal_info["is_image"], "multimodal_info": multimodal_info, "warnings": [] } @@ -554,7 +554,7 @@ def format_dataset( "chat_column": detected["chat_column"], "is_standardized": True, "requires_manual_mapping": False, - "is_multimodal": multimodal_info["is_multimodal"], + "is_image": multimodal_info["is_image"], "multimodal_info": multimodal_info, "warnings": [] } @@ -575,7 +575,7 @@ def format_dataset( "chat_column": detected["chat_column"], "is_standardized": True, "requires_manual_mapping": False, - "is_multimodal": multimodal_info["is_multimodal"], + "is_image": multimodal_info["is_image"], "multimodal_info": multimodal_info, "warnings": warnings } @@ -589,7 +589,7 @@ def format_dataset( "chat_column": detected["chat_column"], "is_standardized": False, "requires_manual_mapping": True, - "is_multimodal": multimodal_info["is_multimodal"], + "is_image": multimodal_info["is_image"], "multimodal_info": multimodal_info, "warnings": warnings } @@ -675,7 +675,7 @@ def format_and_template_dataset( "final_format": "vlm_messages", "chat_column": "messages", "is_vlm": True, - "is_multimodal": True, + "is_image": True, "multimodal_info": multimodal_info, "success": True, "requires_manual_mapping": False, @@ -797,7 +797,7 @@ def format_and_template_dataset( "final_format": "vlm_messages", "chat_column": "messages", "is_vlm": True, - "is_multimodal": multimodal_info["is_multimodal"], + "is_image": multimodal_info["is_image"], "multimodal_info": multimodal_info, "vlm_structure": vlm_structure, "success": True, @@ -826,7 +826,7 @@ def format_and_template_dataset( # Gemma emits a leading that must be stripped for text-only chatml/sharegpt. is_alpaca = format_type == "alpaca" or (format_type == "auto" and dataset_info["detected_format"] == "alpaca") is_gemma = "gemma" in model_name.lower() - if is_gemma and not dataset_info["is_multimodal"] and not is_alpaca: + if is_gemma and not dataset_info["is_image"] and not is_alpaca: remove_bos_prefix = True template_result = apply_chat_template_to_dataset( dataset_info=dataset_info, diff --git a/studio/backend/utils/datasets/format_detection.py b/studio/backend/utils/datasets/format_detection.py index 1a557b5df6..2337833bef 100644 --- a/studio/backend/utils/datasets/format_detection.py +++ b/studio/backend/utils/datasets/format_detection.py @@ -334,7 +334,7 @@ def detect_multimodal_dataset(dataset): Returns: dict: { - "is_multimodal": bool, + "is_image": bool, "multimodal_columns": list of column names containing image data, "modality_types": list of detected types (e.g., ["image", "audio"]), "is_audio": bool, @@ -427,7 +427,7 @@ def detect_multimodal_dataset(dataset): break return { - "is_multimodal": len(multimodal_columns) > 0 or is_audio, + "is_image": len(multimodal_columns) > 0, "multimodal_columns": multimodal_columns, "modality_types": list(modality_types), "is_audio": is_audio, diff --git a/studio/backend/utils/hardware/hardware.py b/studio/backend/utils/hardware/hardware.py index cb2f097b92..b885e130d5 100644 --- a/studio/backend/utils/hardware/hardware.py +++ b/studio/backend/utils/hardware/hardware.py @@ -423,14 +423,13 @@ def safe_num_proc(desired: Optional[int] = None) -> int: """ Return a safe ``num_proc`` for ``dataset.map()`` calls. - Fork-based multiprocessing deadlocks when CUDA has already been - initialized (e.g. after inference). This helper detects that case - and forces ``num_proc=1``. - On multi-GPU machines the NVIDIA driver spawns extra background threads, making ``os.fork()`` prone to deadlocks when many workers are created. This helper caps ``num_proc`` to 4 on such machines. + On single-GPU (or CPU-only) machines the original value is returned + unchanged. + Args: desired: The num_proc you *want*. If None, auto-computes from ``os.cpu_count()``. @@ -443,14 +442,6 @@ def safe_num_proc(desired: Optional[int] = None) -> int: if desired is None or not isinstance(desired, int): desired = max(1, os.cpu_count() // 3) - # After inference, CUDA is initialized — forking will deadlock. - try: - import torch - if torch.cuda.is_initialized(): - return 1 - except ImportError: - pass - if get_physical_gpu_count() > 1: capped = min(4, desired) print( diff --git a/studio/frontend/src/features/studio/sections/dataset-preview-dialog.tsx b/studio/frontend/src/features/studio/sections/dataset-preview-dialog.tsx index 7537534329..abdce2bd47 100644 --- a/studio/frontend/src/features/studio/sections/dataset-preview-dialog.tsx +++ b/studio/frontend/src/features/studio/sections/dataset-preview-dialog.tsx @@ -62,10 +62,10 @@ export function DatasetPreviewDialog({ ); const { isStarting, startError, startTrainingRun } = useTrainingActions(); - // If the backend reports multimodal data, treat as VLM even if the prop - // hasn't caught up yet (isDatasetMultimodal may still be null in the store). + // If the backend reports image data, treat as VLM even if the prop + // hasn't caught up yet (isDatasetImage may still be null in the store). const effectiveIsAudio = !!data?.is_audio; - const effectiveIsVlm = !effectiveIsAudio && (isVlm || !!data?.is_multimodal); + const effectiveIsVlm = isVlm || !!data?.is_image; const hasHeuristicMapping = !data?.requires_manual_mapping && !!data?.suggested_mapping; const mappingEnabled = !!data?.requires_manual_mapping || hasHeuristicMapping; diff --git a/studio/frontend/src/features/studio/sections/params-section.tsx b/studio/frontend/src/features/studio/sections/params-section.tsx index b7ad7aec0b..77680cfd82 100644 --- a/studio/frontend/src/features/studio/sections/params-section.tsx +++ b/studio/frontend/src/features/studio/sections/params-section.tsx @@ -114,7 +114,7 @@ function SliderRow({ export function ParamsSection(): ReactElement { const store = useTrainingConfigStore(); const isLora = store.trainingMethod !== "full"; - const showVisionLora = store.isVisionModel && store.isDatasetMultimodal === true; + const showVisionLora = store.isVisionModel && store.isDatasetImage === true; const [loraOpen, setLoraOpen] = useState(false); const [hyperOpen, setHyperOpen] = useState(false); const maxStepsSliderMax = Math.max(500, store.maxSteps, 30); diff --git a/studio/frontend/src/features/studio/sections/training-section.tsx b/studio/frontend/src/features/studio/sections/training-section.tsx index 160a0808c9..252fe5d2d2 100644 --- a/studio/frontend/src/features/studio/sections/training-section.tsx +++ b/studio/frontend/src/features/studio/sections/training-section.tsx @@ -42,7 +42,7 @@ export function TrainingSection() { const store = useTrainingConfigStore(); const { isStarting, startError, startTrainingRun } = useTrainingActions(); const isIncompatible = - !store.isVisionModel && !store.isDatasetAudio && store.isDatasetMultimodal === true; + !store.isVisionModel && store.isDatasetImage === true; const fileInputRef = useRef(null); const handleFileUpload = (e: React.ChangeEvent) => { diff --git a/studio/frontend/src/features/studio/studio-page.tsx b/studio/frontend/src/features/studio/studio-page.tsx index 75fde25e07..a0cfe5f48e 100644 --- a/studio/frontend/src/features/studio/studio-page.tsx +++ b/studio/frontend/src/features/studio/studio-page.tsx @@ -93,7 +93,7 @@ export function StudioPage(): ReactElement { datasetSplit={config.datasetSplit} mode={dialogMode} initialData={dialogInitial} - isVlm={config.isVisionModel && config.isDatasetMultimodal === true} + isVlm={config.isVisionModel && config.isDatasetImage === true} /> {canGoBack && ( diff --git a/studio/frontend/src/features/training/api/mappers.ts b/studio/frontend/src/features/training/api/mappers.ts index 923306861d..ed8142f6ad 100644 --- a/studio/frontend/src/features/training/api/mappers.ts +++ b/studio/frontend/src/features/training/api/mappers.ts @@ -57,7 +57,7 @@ export function buildTrainingStartPayload( finetune_language_layers: config.finetuneLanguageLayers, finetune_attention_modules: config.finetuneAttentionModules, finetune_mlp_modules: config.finetuneMLPModules, - is_dataset_multimodal: !!config.isDatasetMultimodal, + is_dataset_image: !!config.isDatasetImage, is_dataset_audio: config.isDatasetAudio, enable_wandb: config.enableWandb, wandb_token: config.enableWandb ? config.wandbToken.trim() || null : null, diff --git a/studio/frontend/src/features/training/hooks/use-training-actions.ts b/studio/frontend/src/features/training/hooks/use-training-actions.ts index 10f67493dc..2225857c07 100644 --- a/studio/frontend/src/features/training/hooks/use-training-actions.ts +++ b/studio/frontend/src/features/training/hooks/use-training-actions.ts @@ -36,7 +36,7 @@ export function useTrainingActions() { try { const datasetName = getDatasetName(config); - let isVlm = config.isVisionModel && config.isDatasetMultimodal === true; + let isVlm = config.isVisionModel && config.isDatasetImage === true; if (datasetName) { const check = await checkDatasetFormat({ @@ -47,17 +47,17 @@ export function useTrainingActions() { isVlm, }); - // Backend auto-detects multimodal/audio from dataset content. + // Backend auto-detects image/audio from dataset content. // Sync these flags into the store so buildTrainingStartPayload picks them up. const isAudio = !!check.is_audio; - const isMultimodal = !!check.is_multimodal; + const isImage = !!check.is_image; - if (isMultimodal && config.isVisionModel) { + if (isImage && config.isVisionModel) { isVlm = true; } - if (isMultimodal !== config.isDatasetMultimodal || isAudio !== config.isDatasetAudio) { + if (isImage !== config.isDatasetImage || isAudio !== config.isDatasetAudio) { useTrainingConfigStore.setState({ - isDatasetMultimodal: isMultimodal, + isDatasetImage: isImage, isDatasetAudio: isAudio, }); } diff --git a/studio/frontend/src/features/training/stores/training-config-store.ts b/studio/frontend/src/features/training/stores/training-config-store.ts index 4e50ed6ae3..b9aa2953ef 100644 --- a/studio/frontend/src/features/training/stores/training-config-store.ts +++ b/studio/frontend/src/features/training/stores/training-config-store.ts @@ -35,7 +35,7 @@ const initialState: TrainingConfigState = { modelDefaultsError: null, modelDefaultsAppliedFor: null, isCheckingDataset: false, - isDatasetMultimodal: null, + isDatasetImage: null, isDatasetAudio: false, ...DEFAULT_HYPERPARAMS, }; @@ -57,7 +57,7 @@ const NON_PERSISTED_STATE_KEYS: ReadonlySet = new Set "modelDefaultsError", "modelDefaultsAppliedFor", "isCheckingDataset", - "isDatasetMultimodal", + "isDatasetImage", "isDatasetAudio", "trainOnCompletions", ]); @@ -116,9 +116,9 @@ export const useTrainingConfigStore = create()( _trainOnCompletionsManuallySet = false; const patch = mapBackendModelConfigToTrainingPatch(modelDetails.config); - // If vision model + multimodal dataset already known, override + // If vision model + image dataset already known, override // trainOnCompletions to false regardless of backend default. - if (modelDetails.is_vision && get().isDatasetMultimodal === true) { + if (modelDetails.is_vision && get().isDatasetImage === true) { patch.trainOnCompletions = false; } @@ -174,16 +174,16 @@ export const useTrainingConfigStore = create()( }) .then((res) => { if (controller.signal.aborted) return; - const isMultimodal = !!res.is_multimodal; + const isImage = !!res.is_image; const isAudio = !!res.is_audio; const updates: Record = { - isDatasetMultimodal: isMultimodal, + isDatasetImage: isImage, isDatasetAudio: isAudio, isCheckingDataset: false, }; if (!_trainOnCompletionsManuallySet) { const { isVisionModel } = get(); - if (isVisionModel && isMultimodal) { + if (isVisionModel && isImage) { updates.trainOnCompletions = false; } } @@ -191,7 +191,7 @@ export const useTrainingConfigStore = create()( }) .catch(() => { if (controller.signal.aborted) return; - set({ isDatasetMultimodal: null, isCheckingDataset: false }); + set({ isDatasetImage: null, isCheckingDataset: false }); }); }; @@ -261,7 +261,7 @@ export const useTrainingConfigStore = create()( datasetSplit: null, datasetEvalSplit: null, datasetManualMapping: emptyManualMapping(), - isDatasetMultimodal: null, + isDatasetImage: null, isCheckingDataset: false, }); }, @@ -274,7 +274,7 @@ export const useTrainingConfigStore = create()( datasetSplit: null, datasetEvalSplit: null, datasetManualMapping: emptyManualMapping(), - isDatasetMultimodal: null, + isDatasetImage: null, isCheckingDataset: false, }); }, @@ -282,7 +282,7 @@ export const useTrainingConfigStore = create()( set({ datasetSplit, datasetManualMapping: emptyManualMapping(), - isDatasetMultimodal: null, + isDatasetImage: null, isCheckingDataset: false, }); @@ -298,7 +298,7 @@ export const useTrainingConfigStore = create()( ensureDatasetChecked: () => { const state = get(); if (state.isCheckingDataset) return; - if (state.isDatasetMultimodal !== null) return; + if (state.isDatasetImage !== null) return; const datasetName = state.datasetSource === "huggingface" diff --git a/studio/frontend/src/features/training/types/api.ts b/studio/frontend/src/features/training/types/api.ts index 9f4c849200..e22bd2b9ac 100644 --- a/studio/frontend/src/features/training/types/api.ts +++ b/studio/frontend/src/features/training/types/api.ts @@ -38,7 +38,7 @@ export interface TrainingStartRequest { finetune_language_layers: boolean; finetune_attention_modules: boolean; finetune_mlp_modules: boolean; - is_dataset_multimodal: boolean; + is_dataset_image: boolean; is_dataset_audio: boolean; enable_wandb: boolean; wandb_token: string | null; diff --git a/studio/frontend/src/features/training/types/config.ts b/studio/frontend/src/features/training/types/config.ts index 0cf2129f3a..c66d0ad160 100644 --- a/studio/frontend/src/features/training/types/config.ts +++ b/studio/frontend/src/features/training/types/config.ts @@ -59,7 +59,7 @@ export interface TrainingConfigState { modelDefaultsError: string | null; modelDefaultsAppliedFor: string | null; isCheckingDataset: boolean; - isDatasetMultimodal: boolean | null; + isDatasetImage: boolean | null; isDatasetAudio: boolean; finetuneVisionLayers: boolean; finetuneLanguageLayers: boolean; diff --git a/studio/frontend/src/features/training/types/datasets.ts b/studio/frontend/src/features/training/types/datasets.ts index c1880116ca..1e556a9d43 100644 --- a/studio/frontend/src/features/training/types/datasets.ts +++ b/studio/frontend/src/features/training/types/datasets.ts @@ -9,7 +9,7 @@ export type CheckFormatResponse = { detected_speaker_column?: string | null; preview_samples?: Record[] | null; total_rows?: number | null; - is_multimodal?: boolean; + is_image?: boolean; is_audio?: boolean; multimodal_columns?: string[] | null; };