diff --git a/studio/backend/utils/datasets/dataset_utils.py b/studio/backend/utils/datasets/dataset_utils.py index 9d15b86ca1..7703488878 100644 --- a/studio/backend/utils/datasets/dataset_utils.py +++ b/studio/backend/utils/datasets/dataset_utils.py @@ -80,9 +80,6 @@ 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_image"]: - is_vlm = True # Route to VLM detection for image datasets - # Common audio fields for all return paths audio_fields = { "is_audio": is_audio, @@ -153,8 +150,8 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict: "suggested_mapping": heuristic_mapping, "detected_image_column": None, "detected_text_column": None, - "is_image": False, - "multimodal_columns": None, + "is_image": multimodal_info["is_image"], + "multimodal_columns": multimodal_info.get("multimodal_columns"), **audio_fields, } else: @@ -166,8 +163,8 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict: "suggested_mapping": None, "detected_image_column": None, "detected_text_column": None, - "is_image": False, - "multimodal_columns": None, + "is_image": multimodal_info["is_image"], + "multimodal_columns": multimodal_info.get("multimodal_columns"), "warning": ( f"Could not auto-detect column roles for columns: {columns}. " "Please assign roles manually, or use AI Assist." @@ -183,8 +180,8 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict: "suggested_mapping": None, "detected_image_column": None, "detected_text_column": None, - "is_image": False, - "multimodal_columns": None, + "is_image": multimodal_info["is_image"], + "multimodal_columns": multimodal_info.get("multimodal_columns"), **audio_fields, } diff --git a/studio/frontend/src/features/studio/sections/training-section.tsx b/studio/frontend/src/features/studio/sections/training-section.tsx index d1cec00e22..a47654d290 100644 --- a/studio/frontend/src/features/studio/sections/training-section.tsx +++ b/studio/frontend/src/features/studio/sections/training-section.tsx @@ -46,7 +46,8 @@ export function TrainingSection() { const store = useTrainingConfigStore(); const { isStarting, startError, startTrainingRun } = useTrainingActions(); const isIncompatible = - !store.isVisionModel && store.isDatasetImage === true; + (!store.isVisionModel && store.isDatasetImage === true) || + (!store.isAudioModel && store.isDatasetAudio === true); const configValidation = validateTrainingConfig(store); const fileInputRef = useRef(null); @@ -155,10 +156,10 @@ export function TrainingSection() { data-tour="studio-start" className="w-full cursor-pointer bg-gradient-to-r from-emerald-500 to-teal-500 text-white hover:from-emerald-600 hover:to-teal-600" onClick={() => void startTrainingRun()} - disabled={isStarting || isIncompatible || !configValidation.ok} + disabled={isStarting || isIncompatible || store.isCheckingDataset || !configValidation.ok} > - {isStarting ? "Starting..." : "Start Training"} + {isStarting ? "Starting..." : store.isCheckingDataset ? "Checking dataset..." : "Start Training"} {startError && (

{startError}

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 212dc78e6c..e1b8d4b1a2 100644 --- a/studio/frontend/src/features/training/stores/training-config-store.ts +++ b/studio/frontend/src/features/training/stores/training-config-store.ts @@ -210,6 +210,7 @@ export const useTrainingConfigStore = create()( hfToken: state.hfToken.trim() || null, subset: state.datasetSubset, split, + isVlm: state.isVisionModel, }) .then((res) => { if (controller.signal.aborted) return;