unsloth/studio/backend/routes/data_recipe/seed.py
Roland Tannous 47654cb91c Final cleanup
2026-03-12 18:28:04 +00:00

365 lines
12 KiB
Python

# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""Seed inspect endpoints for data recipe."""
from __future__ import annotations
import base64
import binascii
from itertools import islice
from pathlib import Path
from typing import Any
from uuid import uuid4
from fastapi import APIRouter, HTTPException
from data_designer_unstructured_seed.chunking import (
build_unstructured_preview_rows,
resolve_chunking,
)
from core.data_recipe.jsonable import to_preview_jsonable
from utils.paths import ensure_dir, seed_uploads_root
from models.data_recipe import (
SeedInspectRequest,
SeedInspectResponse,
SeedInspectUploadRequest,
)
router = APIRouter()
DATA_EXTS = (".parquet", ".jsonl", ".json", ".csv")
DEFAULT_SPLIT = "train"
LOCAL_UPLOAD_EXTS = {".csv", ".json", ".jsonl"}
UNSTRUCTURED_UPLOAD_EXTS = {".txt", ".md"}
SEED_UPLOAD_DIR = seed_uploads_root()
def _serialize_preview_value(value: Any) -> Any:
return to_preview_jsonable(value)
def _serialize_preview_rows(rows: list[dict[str, Any]]) -> list[dict[str, Any]]:
return [
{str(key): _serialize_preview_value(value) for key, value in row.items()}
for row in rows
]
def _normalize_optional_text(value: str | None) -> str | None:
if value is None:
return None
trimmed = value.strip()
return trimmed if trimmed else None
def _list_hf_data_files(*, dataset_name: str, token: str | None) -> list[str]:
try:
from huggingface_hub import HfApi
from huggingface_hub.utils import HfHubHTTPError
except ImportError:
return []
try:
api = HfApi()
repo_files = api.list_repo_files(dataset_name, repo_type = "dataset", token = token)
return [file for file in repo_files if file.lower().endswith(DATA_EXTS)]
except (HfHubHTTPError, OSError, ValueError):
return []
def _select_best_file(data_files: list[str], split: str = DEFAULT_SPLIT) -> str | None:
if not data_files:
return None
split_lower = split.lower()
def score(path: str) -> tuple[int, int]:
name = path.lower()
if f"/{split_lower}/" in name:
return (0, len(path))
if (
f"_{split_lower}." in name
or f"-{split_lower}." in name
or f"/{split_lower}." in name
or f"/{split_lower}_" in name
or f"/{split_lower}-" in name
):
return (1, len(path))
return (2, len(path))
return sorted(data_files, key = score)[0]
def _resolve_seed_hf_path(
dataset_name: str, data_files: list[str], split: str = DEFAULT_SPLIT
) -> str | None:
selected = _select_best_file(data_files, split)
if not selected:
return None
ext = Path(selected).suffix.lower()
if ext not in DATA_EXTS:
return f"datasets/{dataset_name}/{selected}"
parent = Path(selected).parent.as_posix()
if not parent or parent == ".":
return f"datasets/{dataset_name}/**/*{ext}"
return f"datasets/{dataset_name}/{parent}/**/*{ext}"
def _build_stream_load_kwargs(
*,
dataset_name: str,
split: str,
subset: str | None,
token: str | None,
data_file: str | None = None,
) -> dict[str, Any]:
kwargs: dict[str, Any] = {
"path": dataset_name,
"split": split,
"streaming": True,
}
if data_file:
kwargs["data_files"] = [data_file]
if subset:
kwargs["name"] = subset
if token:
kwargs["token"] = token
return kwargs
def _load_preview_rows(
*,
load_dataset_fn,
load_kwargs: dict[str, Any],
preview_size: int,
) -> list[dict[str, Any]]:
streamed_ds = load_dataset_fn(**load_kwargs)
return [row for row in islice(streamed_ds, preview_size)]
def _extract_columns(rows: list[dict[str, Any]]) -> list[str]:
columns_seen: dict[str, None] = {}
for row in rows:
for key in row.keys():
columns_seen[str(key)] = None
return list(columns_seen.keys())
def _sanitize_filename(filename: str) -> str:
name = Path(filename).name.strip().replace("\x00", "")
if not name:
return "seed_upload"
return name
def _decode_base64_payload(content_base64: str) -> bytes:
raw = content_base64.strip()
if "," in raw and raw.lower().startswith("data:"):
raw = raw.split(",", 1)[1]
try:
return base64.b64decode(raw, validate = True)
except binascii.Error as exc:
raise HTTPException(status_code = 400, detail = "invalid base64 payload") from exc
def _read_preview_rows_from_local_file(
path: Path, preview_size: int
) -> list[dict[str, Any]]:
try:
import pandas as pd
except ImportError as exc:
raise HTTPException(
status_code = 500, detail = f"seed inspect dependencies unavailable: {exc}"
) from exc
ext = path.suffix.lower()
try:
if ext == ".csv":
df = pd.read_csv(path, nrows = preview_size)
elif ext == ".jsonl":
df = pd.read_json(path, lines = True).head(preview_size)
elif ext == ".json":
try:
df = pd.read_json(path).head(preview_size)
except ValueError:
df = pd.read_json(path, lines = True).head(preview_size)
else:
raise HTTPException(status_code = 422, detail = f"unsupported file type: {ext}")
except HTTPException:
raise
except (ValueError, OSError) as exc:
raise HTTPException(
status_code = 422, detail = f"seed inspect failed: {exc}"
) from exc
rows = df.to_dict(orient = "records")
return _serialize_preview_rows(rows)
def _read_preview_rows_from_unstructured_file(
*,
path: Path,
preview_size: int,
chunk_size: int | None,
chunk_overlap: int | None,
) -> list[dict[str, Any]]:
size, overlap = resolve_chunking(chunk_size, chunk_overlap)
try:
rows = build_unstructured_preview_rows(
source_path = path,
preview_size = preview_size,
chunk_size = size,
chunk_overlap = overlap,
)
except (FileNotFoundError, RuntimeError, ValueError, OSError) as exc:
raise HTTPException(
status_code = 422, detail = f"seed inspect failed: {exc}"
) from exc
return _serialize_preview_rows(rows)
@router.post("/seed/inspect", response_model = SeedInspectResponse)
def inspect_seed_dataset(payload: SeedInspectRequest) -> SeedInspectResponse:
dataset_name = payload.dataset_name.strip()
if not dataset_name or dataset_name.count("/") < 1:
raise HTTPException(
status_code = 400,
detail = "dataset_name must be a Hugging Face repo id like org/repo",
)
try:
from datasets import load_dataset
except ImportError as exc:
raise HTTPException(
status_code = 500, detail = f"seed inspect dependencies unavailable: {exc}"
) from exc
split = _normalize_optional_text(payload.split) or DEFAULT_SPLIT
subset = _normalize_optional_text(payload.subset)
token = _normalize_optional_text(payload.hf_token)
preview_size = int(payload.preview_size)
preview_rows: list[dict[str, Any]] = []
data_files = _list_hf_data_files(dataset_name = dataset_name, token = token)
selected_file = _select_best_file(data_files, split)
if selected_file:
try:
single_file_kwargs = _build_stream_load_kwargs(
dataset_name = dataset_name,
split = split,
subset = subset,
token = token,
data_file = selected_file,
)
preview_rows = _load_preview_rows(
load_dataset_fn = load_dataset,
load_kwargs = single_file_kwargs,
preview_size = preview_size,
)
except (ValueError, OSError, RuntimeError):
preview_rows = []
if not preview_rows:
try:
split_kwargs = _build_stream_load_kwargs(
dataset_name = dataset_name,
split = split,
subset = subset,
token = token,
)
preview_rows = _load_preview_rows(
load_dataset_fn = load_dataset,
load_kwargs = split_kwargs,
preview_size = preview_size,
)
except (ValueError, OSError, RuntimeError) as exc:
raise HTTPException(
status_code = 422, detail = f"seed inspect failed: {exc}"
) from exc
if not preview_rows:
raise HTTPException(
status_code = 422, detail = "dataset appears empty or unreadable"
)
preview_rows = _serialize_preview_rows(preview_rows)
columns = _extract_columns(preview_rows)
if not data_files:
resolved_path = f"datasets/{dataset_name}/**/*.parquet"
else:
resolved_path = _resolve_seed_hf_path(dataset_name, data_files, split)
if not resolved_path:
raise HTTPException(
status_code = 422, detail = "unable to resolve seed dataset path"
)
return SeedInspectResponse(
dataset_name = dataset_name,
resolved_path = resolved_path,
columns = columns,
preview_rows = preview_rows,
split = split,
subset = subset,
)
@router.post("/seed/inspect-upload", response_model = SeedInspectResponse)
def inspect_seed_upload(payload: SeedInspectUploadRequest) -> SeedInspectResponse:
seed_source_type = _normalize_optional_text(payload.seed_source_type) or "local"
filename = _sanitize_filename(payload.filename)
ext = Path(filename).suffix.lower()
if seed_source_type == "unstructured":
if ext not in UNSTRUCTURED_UPLOAD_EXTS:
allowed = ", ".join(sorted(UNSTRUCTURED_UPLOAD_EXTS))
raise HTTPException(
status_code = 400,
detail = f"unsupported file type: {ext}. allowed: {allowed}",
)
else:
if ext not in LOCAL_UPLOAD_EXTS:
allowed = ", ".join(sorted(LOCAL_UPLOAD_EXTS))
raise HTTPException(
status_code = 400,
detail = f"unsupported file type: {ext}. allowed: {allowed}",
)
file_bytes = _decode_base64_payload(payload.content_base64)
if not file_bytes:
raise HTTPException(status_code = 400, detail = "empty upload payload")
max_size_bytes = 50 * 1024 * 1024
if len(file_bytes) > max_size_bytes:
raise HTTPException(status_code = 413, detail = "file too large (max 50MB)")
ensure_dir(SEED_UPLOAD_DIR)
stored_name = f"{uuid4().hex}_{filename}"
stored_path = SEED_UPLOAD_DIR / stored_name
stored_path.write_bytes(file_bytes)
if seed_source_type == "unstructured":
preview_rows = _read_preview_rows_from_unstructured_file(
path = stored_path,
preview_size = int(payload.preview_size),
chunk_size = payload.unstructured_chunk_size,
chunk_overlap = payload.unstructured_chunk_overlap,
)
else:
preview_rows = _read_preview_rows_from_local_file(
stored_path,
int(payload.preview_size),
)
if not preview_rows:
raise HTTPException(
status_code = 422, detail = "dataset appears empty or unreadable"
)
columns = _extract_columns(preview_rows)
return SeedInspectResponse(
dataset_name = filename,
resolved_path = str(stored_path),
columns = columns,
preview_rows = preview_rows,
split = None,
subset = None,
)