unsloth/studio/backend/utils/datasets/format_detection.py
Roland Tannous dd6c38cc7b fix: probe image column candidates when multiple exist
When multiple image columns are found, probes them (HEAD for URLs,
os.path.exists for paths) and picks the first that works.
Skips probing when top candidate is PIL/dict (score >= 75).
2026-03-10 01:38:33 +00:00

799 lines
28 KiB
Python

# 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 <image> 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 <image> placeholder anywhere in the conversation
has_image_placeholder = any(
"<image>" 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,
}