From 90332924debb7d5ce91de0ab9b3471cd45f8e428 Mon Sep 17 00:00:00 2001 From: Manan17 Date: Thu, 26 Feb 2026 22:17:25 +0000 Subject: [PATCH] Changes with audio training --- studio/backend/core/training/trainer.py | 207 +++++++++++------- studio/backend/models/datasets.py | 3 + studio/backend/routes/datasets.py | 3 + .../backend/utils/datasets/dataset_utils.py | 112 ++++++---- .../utils/datasets/format_detection.py | 121 ++++++++-- .../dataset-preview-dialog-mapping.tsx | 34 ++- .../sections/dataset-preview-dialog.tsx | 10 +- .../training/hooks/use-training-actions.ts | 30 ++- .../features/training/lib/model-defaults.ts | 4 - .../training/stores/training-config-store.ts | 2 + .../src/features/training/types/datasets.ts | 3 + 11 files changed, 371 insertions(+), 158 deletions(-) diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index 75b6aac0ef..d77effc0d4 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -111,6 +111,41 @@ class UnslothTrainer: except Exception as e: logger.error(f"Error in progress callback: {e}") + def _resolve_audio_columns(self, dataset, custom_format_mapping: dict = None): + """Resolve audio, text, and speaker columns from user mapping or hardcoded fallback. + + Returns: + dict with keys: audio_col, text_col, speaker_col (speaker_col may be None) + """ + cols = dataset.column_names + + if custom_format_mapping: + audio_col = None + text_col = None + speaker_col = None + for col, role in custom_format_mapping.items(): + if role == "audio": + audio_col = col + elif role == "text": + text_col = col + elif role == "speaker_id": + speaker_col = col + # Use mapping if both required columns exist in the dataset + if audio_col and audio_col in cols and text_col and text_col in cols: + return {"audio_col": audio_col, "text_col": text_col, "speaker_col": speaker_col} + + # Hardcoded fallback (existing behavior) + audio_col = next((c for c in cols if c.lower() in ("audio", "speech")), None) + text_col = next((c for c in cols if c.lower() in ("text", "sentence", "transcript", "transcription")), None) + + speaker_col = None + if "source" in cols: + speaker_col = "source" + elif "speaker_id" in cols: + speaker_col = "speaker_id" + + return {"audio_col": audio_col, "text_col": text_col, "speaker_col": speaker_col} + def _resolve_audio_type(self, model_name: str) -> Optional[str]: """Resolve audio_type from YAML model config. Returns None for non-audio models.""" try: @@ -150,12 +185,11 @@ class UnslothTrainer: self._audio_type = self._resolve_audio_type(model_name) self.is_audio = self._audio_type is not None - # 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 = not self.is_audio and is_vision_model(model_name) and is_dataset_multimodal # Audio VLM: multimodal model (e.g. Gemma 3N) trained on audio data - # Uses FastModel + SFTTrainer with audio collator (same pattern as VLM) + # 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.model_name = model_name logger.info(f"Audio type: {self._audio_type}") @@ -693,7 +727,7 @@ class UnslothTrainer: CsmForConditionalGeneration.forward = _fixed_csm_forward print("Applied CSM forward fix (class + instance level)\n") - def _preprocess_csm_dataset(self, dataset): + def _preprocess_csm_dataset(self, dataset, custom_format_mapping=None): """Preprocess dataset for CSM TTS training (exact notebook copy).""" from transformers import AutoProcessor from datasets import Audio @@ -701,22 +735,20 @@ class UnslothTrainer: processor = AutoProcessor.from_pretrained(self.model_name) - # Resolve speaker key - speaker_key = "source" - if "source" not in dataset.column_names and "speaker_id" not in dataset.column_names: - print("No speaker found, adding default 'source' of 0 for all examples\n") - dataset = dataset.add_column("source", ["0"] * len(dataset)) - elif "source" not in dataset.column_names and "speaker_id" in dataset.column_names: - speaker_key = "speaker_id" - - # Resolve audio and text columns - audio_col = next((c for c in dataset.column_names if c in ("audio", "Audio")), None) - text_col = next((c for c in dataset.column_names if c in ("text", "sentence", "transcript")), None) + # Resolve columns from user mapping or hardcoded fallback + resolved = self._resolve_audio_columns(dataset, custom_format_mapping) + audio_col = resolved["audio_col"] + text_col = resolved["text_col"] + speaker_key = resolved["speaker_col"] if audio_col is None: raise ValueError(f"No audio column found in dataset. Columns: {dataset.column_names}") if text_col is None: raise ValueError(f"No text column found in dataset. Columns: {dataset.column_names}") + if speaker_key is None: + print("No speaker found, adding default 'source' of 0 for all examples\n") + dataset = dataset.add_column("source", ["0"] * len(dataset)) + speaker_key = "source" print(f"CSM preprocessing: audio_col='{audio_col}', text_col='{text_col}', speaker_key='{speaker_key}'\n") @@ -773,7 +805,7 @@ class UnslothTrainer: print(f"CSM preprocessing complete: {len(processed)} examples\n") return processed - def _format_audio_vlm_dataset(self, dataset): + def _format_audio_vlm_dataset(self, dataset, custom_format_mapping=None): """Format dataset as audio chat messages for multimodal models (e.g. Gemma 3N). Expects columns: audio (Audio), text (str). @@ -781,15 +813,17 @@ class UnslothTrainer: """ from datasets import Audio - # Detect audio and text columns - cols = dataset.column_names - audio_col = next((c for c in cols if c.lower() in ("audio", "speech")), None) - text_col = next((c for c in cols if c.lower() in ("text", "sentence", "transcript", "transcription")), None) + resolved = self._resolve_audio_columns(dataset, custom_format_mapping) + audio_col = resolved["audio_col"] + text_col = resolved["text_col"] if not audio_col or not text_col: raise ValueError( - f"Audio VLM dataset needs 'audio' and 'text' columns, got: {cols}" + f"Audio VLM dataset needs 'audio' and 'text' columns, got: {dataset.column_names}" ) + # Store resolved audio column name for the collator closure + self._audio_vlm_audio_col = audio_col + # Cast audio to 16kHz (standard for speech models) dataset = dataset.cast_column(audio_col, Audio(sampling_rate=16000)) @@ -818,7 +852,7 @@ class UnslothTrainer: print(f"Audio VLM dataset formatted: {len(dataset)} examples\n") return dataset - def _preprocess_snac_dataset(self, dataset): + def _preprocess_snac_dataset(self, dataset, custom_format_mapping=None): """Preprocess dataset for Orpheus TTS training with SNAC codec. Mirrors Orpheus_(3B)-TTS.ipynb: encode audio with SNAC (24kHz, 3 hierarchical @@ -844,14 +878,14 @@ class UnslothTrainer: END_OF_TEXT = 128009 AUDIO_OFFSET = 128266 - # Resolve audio and text columns (reuse CSM pattern) - cols = dataset.column_names - audio_col = next((c for c in cols if c.lower() in ("audio", "speech")), None) - text_col = next((c for c in cols if c.lower() in ("text", "sentence", "transcript", "transcription")), None) - has_source = "source" in cols + resolved = self._resolve_audio_columns(dataset, custom_format_mapping) + audio_col = resolved["audio_col"] + text_col = resolved["text_col"] + speaker_col = resolved["speaker_col"] + has_source = speaker_col is not None if not audio_col or not text_col: raise ValueError( - f"SNAC dataset needs 'audio' and 'text' columns, got: {cols}" + f"SNAC dataset needs 'audio' and 'text' columns, got: {dataset.column_names}" ) # Get dataset sample rate from first example @@ -923,7 +957,7 @@ class UnslothTrainer: all_codes = deduped # --- Build text tokens (notebook lines 217-224) --- - text_prompt = f"{example['source']}: {text}" if has_source and example.get("source") else text + text_prompt = f"{example[speaker_col]}: {text}" if has_source and example.get(speaker_col) else text text_ids = tokenizer.encode(text_prompt, add_special_tokens=True) text_ids.append(END_OF_TEXT) @@ -979,7 +1013,7 @@ class UnslothTrainer: f"({skipped} skipped)\n") return result_dataset - def _preprocess_bicodec_dataset(self, dataset): + def _preprocess_bicodec_dataset(self, dataset, custom_format_mapping=None): """Preprocess dataset for Spark-TTS training with BiCodec tokenizer. Mirrors Spark_TTS_(0_5B).ipynb: encode audio with BiCodec (semantic + global tokens), @@ -1013,13 +1047,14 @@ class UnslothTrainer: from sparktts.utils.audio import audio_volume_normalize # Resolve audio and text columns - cols = dataset.column_names - audio_col = next((c for c in cols if c.lower() in ("audio", "speech")), None) - text_col = next((c for c in cols if c.lower() in ("text", "sentence", "transcript", "transcription")), None) - has_source = "source" in cols + resolved = self._resolve_audio_columns(dataset, custom_format_mapping) + audio_col = resolved["audio_col"] + text_col = resolved["text_col"] + speaker_col = resolved["speaker_col"] + has_source = speaker_col is not None if not audio_col or not text_col: raise ValueError( - f"BiCodec dataset needs 'audio' and 'text' columns, got: {cols}" + f"BiCodec dataset needs 'audio' and 'text' columns, got: {dataset.column_names}" ) # Load BiCodec tokenizer @@ -1117,7 +1152,7 @@ class UnslothTrainer: ) # Format text with source prefix if available - text_content = f"{example['source']}: {text}" if has_source and example.get("source") else text + text_content = f"{example[speaker_col]}: {text}" if has_source and example.get(speaker_col) else text formatted = "".join([ "<|task_tts|>", @@ -1166,7 +1201,7 @@ class UnslothTrainer: print(f"Sample text length: {len(sample)} chars\n") return result_dataset - def _preprocess_whisper_dataset(self, dataset, eval_split=None): + def _preprocess_whisper_dataset(self, dataset, eval_split=None, custom_format_mapping=None): """Preprocess dataset for Whisper speech-to-text training. Mirrors Whisper.ipynb: extract audio features with Whisper's feature @@ -1177,13 +1212,12 @@ class UnslothTrainer: WHISPER_SAMPLE_RATE = 16000 - # Resolve audio and text columns - cols = dataset.column_names - audio_col = next((c for c in cols if c.lower() in ("audio", "speech")), None) - text_col = next((c for c in cols if c.lower() in ("text", "sentence", "transcript", "transcription")), None) + resolved = self._resolve_audio_columns(dataset, custom_format_mapping) + audio_col = resolved["audio_col"] + text_col = resolved["text_col"] if not audio_col or not text_col: raise ValueError( - f"Whisper dataset needs 'audio' and 'text' columns, got: {cols}" + f"Whisper dataset needs 'audio' and 'text' columns, got: {dataset.column_names}" ) # Cast audio to 16kHz (Whisper's expected sample rate) @@ -1355,20 +1389,21 @@ class UnslothTrainer: # ========== AUDIO MODELS: custom preprocessing ========== if self._audio_type == 'csm': - processed = self._preprocess_csm_dataset(dataset) - # CSM returns a ready-to-train Dataset (not a dict) with no eval + processed = self._preprocess_csm_dataset(dataset, custom_format_mapping) return (processed, None) elif self._audio_type == 'whisper': - train_data, eval_data = self._preprocess_whisper_dataset(dataset, eval_split=eval_split) + train_data, eval_data = self._preprocess_whisper_dataset( + dataset, eval_split=eval_split, custom_format_mapping=custom_format_mapping + ) return (train_data, eval_data) elif self._audio_type == 'snac': - processed = self._preprocess_snac_dataset(dataset) + processed = self._preprocess_snac_dataset(dataset, custom_format_mapping) return (processed, None) elif self._audio_type == 'bicodec': - processed = self._preprocess_bicodec_dataset(dataset) + processed = self._preprocess_bicodec_dataset(dataset, custom_format_mapping) return (processed, None) elif self._audio_type in ('xcodec2', 'dac'): @@ -1376,8 +1411,7 @@ class UnslothTrainer: raise NotImplementedError(f"Audio dataset preprocessing for '{self._audio_type}' not yet implemented") elif self.is_audio_vlm: - # Audio VLM (e.g. Gemma 3N): format as chat messages with audio content - formatted = self._format_audio_vlm_dataset(dataset) + formatted = self._format_audio_vlm_dataset(dataset, custom_format_mapping) return (formatted, None) # ========== FORMAT FIRST ========== @@ -1438,7 +1472,7 @@ class UnslothTrainer: from datasets import get_dataset_split_names load_kwargs = {"path": dataset_source} if subset: - load_kwargs["name"] = subset + load_kwargs["config_name"] = subset available_splits = get_dataset_split_names(**load_kwargs) print(f"Available splits: {available_splits}\n") @@ -1518,6 +1552,22 @@ class UnslothTrainer: self._update_progress(error="Model not loaded") return False + # Pre-import heavy transformers modules on the main thread. + # Unsloth's patched_import hook (deepseek_v3_moe.py) is not thread-safe + # with Python's importlib cache, causing KeyError: 'size' if these are + # first imported inside the worker thread. + import transformers # noqa: F401 – ensures submodules are cached + from transformers import ( # noqa: F401 + Trainer as _HFTrainer, + TrainingArguments as _TrainingArguments, + TrainerCallback as _TrainerCallback, + ) + if self._audio_type == 'whisper': + from transformers import ( # noqa: F401 + Seq2SeqTrainer as _Seq2SeqTrainer, + Seq2SeqTrainingArguments as _Seq2SeqTrainingArguments, + ) + # Start training in separate thread self.training_thread = threading.Thread( target=self._train_worker, @@ -1970,7 +2020,7 @@ class UnslothTrainer: "model": self.model, "train_dataset": train_ds, "data_collator": data_collator, - "tokenizer": self.tokenizer.feature_extractor, + "processing_class": self.tokenizer.feature_extractor, "args": Seq2SeqTrainingArguments(**whisper_training_args), } if eval_dataset: @@ -2293,21 +2343,14 @@ class UnslothTrainer: self._update_progress(error=error_msg, is_training=False) return - elif self.is_vlm: - # Standard VLM collator - print("Using UnslothVisionDataCollator for vision model\n") - from unsloth.trainer import UnslothVisionDataCollator - - FastVisionModel.for_training(self.model) - data_collator = UnslothVisionDataCollator(self.model, self.tokenizer) - print("Vision data collator configured\n") - elif self.is_audio_vlm: # Audio VLM collator (e.g. Gemma 3N with audio data) # Mirrors the collate_fn from Gemma3N_(4B)-Audio notebook print("Configuring audio VLM data collator...\n") processor = self.tokenizer # FastModel returns processor as tokenizer + audio_col_name = getattr(self, '_audio_vlm_audio_col', 'audio') + def audio_vlm_collate_fn(examples): texts = [] audios = [] @@ -2316,7 +2359,7 @@ class UnslothTrainer: example["messages"], tokenize=False, add_generation_prompt=False ).strip() texts.append(text) - audios.append(example["audio"]["array"]) + audios.append(example[audio_col_name]["array"]) batch = processor( text=texts, audio=audios, return_tensors="pt", padding=True @@ -2335,6 +2378,15 @@ class UnslothTrainer: data_collator = audio_vlm_collate_fn print("Audio VLM data collator configured\n") + elif self.is_vlm: + # Standard VLM collator (images) + print("Using UnslothVisionDataCollator for vision model\n") + from unsloth.trainer import UnslothVisionDataCollator + + FastVisionModel.for_training(self.model) + data_collator = UnslothVisionDataCollator(self.model, self.tokenizer) + print("Vision data collator configured\n") + # ========== TRAINING CONFIGURATION ========== # Handle epochs vs max_steps properly max_steps_val = training_args.get('max_steps', 0) @@ -2443,19 +2495,28 @@ class UnslothTrainer: print("Training configuration prepared\n") # ========== TRAINER INITIALIZATION ========== - if self.is_vlm or self.is_audio_vlm: - # VLM: dataset is dict wrapper from format_and_template_dataset - # Audio VLM: dataset is raw Dataset from _format_audio_vlm_dataset + if self.is_audio_vlm: + # Audio VLM (e.g. Gemma 3N + audio): raw Dataset from _format_audio_vlm_dataset + # Notebook uses processing_class=processor.tokenizer (text tokenizer only) + train_dataset = dataset if isinstance(dataset, Dataset) else dataset['dataset'] + processing_class = self.tokenizer.tokenizer if hasattr(self.tokenizer, 'tokenizer') else self.tokenizer + trainer_kwargs = { + "model": self.model, + "train_dataset": train_dataset, + "processing_class": processing_class, + "data_collator": data_collator, + "args": SFTConfig(**config_args), + } + if eval_dataset is not None: + trainer_kwargs["eval_dataset"] = eval_dataset + self.trainer = SFTTrainer(**trainer_kwargs) + elif self.is_vlm: + # Image VLM: dataset is dict wrapper from format_and_template_dataset train_dataset = dataset['dataset'] if isinstance(dataset, dict) else dataset trainer_kwargs = { "model": self.model, -<<<<<<< HEAD - "train_dataset": dataset['dataset'], - "processing_class": self.tokenizer, -======= "train_dataset": train_dataset, - "processing_class": self.tokenizer.tokenizer, ->>>>>>> 0a7e75e (Adding support for audio llms) + "processing_class": self.tokenizer, "data_collator": data_collator, "args": SFTConfig(**config_args), } @@ -2650,11 +2711,7 @@ class UnslothTrainer: progress_callback = ProgressCallback(self) self.trainer.add_callback(progress_callback) -<<<<<<< HEAD - num_samples = len(self.trainer.train_dataset) -======= num_samples = len(dataset['dataset'] if isinstance(dataset, dict) else dataset) ->>>>>>> 0a7e75e (Adding support for audio llms) batch_size = training_args.get('batch_size', 2) grad_accum = training_args.get('gradient_accumulation_steps', 4) num_epochs = training_args.get('num_epochs', 3) diff --git a/studio/backend/models/datasets.py b/studio/backend/models/datasets.py index 81adef7577..a144cf65cb 100644 --- a/studio/backend/models/datasets.py +++ b/studio/backend/models/datasets.py @@ -28,9 +28,12 @@ class CheckFormatResponse(BaseModel): detected_format: str columns: List[str] is_multimodal: bool = False + is_audio: bool = False multimodal_columns: Optional[List[str]] = None suggested_mapping: Optional[Dict[str, str]] = None detected_image_column: Optional[str] = None + detected_audio_column: Optional[str] = None detected_text_column: Optional[str] = None + detected_speaker_column: Optional[str] = None preview_samples: Optional[List[Dict]] = None total_rows: Optional[int] = None diff --git a/studio/backend/routes/datasets.py b/studio/backend/routes/datasets.py index cb1ea75e33..af50b2b72e 100644 --- a/studio/backend/routes/datasets.py +++ b/studio/backend/routes/datasets.py @@ -212,10 +212,13 @@ def check_format(request: CheckFormatRequest): detected_format=result["detected_format"], columns=result["columns"], is_multimodal=result.get("is_multimodal", False), + is_audio=result.get("is_audio", False), multimodal_columns=result.get("multimodal_columns"), suggested_mapping=result.get("suggested_mapping"), detected_image_column=result.get("detected_image_column"), + detected_audio_column=result.get("detected_audio_column"), detected_text_column=result.get("detected_text_column"), + detected_speaker_column=result.get("detected_speaker_column"), preview_samples=preview_samples, total_rows=total_rows, ) diff --git a/studio/backend/utils/datasets/dataset_utils.py b/studio/backend/utils/datasets/dataset_utils.py index 3a4d54f93f..6ddd2c753c 100644 --- a/studio/backend/utils/datasets/dataset_utils.py +++ b/studio/backend/utils/datasets/dataset_utils.py @@ -65,13 +65,22 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict: # Auto-detect multimodal data regardless of is_vlm flag multimodal_info = detect_multimodal_dataset(dataset) - if multimodal_info["is_multimodal"]: - is_vlm = True # Route to VLM detection automatically - + 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 + + # Common audio fields for all return paths + audio_fields = { + "is_audio": is_audio, + "detected_audio_column": multimodal_info.get("detected_audio_column"), + "detected_speaker_column": multimodal_info.get("detected_speaker_column"), + } + if is_vlm: vlm_structure = detect_vlm_dataset_structure(dataset) requires_mapping = vlm_structure["format"] == "unknown" - + return { "requires_manual_mapping": requires_mapping, "detected_format": vlm_structure["format"], @@ -81,51 +90,70 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict: "detected_text_column": vlm_structure.get("text_column"), "is_multimodal": multimodal_info["is_multimodal"], "multimodal_columns": multimodal_info.get("multimodal_columns"), + **audio_fields, } - else: - # LLM flow - detected = detect_dataset_format(dataset) - - # If format is unknown, try heuristic detection - if detected["format"] == "unknown": - heuristic_mapping = detect_custom_format_heuristic(dataset) - if heuristic_mapping: - # Heuristic succeeded - no manual mapping needed - return { - "requires_manual_mapping": False, - "detected_format": "custom_heuristic", - "columns": columns, - "suggested_mapping": heuristic_mapping, - "detected_image_column": None, - "detected_text_column": None, - "is_multimodal": False, - "multimodal_columns": None, - } - else: - # Both detection and heuristic failed - return { - "requires_manual_mapping": True, - "detected_format": "unknown", - "columns": columns, - "suggested_mapping": None, - "detected_image_column": None, - "detected_text_column": None, - "is_multimodal": False, - "multimodal_columns": None, - } - - # Known format detected + + if is_audio: + # Audio dataset — require manual mapping only when columns can't be auto-detected + detected_audio = multimodal_info.get("detected_audio_column") + detected_text = multimodal_info.get("detected_text_column") + needs_mapping = not detected_audio or not detected_text return { - "requires_manual_mapping": False, - "detected_format": detected["format"], + "requires_manual_mapping": needs_mapping, + "detected_format": "audio", "columns": columns, "suggested_mapping": None, "detected_image_column": None, - "detected_text_column": None, - "is_multimodal": False, - "multimodal_columns": None, + "detected_text_column": multimodal_info.get("detected_text_column"), + "is_multimodal": True, + "multimodal_columns": multimodal_info.get("audio_columns"), + **audio_fields, } + # LLM flow + detected = detect_dataset_format(dataset) + + # If format is unknown, try heuristic detection + if detected["format"] == "unknown": + heuristic_mapping = detect_custom_format_heuristic(dataset) + if heuristic_mapping: + return { + "requires_manual_mapping": False, + "detected_format": "custom_heuristic", + "columns": columns, + "suggested_mapping": heuristic_mapping, + "detected_image_column": None, + "detected_text_column": None, + "is_multimodal": False, + "multimodal_columns": None, + **audio_fields, + } + else: + return { + "requires_manual_mapping": True, + "detected_format": "unknown", + "columns": columns, + "suggested_mapping": None, + "detected_image_column": None, + "detected_text_column": None, + "is_multimodal": False, + "multimodal_columns": None, + **audio_fields, + } + + # Known format detected + return { + "requires_manual_mapping": False, + "detected_format": detected["format"], + "columns": columns, + "suggested_mapping": None, + "detected_image_column": None, + "detected_text_column": None, + "is_multimodal": False, + "multimodal_columns": None, + **audio_fields, + } + # Normalise any format-specific role to canonical chatml (user/assistant/system) _TO_CHATML = { "user": "user", "human": "user", "instruction": "user", diff --git a/studio/backend/utils/datasets/format_detection.py b/studio/backend/utils/datasets/format_detection.py index 9283ea5d55..1a557b5df6 100644 --- a/studio/backend/utils/datasets/format_detection.py +++ b/studio/backend/utils/datasets/format_detection.py @@ -326,45 +326,51 @@ def detect_custom_format_heuristic(dataset): def detect_multimodal_dataset(dataset): """ - Detects if dataset contains multimodal data (images/vision). + Detects if dataset contains multimodal data (images and/or audio). - Two-pass approach: - 1. Column-name heuristic (fast): checks for keywords like 'image', 'img', 'pixel'. - 2. Value-type inspection (reliable): checks if actual values are PIL Images, - bytes with image headers, or HF Image-feature dicts. + Two-pass approach for each modality: + 1. Column-name heuristic (fast): checks for keywords. + 2. Value-type inspection (reliable): checks actual sample values. Returns: dict: { "is_multimodal": bool, "multimodal_columns": list of column names containing image data, - "modality_types": list of detected types (e.g., ["image", "pixel"]) + "modality_types": list of detected types (e.g., ["image", "audio"]), + "is_audio": bool, + "audio_columns": list of column names containing audio data, + "detected_audio_column": str or None, + "detected_text_column": str or None, } """ sample = next(iter(dataset)) column_names = list(sample.keys()) - # Keywords that indicate multimodal/image data - multimodal_keywords = [ + # Keywords that indicate image data + image_keywords = [ 'image', 'img', 'pixel', 'jpg', 'jpeg', 'png', 'webp', 'bmp', 'gif', 'tiff', 'svg', 'photo', 'pic', 'picture', 'visual', ] + # Keywords that indicate audio data + audio_keywords = ['audio', 'speech', 'wav', 'waveform', 'sound'] + multimodal_columns = [] + audio_columns = [] modality_types = set() - # ── Pass 1: column-name heuristic ─────────────────────── + # ── Image detection ───────────────────────────────────── + # Pass 1: column-name heuristic for col_name in column_names: col_lower = col_name.lower() - - for keyword in multimodal_keywords: + for keyword in image_keywords: if keyword in col_lower: multimodal_columns.append(col_name) modality_types.add(keyword) - break # Don't check other keywords for this column + break - # ── Pass 2: inspect actual values ─────────────────────── - # Catches columns with non-obvious names (e.g. "jpg", "photo", "pic") + # Pass 2: inspect actual values already_detected = set(multimodal_columns) for col_name in column_names: if col_name in already_detected: @@ -374,10 +380,61 @@ def detect_multimodal_dataset(dataset): multimodal_columns.append(col_name) modality_types.add("image") + # ── Audio detection ───────────────────────────────────── + # Pass 1: column-name heuristic + for col_name in column_names: + col_lower = col_name.lower() + for keyword in audio_keywords: + if keyword in col_lower: + audio_columns.append(col_name) + modality_types.add("audio") + break + + # Pass 2: inspect actual values (catches non-obvious column names) + already_audio = set(audio_columns) + for col_name in column_names: + if col_name in already_audio: + continue + value = sample[col_name] + if _is_audio_value(value): + audio_columns.append(col_name) + modality_types.add("audio") + + # Filter out columns that are actually audio from the image list + # (e.g. a column named "audio" with {"bytes", "path"} could match _is_image_value) + if audio_columns: + audio_set = set(audio_columns) + multimodal_columns = [c for c in multimodal_columns if c not in audio_set] + + # Detect text column for audio datasets + detected_text_col = None + if audio_columns: + text_keywords = ['text', 'sentence', 'transcript', 'transcription', 'label'] + for col_name in column_names: + if col_name.lower() in text_keywords: + detected_text_col = col_name + break + + is_audio = len(audio_columns) > 0 + + # Detect speaker_id column for TTS datasets (CSM, Orpheus, Spark) + detected_speaker_col = None + if audio_columns: + speaker_keywords = ['source', 'speaker', 'speaker_id'] + for col_name in column_names: + if col_name.lower() in speaker_keywords: + detected_speaker_col = col_name + break + return { - "is_multimodal": len(multimodal_columns) > 0, + "is_multimodal": len(multimodal_columns) > 0 or is_audio, "multimodal_columns": multimodal_columns, - "modality_types": list(modality_types) + "modality_types": list(modality_types), + "is_audio": is_audio, + "audio_columns": audio_columns, + "detected_audio_column": audio_columns[0] if audio_columns else None, + "detected_text_column": detected_text_col, + "detected_speaker_column": detected_speaker_col, } @@ -395,9 +452,16 @@ def _is_image_value(value) -> bool: pass # HF datasets Image feature stores decoded images as PIL or dicts with - # {"bytes": b"...", "path": "..."} when not yet decoded + # {"bytes": b"...", "path": "..."} when not yet decoded. + # Exclude audio dicts (decoded audio has "array" + "sampling_rate"). if isinstance(value, dict): + if "array" in value and "sampling_rate" in value: + return False # This is audio, not image if "bytes" in value and "path" in value: + # Check path extension to exclude audio files + path = value.get("path") or "" + if isinstance(path, str) and any(path.lower().endswith(ext) for ext in _AUDIO_EXTENSIONS): + return False return True # Raw bytes with a known image magic header @@ -407,6 +471,29 @@ def _is_image_value(value) -> bool: return False +_AUDIO_EXTENSIONS = ( + ".wav", ".mp3", ".flac", ".ogg", ".opus", ".m4a", ".aac", ".wma", ".webm", +) + + +def _is_audio_value(value) -> bool: + """Check if a single sample value looks like audio data.""" + if value is None: + return False + + # HF datasets Audio feature: decoded → {"array": np.ndarray, "sampling_rate": int} + if isinstance(value, dict): + if "array" in value and "sampling_rate" in value: + return True + # Undecoded/streaming → {"bytes": b"...", "path": "some.wav"} + if "bytes" in value or "path" in value: + path = value.get("path") or "" + if isinstance(path, str) and any(path.lower().endswith(ext) for ext in _AUDIO_EXTENSIONS): + return True + + return False + + def _has_image_header(data: bytes) -> bool: """Quick magic-byte check for common image formats.""" if len(data) < 4: diff --git a/studio/frontend/src/features/studio/sections/dataset-preview-dialog-mapping.tsx b/studio/frontend/src/features/studio/sections/dataset-preview-dialog-mapping.tsx index c1455c20d2..a0938f1338 100644 --- a/studio/frontend/src/features/studio/sections/dataset-preview-dialog-mapping.tsx +++ b/studio/frontend/src/features/studio/sections/dataset-preview-dialog-mapping.tsx @@ -16,6 +16,7 @@ const CHATML_ROLES = ["system", "user", "assistant"] as const; const ALPACA_ROLES = ["instruction", "input", "output"] as const; const SHAREGPT_ROLES = ["system", "human", "gpt"] as const; const VLM_ROLES = ["image", "text"] as const; +const AUDIO_ROLES = ["audio", "text", "speaker_id"] as const; const ROLE_LABELS: Record = { system: "System", @@ -28,9 +29,12 @@ const ROLE_LABELS: Record = { output: "Output", image: "Image", text: "Text", + audio: "Audio", + speaker_id: "Speaker ID", }; -export function getAvailableRoles(isVlm: boolean, format?: string): readonly string[] { +export function getAvailableRoles(isVlm: boolean, format?: string, isAudio?: boolean): readonly string[] { + if (isAudio) return AUDIO_ROLES; if (isVlm) return VLM_ROLES; if (format === "alpaca") return ALPACA_ROLES; if (format === "sharegpt") return SHAREGPT_ROLES; @@ -41,8 +45,10 @@ export function isMappingComplete( mapping: Record, isVlm: boolean, format?: string, + isAudio?: boolean, ): boolean { const roles = new Set(Object.values(mapping)); + if (isAudio) return roles.has("audio") && roles.has("text"); if (isVlm) return roles.has("image") && roles.has("text"); if (format === "alpaca") return roles.has("instruction") && roles.has("output"); if (format === "sharegpt") return roles.has("human") && roles.has("gpt"); @@ -91,16 +97,19 @@ export function DatasetMappingCard({ mappingOk: boolean; autoDetected?: boolean; isVlm?: boolean; + isAudio?: boolean; format?: string; }) { const entries = Object.entries(mapping); - const requiredLabel = isVlm - ? "image and text" - : format === "alpaca" - ? "instruction and output" - : format === "sharegpt" - ? "human and gpt" - : "user and assistant"; + const requiredLabel = isAudio + ? "audio and text" + : isVlm + ? "image and text" + : format === "alpaca" + ? "instruction and output" + : format === "sharegpt" + ? "human and gpt" + : "user and assistant"; return (
= { instruction: "user", input: "system", output: "assistant", human: "user", gpt: "assistant", image: "image", text: "text", + audio: "audio", speaker_id: "speaker_id", }; /** Chatml → format-specific role names (only for formats that differ). */ @@ -257,10 +267,18 @@ export function deriveDefaultMapping( data: CheckFormatResponse, isVlm: boolean, format?: string, + isAudio?: boolean, ): Record { if (data.suggested_mapping) { return remapRolesForFormat({ ...data.suggested_mapping }, format); } + if (isAudio) { + const result: Record = {}; + if (data.detected_audio_column) result[data.detected_audio_column] = "audio"; + if (data.detected_text_column) result[data.detected_text_column] = "text"; + if (data.detected_speaker_column) result[data.detected_speaker_column] = "speaker_id"; + return result; + } if (isVlm) { const result: Record = {}; if (data.detected_image_column) result[data.detected_image_column] = "image"; 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 00a738bc64..7537534329 100644 --- a/studio/frontend/src/features/studio/sections/dataset-preview-dialog.tsx +++ b/studio/frontend/src/features/studio/sections/dataset-preview-dialog.tsx @@ -64,13 +64,14 @@ export function DatasetPreviewDialog({ // 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). - const effectiveIsVlm = isVlm || !!data?.is_multimodal; + const effectiveIsAudio = !!data?.is_audio; + const effectiveIsVlm = !effectiveIsAudio && (isVlm || !!data?.is_multimodal); const hasHeuristicMapping = !data?.requires_manual_mapping && !!data?.suggested_mapping; const mappingEnabled = !!data?.requires_manual_mapping || hasHeuristicMapping; const showMappingFooter = mode === "mapping" && mappingEnabled; - const mappingOk = isMappingComplete(manualMapping, effectiveIsVlm, datasetFormat); - const availableRoles = getAvailableRoles(effectiveIsVlm, datasetFormat); + const mappingOk = isMappingComplete(manualMapping, effectiveIsVlm, datasetFormat, effectiveIsAudio); + const availableRoles = getAvailableRoles(effectiveIsVlm, datasetFormat, effectiveIsAudio); const isHfDataset = !!datasetName && datasetName.includes("/"); // When format changes, remap existing mapping roles to the new format's role names @@ -150,7 +151,7 @@ export function DatasetPreviewDialog({ if (!data?.requires_manual_mapping && !data?.suggested_mapping) return; // Don't overwrite if mapping already has entries if (Object.keys(manualMapping).length > 0) return; - const derived = deriveDefaultMapping(data, effectiveIsVlm, datasetFormat); + const derived = deriveDefaultMapping(data, effectiveIsVlm, datasetFormat, effectiveIsAudio); if (Object.keys(derived).length === 0) return; setManualMapping(derived); }, [open, datasetName, data, effectiveIsVlm, datasetFormat, manualMapping, setManualMapping]); @@ -346,6 +347,7 @@ export function DatasetPreviewDialog({ mappingOk={mappingOk} autoDetected={hasHeuristicMapping} isVlm={effectiveIsVlm} + isAudio={effectiveIsAudio} format={datasetFormat} /> )} 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 bfe035cf4a..10f67493dc 100644 --- a/studio/frontend/src/features/training/hooks/use-training-actions.ts +++ b/studio/frontend/src/features/training/hooks/use-training-actions.ts @@ -47,12 +47,22 @@ export function useTrainingActions() { isVlm, }); - // Backend auto-detects multimodal even if we didn't know yet - if (check.is_multimodal && config.isVisionModel) { + // Backend auto-detects multimodal/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; + + if (isMultimodal && config.isVisionModel) { isVlm = true; } + if (isMultimodal !== config.isDatasetMultimodal || isAudio !== config.isDatasetAudio) { + useTrainingConfigStore.setState({ + isDatasetMultimodal: isMultimodal, + isDatasetAudio: isAudio, + }); + } - if (check.requires_manual_mapping && !hasManualMapping(config, isVlm)) { + if (check.requires_manual_mapping && !hasManualMapping(config, isVlm, isAudio)) { // Pre-fill from suggested_mapping or VLM detected columns const hint: Record = {}; if (check.suggested_mapping) { @@ -60,6 +70,10 @@ export function useTrainingActions() { for (const [col, role] of Object.entries(check.suggested_mapping)) { hint[col] = table ? (table[role] ?? role) : role; } + } else if (isAudio) { + if (check.detected_audio_column) hint[check.detected_audio_column] = "audio"; + if (check.detected_text_column) hint[check.detected_text_column] = "text"; + if (check.detected_speaker_column) hint[check.detected_speaker_column] = "speaker_id"; } else if (isVlm) { if (check.detected_image_column) hint[check.detected_image_column] = "image"; if (check.detected_text_column) hint[check.detected_text_column] = "text"; @@ -75,7 +89,8 @@ export function useTrainingActions() { } } - const payload = buildTrainingStartPayload(config); + // Re-read config after potential store updates from dataset check + const payload = buildTrainingStartPayload(useTrainingConfigStore.getState()); const response = await startTraining(payload); if (response.status === "error") { @@ -143,12 +158,11 @@ function getDatasetName(config: TrainingConfigState): string | null { : config.uploadedFile; } -function hasManualMapping(config: TrainingConfigState, isVlm = false): boolean { +function hasManualMapping(config: TrainingConfigState, isVlm = false, isAudio = false): boolean { const mapping = config.datasetManualMapping; const roles = new Set(Object.values(mapping)); - if (isVlm) { - return roles.has("image") && roles.has("text"); - } + if (isAudio) return roles.has("audio") && roles.has("text"); + if (isVlm) return roles.has("image") && roles.has("text"); const fmt = config.datasetFormat; if (fmt === "alpaca") return roles.has("instruction") && roles.has("output"); if (fmt === "sharegpt") return roles.has("human") && roles.has("gpt"); diff --git a/studio/frontend/src/features/training/lib/model-defaults.ts b/studio/frontend/src/features/training/lib/model-defaults.ts index e68516e7de..35ce562dbf 100644 --- a/studio/frontend/src/features/training/lib/model-defaults.ts +++ b/studio/frontend/src/features/training/lib/model-defaults.ts @@ -4,7 +4,6 @@ import type { TrainingConfigState } from "../types/config"; type ModelDefaultsPatch = Partial< Pick< TrainingConfigState, - | "isDatasetAudio" | "epochs" | "contextLength" | "learningRate" @@ -80,9 +79,6 @@ export function mapBackendModelConfigToTrainingPatch( const lora = config.lora; const logging = config.logging; - // Audio models: set isDatasetAudio based on audio_type from YAML - patch.isDatasetAudio = typeof config.audio_type === "string" && config.audio_type.length > 0; - const maxSeqLength = toNumber(training?.max_seq_length); if (maxSeqLength !== undefined) patch.contextLength = maxSeqLength; 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 e72acef5d2..4e50ed6ae3 100644 --- a/studio/frontend/src/features/training/stores/training-config-store.ts +++ b/studio/frontend/src/features/training/stores/training-config-store.ts @@ -175,8 +175,10 @@ export const useTrainingConfigStore = create()( .then((res) => { if (controller.signal.aborted) return; const isMultimodal = !!res.is_multimodal; + const isAudio = !!res.is_audio; const updates: Record = { isDatasetMultimodal: isMultimodal, + isDatasetAudio: isAudio, isCheckingDataset: false, }; if (!_trainOnCompletionsManuallySet) { diff --git a/studio/frontend/src/features/training/types/datasets.ts b/studio/frontend/src/features/training/types/datasets.ts index 96cf699f89..c1880116ca 100644 --- a/studio/frontend/src/features/training/types/datasets.ts +++ b/studio/frontend/src/features/training/types/datasets.ts @@ -4,10 +4,13 @@ export type CheckFormatResponse = { columns: string[]; suggested_mapping?: Record | null; detected_image_column?: string | null; + detected_audio_column?: string | null; detected_text_column?: string | null; + detected_speaker_column?: string | null; preview_samples?: Record[] | null; total_rows?: number | null; is_multimodal?: boolean; + is_audio?: boolean; multimodal_columns?: string[] | null; };