# SPDX-License-Identifier: AGPL-3.0-only # Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 """ Format detection utilities for dataset processing. This module contains functions for detecting dataset formats (Alpaca, ShareGPT, ChatML), detecting multimodal/VLM dataset structures, and heuristic-based column mapping. """ import re def _keyword_in_column(keyword: str, col_name: str) -> bool: """Word-boundary keyword match to avoid false positives like 'pic' in 'topic'.""" return ( re.search(r"\b" + re.escape(keyword) + r"\b", col_name, re.IGNORECASE) is not None ) def detect_dataset_format(dataset): """ Detects dataset format by inspecting structure. Returns: dict: { "format": "alpaca" | "sharegpt" | "chatml" | "unknown", "chat_column": "messages" | "conversations" | None, "needs_standardization": bool, "sample_keys": list of keys found in messages (for debugging) } """ column_names = set(next(iter(dataset)).keys()) # Check for Alpaca alpaca_columns = {"instruction", "output"} if alpaca_columns.issubset(column_names): return { "format": "alpaca", "chat_column": None, "needs_standardization": False, "sample_keys": [], } # Check for chat-based formats (messages or conversations) chat_column = None if "messages" in column_names: chat_column = "messages" elif "conversations" in column_names: chat_column = "conversations" elif "texts" in column_names: chat_column = "texts" if chat_column: # Inspect the structure to determine if ShareGPT or ChatML try: sample = next(iter(dataset)) chat_data = sample[chat_column] if chat_data and len(chat_data) > 0: first_msg = chat_data[0] msg_keys = set(first_msg.keys()) # ShareGPT uses "from" and "value" if "from" in msg_keys or "value" in msg_keys: return { "format": "sharegpt", "chat_column": chat_column, "needs_standardization": True, "sample_keys": list(msg_keys), } # ChatML uses "role" and "content" elif "role" in msg_keys and "content" in msg_keys: return { "format": "chatml", "chat_column": chat_column, "needs_standardization": False, "sample_keys": list(msg_keys), } # Unknown structure but has chat column else: return { "format": "unknown", "chat_column": chat_column, "needs_standardization": None, "sample_keys": list(msg_keys), } except Exception as e: return { "format": "unknown", "chat_column": chat_column, "needs_standardization": None, "sample_keys": [], "error": str(e), } # No recognized format return { "format": "unknown", "chat_column": None, "needs_standardization": None, "sample_keys": [], } def detect_custom_format_heuristic(dataset): """ Smart detection with priority scoring. Strategy for ambiguous keywords like 'task': 1. Detect assistant first (unambiguous) 2. Detect user using high-priority keywords first 3. Check REMAINING columns for system keywords (including 'task') 4. Only if no system match, use 'task' as fallback user """ sample = next(iter(dataset)) all_columns = list(sample.keys()) mapping = {} # Keywords assistant_words = [ "output", "answer", "response", "assistant", "completion", "expected", "recommendation", "reply", "result", "target", "solution", "explanation", "solve", ] # Split into high/low priority user_words_high_priority = [ "input", "question", "query", "prompt", "instruction", "request", "snippet", "user", "text", "problem", "exercise", ] user_words_low_priority = ["task"] # Ambiguous - can be user OR system user_words = user_words_high_priority + user_words_low_priority system_words = [ "system", "context", "description", "persona", "role", "template", "task", # Also in system ] # Metadata columns to ignore metadata_exact_match = { "id", "idx", "index", "key", "timestamp", "date", "metadata", "source", "kind", "type", "category", "score", "label", "tag", "inference_mode", } metadata_prefix_patterns = [ "problem_type", "problem_source", "generation_model", "pass_rate", ] priority_patterns = { "generated": 100, "gen_": 90, "model_": 80, "predicted": 70, "completion": 60, } def has_keyword(col_name, keywords): """Check if any keyword appears in column name.""" col_lower = col_name.lower() col_normalized = col_lower.replace("_", "").replace("-", "").replace(" ", "") for keyword in keywords: if keyword in col_lower or keyword in col_normalized: return True return False def is_metadata(col_name): """Check if column is likely metadata.""" col_lower = col_name.lower() if col_lower in metadata_exact_match: return True if col_lower in metadata_prefix_patterns: return True for pattern in metadata_prefix_patterns: if ( col_lower.startswith(pattern.split("_")[0] + "_") and col_lower != pattern ): if "_" in col_lower: prefix = col_lower.split("_")[0] if prefix in ["generation", "pass", "inference"]: return True if len(col_lower) <= 2 and not col_lower in ["qa", "q", "a"]: return True return False def get_priority_score(col_name): """Calculate priority score based on column name patterns.""" col_lower = col_name.lower() score = 0 for pattern, pattern_score in priority_patterns.items(): if pattern in col_lower: score += pattern_score return score def get_content_length(col_name): """Get average content length for this column.""" try: if col_name in sample and sample[col_name]: content = str(sample[col_name]) return len(content) return 0 except: return 0 def score_column(col_name, keywords, role_type, num_candidates): """Score a column for how likely it is to be a particular role.""" if not has_keyword(col_name, keywords): return 0 score = 0 score += 10 # Penalize ambiguous keywords when scoring for user if role_type == "user": col_lower = col_name.lower() # If column is ONLY "task" (or task_xxx), give it lower priority for user role if "task" in col_lower and not any( kw in col_lower for kw in user_words_high_priority ): score -= 15 # Significant penalty so other user columns win priority_bonus = get_priority_score(col_name) score += priority_bonus if role_type in ["assistant", "user"]: avg_length = get_content_length(col_name) if num_candidates > 1: if avg_length > 1000: score += 50 elif avg_length > 200: score += 30 elif avg_length > 50: score += 10 elif avg_length < 50: score -= 20 else: if avg_length > 1000: score += 50 elif avg_length > 200: score += 30 elif avg_length > 50: score += 10 return score # Filter out metadata columns content_columns = [col for col in all_columns if not is_metadata(col)] # Count candidates first assistant_potential = [ col for col in content_columns if has_keyword(col, assistant_words) ] user_potential = [col for col in content_columns if has_keyword(col, user_words)] # STEP 1: Find best ASSISTANT column assistant_candidates = [] for col in assistant_potential: score = score_column( col, assistant_words, "assistant", len(assistant_potential) ) if score > 0: assistant_candidates.append((col, score)) if assistant_candidates: assistant_candidates.sort(key = lambda x: x[1], reverse = True) assistant_col = assistant_candidates[0][0] mapping[assistant_col] = "assistant" else: assistant_col = None # STEP 2: Find best USER column (with penalty for ambiguous keywords) user_candidates = [] for col in user_potential: if col == assistant_col: continue score = score_column(col, user_words, "user", len(user_potential)) if score > 0: user_candidates.append((col, score)) if user_candidates: user_candidates.sort(key = lambda x: x[1], reverse = True) user_col = user_candidates[0][0] mapping[user_col] = "user" else: user_col = None # STEP 3: Check ALL remaining columns for SYSTEM matches (priority check) remaining_columns = [col for col in content_columns if col not in mapping] system_col = None for col in remaining_columns: if has_keyword(col, system_words): # Found a system match in remaining columns mapping[col] = "system" system_col = col break # STEP 4: Handle any additional remaining columns if system_col: remaining_columns = [col for col in remaining_columns if col != system_col] if len(remaining_columns) >= 1: remaining_col = remaining_columns[0] # If no strong keyword match, decide based on what's missing if not has_keyword(remaining_col, user_words + assistant_words): mapping[remaining_col] = "system" elif user_col is None: # No user column yet, assign this as user mapping[remaining_col] = "user" else: # Already have user + assistant, treat as system context mapping[remaining_col] = "system" # VALIDATION: Ensure we have at least user + assistant has_user = any(role == "user" for role in mapping.values()) has_assistant = any(role == "assistant" for role in mapping.values()) if not has_user and len(remaining_columns) > 0: for col in remaining_columns: if col not in mapping: mapping[col] = "user" has_user = True break if has_user and has_assistant: return mapping return None def detect_multimodal_dataset(dataset): """ Detects if dataset contains multimodal data (images and/or audio). 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_image": bool, "multimodal_columns": list of column names containing image data, "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 image data image_keywords = [ "image", "img", "pixel", "jpg", "jpeg", "png", "webp", "bmp", "gif", "tiff", "svg", "photo", "pic", "picture", "visual", "file_name", "filename", ] # Keywords that indicate audio data audio_keywords = ["audio", "speech", "wav", "waveform", "sound"] multimodal_columns = [] audio_columns = [] modality_types = set() # ── Image detection ───────────────────────────────────── # Pass 1: column-name heuristic (word-boundary match to avoid # false positives like 'pic' in 'topic') for col_name in column_names: for keyword in image_keywords: if _keyword_in_column(keyword, col_name): multimodal_columns.append(col_name) modality_types.add(keyword) break # Pass 2: inspect actual values already_detected = set(multimodal_columns) for col_name in column_names: if col_name in already_detected: continue value = sample[col_name] if _is_image_value(value): multimodal_columns.append(col_name) modality_types.add("image") # ── Audio detection ───────────────────────────────────── # Pass 1: column-name heuristic (word-boundary match) for col_name in column_names: for keyword in audio_keywords: if _keyword_in_column(keyword, col_name): 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_image": len(multimodal_columns) > 0, "multimodal_columns": multimodal_columns, "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, } def _is_image_value(value) -> bool: """Check if a single sample value looks like image data.""" if value is None: return False # PIL Image instance try: from PIL.Image import Image as PILImage if isinstance(value, PILImage): return True except ImportError: pass # HF datasets Image feature stores decoded images as PIL or dicts with # {"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 if isinstance(value, (bytes, bytearray)): return _has_image_header(value) # String that looks like an image file path or URL _IMAGE_EXTS = (".png", ".jpg", ".jpeg", ".webp", ".gif", ".bmp", ".tiff", ".svg") if isinstance(value, str) and len(value) < 1000: lower = value.strip().lower() # Image URL (http://... ending in image extension) if lower.startswith(("http://", "https://")) and any( lower.split("?")[0].endswith(ext) for ext in _IMAGE_EXTS ): return True # Image file path (relative or absolute path ending in image extension) if any(lower.endswith(ext) for ext in _IMAGE_EXTS): return True 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: return False # JPEG if data[:2] == b"\xff\xd8": return True # PNG if data[:4] == b"\x89PNG": return True # GIF if data[:3] == b"GIF": return True # WebP if data[:4] == b"RIFF" and len(data) >= 12 and data[8:12] == b"WEBP": return True # BMP if data[:2] == b"BM": return True return False def detect_vlm_dataset_structure(dataset): """ Detects if VLM dataset is: - Standard VLM messages format (image objects in content) - Llava format (image indices + separate images column) - Simple format needing conversion (image + text columns) """ try: sample = next(iter(dataset)) except StopIteration: return { "format": "unknown", "needs_conversion": None, "image_column": None, "text_column": None, "messages_column": None, } column_names = set(sample.keys()) # Check if has messages column if "messages" in column_names: messages = sample["messages"] if messages and len(messages) > 0: first_msg = messages[0] if "content" in first_msg: content = first_msg["content"] if isinstance(content, list) and len(content) > 0: if isinstance(content[0], dict) and "type" in content[0]: # Check for llava format has_index = any( "index" in item for item in content if isinstance(item, dict) ) has_images_column = "images" in column_names if has_index and has_images_column: return { "format": "vlm_messages_llava", "needs_conversion": True, "messages_column": "messages", "image_column": "images", "text_column": None, } # Standard VLM format has_image = any( "image" in item for item in content if isinstance(item, dict) ) if has_image: return { "format": "vlm_messages", "needs_conversion": False, "messages_column": "messages", "image_column": None, "text_column": None, } # Check for ShareGPT/ChatML conversations with placeholder + companion image column # (e.g. Lin-Chen/ShareGPT4V, LLaVA-style datasets) for chat_col in ("conversations", "messages"): if chat_col not in column_names: continue chat_data = sample[chat_col] if not isinstance(chat_data, list) or len(chat_data) == 0: continue first_msg = chat_data[0] if not isinstance(first_msg, dict): continue # Detect ShareGPT (from/value) or ChatML (role/content) keys msg_text = first_msg.get("value") or first_msg.get("content") if not isinstance(msg_text, str): continue # Check for placeholder anywhere in the conversation has_image_placeholder = any( "" in str(m.get("value", "") or m.get("content", "")) for m in chat_data if isinstance(m, dict) ) if not has_image_placeholder: continue # Find companion image column image_col = None for col in column_names: if col == chat_col: continue if _keyword_in_column("image", col) or _keyword_in_column("img", col): image_col = col break if image_col: return { "format": "sharegpt_with_images", "needs_conversion": True, "image_column": image_col, "text_column": None, "messages_column": chat_col, } # Find image and text columns using metadata filtering # Define metadata patterns to EXCLUDE metadata_patterns = { "suffixes": [ "_id", "_url", "_name", "_filename", "_uri", "_link", "_key", "_index", ], "prefixes": [ "id_", "url_", "name_", "filename_", "uri_", "link_", "key_", "index_", ], } # Image-related keywords image_keywords = [ "image", "img", "photo", "picture", "pic", "visual", "scan", "file_name", "filename", ] # Text-related keywords text_keywords = [ "text", "caption", "captions", "description", "answer", "output", "response", "label", ] def is_metadata_column(col_name): """Check if column name looks like metadata.""" col_lower = col_name.lower() # Check suffixes if any(col_lower.endswith(suffix) for suffix in metadata_patterns["suffixes"]): return True # Check prefixes if any( col_lower.startswith(prefix) for prefix in metadata_patterns["prefixes"] ): return True return False def _score_image_candidate(col, sample_value): """Score a candidate image column by how resolvable its value is.""" # PIL Image object (highest priority - already loaded) if hasattr(sample_value, "size") and hasattr(sample_value, "mode"): return 100 # Dict with image data (bytes/path from HF Image feature) if isinstance(sample_value, dict) and ( "bytes" in sample_value or "path" in sample_value ): return 75 if isinstance(sample_value, str): # URL strings if sample_value.startswith(("http://", "https://")): return 70 if not is_metadata_column(col) else 55 # Bare file path if is_metadata_column(col): return 30 return 50 return 0 def _probe_image_candidate(col, sample_value): """Quick probe to check if an image candidate is actually reachable. Returns True if likely valid, False if definitely broken.""" import os # PIL / dict — already loaded, always valid if not isinstance(sample_value, str): return True # Local file — check it exists if not sample_value.startswith(("http://", "https://")): return os.path.exists( sample_value ) # bare filenames return False here, that's OK # URL — quick HEAD request with short timeout try: import urllib.request req = urllib.request.Request(sample_value, method = "HEAD") resp = urllib.request.urlopen(req, timeout = 3) return resp.status < 400 except Exception: return False def find_image_column(): """Find image column by keyword match + value-based fallback. When multiple candidates exist, probes them to find one that works.""" candidates = [] # Pass 1: keyword-matched columns for col in column_names: if any(_keyword_in_column(keyword, col) for keyword in image_keywords): sample_value = sample[col] score = _score_image_candidate(col, sample_value) if score > 0: candidates.append((col, score)) # Pass 2: value-based fallback — find columns with image URLs/paths # even if the column name doesn't match image keywords already = {c[0] for c in candidates} for col in column_names: if col in already: continue sample_value = sample[col] if _is_image_value(sample_value): score = _score_image_candidate(col, sample_value) # Slightly penalise non-keyword columns so keyword matches win on ties candidates.append((col, max(score - 5, 1))) if not candidates: return None candidates.sort(key = lambda x: x[1], reverse = True) # Single candidate or top candidate is PIL/dict — no probing needed if len(candidates) == 1 or candidates[0][1] >= 75: return candidates[0][0] # Multiple string-based candidates — probe to find one that actually works for col, score in candidates: sample_value = sample[col] if _probe_image_candidate(col, sample_value): return col # Nothing probed successfully — return highest-scored anyway and let # conversion handle the error (it may still resolve via hf_hub_download) return candidates[0][0] def find_text_column(): """Find text column by filtering out metadata and checking keywords.""" candidates = [] for col in column_names: # Skip metadata columns if is_metadata_column(col): continue # Check if contains text keywords (word-boundary match) if any(_keyword_in_column(keyword, col) for keyword in text_keywords): # Verify it's actually text sample_value = sample[col] if isinstance(sample_value, str) and len(sample_value) > 0: # Longer text = higher priority (likely content, not just a label) priority = min(len(sample_value), 1000) # Cap at 1000 candidates.append((col, priority)) elif ( isinstance(sample_value, list) and len(sample_value) > 0 and isinstance(sample_value[0], str) ): # List of strings (e.g. captions list) — lower priority than plain strings priority = min(len(sample_value[0]), 1000) // 2 candidates.append((col, priority)) # Return highest priority candidate if candidates: candidates.sort(key = lambda x: x[1], reverse = True) return candidates[0][0] return None found_image = find_image_column() found_text = find_text_column() if found_image and found_text: return { "format": "simple_image_text", "needs_conversion": True, "image_column": found_image, "text_column": found_text, "messages_column": None, } return { "format": "unknown", "needs_conversion": None, "image_column": found_image, "text_column": found_text, "messages_column": None, }