From e30fc87187e6c022b8b9e3ddc1abc54f98bc03f6 Mon Sep 17 00:00:00 2001 From: Shine1i Date: Wed, 4 Mar 2026 20:11:39 +0100 Subject: [PATCH] refactor(studio): add local data-recipe dataset selection + training wiring --- studio/backend/core/training/trainer.py | 31 +- studio/backend/models/datasets.py | 25 +- studio/backend/routes/datasets.py | 194 ++++++- .../sections/dataset-preview-dialog.tsx | 8 +- .../studio/sections/dataset-section.tsx | 497 ++++++++++++++---- .../src/features/studio/studio-page.tsx | 1 + .../src/features/training/api/datasets-api.ts | 13 +- .../src/features/training/api/mappers.ts | 11 +- .../src/features/training/lib/validation.ts | 22 +- .../training/stores/training-config-store.ts | 15 +- .../src/features/training/types/datasets.ts | 18 + .../src/hooks/use-hf-dataset-search.ts | 9 +- .../src/hooks/use-hf-paginated-search.ts | 11 +- 13 files changed, 708 insertions(+), 147 deletions(-) diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index e4e7a474be..9987f0a260 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -17,6 +17,7 @@ import threading import math import logging import time +from pathlib import Path from typing import Optional, Callable from dataclasses import dataclass import pandas as pd @@ -382,17 +383,37 @@ class UnslothTrainer: script_dir = Path(__file__).parent.parent assets_datasets_dir = script_dir / "assets" / "datasets" file_path = assets_datasets_dir / dataset_file - - if str(file_path).endswith('.json'): - with open(file_path, 'r', encoding='utf-8') as f: + + file_path_obj = Path(file_path) + file_path_str = str(file_path_obj) + + if file_path_obj.is_dir(): + parquet_dir = ( + file_path_obj / "parquet-files" + if (file_path_obj / "parquet-files").exists() + else file_path_obj + ) + parquet_files = sorted(parquet_dir.glob("*.parquet")) + if parquet_files: + for parquet_file in parquet_files: + df = pd.read_parquet(parquet_file) + all_data.extend(df.to_dict("records")) + continue + + if file_path_str.endswith('.json'): + with open(file_path_obj, 'r', encoding='utf-8') as f: data = json.load(f) if isinstance(data, list): all_data.extend(data) else: all_data.append(data) - elif str(file_path).endswith('.csv'): - df = pd.read_csv(file_path) + elif file_path_str.endswith('.csv'): + df = pd.read_csv(file_path_obj) all_data.extend(df.to_dict('records')) + elif file_path_str.endswith('.parquet'): + df = pd.read_parquet(file_path_obj) + all_data.extend(df.to_dict('records')) + continue if all_data: dataset = Dataset.from_list(all_data) diff --git a/studio/backend/models/datasets.py b/studio/backend/models/datasets.py index 81adef7577..9e0643415b 100644 --- a/studio/backend/models/datasets.py +++ b/studio/backend/models/datasets.py @@ -1,8 +1,9 @@ """ Dataset-related Pydantic models for API requests and responses. """ -from pydantic import BaseModel, model_validator -from typing import Any, Optional, Dict, List +from typing import Any, Dict, List, Optional + +from pydantic import BaseModel, Field, model_validator class CheckFormatRequest(BaseModel): @@ -34,3 +35,23 @@ class CheckFormatResponse(BaseModel): detected_text_column: Optional[str] = None preview_samples: Optional[List[Dict]] = None total_rows: Optional[int] = None + + +class LocalDatasetItem(BaseModel): + class Metadata(BaseModel): + actual_num_records: Optional[int] = None + target_num_records: Optional[int] = None + total_num_batches: Optional[int] = None + num_completed_batches: Optional[int] = None + columns: Optional[List[str]] = None + + id: str + label: str + path: str + rows: Optional[int] = None + updated_at: Optional[float] = None + metadata: Optional[Metadata] = None + + +class LocalDatasetsResponse(BaseModel): + datasets: List[LocalDatasetItem] = Field(default_factory=list) diff --git a/studio/backend/routes/datasets.py b/studio/backend/routes/datasets.py index cb1ea75e33..61f1bf94da 100644 --- a/studio/backend/routes/datasets.py +++ b/studio/backend/routes/datasets.py @@ -3,6 +3,7 @@ Datasets API routes """ import base64 import io +import json import sys from pathlib import Path from fastapi import APIRouter, HTTPException @@ -29,7 +30,12 @@ if not logger.handlers: logger.setLevel(logging.INFO) -from models.datasets import CheckFormatRequest, CheckFormatResponse +from models.datasets import ( + CheckFormatRequest, + CheckFormatResponse, + LocalDatasetItem, + LocalDatasetsResponse, +) def _serialize_preview_value(value): @@ -81,6 +87,173 @@ DATA_EXTS = ( '.gz', '.zst', '.zip', ) +LOCAL_FILE_EXTS = ('.json', '.jsonl', '.csv', '.parquet') +BACKEND_ROOT = Path(__file__).resolve().parents[1] +LOCAL_DATASETS_ROOT = BACKEND_ROOT / "assets" / "datasets" + + +def _safe_read_metadata(path: Path) -> dict | None: + try: + payload = json.loads(path.read_text(encoding="utf-8")) + except (OSError, ValueError, TypeError): + return None + if not isinstance(payload, dict): + return None + return payload + + +def _safe_read_rows_from_metadata(payload: dict | None) -> int | None: + if not payload: + return None + for key in ("actual_num_records", "target_num_records"): + value = payload.get(key) + if isinstance(value, int): + return value + return None + + +def _safe_read_metadata_summary(payload: dict | None) -> dict | None: + if not payload: + return None + + actual_num_records = ( + payload.get("actual_num_records") + if isinstance(payload.get("actual_num_records"), int) + else None + ) + target_num_records = ( + payload.get("target_num_records") + if isinstance(payload.get("target_num_records"), int) + else actual_num_records + ) + + columns: list[str] | None = None + schema = payload.get("schema") + if isinstance(schema, dict): + columns = [str(key) for key in schema.keys()] + if not columns: + stats = payload.get("column_statistics") + if isinstance(stats, list): + derived = [ + str(item.get("column_name")) + for item in stats + if isinstance(item, dict) and item.get("column_name") + ] + columns = derived or None + + parquet_files_count = None + file_paths = payload.get("file_paths") + if isinstance(file_paths, dict): + parquet_files = file_paths.get("parquet-files") + if isinstance(parquet_files, list): + parquet_files_count = len(parquet_files) + + total_num_batches = ( + payload.get("total_num_batches") + if isinstance(payload.get("total_num_batches"), int) + else parquet_files_count + ) + num_completed_batches = ( + payload.get("num_completed_batches") + if isinstance(payload.get("num_completed_batches"), int) + else total_num_batches + ) + + return { + "actual_num_records": actual_num_records, + "target_num_records": target_num_records, + "total_num_batches": total_num_batches, + "num_completed_batches": num_completed_batches, + "columns": columns, + } + + +def _build_local_dataset_items() -> list[LocalDatasetItem]: + if not LOCAL_DATASETS_ROOT.exists(): + return [] + + items: list[LocalDatasetItem] = [] + for entry in LOCAL_DATASETS_ROOT.iterdir(): + if not entry.is_dir() or not entry.name.startswith("recipe_"): + continue + parquet_dir = entry / "parquet-files" + if not parquet_dir.exists() or not any(parquet_dir.glob("*.parquet")): + continue + + rows = None + metadata_summary = None + metadata_path = entry / "metadata.json" + if metadata_path.exists(): + metadata_payload = _safe_read_metadata(metadata_path) + rows = _safe_read_rows_from_metadata(metadata_payload) + metadata_summary = _safe_read_metadata_summary(metadata_payload) + + try: + updated_at = entry.stat().st_mtime + except OSError: + updated_at = None + + items.append( + LocalDatasetItem( + id=entry.name, + label=entry.name, + path=str(parquet_dir.resolve()), + rows=rows, + updated_at=updated_at, + metadata=metadata_summary, + ) + ) + + items.sort(key=lambda item: item.updated_at or 0, reverse=True) + return items + + +def _load_local_preview_slice(*, dataset_path: Path, train_split: str, preview_size: int): + from datasets import load_dataset + + if dataset_path.is_dir(): + parquet_dir = dataset_path / "parquet-files" if (dataset_path / "parquet-files").exists() else dataset_path + parquet_files = sorted(parquet_dir.glob("*.parquet")) + if parquet_files: + dataset = load_dataset( + "parquet", + data_files=[str(path) for path in parquet_files], + split=train_split, + ) + total_rows = len(dataset) + preview_slice = dataset.select(range(min(preview_size, total_rows))) + return preview_slice, total_rows + else: + candidate_files: list[Path] = [] + for ext in LOCAL_FILE_EXTS: + candidate_files.extend(sorted(dataset_path.glob(f"*{ext}"))) + if not candidate_files: + raise HTTPException( + status_code=400, + detail="Unsupported local dataset directory (expected parquet/json/jsonl/csv files)", + ) + dataset_path = candidate_files[0] + + if dataset_path.suffix in ['.json', '.jsonl']: + dataset = load_dataset('json', data_files=str(dataset_path), split=train_split) + elif dataset_path.suffix == '.csv': + dataset = load_dataset('csv', data_files=str(dataset_path), split=train_split) + elif dataset_path.suffix == '.parquet': + dataset = load_dataset('parquet', data_files=str(dataset_path), split=train_split) + else: + raise HTTPException( + status_code=400, + detail=f"Unsupported file format: {dataset_path.suffix}" + ) + + total_rows = len(dataset) + preview_slice = dataset.select(range(min(preview_size, total_rows))) + return preview_slice, total_rows + + +@router.get("/local", response_model=LocalDatasetsResponse) +def list_local_datasets() -> LocalDatasetsResponse: + return LocalDatasetsResponse(datasets=_build_local_dataset_items()) @router.post("/check-format", response_model=CheckFormatResponse) @@ -112,19 +285,12 @@ def check_format(request: CheckFormatRequest): if dataset_path.exists(): # ── Local file ────────────────────────────────────────── - if dataset_path.suffix in ['.json', '.jsonl']: - dataset = load_dataset('json', data_files=str(dataset_path), split=request.train_split) - elif dataset_path.suffix == '.csv': - dataset = load_dataset('csv', data_files=str(dataset_path), split=request.train_split) - elif dataset_path.suffix == '.parquet': - dataset = load_dataset('parquet', data_files=str(dataset_path), split=request.train_split) - else: - raise HTTPException( - status_code=400, - detail=f"Unsupported file format: {dataset_path.suffix}" - ) - total_rows = len(dataset) - preview_slice = dataset.select(range(min(PREVIEW_SIZE, total_rows))) + train_split = request.train_split or "train" + preview_slice, total_rows = _load_local_preview_slice( + dataset_path=dataset_path, + train_split=train_split, + preview_size=PREVIEW_SIZE, + ) else: # ── HuggingFace dataset ───────────────────────────────── # Tier 1: list_repo_files → load only the first data file 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 00a738bc64..3135087383 100644 --- a/studio/frontend/src/features/studio/sections/dataset-preview-dialog.tsx +++ b/studio/frontend/src/features/studio/sections/dataset-preview-dialog.tsx @@ -30,6 +30,7 @@ type DatasetPreviewDialogProps = { open: boolean; onOpenChange: (open: boolean) => void; datasetName: string | null; + datasetSource?: "huggingface" | "upload"; hfToken: string | null; datasetSubset?: string | null; datasetSplit?: string | null; @@ -42,6 +43,7 @@ export function DatasetPreviewDialog({ open, onOpenChange, datasetName, + datasetSource, hfToken, datasetSubset, datasetSplit, @@ -71,7 +73,7 @@ export function DatasetPreviewDialog({ const showMappingFooter = mode === "mapping" && mappingEnabled; const mappingOk = isMappingComplete(manualMapping, effectiveIsVlm, datasetFormat); const availableRoles = getAvailableRoles(effectiveIsVlm, datasetFormat); - const isHfDataset = !!datasetName && datasetName.includes("/"); + const isHfDataset = datasetSource === "huggingface"; // When format changes, remap existing mapping roles to the new format's role names const prevFormatRef = useRef(datasetFormat); @@ -161,7 +163,7 @@ export function DatasetPreviewDialog({ // Determine source label const sourceLabel = useMemo(() => { if (!datasetName) return ""; - if (datasetName.includes("/")) { + if (datasetSource === "huggingface") { let label = `Hugging Face (${datasetName}`; if (datasetSubset) label += ` / ${datasetSubset}`; if (datasetSplit) label += ` / ${datasetSplit}`; @@ -169,7 +171,7 @@ export function DatasetPreviewDialog({ return label; } return `Local Files (${datasetName})`; - }, [datasetName, datasetSubset, datasetSplit]); + }, [datasetName, datasetSource, datasetSubset, datasetSplit]); // Build TanStack Table columns from the column names const tableColumns = useMemo>[]>(() => { diff --git a/studio/frontend/src/features/studio/sections/dataset-section.tsx b/studio/frontend/src/features/studio/sections/dataset-section.tsx index 57e50bced5..9fff45cb09 100644 --- a/studio/frontend/src/features/studio/sections/dataset-section.tsx +++ b/studio/frontend/src/features/studio/sections/dataset-section.tsx @@ -22,6 +22,7 @@ import { SelectValue, } from "@/components/ui/select"; import { Spinner } from "@/components/ui/spinner"; +import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; import { Tooltip, TooltipContent, @@ -38,6 +39,8 @@ import { useDatasetPreviewDialogStore, useTrainingConfigStore, } from "@/features/training"; +import { listLocalDatasets } from "@/features/training/api/datasets-api"; +import type { LocalDatasetInfo } from "@/features/training/types/datasets"; import { ArrowDown01Icon, CloudUploadIcon, @@ -48,7 +51,7 @@ import { ViewIcon, } from "@hugeicons/core-free-icons"; import { HugeiconsIcon } from "@hugeicons/react"; -import { useMemo, useRef, useState } from "react"; +import { useCallback, useEffect, useMemo, useRef, useState } from "react"; import { useShallow } from "zustand/react/shallow"; function isLikelyLocalDatasetRef(value: string) { @@ -61,10 +64,20 @@ function isLikelyLocalDatasetRef(value: string) { ); } +function deriveLocalDatasetName(path: string): string { + const normalized = path.replaceAll("\\", "/"); + const parts = normalized.split("/").filter(Boolean); + const parquetIndex = parts.lastIndexOf("parquet-files"); + if (parquetIndex > 0) return parts[parquetIndex - 1]; + return parts[parts.length - 1] ?? path; +} + export function DatasetSection() { const { dataset, setDataset, + datasetSource, + setDatasetSource, datasetFormat, setDatasetFormat, datasetSubset, @@ -73,12 +86,16 @@ export function DatasetSection() { setDatasetSplit, datasetEvalSplit, setDatasetEvalSplit, + uploadedFile, + setUploadedFile, hfToken, modelType, } = useTrainingConfigStore( useShallow((s) => ({ dataset: s.dataset, setDataset: s.setDataset, + datasetSource: s.datasetSource, + setDatasetSource: s.setDatasetSource, datasetFormat: s.datasetFormat, setDatasetFormat: s.setDatasetFormat, datasetSubset: s.datasetSubset, @@ -87,6 +104,8 @@ export function DatasetSection() { setDatasetSplit: s.setDatasetSplit, datasetEvalSplit: s.datasetEvalSplit, setDatasetEvalSplit: s.setDatasetEvalSplit, + uploadedFile: s.uploadedFile, + setUploadedFile: s.setUploadedFile, hfToken: s.hfToken, modelType: s.modelType, })), @@ -94,13 +113,56 @@ export function DatasetSection() { const [inputValue, setInputValue] = useState(""); const [advancedOpen, setAdvancedOpen] = useState(false); + const [pickerTab, setPickerTab] = useState<"huggingface" | "local">( + datasetSource === "upload" ? "local" : "huggingface", + ); + const [localDatasets, setLocalDatasets] = useState([]); + const [localLoading, setLocalLoading] = useState(false); + const [localError, setLocalError] = useState(null); + const [localSearchActive, setLocalSearchActive] = useState(false); const openPreview = useDatasetPreviewDialogStore((s) => s.openPreview); const selectingRef = useRef(false); const debouncedQuery = useDebouncedValue(inputValue); + useEffect(() => { + setPickerTab(datasetSource === "upload" ? "local" : "huggingface"); + }, [datasetSource]); + + const refreshLocalDatasets = useCallback(async () => { + setLocalLoading(true); + setLocalError(null); + try { + const response = await listLocalDatasets(); + setLocalDatasets(response.datasets ?? []); + } catch (error) { + setLocalError( + error instanceof Error ? error.message : "Failed to load local datasets.", + ); + } finally { + setLocalLoading(false); + } + }, []); + + useEffect(() => { + if (pickerTab !== "local") return; + void refreshLocalDatasets(); + }, [pickerTab, refreshLocalDatasets]); + function handleDatasetSelect(id: string | null) { selectingRef.current = true; + setDatasetSource("huggingface"); setDataset(id); + setInputValue(id ?? ""); + setLocalSearchActive(false); + } + + function handleLocalDatasetSelect(path: string) { + selectingRef.current = true; + setDatasetSource("upload"); + setUploadedFile(path); + const label = localDatasets.find((item) => item.path === path)?.label; + setInputValue(label ?? deriveLocalDatasetName(path)); + setLocalSearchActive(false); } function handleInputChange(val: string) { @@ -108,6 +170,9 @@ export function DatasetSection() { selectingRef.current = false; return; } + if (pickerTab === "local") { + setLocalSearchActive(true); + } setInputValue(val); } const { @@ -116,15 +181,16 @@ export function DatasetSection() { isLoadingMore, fetchMore, error: hfSearchError, - } = useHfDatasetSearch(debouncedQuery, { + } = useHfDatasetSearch(pickerTab === "huggingface" ? debouncedQuery : "", { modelType, accessToken: hfToken || undefined, + enabled: pickerTab === "huggingface", }); const { error: tokenValidationError, isChecking: isCheckingToken } = useHfTokenValidation(hfToken); - const resultIds = useMemo(() => { + const hfResultIds = useMemo(() => { const ids = hfResults.map((r) => r.id); if (dataset && !ids.includes(dataset)) { ids.push(dataset); @@ -132,6 +198,57 @@ export function DatasetSection() { return ids; }, [hfResults, dataset]); + const localFilteredDatasets = useMemo(() => { + const query = localSearchActive ? inputValue.trim().toLowerCase() : ""; + if (!query) return localDatasets; + return localDatasets.filter( + (item) => + item.label.toLowerCase().includes(query) || + item.path.toLowerCase().includes(query), + ); + }, [localDatasets, inputValue, localSearchActive]); + + const localPathById = useMemo(() => { + return new Map(localDatasets.map((item) => [item.id, item.path])); + }, [localDatasets]); + + const localLabelById = useMemo(() => { + return new Map(localDatasets.map((item) => [item.id, item.label])); + }, [localDatasets]); + + const selectedLocalId = useMemo(() => { + if (!uploadedFile) return null; + const item = localDatasets.find((entry) => entry.path === uploadedFile); + return item?.id ?? deriveLocalDatasetName(uploadedFile); + }, [localDatasets, uploadedFile]); + + const localResultIds = useMemo(() => { + const ids = localFilteredDatasets.map((item) => item.id); + if (selectedLocalId && !ids.includes(selectedLocalId)) { + ids.push(selectedLocalId); + } + return ids; + }, [localFilteredDatasets, selectedLocalId]); + + const comboboxItems = pickerTab === "huggingface" ? hfResultIds : localResultIds; + const comboboxValue = + pickerTab === "huggingface" ? dataset : selectedLocalId; + const isHfDatasetSelected = + datasetSource === "huggingface" && + !!dataset && + !isLikelyLocalDatasetRef(dataset); + + const selectedDatasetName = datasetSource === "upload" ? uploadedFile : dataset; + const selectedLocalDataset = useMemo(() => { + if (!uploadedFile) return null; + return localDatasets.find((item) => item.path === uploadedFile) ?? null; + }, [localDatasets, uploadedFile]); + const selectedLocalMetadata = selectedLocalDataset?.metadata ?? null; + const selectedLocalColumns = selectedLocalMetadata?.columns ?? []; + const selectedLocalRows = + selectedLocalDataset?.rows ?? selectedLocalMetadata?.actual_num_records ?? null; + const selectedLocalUpdatedAt = selectedLocalDataset?.updated_at ?? null; + const comboboxAnchorRef = useRef(null); const { scrollRef, sentinelRef } = useInfiniteScroll( fetchMore, @@ -147,7 +264,7 @@ export function DatasetSection() { accent="indigo" className="md:min-h-[470px] dark:shadow-border" > -
+
Load from Hub @@ -183,22 +300,47 @@ export function DatasetSection() { if (event.key !== "Enter") return; if (!(event.target instanceof HTMLInputElement)) return; event.preventDefault(); - if (hfResults.length > 0) { - handleDatasetSelect(hfResults[0].id); - } else { - const text = event.target.value.trim(); - if (text) handleDatasetSelect(text); + if (pickerTab === "huggingface") { + if (hfResults.length > 0) { + handleDatasetSelect(hfResults[0].id); + } else { + const text = event.target.value.trim(); + if (text) handleDatasetSelect(text); + } + return; + } + + if (localResultIds.length > 0) { + const selectedId = localResultIds[0]; + const path = localPathById.get(selectedId); + if (path) { + handleLocalDatasetSelect(path); + } } }} > { + if (!value) return; + if (pickerTab === "huggingface") { + handleDatasetSelect(value); + return; + } + const path = localPathById.get(value); + if (path) { + handleLocalDatasetSelect(path); + } + }} onInputValueChange={handleInputChange} - itemToStringValue={(id) => id} + itemToStringValue={(id) => + pickerTab === "local" + ? localLabelById.get(id) ?? id + : id + } autoHighlight={true} > - {isLoading ? ( -
- Searching... -
- ) : ( - No datasets found - )} -
- - {(id: string) => { - return ( - - - - - {id} - - - - {id} - - - - ); +
+ { + setPickerTab(value as "huggingface" | "local"); + setInputValue(""); + setLocalSearchActive(false); }} - -
- {isLoadingMore && ( -
- -
- )} + className="w-full" + > + + Hugging Face + Local + + + + {isLoading ? ( +
+ Searching... +
+ ) : ( + No datasets found + )} +
+ + {(id: string) => { + return ( + + + + + {id} + + + + {id} + + + + ); + }} + +
+ {isLoadingMore && ( +
+ +
+ )} +
+ + + + {localLoading ? ( +
+ Loading local datasets... +
+ ) : ( + <> + {localError ? ( +

{localError}

+ ) : ( + +
+

+ {localDatasets.length === 0 + ? "No local datasets yet." + : "No local datasets match search."} +

+ {localDatasets.length === 0 ? ( + + ) : null} +
+
+ )} +
+ + {(id: string) => { + const label = localLabelById.get(id) ?? id; + return ( + + + + + {label} + + + + {label} + + + + ); + }} + +
+ + )} +
+
@@ -273,8 +487,8 @@ export function DatasetSection() { + {datasetSource === "upload" && ( +
+
+

Dataset Metadata

+

Data Recipe Output

+
+ + {selectedLocalDataset ? ( + <> +
+ + 0 + ? String(selectedLocalColumns.length) + : "--" + } + /> + + +
+ + ) : ( +

+ Select a local dataset to view metadata. +

+ )} +
+ )} + - {dataset ? ( -
-
- -
-
-

- {dataset} -

-

- Hugging Face Dataset - {datasetSubset && ` / ${datasetSubset}`} - {datasetSplit && ` / ${datasetSplit}`} -

-
-
- ) : ( -
- - - No dataset selected - -
- )} +
+ {selectedDatasetName ? ( +
+
+ +
+
+

+ {datasetSource === "upload" + ? selectedLocalDataset?.label ?? + deriveLocalDatasetName(selectedDatasetName) + : selectedDatasetName} +

+

+ {datasetSource === "upload" ? ( + selectedLocalDataset && typeof selectedLocalDataset.rows === "number" ? ( + `${selectedLocalDataset.rows.toLocaleString()} rows` + ) : ( + "Local dataset" + ) + ) : ( + <> + Hugging Face Dataset + {datasetSubset && ` / ${datasetSubset}`} + {datasetSplit && ` / ${datasetSplit}`} + + )} +

+
+
+ ) : ( +
+ + + No dataset selected + +
+ )} -
- - -
+
+ + +
+
); } + +function MetadataRow({ label, value }: { label: string; value: string }) { + return ( +
+ {label} + {value} +
+ ); +} diff --git a/studio/frontend/src/features/studio/studio-page.tsx b/studio/frontend/src/features/studio/studio-page.tsx index 75fde25e07..54a3f2023f 100644 --- a/studio/frontend/src/features/studio/studio-page.tsx +++ b/studio/frontend/src/features/studio/studio-page.tsx @@ -85,6 +85,7 @@ export function StudioPage(): ReactElement { onOpenChange={(open) => { if (!open) closeDialog(); }} + datasetSource={config.datasetSource} datasetName={ config.datasetSource === "huggingface" ? config.dataset : config.uploadedFile } diff --git a/studio/frontend/src/features/training/api/datasets-api.ts b/studio/frontend/src/features/training/api/datasets-api.ts index 7bba75ca38..7149271c0e 100644 --- a/studio/frontend/src/features/training/api/datasets-api.ts +++ b/studio/frontend/src/features/training/api/datasets-api.ts @@ -1,4 +1,7 @@ -import type { CheckFormatResponse } from "../types/datasets"; +import type { + CheckFormatResponse, + LocalDatasetsResponse, +} from "../types/datasets"; type CheckDatasetFormatArgs = { datasetName: string; @@ -35,3 +38,11 @@ export async function checkDatasetFormat({ return res.json(); } +export async function listLocalDatasets(): Promise { + const res = await fetch("/api/datasets/local"); + if (!res.ok) { + const body = await res.json().catch(() => null); + throw new Error(body?.detail || `Request failed (${res.status})`); + } + return res.json(); +} diff --git a/studio/frontend/src/features/training/api/mappers.ts b/studio/frontend/src/features/training/api/mappers.ts index 1adfbd8b6d..80a2b3aaf8 100644 --- a/studio/frontend/src/features/training/api/mappers.ts +++ b/studio/frontend/src/features/training/api/mappers.ts @@ -12,8 +12,12 @@ export function buildTrainingStartPayload( config: TrainingConfigState, ): TrainingStartRequest { const adapterMethod = config.trainingMethod !== "full"; - const isQlorMethod = config.trainingMethod === "qlora"; + const isQloraMethod = config.trainingMethod === "qlora"; const hfDataset = config.datasetSource === "huggingface" ? config.dataset : null; + const localDatasets = + config.datasetSource === "upload" && config.uploadedFile + ? [config.uploadedFile] + : []; const customFormatMapping = Object.keys(config.datasetManualMapping).length > 0 ? config.datasetManualMapping : undefined; @@ -21,13 +25,13 @@ export function buildTrainingStartPayload( model_name: config.selectedModel ?? "", training_type: toBackendTrainingType(config.trainingMethod), hf_token: config.hfToken.trim() || null, - load_in_4bit: adapterMethod ? isQlorMethod : false, + load_in_4bit: adapterMethod ? isQloraMethod : false, max_seq_length: config.contextLength, hf_dataset: hfDataset, subset: hfDataset ? config.datasetSubset : null, train_split: hfDataset ? config.datasetSplit : null, eval_split: hfDataset ? config.datasetEvalSplit : null, - local_datasets: [], + local_datasets: localDatasets, format_type: config.datasetFormat, custom_format_mapping: customFormatMapping, num_epochs: config.epochs, @@ -69,4 +73,3 @@ export function buildTrainingStartPayload( : null, }; } - diff --git a/studio/frontend/src/features/training/lib/validation.ts b/studio/frontend/src/features/training/lib/validation.ts index 8e966153d3..ae89cdaf51 100644 --- a/studio/frontend/src/features/training/lib/validation.ts +++ b/studio/frontend/src/features/training/lib/validation.ts @@ -12,16 +12,22 @@ export function validateTrainingConfig( return { ok: false, message: "Select a base model first." }; } - if (config.datasetSource !== "huggingface") { - return { - ok: false, - message: "Only Hugging Face dataset source is enabled right now.", - }; + if (config.datasetSource === "huggingface") { + if (!config.dataset) { + return { ok: false, message: "Select a Hugging Face dataset first." }; + } + return { ok: true, message: null }; } - if (!config.dataset) { - return { ok: false, message: "Select a Hugging Face dataset first." }; + if (config.datasetSource === "upload") { + if (!config.uploadedFile) { + return { ok: false, message: "Select a local dataset first." }; + } + return { ok: true, message: null }; } - return { ok: true, message: null }; + return { + ok: false, + message: "Unsupported dataset source.", + }; } diff --git a/studio/frontend/src/features/training/stores/training-config-store.ts b/studio/frontend/src/features/training/stores/training-config-store.ts index b2d1858716..22218e62a9 100644 --- a/studio/frontend/src/features/training/stores/training-config-store.ts +++ b/studio/frontend/src/features/training/stores/training-config-store.ts @@ -311,7 +311,20 @@ export const useTrainingConfigStore = create()( }, setDatasetManualMapping: (datasetManualMapping) => set({ datasetManualMapping }), - setUploadedFile: (uploadedFile) => set({ uploadedFile }), + setUploadedFile: (uploadedFile) => { + _datasetCheckController?.abort(); + _datasetCheckController = null; + _trainOnCompletionsManuallySet = false; + set({ + uploadedFile, + datasetSubset: null, + datasetSplit: null, + datasetEvalSplit: null, + datasetManualMapping: emptyManualMapping(), + isDatasetMultimodal: null, + isCheckingDataset: false, + }); + }, setEpochs: (epochs) => set({ epochs }), setContextLength: (contextLength) => set({ contextLength }), setLearningRate: (learningRate) => set({ learningRate }), diff --git a/studio/frontend/src/features/training/types/datasets.ts b/studio/frontend/src/features/training/types/datasets.ts index 96cf699f89..15f9a2d0b8 100644 --- a/studio/frontend/src/features/training/types/datasets.ts +++ b/studio/frontend/src/features/training/types/datasets.ts @@ -11,3 +11,21 @@ export type CheckFormatResponse = { multimodal_columns?: string[] | null; }; +export type LocalDatasetInfo = { + metadata?: { + actual_num_records?: number | null; + target_num_records?: number | null; + total_num_batches?: number | null; + num_completed_batches?: number | null; + columns?: string[] | null; + } | null; + id: string; + label: string; + path: string; + rows?: number | null; + updated_at?: number | null; +}; + +export type LocalDatasetsResponse = { + datasets: LocalDatasetInfo[]; +}; diff --git a/studio/frontend/src/hooks/use-hf-dataset-search.ts b/studio/frontend/src/hooks/use-hf-dataset-search.ts index ec64ffd7e6..f051ec1a68 100644 --- a/studio/frontend/src/hooks/use-hf-dataset-search.ts +++ b/studio/frontend/src/hooks/use-hf-dataset-search.ts @@ -277,9 +277,9 @@ function isOcrOrVisionTextDataset(dataset: HfDatasetResult): boolean { export function useHfDatasetSearch( query: string, - options?: { modelType?: ModelType | null; accessToken?: string }, + options?: { modelType?: ModelType | null; accessToken?: string; enabled?: boolean }, ) { - const { modelType, accessToken } = options ?? {}; + const { modelType, accessToken, enabled = true } = options ?? {}; const createIter = useCallback( () => listDatasets({ @@ -291,9 +291,10 @@ export function useHfDatasetSearch( [query, accessToken], ); - const search = useHfPaginatedSearch(createIter, mapDataset); + const search = useHfPaginatedSearch(createIter, mapDataset, { enabled }); const results = useMemo(() => { + if (!enabled) return []; const hideOcr = modelType !== "vision"; const baseResults = hideOcr ? search.results.filter((ds) => !isOcrOrVisionTextDataset(ds)) @@ -311,7 +312,7 @@ export function useHfDatasetSearch( } return [...boosted, ...neutral]; - }, [search.results, modelType]); + }, [enabled, search.results, modelType]); return { ...search, results }; } diff --git a/studio/frontend/src/hooks/use-hf-paginated-search.ts b/studio/frontend/src/hooks/use-hf-paginated-search.ts index a57b0664c0..f78f560dd6 100644 --- a/studio/frontend/src/hooks/use-hf-paginated-search.ts +++ b/studio/frontend/src/hooks/use-hf-paginated-search.ts @@ -39,14 +39,16 @@ async function pullBatch( export function useHfPaginatedSearch( createIter: () => AsyncGenerator, mapItem: (raw: unknown) => T | null, + options?: { enabled?: boolean }, ): HfPaginatedState & { fetchMore: () => void } { + const enabled = options?.enabled ?? true; const [state, setState] = useState>( INITIAL as HfPaginatedState, ); const stateRef = useRef(state); useEffect(() => { stateRef.current = state; - }); + }, [state]); const iterRef = useRef | null>(null); const versionRef = useRef(0); @@ -55,6 +57,11 @@ export function useHfPaginatedSearch( const v = ++versionRef.current; iterRef.current = null; + if (!enabled) { + setState(INITIAL as HfPaginatedState); + return; + } + setState({ ...(INITIAL as HfPaginatedState), isLoading: true, @@ -88,7 +95,7 @@ export function useHfPaginatedSearch( error: err instanceof Error ? err.message : "Search failed", }); }); - }, [createIter, mapItem]); + }, [createIter, mapItem, enabled]); const fetchMore = useCallback(() => { const iter = iterRef.current;