Studio fix recipe dataset preview (#6031)
* Studio: fix recipe dataset preview * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
This commit is contained in:
parent
aec41d17ed
commit
33f4397b78
7 changed files with 506 additions and 165 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.")
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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({
|
|||
/>
|
||||
</div>
|
||||
|
||||
{readyForTraining && (
|
||||
<div className="mb-4 flex items-start gap-2.5 rounded-lg border border-emerald-200 bg-emerald-50 px-4 py-3 text-xs text-emerald-800 dark:border-emerald-800 dark:bg-emerald-950/70 dark:text-emerald-300">
|
||||
<HugeiconsIcon icon={CheckmarkCircle02Icon} className="mt-0.5 size-4 shrink-0" />
|
||||
<div className="space-y-0.5">
|
||||
<p className="font-medium">Ready for training</p>
|
||||
{readyDetail && <p>{readyDetail}</p>}
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{data.warning && !isRawFormat && (
|
||||
<div className="rounded-lg border border-amber-200 bg-amber-50 px-4 py-3 text-xs text-amber-700 dark:border-amber-800 dark:bg-amber-950 dark:text-amber-400 mb-4 flex items-start gap-2.5">
|
||||
<HugeiconsIcon icon={AlertCircleIcon} className="size-4 shrink-0 mt-0.5" />
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ export type CheckFormatResponse = {
|
|||
detected_audio_column?: string | null;
|
||||
detected_text_column?: string | null;
|
||||
detected_speaker_column?: string | null;
|
||||
chat_column?: string | null;
|
||||
preview_samples?: Record<string, unknown>[] | null;
|
||||
total_rows?: number | null;
|
||||
is_image?: boolean;
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue