diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index cee84a8809..ecc11c30d5 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -2200,11 +2200,59 @@ class UnslothTrainer: return (train_data, eval_data) + @staticmethod + def _resolve_local_files(file_paths: list) -> list[str]: + """Resolve a list of local dataset paths to concrete file paths.""" + all_files: list[str] = [] + for dataset_file in file_paths: + if os.path.isabs(dataset_file): + file_path = dataset_file + else: + file_path = str(resolve_dataset_path(dataset_file)) + + file_path_obj = Path(file_path) + + 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: + all_files.extend(str(p) for p in parquet_files) + continue + candidates: list[Path] = [] + for ext in (".json", ".jsonl", ".csv", ".parquet"): + candidates.extend(sorted(file_path_obj.glob(f"*{ext}"))) + if candidates: + all_files.extend(str(c) for c in candidates) + continue + raise ValueError( + f"No supported data files in directory: {file_path_obj}" + ) + else: + all_files.append(str(file_path_obj)) + return all_files + + @staticmethod + def _loader_for_files(files: list[str]) -> str: + """Determine the HF datasets loader type from file extensions.""" + first_ext = Path(files[0]).suffix.lower() + if first_ext in (".json", ".jsonl"): + return "json" + elif first_ext == ".csv": + return "csv" + elif first_ext == ".parquet": + return "parquet" + raise ValueError(f"Unsupported dataset format: {files[0]}") + def load_and_format_dataset( self, dataset_source: str, format_type: str = "auto", local_datasets: list = None, + local_eval_datasets: list = None, custom_format_mapping: dict = None, subset: str = None, train_split: str = "train", @@ -2236,54 +2284,10 @@ class UnslothTrainer: # Arrow-backed (has cache files). Dataset.from_list() creates # an in-memory dataset with no cache, which forces num_proc=1 # during tokenization/map because sharding requires Arrow files. - all_files: list[str] = [] - for dataset_file in local_datasets: - # dataset_file may already be an absolute path from routes/training.py - if os.path.isabs(dataset_file): - file_path = dataset_file - else: - # Fallback: try relative to assets/datasets - file_path = str(resolve_dataset_path(dataset_file)) - - file_path_obj = Path(file_path) - - 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: - all_files.extend(str(p) for p in parquet_files) - continue - # Fall through to single-file detection for dirs with json/csv - candidates: list[Path] = [] - for ext in (".json", ".jsonl", ".csv", ".parquet"): - candidates.extend(sorted(file_path_obj.glob(f"*{ext}"))) - if candidates: - all_files.extend(str(c) for c in candidates) - continue - raise ValueError( - f"No supported data files in directory: {file_path_obj}" - ) - else: - all_files.append(str(file_path_obj)) + all_files = self._resolve_local_files(local_datasets) if all_files: - # Determine loader type from the first file extension - first_ext = Path(all_files[0]).suffix.lower() - if first_ext in (".json", ".jsonl"): - loader = "json" - elif first_ext == ".csv": - loader = "csv" - elif first_ext == ".parquet": - loader = "parquet" - else: - raise ValueError( - f"Unsupported local dataset format: {all_files[0]}" - ) - + loader = self._loader_for_files(all_files) dataset = load_dataset(loader, data_files = all_files, split = "train") # Check if stopped during dataset loading @@ -2297,6 +2301,19 @@ class UnslothTrainer: logger.info(f"Loaded {len(dataset)} samples from local files\n") logger.info(f"[DEBUG] Dataset cache_files: {dataset.cache_files}\n") + # Load local eval datasets if provided + if local_eval_datasets and eval_enabled: + eval_all_files = self._resolve_local_files(local_eval_datasets) + if eval_all_files: + eval_loader = self._loader_for_files(eval_all_files) + eval_dataset = load_dataset( + eval_loader, data_files = eval_all_files, split = "train" + ) + has_separate_eval_source = True + logger.info( + f"Loaded {len(eval_dataset)} eval samples from local eval files\n" + ) + elif dataset_source: # Load from Hugging Face split_name = train_split or "train" diff --git a/studio/backend/core/training/training.py b/studio/backend/core/training/training.py index 96a92af942..9626f9df2e 100644 --- a/studio/backend/core/training/training.py +++ b/studio/backend/core/training/training.py @@ -141,6 +141,7 @@ class TrainingBackend: "max_seq_length": kwargs.get("max_seq_length", 2048), "hf_dataset": kwargs.get("hf_dataset", ""), "local_datasets": kwargs.get("local_datasets"), + "local_eval_datasets": kwargs.get("local_eval_datasets"), "format_type": kwargs.get("format_type", ""), "subset": kwargs.get("subset"), "train_split": kwargs.get("train_split", "train"), diff --git a/studio/backend/core/training/worker.py b/studio/backend/core/training/worker.py index b5cf984bf9..f0e8087672 100644 --- a/studio/backend/core/training/worker.py +++ b/studio/backend/core/training/worker.py @@ -280,6 +280,7 @@ def run_training_process( dataset_source = hf_dataset if hf_dataset and hf_dataset.strip() else None, format_type = config.get("format_type", ""), local_datasets = config.get("local_datasets") or None, + local_eval_datasets = config.get("local_eval_datasets") or None, custom_format_mapping = config.get("custom_format_mapping"), subset = config.get("subset"), train_split = config.get("train_split", "train"), diff --git a/studio/backend/models/training.py b/studio/backend/models/training.py index 35bd9744a9..342e44cc09 100644 --- a/studio/backend/models/training.py +++ b/studio/backend/models/training.py @@ -34,6 +34,9 @@ class TrainingStartRequest(BaseModel): local_datasets: List[str] = Field( default_factory = list, description = "List of local dataset paths" ) + local_eval_datasets: List[str] = Field( + default_factory = list, description = "List of local eval dataset paths" + ) format_type: str = Field(..., description = "Dataset format type") subset: Optional[str] = None train_split: Optional[str] = Field("train", description = "Training split name") diff --git a/studio/backend/routes/datasets.py b/studio/backend/routes/datasets.py index 415c30cd72..611a34e39d 100644 --- a/studio/backend/routes/datasets.py +++ b/studio/backend/routes/datasets.py @@ -284,25 +284,17 @@ async def upload_dataset( detail = f"Unsupported file type: {ext}. Allowed: {allowed}", ) - max_size_bytes = 512 * 1024 * 1024 ensure_dir(DATASET_UPLOAD_DIR) stem = Path(filename).stem stored_name = f"{uuid4().hex}_{stem}{ext}" stored_path = DATASET_UPLOAD_DIR / stored_name # Stream file to disk in chunks to avoid holding entire file in memory - size = 0 with open(stored_path, "wb") as f: while chunk := await file.read(1024 * 1024): - size += len(chunk) - if size > max_size_bytes: - stored_path.unlink(missing_ok = True) - raise HTTPException( - status_code = 413, detail = "File too large (max 512MB)" - ) f.write(chunk) - if size == 0: + if stored_path.stat().st_size == 0: stored_path.unlink(missing_ok = True) raise HTTPException(status_code = 400, detail = "Empty upload payload") diff --git a/studio/backend/routes/training.py b/studio/backend/routes/training.py index 43b5212010..4f8054f80e 100644 --- a/studio/backend/routes/training.py +++ b/studio/backend/routes/training.py @@ -56,6 +56,29 @@ router = APIRouter() logger = get_logger(__name__) +def _validate_local_dataset_paths( + paths: list[str], label: str = "Local dataset" +) -> list[str]: + """Resolve and validate a list of local dataset paths. Returns validated absolute paths.""" + validated = [] + missing = [] + for dataset_path in paths: + dataset_file = resolve_dataset_path(dataset_path) + if not dataset_file.exists(): + missing.append(f"{dataset_path} (resolved: {dataset_file})") + continue + logger.info(f"Found {label.lower()} file: {dataset_file}") + validated.append(str(dataset_file)) + + if missing: + missing_detail = "; ".join(missing[:3]) + raise HTTPException( + status_code = 400, + detail = f"{label} not found: {missing_detail}", + ) + return validated + + @router.get("/hardware") async def get_hardware_utilization( current_subject: str = Depends(get_current_subject), @@ -111,27 +134,13 @@ async def start_training( # Validate dataset paths if provided if request.local_datasets: - validated_datasets = [] - missing_datasets = [] - for dataset_path in request.local_datasets: - dataset_file = resolve_dataset_path(dataset_path) - - if not dataset_file.exists(): - missing_datasets.append( - f"{dataset_path} (resolved: {dataset_file})" - ) - continue - - logger.info(f"Found dataset file: {dataset_file}") - validated_datasets.append(str(dataset_file)) - - if missing_datasets: - missing_detail = "; ".join(missing_datasets[:3]) - raise HTTPException( - status_code = 400, - detail = f"Local dataset not found: {missing_detail}", - ) - request.local_datasets = validated_datasets + request.local_datasets = _validate_local_dataset_paths( + request.local_datasets, "Local dataset" + ) + if request.local_eval_datasets and request.eval_steps > 0: + request.local_eval_datasets = _validate_local_dataset_paths( + request.local_eval_datasets, "Local eval dataset" + ) # Convert request to kwargs for backend training_kwargs = { @@ -142,6 +151,7 @@ async def start_training( "max_seq_length": request.max_seq_length, "hf_dataset": request.hf_dataset or "", "local_datasets": request.local_datasets, + "local_eval_datasets": request.local_eval_datasets, "format_type": request.format_type, "subset": request.subset, "train_split": request.train_split, diff --git a/studio/frontend/src/features/studio/sections/dataset-section.tsx b/studio/frontend/src/features/studio/sections/dataset-section.tsx index a7e2c6555d..ab05d8e70b 100644 --- a/studio/frontend/src/features/studio/sections/dataset-section.tsx +++ b/studio/frontend/src/features/studio/sections/dataset-section.tsx @@ -49,6 +49,7 @@ import type { LocalDatasetInfo } from "@/features/training/types/datasets"; import { useNavigate } from "@tanstack/react-router"; import { ArrowDown01Icon, + Cancel01Icon, CloudUploadIcon, Database02Icon, FileAttachmentIcon, @@ -118,6 +119,8 @@ export function DatasetSection() { datasetEvalSplit, setDatasetEvalSplit, uploadedFile, + uploadedEvalFile, + setUploadedEvalFile, hfToken, modelType, datasetSliceStart, @@ -139,6 +142,8 @@ export function DatasetSection() { datasetEvalSplit: s.datasetEvalSplit, setDatasetEvalSplit: s.setDatasetEvalSplit, uploadedFile: s.uploadedFile, + uploadedEvalFile: s.uploadedEvalFile, + setUploadedEvalFile: s.setUploadedEvalFile, hfToken: s.hfToken, modelType: s.modelType, datasetSliceStart: s.datasetSliceStart, @@ -342,6 +347,7 @@ export function DatasetSection() { const comboboxAnchorRef = useRef(null); const fileInputRef = useRef(null); + const evalFileInputRef = useRef(null); const { scrollRef, sentinelRef } = useInfiniteScroll( fetchMore, hfResults.length, @@ -355,6 +361,25 @@ export function DatasetSection() { fileInputRef.current?.click(); }; + const handleFileUpload = async ( + file: File, + onSuccess: (storedPath: string) => void, + successMessage: string, + ) => { + setIsUploading(true); + try { + const uploaded = await uploadTrainingDataset(file); + onSuccess(uploaded.stored_path); + toast.success(successMessage, { description: uploaded.filename }); + } catch (error) { + toast.error("Upload failed", { + description: error instanceof Error ? error.message : "Unknown error", + }); + } finally { + setIsUploading(false); + } + }; + const handleDatasetFileChange = async (event: ChangeEvent) => { const file = event.target.files?.[0]; event.target.value = ""; @@ -367,30 +392,15 @@ export function DatasetSection() { return; } - const MAX_SIZE_BYTES = 512 * 1024 * 1024; - if (file.size > MAX_SIZE_BYTES) { - toast.error("File too large", { - description: "Maximum upload size is 512 MB.", - }); - return; - } + await handleFileUpload(file, selectLocalDataset, "Dataset uploaded"); + }; - setIsUploading(true); - try { - const uploaded = await uploadTrainingDataset(file); + const handleEvalFileChange = async (event: ChangeEvent) => { + const file = event.target.files?.[0]; + event.target.value = ""; + if (!file) return; - selectLocalDataset(uploaded.stored_path); - - toast.success("Dataset uploaded", { - description: uploaded.filename, - }); - } catch (error) { - toast.error("Upload failed", { - description: error instanceof Error ? error.message : "Unknown error", - }); - } finally { - setIsUploading(false); - } + await handleFileUpload(file, setUploadedEvalFile, "Eval dataset uploaded"); }; const handleOpenLearningRecipes = useCallback(() => { @@ -732,6 +742,52 @@ export function DatasetSection() { ) : null} + {datasetSource === "upload" && uploadedFile && ( +
+

+ Eval dataset +

+ {uploadedEvalFile ? ( +
+
+ + + {deriveLocalDatasetName(uploadedEvalFile)} + +
+ +
+ ) : ( +
+ +

+ Optional. If not provided, a small portion will be split from the training data. +

+
+ )} +
+ )} + + { + void handleEvalFileChange(event); + }} + /> 0 && !config.datasetEvalSplit) { - return { - ok: false, - message: - "Eval Steps is set but no Eval Split is selected. Choose an Eval Split or set Eval Steps to 0.", - }; - } return { ok: true, message: null }; } 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 19d0065713..212dc78e6c 100644 --- a/studio/frontend/src/features/training/stores/training-config-store.ts +++ b/studio/frontend/src/features/training/stores/training-config-store.ts @@ -39,6 +39,7 @@ const initialState: TrainingConfigState = { datasetSliceStart: null, datasetSliceEnd: null, uploadedFile: null, + uploadedEvalFile: null, isCheckingVision: false, isVisionModel: false, isEmbeddingModel: false, @@ -253,6 +254,7 @@ export const useTrainingConfigStore = create()( datasetAdvisorNotification: null, datasetSliceStart: null, datasetSliceEnd: null, + uploadedEvalFile: null, isDatasetImage: null, isDatasetAudio: false, isCheckingDataset: false, @@ -444,11 +446,16 @@ export const useTrainingConfigStore = create()( datasetManualMapping: emptyManualMapping(), datasetSliceStart: null, datasetSliceEnd: null, + uploadedEvalFile: null, isDatasetImage: null, isDatasetAudio: false, isCheckingDataset: false, }); }, + setUploadedEvalFile: (uploadedEvalFile) => set({ + uploadedEvalFile, + evalSteps: uploadedEvalFile ? 0.1 : 0, + }), setEpochs: (epochs) => set({ epochs }), setContextLength: (contextLength) => set({ contextLength }), setLearningRate: (learningRate) => set({ learningRate }), diff --git a/studio/frontend/src/features/training/types/api.ts b/studio/frontend/src/features/training/types/api.ts index 8f7ccb9463..1d72e9d9d8 100644 --- a/studio/frontend/src/features/training/types/api.ts +++ b/studio/frontend/src/features/training/types/api.ts @@ -16,6 +16,7 @@ export interface TrainingStartRequest { dataset_slice_start: number | null; dataset_slice_end: number | null; local_datasets: string[]; + local_eval_datasets: string[]; format_type: string; custom_format_mapping?: Record | null; num_epochs: number; diff --git a/studio/frontend/src/features/training/types/config.ts b/studio/frontend/src/features/training/types/config.ts index 40dae7bb3f..635ef8671a 100644 --- a/studio/frontend/src/features/training/types/config.ts +++ b/studio/frontend/src/features/training/types/config.ts @@ -37,6 +37,7 @@ export interface TrainingConfigState { datasetSliceStart: string | null; datasetSliceEnd: string | null; uploadedFile: string | null; + uploadedEvalFile: string | null; epochs: number; contextLength: number; learningRate: number; @@ -109,6 +110,7 @@ export interface TrainingConfigActions { setDatasetSliceStart: (value: string | null) => void; setDatasetSliceEnd: (value: string | null) => void; setUploadedFile: (file: string | null) => void; + setUploadedEvalFile: (file: string | null) => void; setEpochs: (epochs: number) => void; setContextLength: (length: number) => void; setLearningRate: (rate: number) => void;