# SPDX-License-Identifier: AGPL-3.0-only - See /studio/LICENSE.AGPL-3.0 # Copyright © 2025 Unsloth AI """ 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, }