diff --git a/studio/backend/models/datasets.py b/studio/backend/models/datasets.py index 9a2d8e91cb..b40a8416ec 100644 --- a/studio/backend/models/datasets.py +++ b/studio/backend/models/datasets.py @@ -40,6 +40,7 @@ class CheckFormatResponse(BaseModel): detected_audio_column: Optional[str] = None detected_text_column: Optional[str] = None detected_speaker_column: Optional[str] = None + chat_column: Optional[str] = None preview_samples: Optional[List[Dict]] = None total_rows: Optional[int] = None warning: Optional[str] = None diff --git a/studio/backend/routes/datasets.py b/studio/backend/routes/datasets.py index b755bf9d4a..859f2958f9 100644 --- a/studio/backend/routes/datasets.py +++ b/studio/backend/routes/datasets.py @@ -122,12 +122,12 @@ def _serialize_preview_rows(rows): ] -# Data-file extensions for the single-file fallback. Tier 1 preview prefers -# tabular over archives: archives (e.g. images.zip) load as ImageFolder with -# synthetic image/label columns that don't match the real schema. -_TABULAR_EXTS = (".parquet", ".json", ".jsonl", ".csv", ".tsv", ".arrow") -_ARCHIVE_EXTS = (".tar", ".tar.gz", ".tgz", ".gz", ".zst", ".zip", ".txt") -DATA_EXTS = _TABULAR_EXTS + _ARCHIVE_EXTS +# Data-file extensions for single-file preview. Tier 1 only uses tabular +# files; archives/text/config fall through to full load_dataset. +_COLUMNAR_EXTS = (".parquet", ".arrow") +_RECORD_EXTS = (".jsonl", ".csv", ".tsv") +_JSON_EXTS = (".json",) +_TABULAR_EXTS = _COLUMNAR_EXTS + _RECORD_EXTS + _JSON_EXTS LOCAL_FILE_EXTS = (".json", ".jsonl", ".csv", ".parquet") LOCAL_UPLOAD_EXTS = {".csv", ".json", ".jsonl", ".parquet"} # sync: training dataset upload limits are exposed by /api/settings/upload-limit @@ -145,6 +145,166 @@ def _safe_read_metadata(path: Path) -> dict | None: return payload +_HF_PREVIEW_EXT_PRIORITY = { + ".parquet": 0, + ".arrow": 0, + ".jsonl": 1, + ".csv": 2, + ".tsv": 3, + ".json": 4, +} +_HF_NON_DATA_EXACT_FILENAMES = { + ".gitattributes", + "builder_config.json", + "config.json", + "dataset_info.json", + "dataset_infos.json", + "metadata.json", +} +_HF_NON_DATA_CARD_FILENAMES = {"card.json", "dataset_card.json"} + + +def _normalize_hf_repo_path(path: str) -> str: + return path.strip().replace("\\", "/").lstrip("./") + + +def _hf_preview_extension(path: str) -> str | None: + lower = path.lower() + for ext in _HF_PREVIEW_EXT_PRIORITY: + if lower.endswith(ext): + return ext + return None + + +def _is_known_hf_non_data_file(path: str) -> bool: + name = Path(path).name.lower() + if name in _HF_NON_DATA_EXACT_FILENAMES: + return True + if name in _HF_NON_DATA_CARD_FILENAMES: + return True + if name == "readme" or name.startswith("readme."): + return True + if name.endswith("_config.json") or name.endswith("-config.json"): + return True + if name.endswith("_card.json") or name.endswith("-card.json"): + return True + return False + + +def _is_hf_preview_data_file(path: str) -> bool: + normalized = _normalize_hf_repo_path(path) + if not normalized or _is_known_hf_non_data_file(normalized): + return False + return _hf_preview_extension(normalized) is not None + + +def _extract_hf_metadata_data_paths(metadata: dict | None) -> list[str]: + if not metadata: + return [] + file_paths = metadata.get("file_paths") + if not isinstance(file_paths, dict): + return [] + raw_data_paths = file_paths.get("data") + if isinstance(raw_data_paths, str): + values = [raw_data_paths] + elif isinstance(raw_data_paths, list): + values = raw_data_paths + else: + return [] + + paths: list[str] = [] + for value in values: + if not isinstance(value, str): + continue + normalized = _normalize_hf_repo_path(value) + if normalized: + paths.append(normalized) + return paths + + +def _select_best_hf_preview_candidate( + candidates: list[str], *, subset: str | None, split: str | None +) -> str | None: + if not candidates: + return None + subset_lower = subset.lower() if subset else None + split_lower = split.lower() if split else None + + def score(path: str) -> tuple[int, int, int, int, str]: + ext = _hf_preview_extension(path) + ext_priority = _HF_PREVIEW_EXT_PRIORITY[ext] if ext else 99 + stem = Path(path).stem.lower() + path_lower = path.lower() + + subset_miss = 0 + if subset_lower: + subset_miss = 0 if subset_lower in stem or subset_lower in path_lower else 1 + + split_miss = 0 + if split_lower: + split_hit = ( + stem == split_lower + or stem.startswith(f"{split_lower}_") + or stem.startswith(f"{split_lower}-") + or f"/{split_lower}/" in path_lower + or f"_{split_lower}." in path_lower + or f"-{split_lower}." in path_lower + or f"/{split_lower}." in path_lower + or f"/{split_lower}_" in path_lower + or f"/{split_lower}-" in path_lower + ) + split_miss = 0 if split_hit else 1 + + return (subset_miss, split_miss, ext_priority, len(path), path) + + return sorted(candidates, key = score)[0] + + +def _select_hf_preview_file( + repo_files: list[str], *, metadata: dict | None, subset: str | None, split: str | None +) -> str | None: + normalized_repo_files = [_normalize_hf_repo_path(path) for path in repo_files] + repo_file_set = set(normalized_repo_files) + + metadata_candidates = [ + path + for path in _extract_hf_metadata_data_paths(metadata) + if path in repo_file_set and _is_hf_preview_data_file(path) + ] + if metadata_candidates: + return _select_best_hf_preview_candidate(metadata_candidates, subset = subset, split = split) + + data_candidates = [path for path in normalized_repo_files if _is_hf_preview_data_file(path)] + return _select_best_hf_preview_candidate(data_candidates, subset = subset, split = split) + + +def _download_hf_metadata(*, repo_id: str, repo_files: list[str], token: str | None) -> dict | None: + metadata_file = next( + ( + path + for path in repo_files + if Path(_normalize_hf_repo_path(path)).name.lower() == "metadata.json" + ), + None, + ) + if not metadata_file: + return None + + try: + from huggingface_hub import hf_hub_download + local_path = hf_hub_download( + repo_id = repo_id, + filename = metadata_file, + repo_type = "dataset", + token = token, + ) + except Exception as exc: + logger.warning(f"Could not read HF dataset metadata for {repo_id}: {exc}") + return None + + return _safe_read_metadata(Path(local_path)) + + def _safe_read_rows_from_metadata(payload: dict | None) -> int | None: if not payload: return None @@ -442,7 +602,8 @@ def check_format(request: CheckFormatRequest, current_subject: str = Depends(get """Check if a dataset requires manual column mapping. HuggingFace strategy: - 1. list_repo_files -> first data file -> load_dataset (avoids resolving thousands of files; ~2-4 s). + 1. list_repo_files -> select one tabular data file -> load_dataset + (avoids resolving thousands of files; ~2-4 s). 2. Full streaming load_dataset as a last-resort fallback. Local files load directly. Plain `def` (not async) so FastAPI runs it in a @@ -470,7 +631,7 @@ def check_format(request: CheckFormatRequest, current_subject: str = Depends(get ) else: # ── HuggingFace dataset ───────────────────────────────── - # Tier 1: load only the first data file from list_repo_files + # Tier 1: list_repo_files -> load one selected tabular data file preview_slice = None try: @@ -482,27 +643,23 @@ def check_format(request: CheckFormatRequest, current_subject: str = Depends(get repo_type = "dataset", token = request.hf_token or None, ) - data_files = [f for f in repo_files if any(f.endswith(ext) for ext in DATA_EXTS)] + metadata = _download_hf_metadata( + repo_id = request.dataset_name, + repo_files = repo_files, + token = request.hf_token or None, + ) + selected_file = _select_hf_preview_file( + repo_files, + metadata = metadata, + subset = request.subset, + split = request.train_split or "train", + ) - # Prefer tabular over archives (e.g. images.zip -> ImageFolder - # with synthetic columns not in the real schema). - tabular_files = [ - f for f in data_files if any(f.endswith(ext) for ext in _TABULAR_EXTS) - ] - candidates = tabular_files or data_files - - # With a subset, narrow to files whose name matches it. - if request.subset and candidates: - subset_matches = [f for f in candidates if request.subset in Path(f).stem] - if subset_matches: - candidates = subset_matches - - if candidates: - first_file = candidates[0] - logger.info(f"Tier 1: loading single file {first_file}") + if selected_file: + logger.info(f"Tier 1: loading single file {selected_file}") load_kwargs = { "path": request.dataset_name, - "data_files": [first_file], + "data_files": [selected_file], "split": "train", "streaming": True, } @@ -596,6 +753,7 @@ def check_format(request: CheckFormatRequest, current_subject: str = Depends(get detected_audio_column = result.get("detected_audio_column"), detected_text_column = result.get("detected_text_column"), detected_speaker_column = result.get("detected_speaker_column"), + chat_column = result.get("chat_column"), preview_samples = preview_samples, total_rows = total_rows, warning = warning, diff --git a/studio/backend/utils/datasets/dataset_utils.py b/studio/backend/utils/datasets/dataset_utils.py index 2d65b6ecbc..faa3deac70 100644 --- a/studio/backend/utils/datasets/dataset_utils.py +++ b/studio/backend/utils/datasets/dataset_utils.py @@ -180,6 +180,7 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict: "suggested_mapping": None, "detected_image_column": None, "detected_text_column": None, + "chat_column": detected.get("chat_column"), "is_image": multimodal_info["is_image"], "multimodal_columns": multimodal_info.get("multimodal_columns"), **audio_fields, @@ -199,6 +200,17 @@ _TO_CHATML = { } _CHATML_ROLE_ORDER = ("system", "user", "assistant") _CHATML_TO_ALPACA = {"user": "instruction", "system": "input", "assistant": "output"} +_KNOWN_CHAT_COLUMNS = {"messages", "conversations", "texts"} + + +def _chatml_final_format(chat_column: str | None) -> str: + return "chatml_messages" if chat_column == "messages" else "chatml_conversations" + + +def _chatml_detected_format_label(chat_column: str | None) -> str: + if chat_column in _KNOWN_CHAT_COLUMNS: + return f"chatml_{chat_column}" + return "chatml_conversations" def _apply_user_mapping( @@ -522,7 +534,7 @@ def format_dataset( } # ShareGPT - needs standardization - elif detected["format"] == "sharegpt": + elif detected["format"] == "sharegpt" and detected.get("chat_column"): try: standardized = standardize_chat_format( dataset, @@ -532,11 +544,12 @@ def format_dataset( aliases_for_assistant, batch_size, num_proc, + chat_column = detected["chat_column"], ) return { "dataset": standardized, "detected_format": "sharegpt", - "final_format": f"chatml_{detected['chat_column']}", + "final_format": _chatml_final_format(detected["chat_column"]), "chat_column": detected["chat_column"], "is_standardized": True, "requires_manual_mapping": False, @@ -558,15 +571,11 @@ def format_dataset( "warnings": warnings, } - elif detected["format"] == "chatml" and detected["chat_column"] in [ - "conversations", - "messages", - "texts", - ]: + elif detected["format"] == "chatml" and detected.get("chat_column"): return { "dataset": dataset, - "detected_format": f"chatml_{detected['chat_column']}", - "final_format": f"chatml_{detected['chat_column']}", + "detected_format": _chatml_detected_format_label(detected["chat_column"]), + "final_format": _chatml_final_format(detected["chat_column"]), "chat_column": detected["chat_column"], "is_standardized": True, "requires_manual_mapping": False, @@ -637,12 +646,13 @@ def format_dataset( aliases_for_assistant, batch_size, num_proc, + chat_column = detected["chat_column"], ) warnings.append("Successfully standardized unknown format") return { "dataset": standardized, "detected_format": "unknown", - "final_format": f"chatml_{detected['chat_column']}", + "final_format": _chatml_final_format(detected["chat_column"]), "chat_column": detected["chat_column"], "is_standardized": True, "requires_manual_mapping": False, @@ -681,32 +691,52 @@ def format_dataset( "warnings": [], } - elif detected["format"] in ["sharegpt", "chatml"]: - # First standardize if ShareGPT - if detected["format"] == "sharegpt": - dataset = standardize_chat_format( + elif detected["format"] in ["sharegpt", "chatml"] and detected.get("chat_column"): + try: + # First standardize if ShareGPT + if detected["format"] == "sharegpt": + dataset = standardize_chat_format( + dataset, + tokenizer, + aliases_for_system, + aliases_for_user, + aliases_for_assistant, + batch_size, + num_proc, + chat_column = detected["chat_column"], + ) + + # Then convert to Alpaca + converted = convert_chatml_to_alpaca( dataset, - tokenizer, - aliases_for_system, - aliases_for_user, - aliases_for_assistant, batch_size, num_proc, + chat_column = detected["chat_column"], ) - - # Then convert to Alpaca - converted = convert_chatml_to_alpaca(dataset, batch_size, num_proc) - return { - "dataset": converted, - "detected_format": detected["format"], - "final_format": "alpaca", - "chat_column": None, - "is_standardized": True, - "requires_manual_mapping": False, - "is_image": multimodal_info["is_image"], - "multimodal_info": multimodal_info, - "warnings": [], - } + return { + "dataset": converted, + "detected_format": detected["format"], + "final_format": "alpaca", + "chat_column": None, + "is_standardized": True, + "requires_manual_mapping": False, + "is_image": multimodal_info["is_image"], + "multimodal_info": multimodal_info, + "warnings": [], + } + except Exception as e: + warnings.append(f"Failed to convert chat dataset to Alpaca: {e}") + return { + "dataset": dataset, + "detected_format": detected["format"], + "final_format": "unknown", + "chat_column": detected["chat_column"], + "is_standardized": False, + "requires_manual_mapping": True, + "is_image": multimodal_info["is_image"], + "multimodal_info": multimodal_info, + "warnings": warnings, + } else: warnings.append(f"Cannot convert unknown format to Alpaca") @@ -738,33 +768,48 @@ def format_dataset( "warnings": [], } - elif detected["format"] == "sharegpt": - standardized = standardize_chat_format( - dataset, - tokenizer, - aliases_for_system, - aliases_for_user, - aliases_for_assistant, - batch_size, - num_proc, - ) - return { - "dataset": standardized, - "detected_format": "sharegpt", - "final_format": f"chatml_{detected['chat_column']}", - "chat_column": detected["chat_column"], - "is_standardized": True, - "requires_manual_mapping": False, - "is_image": multimodal_info["is_image"], - "multimodal_info": multimodal_info, - "warnings": [], - } + elif detected["format"] == "sharegpt" and detected.get("chat_column"): + try: + standardized = standardize_chat_format( + dataset, + tokenizer, + aliases_for_system, + aliases_for_user, + aliases_for_assistant, + batch_size, + num_proc, + chat_column = detected["chat_column"], + ) + return { + "dataset": standardized, + "detected_format": "sharegpt", + "final_format": _chatml_final_format(detected["chat_column"]), + "chat_column": detected["chat_column"], + "is_standardized": True, + "requires_manual_mapping": False, + "is_image": multimodal_info["is_image"], + "multimodal_info": multimodal_info, + "warnings": [], + } + except Exception as e: + warnings.append(f"Failed to standardize ShareGPT format: {e}") + return { + "dataset": dataset, + "detected_format": "sharegpt", + "final_format": "sharegpt", + "chat_column": detected["chat_column"], + "is_standardized": False, + "requires_manual_mapping": True, + "is_image": multimodal_info["is_image"], + "multimodal_info": multimodal_info, + "warnings": warnings, + } - elif detected["format"] == "chatml": + elif detected["format"] == "chatml" and detected.get("chat_column"): return { "dataset": dataset, - "detected_format": f"chatml_{detected['chat_column']}", - "final_format": f"chatml_{detected['chat_column']}", + "detected_format": _chatml_detected_format_label(detected["chat_column"]), + "final_format": _chatml_final_format(detected["chat_column"]), "chat_column": detected["chat_column"], "is_standardized": True, "requires_manual_mapping": False, @@ -785,11 +830,12 @@ def format_dataset( aliases_for_assistant, batch_size, num_proc, + chat_column = detected["chat_column"], ) return { "dataset": standardized, "detected_format": "unknown", - "final_format": f"chatml_{detected['chat_column']}", + "final_format": _chatml_final_format(detected["chat_column"]), "chat_column": detected["chat_column"], "is_standardized": True, "requires_manual_mapping": False, diff --git a/studio/backend/utils/datasets/format_conversion.py b/studio/backend/utils/datasets/format_conversion.py index 7a556a9745..cb24bd96ba 100644 --- a/studio/backend/utils/datasets/format_conversion.py +++ b/studio/backend/utils/datasets/format_conversion.py @@ -29,6 +29,7 @@ def standardize_chat_format( ], batch_size = 1000, num_proc = None, + chat_column: str | None = None, ): """ Standardize BOTH messages and conversations: map non-standard role @@ -46,9 +47,10 @@ def standardize_chat_format( column_names = set(next(iter(dataset)).keys()) - # Find the chat column - chat_column = None - if "conversations" in column_names: + if chat_column: + if chat_column not in column_names: + return dataset + elif "conversations" in column_names: chat_column = "conversations" elif "messages" in column_names: chat_column = "messages" @@ -57,30 +59,48 @@ def standardize_chat_format( else: return dataset # No chat column found - # Inspect structure - examples = itertools.islice(dataset, 10) + def _iter_probe_rows(): + try: + total = min(len(dataset), 100) + for index in range(total): + yield dataset[index] + return + except Exception: + pass + for example in itertools.islice(dataset, 100): + yield example + uniques = collections.defaultdict(list) - for example in examples: - for message in example[chat_column]: + for example in _iter_probe_rows(): + chat_data = example.get(chat_column) + if not isinstance(chat_data, list) or len(chat_data) == 0: + continue + for message in chat_data: + if not isinstance(message, dict): + continue for key, value in message.items(): if type(value) is not str: continue # Skip non-strings uniques[key].append(value) - if len(uniques.keys()) != 2: - return dataset # Unexpected structure - - keys = list(uniques.keys()) - length_first = len(set(uniques[keys[0]])) - length_second = len(set(uniques[keys[1]])) - - # Fewer unique values => role; the other => content - if length_first < length_second: - role_key = keys[0] - content_key = keys[1] + if "from" in uniques and "value" in uniques: + role_key = "from" + content_key = "value" + elif "role" in uniques and "content" in uniques: + role_key = "role" + content_key = "content" + elif len(uniques.keys()) == 2: + keys = list(uniques.keys()) + length_first = len(set(uniques[keys[0]])) + length_second = len(set(uniques[keys[1]])) + if length_first < length_second: + role_key = keys[0] + content_key = keys[1] + else: + role_key = keys[1] + content_key = keys[0] else: - role_key = keys[1] - content_key = keys[0] + raise ValueError(f"Could not infer role/content keys for chat column '{chat_column}'") # Mapping for aliases aliases_mapping = {} @@ -95,10 +115,23 @@ def standardize_chat_format( convos = examples[chat_column] all_convos = [] for convo in convos: + if not isinstance(convo, list): + all_convos.append([]) + continue + new_convo = [] for message in convo: - original_role = message.get(role_key, "") - original_content = message.get(content_key, "") + if not isinstance(message, dict): + continue + + # Use the inferred keys first; fall back per-message so mixed + # ShareGPT/ChatML rows keep valid turns. + original_role = message.get(role_key) + original_content = message.get(content_key) + if original_role is None: + original_role = message.get("role") or message.get("from") or "" + if original_content is None: + original_content = message.get("content") or message.get("value") or "" standard_role = aliases_mapping.get(original_role, original_role) @@ -136,6 +169,7 @@ def convert_chatml_to_alpaca( dataset, batch_size = 1000, num_proc = None, + chat_column: str | None = None, ): """ Convert ChatML (messages OR conversations) to Alpaca format. @@ -151,10 +185,11 @@ def convert_chatml_to_alpaca( _is_torch_iterable = False def _convert(examples): - # Auto-detect the column name - chatml_data = ( - examples.get("messages") or examples.get("conversations") or examples.get("texts") - ) + chatml_data = examples.get(chat_column) if chat_column else None + if chatml_data is None: + chatml_data = ( + examples.get("messages") or examples.get("conversations") or examples.get("texts") + ) if chatml_data is None: raise ValueError("No 'messages' or 'conversations' or 'texts' column found.") diff --git a/studio/backend/utils/datasets/format_detection.py b/studio/backend/utils/datasets/format_detection.py index b3a4e3d6d5..f5ea5ca138 100644 --- a/studio/backend/utils/datasets/format_detection.py +++ b/studio/backend/utils/datasets/format_detection.py @@ -11,22 +11,146 @@ def _keyword_in_column(keyword: str, col_name: str) -> bool: return re.search(r"\b" + re.escape(keyword) + r"\b", col_name, re.IGNORECASE) is not None +CONVERSATION_COLUMNS = ("messages", "conversations", "texts") +_CHATML_KEYS = frozenset({"role", "content"}) +_SHAREGPT_KEYS = frozenset({"from", "value"}) +_TRACE_SUFFIXES = ("__trace", "_trace") + + +def _sample_dataset_rows(dataset, limit: int = 100) -> list[dict]: + try: + total = min(len(dataset), limit) + return [dataset[index] for index in range(total)] + except Exception: + rows = [] + try: + for index, row in enumerate(dataset): + if index >= limit: + break + rows.append(row) + except Exception: + return [] + return rows + + +def _get_dataset_column_names(dataset, sample: dict) -> list[str]: + column_names = getattr(dataset, "column_names", None) + if isinstance(column_names, list): + return [str(column) for column in column_names] + return [str(column) for column in sample.keys()] + + +def _is_trace_conversation_name(column_name: str) -> bool: + return column_name.lower().endswith(_TRACE_SUFFIXES) + + +def _inspect_conversation_column(rows: list[dict], column_name: str) -> dict | None: + turn_keys: set[str] = set() + has_chatml = False + has_sharegpt = False + + for row in rows: + if not isinstance(row, dict) or column_name not in row: + continue + chat_data = row[column_name] + if not isinstance(chat_data, list) or len(chat_data) == 0: + continue + for turn in chat_data: + if not isinstance(turn, dict): + continue + keys = {str(key) for key in turn.keys()} + turn_keys.update(keys) + if _SHAREGPT_KEYS.issubset(keys): + has_sharegpt = True + if _CHATML_KEYS.issubset(keys): + has_chatml = True + + if has_sharegpt: + return { + "format": "sharegpt", + "chat_column": column_name, + "needs_standardization": True, + "sample_keys": sorted(turn_keys), + } + if has_chatml: + return { + "format": "chatml", + "chat_column": column_name, + "needs_standardization": False, + "sample_keys": sorted(turn_keys), + } + if turn_keys: + return { + "format": "unknown", + "chat_column": column_name, + "needs_standardization": None, + "sample_keys": sorted(turn_keys), + } + return None + + +def _detect_conversation_column(rows: list[dict], column_names: list[str]) -> dict | None: + column_name_set = set(column_names) + unknown_exact = None + for column_name in CONVERSATION_COLUMNS: + if column_name not in column_name_set: + continue + inspected = _inspect_conversation_column(rows, column_name) + if inspected and inspected["format"] in {"sharegpt", "chatml"}: + return inspected + if inspected and unknown_exact is None: + unknown_exact = inspected + + structural_candidates = [] + for column_name in column_names: + if column_name in CONVERSATION_COLUMNS: + continue + inspected = _inspect_conversation_column(rows, column_name) + if inspected and inspected["format"] in {"sharegpt", "chatml"}: + structural_candidates.append(inspected) + + trace_candidates = [ + candidate + for candidate in structural_candidates + if _is_trace_conversation_name(candidate["chat_column"]) + ] + if len(trace_candidates) == 1: + return trace_candidates[0] + if len(trace_candidates) > 1: + return unknown_exact + if len(structural_candidates) == 1: + return structural_candidates[0] + if unknown_exact is not None: + return unknown_exact + return None + + def detect_dataset_format(dataset): """Detect dataset format by inspecting structure. Returns: dict: { "format": "alpaca" | "sharegpt" | "chatml" | "unknown", - "chat_column": "messages" | "conversations" | None, + "chat_column": str | None, "needs_standardization": bool, "sample_keys": list of keys found in messages (for debugging) } """ - column_names = set(next(iter(dataset)).keys()) + sample_rows = _sample_dataset_rows(dataset) + if not sample_rows: + return { + "format": "unknown", + "chat_column": None, + "needs_standardization": None, + "sample_keys": [], + } + + column_names = _get_dataset_column_names(dataset, sample_rows[0]) + column_name_set = set(column_names) # Alpaca alpaca_columns = {"instruction", "output"} - if alpaca_columns.issubset(column_names): + if alpaca_columns.issubset(column_name_set): return { "format": "alpaca", "chat_column": None, @@ -34,57 +158,9 @@ def detect_dataset_format(dataset): "sample_keys": [], } - # 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: - 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: "from"/"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: "role"/"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), - } - - 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), - } + conversation = _detect_conversation_column(sample_rows, column_names) + if conversation: + return conversation return { "format": "unknown", 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 6710099400..f337d51004 100644 --- a/studio/frontend/src/features/studio/sections/dataset-preview-dialog.tsx +++ b/studio/frontend/src/features/studio/sections/dataset-preview-dialog.tsx @@ -17,7 +17,11 @@ import { useTrainingActions, useTrainingConfigStore } from "@/features/training" import { checkDatasetFormat } from "@/features/training/api/datasets-api"; import { isRawTextDatasetFormat } from "@/features/training/lib/training-methods"; import type { CheckFormatResponse } from "@/features/training/types/datasets"; -import { Database02Icon, AlertCircleIcon } from "@hugeicons/core-free-icons"; +import { + AlertCircleIcon, + CheckmarkCircle02Icon, + Database02Icon, +} from "@hugeicons/core-free-icons"; import { HugeiconsIcon } from "@hugeicons/react"; import { useShallow } from "zustand/react/shallow"; import { collectPreviewImages, formatCell } from "./dataset-preview-dialog-utils"; @@ -97,6 +101,16 @@ export function DatasetPreviewDialog({ const mappingOk = isRawFormat || isMappingComplete(manualMapping, effectiveIsVlm, datasetFormat, effectiveIsAudio); const availableRoles = getAvailableRoles(effectiveIsVlm, datasetFormat, effectiveIsAudio); const isHfDataset = datasetSource === "huggingface"; + const readyForTraining = + !(isRawFormat || mappingEnabled) && + !data?.requires_manual_mapping && + !!data?.detected_format && + data.detected_format !== "unknown"; + const readyDetail = data?.chat_column && data.detected_format === "chatml" + ? `Detected ChatML conversation column: ${data.chat_column}` + : data?.detected_format + ? `Detected ${data.detected_format} format. No manual column mapping needed.` + : null; // ── AI Assist ────────────────────────────────────────────────────── const [isAiLoading, setIsAiLoading] = useState(false); @@ -438,6 +452,16 @@ export function DatasetPreviewDialog({ /> + {readyForTraining && ( +
Ready for training
+ {readyDetail &&{readyDetail}
} +