From 741b1f2965fc91404b277ba272dc2192a6767df0 Mon Sep 17 00:00:00 2001 From: sanatb187 Date: Fri, 10 Apr 2026 00:14:07 -0500 Subject: [PATCH 01/14] Add HF dataset streaming mode to Studio --- studio/backend/core/training/trainer.py | 272 ++++++++++-------- studio/backend/core/training/training.py | 1 + studio/backend/core/training/worker.py | 1 + studio/backend/models/training.py | 4 + studio/backend/routes/training.py | 1 + .../studio/sections/dataset-section.tsx | 41 +++ .../src/features/training/api/mappers.ts | 1 + .../training/stores/training-config-store.ts | 7 +- .../src/features/training/types/api.ts | 1 + .../src/features/training/types/config.ts | 2 + 10 files changed, 204 insertions(+), 127 deletions(-) diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index 77cbda6b45c..2def28d6a76 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -53,10 +53,10 @@ from loggers import get_logger import time from pathlib import Path -from typing import Optional, Callable +from typing import Any, Dict, List, Optional, Callable from dataclasses import dataclass import pandas as pd -from datasets import Dataset, load_dataset +from datasets import Dataset, IterableDataset, load_dataset from utils.models import is_vision_model, detect_audio_type from utils.datasets import format_and_template_dataset @@ -2321,18 +2321,19 @@ def _loader_for_files(files: list[str]) -> str: 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", - eval_split: str = None, - eval_steps: float = 0.00, - dataset_slice_start: int = None, - dataset_slice_end: int = None, + self, + dataset_source: Optional[str], + format_type: str = "auto", + local_datasets: Optional[List[str]] = None, + local_eval_datasets: Optional[List[str]] = None, + custom_format_mapping: Optional[Dict[str, Any]] = None, + subset: Optional[str] = None, + train_split: str = "train", + eval_split: Optional[str] = None, + dataset_streaming: bool = False, + eval_steps: float = 0.00, + dataset_slice_start: Optional[int] = None, + dataset_slice_end: Optional[int] = None, ) -> Optional[tuple]: """ Load and prepare dataset for training. @@ -2394,28 +2395,30 @@ def load_and_format_dataset( if subset: load_kwargs["name"] = subset - _slice_start = dataset_slice_start or 0 - if ( - dataset_slice_end is not None - and dataset_slice_end >= 0 - and dataset_slice_end >= _slice_start - ): - # Manual slice — stream only the rows we need instead of - # downloading the entire dataset. - rows_to_stream = dataset_slice_end + 1 - logger.info( - f"[dataset-slice] Manual slice specified " - f"(start={dataset_slice_start}, end={dataset_slice_end}), " - f"streaming {rows_to_stream} rows\n" + if dataset_streaming: + self._update_progress( + status_message = f"Streaming dataset: {dataset_source}..." ) - stream = load_dataset(**load_kwargs, streaming = True) - dataset = Dataset.from_list(list(stream.take(rows_to_stream))) + dataset = load_dataset(**load_kwargs, streaming = True) + + # Optional iterable slicing + if dataset_slice_start is not None and dataset_slice_start > 0: + dataset = dataset.skip(dataset_slice_start) + + if dataset_slice_end is not None: + slice_start = dataset_slice_start or 0 + take_count = dataset_slice_end - slice_start + 1 + if take_count <= 0: + raise ValueError( + "Train Split End must be greater than or equal to Train Split Start." + ) + dataset = dataset.take(take_count) + logger.info( - f"[dataset-slice] Downloaded {len(dataset)} rows " - f"(requested {rows_to_stream})\n" + f"Loaded Hugging Face dataset in streaming mode: {dataset_source}\n" ) self._update_progress( - status_message = f"Streamed {len(dataset)} rows from HuggingFace" + status_message = f"Streaming {dataset_source}" ) else: self._update_progress( @@ -2423,40 +2426,56 @@ def load_and_format_dataset( ) dataset = load_dataset(**load_kwargs) + n_rows = len(dataset) if hasattr(dataset, "__len__") else 0 + self._update_progress( + status_message = f"Downloaded {dataset_source} ({n_rows:,} rows)" + ) + logger.info( + f"Loaded dataset from Hugging Face: {dataset_source} ({n_rows:,} rows)\n" + ) + # Check if stopped during dataset loading if self.should_stop: logger.info("Stopped during dataset loading\n") return None - n_rows = len(dataset) if hasattr(dataset, "__len__") else 0 - self._update_progress( - status_message = f"Downloaded {dataset_source} ({n_rows:,} rows)" - ) - logger.info( - f"Loaded dataset from Hugging Face: {dataset_source} ({n_rows:,} rows)\n" - ) - - # Resolve eval split from a separate HF split (explicit or auto-detected) + # Resolve eval split from a separate HF split if eval_enabled: effective_train = train_split or "train" if eval_split and eval_split != effective_train: - # Explicit eval split provided - load it directly logger.info(f"Loading explicit eval split: '{eval_split}'\n") eval_load_kwargs = {"path": dataset_source, "split": eval_split} if subset: eval_load_kwargs["name"] = subset - eval_dataset = load_dataset(**eval_load_kwargs) + + if dataset_streaming: + eval_dataset = load_dataset(**eval_load_kwargs, streaming = True) + else: + eval_dataset = load_dataset(**eval_load_kwargs) + has_separate_eval_source = True - logger.info( - f"Loaded eval split '{eval_split}' with {len(eval_dataset)} rows\n" - ) + if hasattr(eval_dataset, "__len__"): + logger.info( + f"Loaded eval split '{eval_split}' with {len(eval_dataset)} rows\n" + ) + else: + logger.info( + f"Loaded eval split '{eval_split}' in streaming mode\n" + ) elif eval_split and eval_split == effective_train: - # Same split as training — will do 80/20 split after formatting + if dataset_streaming: + raise ValueError( + "Streaming mode does not support using the same split for both train and eval. " + "Please provide a separate eval split or disable evaluation." + ) logger.info( f"Eval split '{eval_split}' is the same as train split — will split 80/20\n" ) else: - # Auto-detect eval split from HF (returns a separate dataset, or None) + if dataset_streaming: + raise ValueError( + "Streaming mode currently requires an explicit eval split when evaluation is enabled." + ) eval_dataset = self._auto_detect_eval_split_from_hf( dataset_source = dataset_source, subset = subset, @@ -2471,8 +2490,10 @@ def load_and_format_dataset( if dataset is None: raise ValueError("No dataset provided") - # Apply index range slicing if requested (inclusive on both ends) - if dataset_slice_start is not None or dataset_slice_end is not None: + # Apply eager-only index range slicing if requested (inclusive on both ends) + if (not dataset_streaming) and ( + dataset_slice_start is not None or dataset_slice_end is not None + ): total_rows = len(dataset) start = dataset_slice_start if dataset_slice_start is not None else 0 end = ( @@ -2563,6 +2584,8 @@ def load_and_format_dataset( final_n = len(final_ds) if hasattr(final_ds, "__len__") else "?" self._update_progress( status_message = f"Dataset ready ({final_n:,} samples, {detected} format)" + if isinstance(final_n, int) + else f"Dataset ready ({final_n} samples, {detected} format)" ) logger.info( f"Dataset formatted successfully ({final_n} samples, {detected})\n" @@ -2571,7 +2594,8 @@ def load_and_format_dataset( # ========== THEN SPLIT ========== if has_separate_eval_source and eval_dataset is not None: # Eval came from a separate HF split — format it too - logger.info(f"Formatting eval dataset ({len(eval_dataset)} rows)...\n") + eval_n = len(eval_dataset) if hasattr(eval_dataset, "__len__") else "?" + logger.info(f"Formatting eval dataset ({eval_n} rows)...\n") eval_info = format_and_template_dataset( eval_dataset, model_name = self.model_name, @@ -2582,8 +2606,8 @@ def load_and_format_dataset( custom_format_mapping = custom_format_mapping, ) eval_dataset = eval_info["dataset"] - logger.info(f"Eval dataset formatted successfully\n") - elif eval_enabled and not has_separate_eval_source: + logger.info("Eval dataset formatted successfully\n") + elif eval_enabled and not has_separate_eval_source and not dataset_streaming: # No separate eval source — split the already-formatted dataset formatted_dataset = dataset_info["dataset"] split_result = self._resolve_eval_split_from_dataset(formatted_dataset) @@ -2597,7 +2621,7 @@ def load_and_format_dataset( logger.error(f"Error loading dataset: {e}") self._update_progress(error = str(e)) return None - + def _auto_detect_eval_split_from_hf( self, dataset_source: str, subset: str ) -> Optional[Dataset]: @@ -2750,7 +2774,7 @@ def start_training( logger.error(f"Failed to start training thread: {e}") return False - def _train_worker(self, dataset: Dataset, **training_args): + def _train_worker(self, dataset: Dataset | IterableDataset | dict, **training_args): """Worker function for training (runs in separate thread)""" try: # On spawn-based platforms (Windows, macOS), register all known @@ -3081,7 +3105,7 @@ def audio_vlm_collate_fn(examples): else: # Default to warmup_steps if neither provided config_args["warmup_steps"] = 5 - logger.info(f"Using default warmup_steps: 5\n") + logger.info("Using default warmup_steps: 5\n") # Add save_steps if specified save_steps_val = training_args.get("save_steps", 0) @@ -3089,7 +3113,7 @@ def audio_vlm_collate_fn(examples): config_args["save_steps"] = save_steps_val config_args["save_strategy"] = "steps" - # If max_steps is specified, use it instead of epochs + # If max_steps is specified, use it instead of epochs max_steps_val = training_args.get("max_steps", 0) if max_steps_val and max_steps_val > 0: del config_args["num_train_epochs"] # Remove epochs @@ -3108,7 +3132,10 @@ def audio_vlm_collate_fn(examples): logger.info( f"✅ Evaluation enabled: eval_steps={eval_steps_val} (fraction of total steps)\n" ) - logger.info(f"Eval dataset: {len(eval_dataset)} rows\n") + if hasattr(eval_dataset, "__len__"): + logger.info(f"Eval dataset: {len(eval_dataset)} rows\n") + else: + logger.info("Eval dataset is streaming / length unknown\n") else: logger.info( f"⚠️ Eval dataset provided but eval_steps={eval_steps_val} (disabled)\n" @@ -3178,7 +3205,7 @@ def audio_vlm_collate_fn(examples): # Audio VLM (e.g. Gemma 3N + audio): raw Dataset from _format_audio_vlm_dataset # Notebook uses processing_class=processor.tokenizer (text tokenizer only) train_dataset = ( - dataset if isinstance(dataset, Dataset) else dataset["dataset"] + dataset["dataset"] if isinstance(dataset, dict) else dataset ) processing_class = ( self.tokenizer.tokenizer @@ -3223,7 +3250,7 @@ def audio_vlm_collate_fn(examples): self.tokenizer, "tokenizer" ): logger.info( - f" ⚠️ Unwrapping Processor → raw tokenizer for text-only SFTTrainer" + "Unwrapping Processor → raw tokenizer for text-only SFTTrainer" ) sft_tokenizer = self.tokenizer.tokenizer @@ -3317,63 +3344,46 @@ def audio_vlm_collate_fn(examples): logger.info("Train on responses only configured successfully\n") # ── Safety net: check if all samples were filtered out ── - # Unsloth's train_on_responses_only masks non-response - # tokens with -100. If max_seq_length is too short and the - # response portion gets truncated away, EVERY sample ends - # up with all labels == -100 and Unsloth removes them, - # leaving 0 usable training samples. - filtered_len = len(self.trainer.train_dataset) - original_len = len(dataset["dataset"]) - dropped = original_len - filtered_len - drop_pct = ( - round(100 * dropped / original_len, 1) - if original_len > 0 - else 0 - ) - - if filtered_len == 0 or drop_pct > 30: - max_seq = training_args.get("max_seq_length", 2048) - error_msg = ( - f"{dropped}/{original_len} samples ({drop_pct}%) " - f"were dropped after applying 'train on responses " - f"only' — only {filtered_len} remain. This usually " - f"means max_seq_length ({max_seq}) is too short " - f"and the response portion is being truncated " - f"away. Try increasing max_seq_length (e.g. 8192) " - f"or disabling 'Train on completions'." - ) - logger.error(error_msg) - self._update_progress(error = error_msg, is_training = False) - return - - if dropped > 0: + # Skip post-filter length checks for streaming datasets. + if isinstance(self.trainer.train_dataset, IterableDataset): logger.info( - f"⚠️ {dropped}/{original_len} samples " - f"({drop_pct}%) were dropped (all labels " - f"masked). {filtered_len} samples remain.\n" + "Skipping post-filter length check for streaming dataset\n" ) - logger.info(f"Post-filter dataset size: {filtered_len} samples\n") - - # [DEBUG] Decode first sample AFTER train_on_completions applied - # try: - # _row = self.trainer.train_dataset[0] - # _space = self.tokenizer( - # " ", add_special_tokens = False - # ).input_ids[0] - # print("[DEBUG] === After train_on_completions ===", flush = True) - # print( - # f"[DEBUG] input_ids decoded:\n{self.tokenizer.decode(_row['input_ids'])}\n", - # flush = True, - # ) - # print( - # f"[DEBUG] labels decoded (-100 → space):\n{self.tokenizer.decode([_space if x == -100 else x for x in _row['labels']])}\n", - # flush = True, - # ) - # except Exception as _dbg_e: - # print( - # f"[DEBUG] Could not decode post-completions sample: {_dbg_e}", - # flush = True, - # ) + else: + filtered_len = len(self.trainer.train_dataset) + original_dataset_obj = ( + dataset["dataset"] if isinstance(dataset, dict) else dataset + ) + original_len = len(original_dataset_obj) + dropped = original_len - filtered_len + drop_pct = ( + round(100 * dropped / original_len, 1) + if original_len > 0 + else 0 + ) + + if filtered_len == 0 or drop_pct > 30: + max_seq = training_args.get("max_seq_length", 2048) + error_msg = ( + f"{dropped}/{original_len} samples ({drop_pct}%) " + f"were dropped after applying 'train on responses " + f"only' — only {filtered_len} remain. This usually " + f"means max_seq_length ({max_seq}) is too short " + f"and the response portion is being truncated " + f"away. Try increasing max_seq_length (e.g. 8192) " + f"or disabling 'Train on completions'." + ) + logger.error(error_msg) + self._update_progress(error = error_msg, is_training = False) + return + + if dropped > 0: + logger.info( + f"⚠️ {dropped}/{original_len} samples " + f"({drop_pct}%) were dropped (all labels " + f"masked). {filtered_len} samples remain.\n" + ) + logger.info(f"Post-filter dataset size: {filtered_len} samples\n") except Exception as e: logger.warning(f"Failed to apply train on responses only: {e}") @@ -3387,17 +3397,27 @@ def audio_vlm_collate_fn(examples): # ========== PROGRESS TRACKING ========== self.trainer.add_callback(self._create_progress_callback()) - num_samples = len( - dataset["dataset"] if isinstance(dataset, dict) else dataset - ) - batch_size = training_args.get("batch_size", 2) - total_steps = self._calculate_total_steps( - num_samples, - batch_size, - training_args.get("gradient_accumulation_steps", 4), - training_args.get("num_epochs", 3), - training_args.get("max_steps", 0), - ) + train_dataset_obj = dataset["dataset"] if isinstance(dataset, dict) else dataset + is_streaming_dataset = isinstance(train_dataset_obj, IterableDataset) + + if is_streaming_dataset and training_args.get("max_steps", 0) <= 0: + raise ValueError( + "Streaming mode requires max_steps > 0 because the training dataset has no length." + ) + + if is_streaming_dataset: + total_steps = training_args.get("max_steps", 0) + else: + num_samples = len(train_dataset_obj) + batch_size = training_args.get("batch_size", 2) + total_steps = self._calculate_total_steps( + num_samples, + batch_size, + training_args.get("gradient_accumulation_steps", 4), + training_args.get("num_epochs", 3), + training_args.get("max_steps", 0), + ) + self._update_progress(total_steps = total_steps) # ========== START TRAINING ========== diff --git a/studio/backend/core/training/training.py b/studio/backend/core/training/training.py index f35c7e8ad3e..6f60cfbe780 100644 --- a/studio/backend/core/training/training.py +++ b/studio/backend/core/training/training.py @@ -146,6 +146,7 @@ def start_training(self, job_id: str, **kwargs) -> bool: "train_split": kwargs.get("train_split", "train"), "eval_split": kwargs.get("eval_split"), "eval_steps": kwargs.get("eval_steps", 0.00), + "dataset_streaming": kwargs.get("dataset_streaming", False), "dataset_slice_start": kwargs.get("dataset_slice_start"), "dataset_slice_end": kwargs.get("dataset_slice_end"), "custom_format_mapping": kwargs.get("custom_format_mapping"), diff --git a/studio/backend/core/training/worker.py b/studio/backend/core/training/worker.py index a461972eca0..ea250528ee4 100644 --- a/studio/backend/core/training/worker.py +++ b/studio/backend/core/training/worker.py @@ -564,6 +564,7 @@ def _poll_stop(): subset = config.get("subset"), train_split = config.get("train_split", "train"), eval_split = config.get("eval_split"), + dataset_streaming = config.get("dataset_streaming", False), eval_steps = config.get("eval_steps", 0.00), dataset_slice_start = config.get("dataset_slice_start"), dataset_slice_end = config.get("dataset_slice_end"), diff --git a/studio/backend/models/training.py b/studio/backend/models/training.py index 07a306ca397..37ddec8dca8 100644 --- a/studio/backend/models/training.py +++ b/studio/backend/models/training.py @@ -43,6 +43,10 @@ class TrainingStartRequest(BaseModel): eval_split: Optional[str] = Field( None, description = "Eval split name. None = auto-detect" ) + dataset_streaming: bool = Field( + False, + description = "Whether to load the Hugging Face dataset in streaming mode", + ) eval_steps: float = Field( 0.00, description = "Fraction of total steps between evals (0-1)" ) diff --git a/studio/backend/routes/training.py b/studio/backend/routes/training.py index e625408bad8..701e043642d 100644 --- a/studio/backend/routes/training.py +++ b/studio/backend/routes/training.py @@ -166,6 +166,7 @@ async def start_training( "format_type": request.format_type, "subset": request.subset, "train_split": request.train_split, + "dataset_streaming": request.dataset_streaming, "eval_split": request.eval_split, "eval_steps": request.eval_steps, "dataset_slice_start": request.dataset_slice_start, diff --git a/studio/frontend/src/features/studio/sections/dataset-section.tsx b/studio/frontend/src/features/studio/sections/dataset-section.tsx index b12bdb09f0a..96fa3bd0845 100644 --- a/studio/frontend/src/features/studio/sections/dataset-section.tsx +++ b/studio/frontend/src/features/studio/sections/dataset-section.tsx @@ -118,6 +118,8 @@ export function DatasetSection() { setDatasetSplit, datasetEvalSplit, setDatasetEvalSplit, + datasetStreaming, + setDatasetStreaming, uploadedFile, uploadedEvalFile, setUploadedEvalFile, @@ -141,6 +143,8 @@ export function DatasetSection() { setDatasetSplit: s.setDatasetSplit, datasetEvalSplit: s.datasetEvalSplit, setDatasetEvalSplit: s.setDatasetEvalSplit, + datasetStreaming: s.datasetStreaming, + setDatasetStreaming: s.setDatasetStreaming, uploadedFile: s.uploadedFile, uploadedEvalFile: s.uploadedEvalFile, setUploadedEvalFile: s.setUploadedEvalFile, @@ -849,6 +853,43 @@ export function DatasetSection() { + {datasetSource === "huggingface" && ( +
+ + Streaming Mode + + + + + + Load Hugging Face datasets using streaming mode when supported. + + + + + + +

+ Only applies to Hugging Face datasets. +

+
+ )}
diff --git a/studio/frontend/src/features/training/api/mappers.ts b/studio/frontend/src/features/training/api/mappers.ts index 561dbe14083..cd196f5cbc5 100644 --- a/studio/frontend/src/features/training/api/mappers.ts +++ b/studio/frontend/src/features/training/api/mappers.ts @@ -60,6 +60,7 @@ export function buildTrainingStartPayload( subset: hfDataset ? config.datasetSubset : null, train_split: hfDataset ? config.datasetSplit : null, eval_split: hfDataset ? config.datasetEvalSplit : null, + dataset_streaming: hfDataset ? config.datasetStreaming : false, dataset_slice_start: parseSliceValue(config.datasetSliceStart), dataset_slice_end: parseSliceValue(config.datasetSliceEnd), local_datasets: localDatasets, 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 8214b0eb2ab..22abcaa36af 100644 --- a/studio/frontend/src/features/training/stores/training-config-store.ts +++ b/studio/frontend/src/features/training/stores/training-config-store.ts @@ -65,6 +65,7 @@ const initialState: TrainingConfigState = { datasetSubset: null, datasetSplit: null, datasetEvalSplit: null, + datasetStreaming: false, datasetManualMapping: emptyManualMapping(), datasetSystemPrompt: "", datasetUserTemplate: "", @@ -518,6 +519,7 @@ export const useTrainingConfigStore = create()( evalSteps: datasetEvalSplit ? 0.1 : 0, }); }, + setDatasetStreaming: (datasetStreaming) => set({datasetStreaming}), setDatasetManualMapping: (datasetManualMapping) => set({ datasetManualMapping }), setDatasetAdvisorFields: (fields) => @@ -629,7 +631,7 @@ export const useTrainingConfigStore = create()( }, { name: "unsloth_training_config_v1", - version: 9, + version: 10, migrate: (persisted, version) => { const s = persisted as Record; if (version < 2 && s.datasetSubset == null && s.datasetConfig != null) { @@ -665,6 +667,9 @@ export const useTrainingConfigStore = create()( s.weightDecay = DEFAULT_HYPERPARAMS.weightDecay; } } + if (version < 10) { + s.datasetStreaming + } return s as unknown as TrainingConfigStore; }, partialize: partializePersistedState, diff --git a/studio/frontend/src/features/training/types/api.ts b/studio/frontend/src/features/training/types/api.ts index 1d72e9d9d86..5b3af7e3299 100644 --- a/studio/frontend/src/features/training/types/api.ts +++ b/studio/frontend/src/features/training/types/api.ts @@ -13,6 +13,7 @@ export interface TrainingStartRequest { subset: string | null; train_split: string | null; eval_split: string | null; + dataset_streaming: boolean; dataset_slice_start: number | null; dataset_slice_end: number | null; local_datasets: string[]; diff --git a/studio/frontend/src/features/training/types/config.ts b/studio/frontend/src/features/training/types/config.ts index 2d19dea8749..78da158550d 100644 --- a/studio/frontend/src/features/training/types/config.ts +++ b/studio/frontend/src/features/training/types/config.ts @@ -28,6 +28,7 @@ export interface TrainingConfigState { datasetSubset: string | null; datasetSplit: string | null; datasetEvalSplit: string | null; + datasetStreaming: boolean; datasetManualMapping: DatasetManualMapping; datasetSystemPrompt: string; datasetUserTemplate: string; @@ -101,6 +102,7 @@ export interface TrainingConfigActions { setDatasetSubset: (subset: string | null) => void; setDatasetSplit: (split: string | null) => void; setDatasetEvalSplit: (split: string | null) => void; + setDatasetStreaming: (value: boolean) => void; setDatasetManualMapping: (mapping: DatasetManualMapping) => void; setDatasetAdvisorFields: (fields: { systemPrompt?: string; From 74823123e626348a2f4fef0bbb2d2a1eeb9f58b2 Mon Sep 17 00:00:00 2001 From: sanatb187 Date: Fri, 10 Apr 2026 03:05:27 -0500 Subject: [PATCH 02/14] Added default value for datasetStreaming in training-config-store.ts --- .../src/features/training/stores/training-config-store.ts | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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 22abcaa36af..3d73fccc428 100644 --- a/studio/frontend/src/features/training/stores/training-config-store.ts +++ b/studio/frontend/src/features/training/stores/training-config-store.ts @@ -668,7 +668,7 @@ export const useTrainingConfigStore = create()( } } if (version < 10) { - s.datasetStreaming + s.datasetStreaming ??= false; } return s as unknown as TrainingConfigStore; }, From e1e721824160a009521cf87c21a0575c48869c15 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Fri, 10 Apr 2026 08:06:57 +0000 Subject: [PATCH 03/14] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- studio/backend/core/training/trainer.py | 48 ++++++++++++++----------- 1 file changed, 27 insertions(+), 21 deletions(-) diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index 2def28d6a76..db123535261 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -2321,19 +2321,19 @@ def _loader_for_files(files: list[str]) -> str: raise ValueError(f"Unsupported dataset format: {files[0]}") def load_and_format_dataset( - self, - dataset_source: Optional[str], - format_type: str = "auto", - local_datasets: Optional[List[str]] = None, - local_eval_datasets: Optional[List[str]] = None, - custom_format_mapping: Optional[Dict[str, Any]] = None, - subset: Optional[str] = None, - train_split: str = "train", - eval_split: Optional[str] = None, - dataset_streaming: bool = False, - eval_steps: float = 0.00, - dataset_slice_start: Optional[int] = None, - dataset_slice_end: Optional[int] = None, + self, + dataset_source: Optional[str], + format_type: str = "auto", + local_datasets: Optional[List[str]] = None, + local_eval_datasets: Optional[List[str]] = None, + custom_format_mapping: Optional[Dict[str, Any]] = None, + subset: Optional[str] = None, + train_split: str = "train", + eval_split: Optional[str] = None, + dataset_streaming: bool = False, + eval_steps: float = 0.00, + dataset_slice_start: Optional[int] = None, + dataset_slice_end: Optional[int] = None, ) -> Optional[tuple]: """ Load and prepare dataset for training. @@ -2417,9 +2417,7 @@ def load_and_format_dataset( logger.info( f"Loaded Hugging Face dataset in streaming mode: {dataset_source}\n" ) - self._update_progress( - status_message = f"Streaming {dataset_source}" - ) + self._update_progress(status_message = f"Streaming {dataset_source}") else: self._update_progress( status_message = f"Downloading dataset: {dataset_source}..." @@ -2449,7 +2447,9 @@ def load_and_format_dataset( eval_load_kwargs["name"] = subset if dataset_streaming: - eval_dataset = load_dataset(**eval_load_kwargs, streaming = True) + eval_dataset = load_dataset( + **eval_load_kwargs, streaming = True + ) else: eval_dataset = load_dataset(**eval_load_kwargs) @@ -2607,7 +2607,9 @@ def load_and_format_dataset( ) eval_dataset = eval_info["dataset"] logger.info("Eval dataset formatted successfully\n") - elif eval_enabled and not has_separate_eval_source and not dataset_streaming: + elif ( + eval_enabled and not has_separate_eval_source and not dataset_streaming + ): # No separate eval source — split the already-formatted dataset formatted_dataset = dataset_info["dataset"] split_result = self._resolve_eval_split_from_dataset(formatted_dataset) @@ -2621,7 +2623,7 @@ def load_and_format_dataset( logger.error(f"Error loading dataset: {e}") self._update_progress(error = str(e)) return None - + def _auto_detect_eval_split_from_hf( self, dataset_source: str, subset: str ) -> Optional[Dataset]: @@ -3383,7 +3385,9 @@ def audio_vlm_collate_fn(examples): f"({drop_pct}%) were dropped (all labels " f"masked). {filtered_len} samples remain.\n" ) - logger.info(f"Post-filter dataset size: {filtered_len} samples\n") + logger.info( + f"Post-filter dataset size: {filtered_len} samples\n" + ) except Exception as e: logger.warning(f"Failed to apply train on responses only: {e}") @@ -3397,7 +3401,9 @@ def audio_vlm_collate_fn(examples): # ========== PROGRESS TRACKING ========== self.trainer.add_callback(self._create_progress_callback()) - train_dataset_obj = dataset["dataset"] if isinstance(dataset, dict) else dataset + train_dataset_obj = ( + dataset["dataset"] if isinstance(dataset, dict) else dataset + ) is_streaming_dataset = isinstance(train_dataset_obj, IterableDataset) if is_streaming_dataset and training_args.get("max_steps", 0) <= 0: From 169ab6593e38318210f8dc429bf59ca91e35ac5b Mon Sep 17 00:00:00 2001 From: sanatb187 Date: Fri, 10 Apr 2026 11:34:51 -0500 Subject: [PATCH 04/14] Handle None max_steps for streaming validation --- studio/backend/core/training/trainer.py | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index db123535261..fabe385bc10 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -3406,13 +3406,16 @@ def audio_vlm_collate_fn(examples): ) is_streaming_dataset = isinstance(train_dataset_obj, IterableDataset) - if is_streaming_dataset and training_args.get("max_steps", 0) <= 0: + max_steps_value = training_args.get("max_steps") + max_steps = 0 if max_steps_value is None else int(max_steps_value) + + if is_streaming_dataset and max_steps <= 0: raise ValueError( "Streaming mode requires max_steps > 0 because the training dataset has no length." ) if is_streaming_dataset: - total_steps = training_args.get("max_steps", 0) + total_steps = max_steps else: num_samples = len(train_dataset_obj) batch_size = training_args.get("batch_size", 2) @@ -3421,7 +3424,7 @@ def audio_vlm_collate_fn(examples): batch_size, training_args.get("gradient_accumulation_steps", 4), training_args.get("num_epochs", 3), - training_args.get("max_steps", 0), + max_steps, ) self._update_progress(total_steps = total_steps) From 682bc3bade4f5aaa7d677bda3ae39f36419e14b2 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Fri, 10 Apr 2026 16:57:05 +0000 Subject: [PATCH 05/14] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- studio/backend/core/training/trainer.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index fabe385bc10..8ec39b86560 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -3408,7 +3408,7 @@ def audio_vlm_collate_fn(examples): max_steps_value = training_args.get("max_steps") max_steps = 0 if max_steps_value is None else int(max_steps_value) - + if is_streaming_dataset and max_steps <= 0: raise ValueError( "Streaming mode requires max_steps > 0 because the training dataset has no length." From 69c85c0aae0b2f50c6fa2835b9ebb476b379f2ae Mon Sep 17 00:00:00 2001 From: Roland Tannous Date: Tue, 14 Apr 2026 22:56:18 +0400 Subject: [PATCH 06/14] studio: fast-fail streaming validation and guard incompatible modes Reject dataset_streaming at the API boundary when hf_dataset is empty, the dataset is vision/audio, or max_steps is not set. Probe eval split with get_dataset_split_names before the streaming load so typos fail immediately instead of mid-training. Guard column_names=None after map on iterables. Hide the UI toggle for non-text configurations and clear the stale flag when config becomes incompatible. --- studio/backend/core/training/trainer.py | 44 +++++++++++++++++-- studio/backend/routes/training.py | 19 ++++++++ .../backend/utils/datasets/dataset_utils.py | 5 ++- .../studio/sections/dataset-section.tsx | 34 ++++++++++++-- 4 files changed, 95 insertions(+), 7 deletions(-) diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index 8ec39b86560..16f5ddffc62 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -2413,6 +2413,13 @@ def load_and_format_dataset( "Train Split End must be greater than or equal to Train Split Start." ) dataset = dataset.take(take_count) + # IterableDataset.take(N) yields *at most* N samples — if + # the source is shorter, the user silently gets fewer rows. + logger.warning( + f"Streaming slice requested up to {take_count} rows " + f"[{slice_start}, {dataset_slice_end}]; actual yield " + f"may be smaller if the dataset has fewer rows." + ) logger.info( f"Loaded Hugging Face dataset in streaming mode: {dataset_source}\n" @@ -2447,6 +2454,30 @@ def load_and_format_dataset( eval_load_kwargs["name"] = subset if dataset_streaming: + # Probe available splits before the streaming load. + # load_dataset(streaming=True) returns an IterableDataset + # without validating the split name — a typo would only + # surface on the first eval batch mid-training. + from datasets import get_dataset_split_names + + probe_kwargs = {"path": dataset_source} + if subset: + probe_kwargs["config_name"] = subset + try: + available_splits = get_dataset_split_names( + **probe_kwargs + ) + except Exception as probe_err: + raise ValueError( + f"Could not list splits for '{dataset_source}' " + f"to validate eval_split='{eval_split}': {probe_err}" + ) + if eval_split not in available_splits: + raise ValueError( + f"Requested eval split '{eval_split}' not found in " + f"dataset '{dataset_source}'. Available splits: " + f"{available_splits}" + ) eval_dataset = load_dataset( **eval_load_kwargs, streaming = True ) @@ -2466,7 +2497,7 @@ def load_and_format_dataset( if dataset_streaming: raise ValueError( "Streaming mode does not support using the same split for both train and eval. " - "Please provide a separate eval split or disable evaluation." + "Please provide a separate eval split or set eval_steps to 0." ) logger.info( f"Eval split '{eval_split}' is the same as train split — will split 80/20\n" @@ -2776,8 +2807,15 @@ def start_training( logger.error(f"Failed to start training thread: {e}") return False - def _train_worker(self, dataset: Dataset | IterableDataset | dict, **training_args): - """Worker function for training (runs in separate thread)""" + def _train_worker(self, dataset: Dataset | dict, **training_args): + """Worker function for training (runs in separate thread). + + ``dataset`` is either a raw ``datasets.Dataset`` (audio preprocessing + paths such as CSM / Whisper / SNAC / Audio-VLM) or a ``dict`` wrapper + returned by ``format_and_template_dataset`` (text and image VLM paths). + Streaming HF datasets arrive wrapped in the latter ``dict`` — they are + never passed as a bare ``IterableDataset``. + """ try: # On spawn-based platforms (Windows, macOS), register all known # compiled-cache directories on sys.path and PYTHONPATH before any diff --git a/studio/backend/routes/training.py b/studio/backend/routes/training.py index 701e043642d..45eb6014891 100644 --- a/studio/backend/routes/training.py +++ b/studio/backend/routes/training.py @@ -153,6 +153,25 @@ async def start_training( request.local_eval_datasets, "Local eval dataset" ) + # Validate streaming-mode compatibility before any expensive work. + # Streaming is supported only for Hugging Face text datasets. + if request.dataset_streaming: + if not request.hf_dataset: + raise HTTPException( + status_code = 400, + detail = "dataset_streaming requires hf_dataset; streaming is not supported for local datasets.", + ) + if request.is_dataset_image or request.is_dataset_audio: + raise HTTPException( + status_code = 400, + detail = "dataset_streaming is not supported for vision or audio datasets.", + ) + if request.max_steps is None or request.max_steps <= 0: + raise HTTPException( + status_code = 422, + detail = "dataset_streaming requires max_steps > 0 because streaming datasets have no known length.", + ) + # Convert request to kwargs for backend training_kwargs = { "model_name": request.model_name, diff --git a/studio/backend/utils/datasets/dataset_utils.py b/studio/backend/utils/datasets/dataset_utils.py index fac8c3d2951..0a574f437b7 100644 --- a/studio/backend/utils/datasets/dataset_utils.py +++ b/studio/backend/utils/datasets/dataset_utils.py @@ -1147,7 +1147,10 @@ def format_and_template_dataset( requires_manual = dataset_info.get("requires_manual_mapping", False) if final_format == "unknown" and template_result["success"]: out_ds = template_result["dataset"] - if hasattr(out_ds, "column_names") and "text" in out_ds.column_names: + # IterableDataset.column_names can be None after .map() loses features; + # guard to avoid `"text" in None` -> TypeError on streaming datasets. + out_columns = getattr(out_ds, "column_names", None) + if out_columns is not None and "text" in out_columns: final_format = "chatml_conversations" requires_manual = False diff --git a/studio/frontend/src/features/studio/sections/dataset-section.tsx b/studio/frontend/src/features/studio/sections/dataset-section.tsx index 96fa3bd0845..fc3ebf4e817 100644 --- a/studio/frontend/src/features/studio/sections/dataset-section.tsx +++ b/studio/frontend/src/features/studio/sections/dataset-section.tsx @@ -120,6 +120,10 @@ export function DatasetSection() { setDatasetEvalSplit, datasetStreaming, setDatasetStreaming, + isVisionModel, + isAudioModel, + isDatasetImage, + isDatasetAudio, uploadedFile, uploadedEvalFile, setUploadedEvalFile, @@ -145,6 +149,10 @@ export function DatasetSection() { setDatasetEvalSplit: s.setDatasetEvalSplit, datasetStreaming: s.datasetStreaming, setDatasetStreaming: s.setDatasetStreaming, + isVisionModel: s.isVisionModel, + isAudioModel: s.isAudioModel, + isDatasetImage: s.isDatasetImage, + isDatasetAudio: s.isDatasetAudio, uploadedFile: s.uploadedFile, uploadedEvalFile: s.uploadedEvalFile, setUploadedEvalFile: s.setUploadedEvalFile, @@ -157,6 +165,26 @@ export function DatasetSection() { })), ); + // Streaming is only supported for Hugging Face text datasets. + // Hide the toggle for vision/audio models or datasets detected as multimodal, + // since downstream preprocessing (convert_to_vlm_format, audio collators) + // requires random access and would crash on an IterableDataset. + const isStreamingSupported = + datasetSource === "huggingface" && + !isVisionModel && + !isAudioModel && + !isDatasetImage && + !isDatasetAudio; + + // If streaming was previously enabled but the config became incompatible + // (model switched to vision, dataset detected as image, etc.), clear it so + // the backend never receives a stale flag. + useEffect(() => { + if (datasetStreaming && !isStreamingSupported) { + setDatasetStreaming(false); + } + }, [datasetStreaming, isStreamingSupported, setDatasetStreaming]); + const [searchQuery, setSearchQuery] = useState(""); const [advancedOpen, setAdvancedOpen] = useState(false); const [pickerTab, setPickerTab] = useState<"huggingface" | "local">( @@ -853,7 +881,7 @@ export function DatasetSection() {
- {datasetSource === "huggingface" && ( + {isStreamingSupported && (
Streaming Mode @@ -870,7 +898,7 @@ export function DatasetSection() { - Load Hugging Face datasets using streaming mode when supported. + Load Hugging Face datasets using streaming mode. Requires max_steps > 0 and a separate eval split when evaluation is enabled. @@ -886,7 +914,7 @@ export function DatasetSection() {

- Only applies to Hugging Face datasets. + Only applies to Hugging Face text datasets.

)} From ebea1a616e8960161f32e30ee16633cb3cea97a5 Mon Sep 17 00:00:00 2001 From: Etherll <61019402+Etherll@users.noreply.github.com> Date: Fri, 19 Jun 2026 22:47:46 +0300 Subject: [PATCH 07/14] studio: add streaming dataset tests, iterable helper, and streaming template/format support (WIP) Work-in-progress on top of feat/studio-dataset-streaming-mode (PR #4946): - new test_training_streaming.py and iterable.py dataset helper - streaming support in chat_templates.py and format_conversion.py - additional streaming guards in trainer.py / models / routes - frontend streaming wiring in params-section and training-config-store Committed to preserve uncommitted work before merging latest main. --- studio/backend/core/training/trainer.py | 13 +- studio/backend/models/training.py | 25 +- studio/backend/routes/training.py | 14 + .../backend/tests/test_training_streaming.py | 362 ++++++++++++++++++ .../backend/utils/datasets/chat_templates.py | 19 +- .../utils/datasets/format_conversion.py | 24 +- studio/backend/utils/datasets/iterable.py | 22 ++ .../studio/sections/dataset-section.tsx | 10 + .../studio/sections/params-section.tsx | 13 +- .../training/stores/training-config-store.ts | 123 +++++- 10 files changed, 583 insertions(+), 42 deletions(-) create mode 100644 studio/backend/tests/test_training_streaming.py create mode 100644 studio/backend/utils/datasets/iterable.py diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index 16f5ddffc62..0a6cceff5df 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -61,6 +61,7 @@ from utils.models import is_vision_model, detect_audio_type from utils.datasets import format_and_template_dataset from utils.datasets import MODEL_TO_TEMPLATE_MAPPER, TEMPLATE_TO_RESPONSES_MAPPER +from utils.datasets.iterable import is_streaming_dataset as detect_streaming_dataset from utils.paths import ( ensure_dir, resolve_dataset_path, @@ -2420,6 +2421,14 @@ def load_and_format_dataset( f"[{slice_start}, {dataset_slice_end}]; actual yield " f"may be smaller if the dataset has fewer rows." ) + if take_count == 1: + # start == end is a valid slice but produces a single + # training row, which is almost always user error. + logger.warning( + "Dataset slice resolves to a single row " + f"(start == end == {slice_start}); training on 1 " + "sample is likely unintended." + ) logger.info( f"Loaded Hugging Face dataset in streaming mode: {dataset_source}\n" @@ -3385,7 +3394,7 @@ def audio_vlm_collate_fn(examples): # ── Safety net: check if all samples were filtered out ── # Skip post-filter length checks for streaming datasets. - if isinstance(self.trainer.train_dataset, IterableDataset): + if detect_streaming_dataset(self.trainer.train_dataset): logger.info( "Skipping post-filter length check for streaming dataset\n" ) @@ -3442,7 +3451,7 @@ def audio_vlm_collate_fn(examples): train_dataset_obj = ( dataset["dataset"] if isinstance(dataset, dict) else dataset ) - is_streaming_dataset = isinstance(train_dataset_obj, IterableDataset) + is_streaming_dataset = detect_streaming_dataset(train_dataset_obj) max_steps_value = training_args.get("max_steps") max_steps = 0 if max_steps_value is None else int(max_steps_value) diff --git a/studio/backend/models/training.py b/studio/backend/models/training.py index 37ddec8dca8..0db7d1374d2 100644 --- a/studio/backend/models/training.py +++ b/studio/backend/models/training.py @@ -51,10 +51,14 @@ class TrainingStartRequest(BaseModel): 0.00, description = "Fraction of total steps between evals (0-1)" ) dataset_slice_start: Optional[int] = Field( - None, description = "Inclusive start row index for dataset slicing" + None, + ge = 0, + description = "Inclusive start row index for dataset slicing", ) dataset_slice_end: Optional[int] = Field( - None, description = "Inclusive end row index for dataset slicing" + None, + ge = 0, + description = "Inclusive end row index for dataset slicing", ) @model_validator(mode = "before") @@ -65,6 +69,23 @@ def _compat_split(cls, values: Any) -> Any: values.setdefault("train_split", values.pop("split")) return values + @model_validator(mode = "after") + def _validate_dataset_slice(self) -> "TrainingStartRequest": + # Only the ordering is validated here. No upper bound is enforced on the + # indices: the trainer slices via datasets `.take()` / `.select()`, which + # clamp gracefully when the end index exceeds the dataset length. + # start == end is intentionally allowed (deliberate single-row slice, + # e.g. for debugging); the trainer logs a warning for that 1-row case. + if ( + self.dataset_slice_start is not None + and self.dataset_slice_end is not None + and self.dataset_slice_end < self.dataset_slice_start + ): + raise ValueError( + "dataset_slice_end must be greater than or equal to dataset_slice_start" + ) + return self + custom_format_mapping: Optional[Dict[str, Any]] = Field( None, description = ( diff --git a/studio/backend/routes/training.py b/studio/backend/routes/training.py index 45eb6014891..7bcdb6c1f9d 100644 --- a/studio/backend/routes/training.py +++ b/studio/backend/routes/training.py @@ -171,6 +171,18 @@ async def start_training( status_code = 422, detail = "dataset_streaming requires max_steps > 0 because streaming datasets have no known length.", ) + if request.train_on_completions: + raise HTTPException( + status_code = 422, + detail = "dataset_streaming is not supported with train_on_completions yet.", + ) + if request.eval_steps > 0: + train_split = request.train_split or "train" + if not request.eval_split or request.eval_split == train_split: + raise HTTPException( + status_code = 422, + detail = "dataset_streaming with evaluation requires a separate eval_split.", + ) # Convert request to kwargs for backend training_kwargs = { @@ -298,6 +310,8 @@ async def start_training( error = None, ) + except HTTPException: + raise except ValueError as e: logger.warning("Rejected training GPU selection: %s", e) raise HTTPException(status_code = 400, detail = str(e)) diff --git a/studio/backend/tests/test_training_streaming.py b/studio/backend/tests/test_training_streaming.py new file mode 100644 index 00000000000..1652ca42d42 --- /dev/null +++ b/studio/backend/tests/test_training_streaming.py @@ -0,0 +1,362 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +from __future__ import annotations + +import asyncio +import importlib.util +import sys +from pathlib import Path +from types import SimpleNamespace +from unittest.mock import patch + +import pytest +from fastapi import HTTPException +from pydantic import ValidationError + +from models.training import TrainingStartRequest +from utils.datasets.chat_templates import apply_chat_template_to_dataset +from utils.datasets.format_conversion import ( + convert_alpaca_to_chatml, + convert_chatml_to_alpaca, + standardize_chat_format, +) +from utils.datasets.iterable import is_streaming_dataset + +datasets = pytest.importorskip("datasets") + +_BACKEND_ROOT = Path(__file__).resolve().parent.parent + + +def _load_route_module(name: str, relative_path: str): + spec = importlib.util.spec_from_file_location(name, _BACKEND_ROOT / relative_path) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +class _Tokenizer: + eos_token = "" + chat_template = "{{ messages }}" + + def apply_chat_template( + self, + conversation, + *, + tokenize = False, + add_generation_prompt = False, + ): + assert tokenize is False + assert add_generation_prompt is False + return "\n".join( + f"{message['role']}: {message['content']}" for message in conversation + ) + + +def _iterable_dataset(rows): + return datasets.IterableDataset.from_generator(lambda: iter(rows)) + + +def test_chat_template_mapping_omits_eager_kwargs_for_streaming(monkeypatch): + seen_kwargs = [] + original_map = datasets.IterableDataset.map + + def spy_map(self, *args, **kwargs): + seen_kwargs.append(dict(kwargs)) + return original_map(self, *args, **kwargs) + + monkeypatch.setattr(datasets.IterableDataset, "map", spy_map) + + dataset = _iterable_dataset( + [ + { + "conversations": [ + {"role": "user", "content": "Hi"}, + {"role": "assistant", "content": "Hello"}, + ] + } + ] + ) + result = apply_chat_template_to_dataset( + { + "dataset": dataset, + "final_format": "chatml_conversations", + "chat_column": "conversations", + "is_standardized": True, + }, + tokenizer = _Tokenizer(), + batch_size = 1, + num_proc = 2, + ) + + assert result["success"] is True + row = next(iter(result["dataset"])) + assert "user: Hi" in row["text"] + assert seen_kwargs + assert all("num_proc" not in kwargs for kwargs in seen_kwargs) + assert all("desc" not in kwargs for kwargs in seen_kwargs) + + +def test_format_conversion_omits_eager_kwargs_for_streaming(monkeypatch): + seen_kwargs = [] + original_map = datasets.IterableDataset.map + + def spy_map(self, *args, **kwargs): + seen_kwargs.append(dict(kwargs)) + return original_map(self, *args, **kwargs) + + monkeypatch.setattr(datasets.IterableDataset, "map", spy_map) + + converted = convert_chatml_to_alpaca( + _iterable_dataset( + [ + { + "conversations": [ + {"from": "human", "value": "Question"}, + {"from": "gpt", "value": "Answer"}, + ] + } + ] + ), + batch_size = 1, + num_proc = 2, + ) + + row = next(iter(converted)) + assert row["instruction"] == "Question" + assert row["output"] == "Answer" + assert seen_kwargs + assert all("num_proc" not in kwargs for kwargs in seen_kwargs) + assert all("desc" not in kwargs for kwargs in seen_kwargs) + + +def test_streaming_start_rejects_train_on_completions_before_backend_start(): + training_route = _load_route_module( + "training_route_module_for_streaming_completion_test", + "routes/training.py", + ) + request = TrainingStartRequest( + model_name = "unsloth/test", + training_type = "LoRA/QLoRA", + hf_dataset = "org/dataset", + format_type = "chatml", + dataset_streaming = True, + train_on_completions = True, + max_steps = 10, + ) + + backend = SimpleNamespace( + current_job_id = None, + is_training_active = lambda: False, + start_training = lambda **kwargs: pytest.fail("backend should not start"), + ) + + with patch.object(training_route, "get_training_backend", return_value = backend): + with pytest.raises(HTTPException) as exc_info: + asyncio.run( + training_route.start_training(request, current_subject = "test-user") + ) + + assert exc_info.value.status_code == 422 + assert "train_on_completions" in exc_info.value.detail + + +@pytest.mark.parametrize("eval_split", [None, "train"]) +def test_streaming_start_requires_separate_eval_split(eval_split): + training_route = _load_route_module( + "training_route_module_for_streaming_eval_test", + "routes/training.py", + ) + request = TrainingStartRequest( + model_name = "unsloth/test", + training_type = "LoRA/QLoRA", + hf_dataset = "org/dataset", + format_type = "chatml", + dataset_streaming = True, + train_split = "train", + eval_split = eval_split, + eval_steps = 0.1, + max_steps = 10, + ) + + backend = SimpleNamespace( + current_job_id = None, + is_training_active = lambda: False, + start_training = lambda **kwargs: pytest.fail("backend should not start"), + ) + + with patch.object(training_route, "get_training_backend", return_value = backend): + with pytest.raises(HTTPException) as exc_info: + asyncio.run( + training_route.start_training(request, current_subject = "test-user") + ) + + assert exc_info.value.status_code == 422 + assert "separate eval_split" in exc_info.value.detail + + +def test_dataset_slice_bounds_are_non_negative(): + with pytest.raises(ValidationError): + TrainingStartRequest( + model_name = "unsloth/test", + training_type = "LoRA/QLoRA", + format_type = "alpaca", + dataset_slice_start = -1, + ) + + with pytest.raises(ValidationError): + TrainingStartRequest( + model_name = "unsloth/test", + training_type = "LoRA/QLoRA", + format_type = "alpaca", + dataset_slice_start = 5, + dataset_slice_end = 4, + ) + + +def test_dataset_slice_accepts_equal_and_ordered_bounds(): + # start == end is intentionally allowed (single-row slice); start < end too. + equal = TrainingStartRequest( + model_name = "unsloth/test", + training_type = "LoRA/QLoRA", + format_type = "alpaca", + dataset_slice_start = 5, + dataset_slice_end = 5, + ) + assert equal.dataset_slice_start == 5 + assert equal.dataset_slice_end == 5 + + ordered = TrainingStartRequest( + model_name = "unsloth/test", + training_type = "LoRA/QLoRA", + format_type = "alpaca", + dataset_slice_start = 2, + dataset_slice_end = 9, + ) + assert ordered.dataset_slice_end == 9 + + +def test_is_streaming_dataset_detects_hf_iterable(): + assert is_streaming_dataset(_iterable_dataset([{"a": 1}])) is True + + +def test_is_streaming_dataset_false_for_plain_list(): + assert is_streaming_dataset([{"a": 1}]) is False + + +def test_is_streaming_dataset_torch_branch_when_datasets_unavailable(): + torch = pytest.importorskip("torch") + + class _TorchIterable(torch.utils.data.IterableDataset): + def __iter__(self): + return iter([1, 2, 3]) + + # Force the `from datasets import IterableDataset` import to fail so the + # torch detection branch is exercised. + with patch.dict(sys.modules, {"datasets": None}): + assert is_streaming_dataset(_TorchIterable()) is True + + +def test_is_streaming_dataset_false_when_both_backends_unavailable(): + # Both `datasets` and `torch.utils.data` imports fail -> graceful False. + with patch.dict(sys.modules, {"datasets": None, "torch.utils.data": None}): + assert is_streaming_dataset(object()) is False + + +def test_standardize_chat_format_omits_eager_kwargs_for_streaming(monkeypatch): + seen_kwargs = [] + original_map = datasets.IterableDataset.map + + def spy_map(self, *args, **kwargs): + seen_kwargs.append(dict(kwargs)) + return original_map(self, *args, **kwargs) + + monkeypatch.setattr(datasets.IterableDataset, "map", spy_map) + + dataset = _iterable_dataset( + [ + { + "conversations": [ + {"from": "human", "value": "Hi"}, + {"from": "gpt", "value": "Hello"}, + ] + } + ] + ) + result = standardize_chat_format( + dataset, + tokenizer = _Tokenizer(), + batch_size = 1, + num_proc = 2, + ) + + next(iter(result)) + assert seen_kwargs + assert all("num_proc" not in kwargs for kwargs in seen_kwargs) + assert all("desc" not in kwargs for kwargs in seen_kwargs) + + +def test_convert_alpaca_to_chatml_omits_eager_kwargs_for_streaming(monkeypatch): + seen_kwargs = [] + original_map = datasets.IterableDataset.map + + def spy_map(self, *args, **kwargs): + seen_kwargs.append(dict(kwargs)) + return original_map(self, *args, **kwargs) + + monkeypatch.setattr(datasets.IterableDataset, "map", spy_map) + + converted = convert_alpaca_to_chatml( + _iterable_dataset( + [{"instruction": "Question", "input": "", "output": "Answer"}] + ), + batch_size = 1, + num_proc = 2, + ) + + row = next(iter(converted)) + assert "conversations" in row + assert seen_kwargs + assert all("num_proc" not in kwargs for kwargs in seen_kwargs) + assert all("desc" not in kwargs for kwargs in seen_kwargs) + + +def test_streaming_start_happy_path_reaches_backend(): + training_route = _load_route_module( + "training_route_module_for_streaming_happy_path_test", + "routes/training.py", + ) + request = TrainingStartRequest( + model_name = "unsloth/test", + training_type = "LoRA/QLoRA", + hf_dataset = "org/dataset", + format_type = "chatml", + dataset_streaming = True, + train_split = "train", + eval_split = "validation", + eval_steps = 0.1, + max_steps = 10, + ) + + captured = {} + + def _start_training(**kwargs): + captured.update(kwargs) + return True + + backend = SimpleNamespace( + current_job_id = "job_test", + is_training_active = lambda: False, + start_training = _start_training, + ) + + with patch.object(training_route, "get_training_backend", return_value = backend): + with patch.object(training_route, "load_model_defaults", return_value = {}): + response = asyncio.run( + training_route.start_training(request, current_subject = "test-user") + ) + + assert response.status == "queued" + assert captured["dataset_streaming"] is True + assert captured["max_steps"] == 10 + assert captured["eval_split"] == "validation" diff --git a/studio/backend/utils/datasets/chat_templates.py b/studio/backend/utils/datasets/chat_templates.py index 35fbaba8f0c..0787225e451 100644 --- a/studio/backend/utils/datasets/chat_templates.py +++ b/studio/backend/utils/datasets/chat_templates.py @@ -9,6 +9,7 @@ """ from .format_detection import detect_dataset_format, detect_multimodal_dataset, detect_custom_format_heuristic +from .iterable import is_streaming_dataset from .model_mappings import MODEL_TO_TEMPLATE_MAPPER from loggers import get_logger logger = get_logger(__name__) @@ -288,13 +289,9 @@ def _format_alpaca_custom(examples): 'batch_size': batch_size, } - try: - from torch.utils.data import IterableDataset - _is_torch_iterable = isinstance(dataset, IterableDataset) - except ImportError: - _is_torch_iterable = False + is_iterable = is_streaming_dataset(dataset) - if not _is_torch_iterable: + if not is_iterable: from utils.hardware import dataset_map_num_proc if num_proc is None or type(num_proc) is not int: num_proc = dataset_map_num_proc() @@ -355,18 +352,14 @@ def _format_chatml(examples): return {"text": texts} try: - try: - from torch.utils.data import IterableDataset - _is_torch_iterable = isinstance(dataset, IterableDataset) - except ImportError: - _is_torch_iterable = False + is_iterable = is_streaming_dataset(dataset) dataset_map_kwargs = { 'batched': True, 'batch_size': batch_size, } - if not _is_torch_iterable: + if not is_iterable: from utils.hardware import dataset_map_num_proc if num_proc is None or type(num_proc) is not int: num_proc = dataset_map_num_proc() @@ -377,7 +370,7 @@ def _format_chatml(examples): # Monitor tqdm progress from dataset.map() and relay to callback _tqdm_monitor_stop = None - if progress_callback and not _is_torch_iterable: + if progress_callback and not is_iterable: import threading from tqdm.auto import tqdm as _tqdm_cls diff --git a/studio/backend/utils/datasets/format_conversion.py b/studio/backend/utils/datasets/format_conversion.py index 289b30e55ea..7a3bad97ef2 100644 --- a/studio/backend/utils/datasets/format_conversion.py +++ b/studio/backend/utils/datasets/format_conversion.py @@ -10,7 +10,7 @@ import os -from datasets import IterableDataset +from .iterable import is_streaming_dataset from loggers import get_logger logger = get_logger(__name__) @@ -41,8 +41,6 @@ def standardize_chat_format( """ import collections import itertools - from datasets import IterableDataset - # Check if vision tokenizer is used is_vlm = False if tokenizer is not None: @@ -126,7 +124,7 @@ def _standardize_dataset(examples): "batch_size": batch_size, } - if not isinstance(dataset, IterableDataset): + if not is_streaming_dataset(dataset): from utils.hardware import dataset_map_num_proc if num_proc is None or type(num_proc) is not int: @@ -149,12 +147,7 @@ def convert_chatml_to_alpaca(dataset, batch_size = 1000, num_proc = None): - "messages" or "conversations" column - "role"/"content" (standard) or "from"/"value" (ShareGPT) """ - try: - from torch.utils.data import IterableDataset - - _is_torch_iterable = isinstance(dataset, IterableDataset) - except ImportError: - _is_torch_iterable = False + is_iterable = is_streaming_dataset(dataset) def _convert(examples): # Auto-detect which column name is used @@ -201,7 +194,7 @@ def _convert(examples): "batch_size": batch_size, } - if not _is_torch_iterable: + if not is_iterable: from utils.hardware import dataset_map_num_proc if num_proc is None or type(num_proc) is not int: @@ -221,12 +214,7 @@ def convert_alpaca_to_chatml(dataset, batch_size = 1000, num_proc = None): Output format: Uses 'conversations' column with standard 'role'/'content' structure. """ - try: - from torch.utils.data import IterableDataset - - _is_torch_iterable = isinstance(dataset, IterableDataset) - except ImportError: - _is_torch_iterable = False + is_iterable = is_streaming_dataset(dataset) def _convert(examples): conversations = [] @@ -256,7 +244,7 @@ def _convert(examples): "batch_size": batch_size, } - if not _is_torch_iterable: + if not is_iterable: from utils.hardware import dataset_map_num_proc if num_proc is None or type(num_proc) is not int: diff --git a/studio/backend/utils/datasets/iterable.py b/studio/backend/utils/datasets/iterable.py new file mode 100644 index 00000000000..d2c5eed8996 --- /dev/null +++ b/studio/backend/utils/datasets/iterable.py @@ -0,0 +1,22 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Helpers for dataset iterable detection.""" + + +def is_streaming_dataset(dataset) -> bool: + """Return True for iterable datasets that do not support eager map kwargs.""" + try: + from datasets import IterableDataset as HfIterableDataset + + if isinstance(dataset, HfIterableDataset): + return True + except ImportError: + pass + + try: + from torch.utils.data import IterableDataset as TorchIterableDataset + + return isinstance(dataset, TorchIterableDataset) + except ImportError: + return False diff --git a/studio/frontend/src/features/studio/sections/dataset-section.tsx b/studio/frontend/src/features/studio/sections/dataset-section.tsx index fc3ebf4e817..1f4686cf507 100644 --- a/studio/frontend/src/features/studio/sections/dataset-section.tsx +++ b/studio/frontend/src/features/studio/sections/dataset-section.tsx @@ -45,6 +45,7 @@ import { useTrainingConfigStore, } from "@/features/training"; import { listLocalDatasets } from "@/features/training/api/datasets-api"; +import { hasSeparateStreamingEvalSplit } from "@/features/training/stores/training-config-store"; import type { LocalDatasetInfo } from "@/features/training/types/datasets"; import { useNavigate } from "@tanstack/react-router"; import { @@ -120,6 +121,9 @@ export function DatasetSection() { setDatasetEvalSplit, datasetStreaming, setDatasetStreaming, + trainOnCompletions, + maxSteps, + evalSteps, isVisionModel, isAudioModel, isDatasetImage, @@ -149,6 +153,9 @@ export function DatasetSection() { setDatasetEvalSplit: s.setDatasetEvalSplit, datasetStreaming: s.datasetStreaming, setDatasetStreaming: s.setDatasetStreaming, + trainOnCompletions: s.trainOnCompletions, + maxSteps: s.maxSteps, + evalSteps: s.evalSteps, isVisionModel: s.isVisionModel, isAudioModel: s.isAudioModel, isDatasetImage: s.isDatasetImage, @@ -171,6 +178,9 @@ export function DatasetSection() { // requires random access and would crash on an IterableDataset. const isStreamingSupported = datasetSource === "huggingface" && + maxSteps > 0 && + !trainOnCompletions && + hasSeparateStreamingEvalSplit({ evalSteps, datasetSplit, datasetEvalSplit }) && !isVisionModel && !isAudioModel && !isDatasetImage && diff --git a/studio/frontend/src/features/studio/sections/params-section.tsx b/studio/frontend/src/features/studio/sections/params-section.tsx index 6566303a818..cadb3cad24a 100644 --- a/studio/frontend/src/features/studio/sections/params-section.tsx +++ b/studio/frontend/src/features/studio/sections/params-section.tsx @@ -907,11 +907,22 @@ export function ParamsSection(): ReactElement { store.setTrainOnCompletions(!!v)} /> 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 3d73fccc428..7c95f36491b 100644 --- a/studio/frontend/src/features/training/stores/training-config-store.ts +++ b/studio/frontend/src/features/training/stores/training-config-store.ts @@ -5,6 +5,7 @@ import { DEFAULT_HYPERPARAMS, LR_DEFAULT_FULL, LR_DEFAULT_LORA, STEPS } from "@/ import { authFetch } from "@/features/auth"; import { isAdapterMethod } from "@/types/training"; import type { ModelType, StepNumber, TrainingMethod } from "@/types/training"; +import { toast } from "sonner"; import { create } from "zustand"; import { persist } from "zustand/middleware"; import { checkDatasetFormat } from "../api/datasets-api"; @@ -157,6 +158,49 @@ function canProceedForStep(state: TrainingConfigState): boolean { } } +// Single source of truth for the "streaming + eval needs a distinct split" +// rule. Shared between the store's compatibility patch and the UI gate +// (DatasetSection) so the two never drift apart. +export function hasSeparateStreamingEvalSplit( + state: Pick< + TrainingConfigState, + "evalSteps" | "datasetSplit" | "datasetEvalSplit" + >, +): boolean { + if (state.evalSteps <= 0) return true; + const trainSplit = state.datasetSplit || "train"; + return !!state.datasetEvalSplit && state.datasetEvalSplit !== trainSplit; +} + +function streamingCompatiblePatch( + state: TrainingConfigState, +): Partial { + const patch: Partial = {}; + + if (state.datasetStreaming && state.maxSteps <= 0) { + patch.datasetStreaming = false; + } + + // Evaluate the remaining streaming constraints against the *post-patch* + // streaming value. If streaming is being turned off in this same patch + // (e.g. maxSteps dropped to 0), its other constraints are moot and we must + // NOT clobber unrelated user preferences like trainOnCompletions/evalSteps. + const willStream = + patch.datasetStreaming !== undefined + ? patch.datasetStreaming + : state.datasetStreaming; + + if (willStream && state.trainOnCompletions) { + patch.trainOnCompletions = false; + } + + if (willStream && !hasSeparateStreamingEvalSplit(state)) { + patch.evalSteps = 0; + } + + return patch; +} + export const useTrainingConfigStore = create()( persist( (set, get) => { @@ -482,15 +526,17 @@ export const useTrainingConfigStore = create()( }); }, setDatasetSplit: (datasetSplit) => { + const state = get(); + const nextState = { ...state, datasetSplit }; set({ datasetSplit, datasetManualMapping: emptyManualMapping(), isDatasetImage: null, isDatasetAudio: false, isCheckingDataset: false, + ...streamingCompatiblePatch(nextState), }); - const state = get(); const datasetName = state.datasetSource === "huggingface" ? state.dataset @@ -514,12 +560,48 @@ export const useTrainingConfigStore = create()( runDatasetCheck(datasetName, split); }, setDatasetEvalSplit: (datasetEvalSplit) => { + const state = get(); + const evalSteps = datasetEvalSplit ? 0.1 : 0; set({ datasetEvalSplit, - evalSteps: datasetEvalSplit ? 0.1 : 0, + evalSteps, + ...streamingCompatiblePatch({ ...state, datasetEvalSplit, evalSteps }), }); }, - setDatasetStreaming: (datasetStreaming) => set({datasetStreaming}), + setDatasetStreaming: (datasetStreaming) => { + if (!datasetStreaming) { + set({ datasetStreaming: false }); + return; + } + + const state = get(); + if (state.maxSteps <= 0) { + set({ datasetStreaming: false }); + toast.warning( + "Streaming needs a fixed Max Steps (streaming datasets have no known length). Set Max Steps > 0 first.", + ); + return; + } + + const dropsTrainOnCompletions = state.trainOnCompletions; + const dropsEval = !hasSeparateStreamingEvalSplit(state); + + set({ + datasetStreaming: true, + trainOnCompletions: false, + evalSteps: dropsEval ? 0 : state.evalSteps, + }); + + if (dropsTrainOnCompletions || dropsEval) { + const disabled = [ + dropsTrainOnCompletions && "assistant-completions-only", + dropsEval && "evaluation (needs a separate eval split)", + ].filter(Boolean); + toast.info( + `Streaming enabled. Disabled incompatible options: ${disabled.join(", ")}.`, + ); + } + }, setDatasetManualMapping: (datasetManualMapping) => set({ datasetManualMapping }), setDatasetAdvisorFields: (fields) => @@ -579,13 +661,29 @@ export const useTrainingConfigStore = create()( set({ gradientAccumulation }), setWeightDecay: (weightDecay) => set({ weightDecay }), setWarmupSteps: (warmupSteps) => set({ warmupSteps }), - setMaxSteps: (maxSteps) => set({ maxSteps }), + setMaxSteps: (maxSteps) => { + const state = get(); + set({ + maxSteps, + ...(maxSteps > 0 ? {} : { datasetStreaming: false }), + ...streamingCompatiblePatch({ ...state, maxSteps }), + }); + }, setSaveSteps: (saveSteps) => set({ saveSteps }), - setEvalSteps: (evalSteps) => set({ evalSteps }), + setEvalSteps: (evalSteps) => { + const state = get(); + set({ + evalSteps, + ...streamingCompatiblePatch({ ...state, evalSteps }), + }); + }, setPacking: (packing) => set({ packing }), setTrainOnCompletions: (trainOnCompletions) => { _trainOnCompletionsManuallySet = true; - set({ trainOnCompletions }); + set({ + trainOnCompletions, + ...(trainOnCompletions ? { datasetStreaming: false } : {}), + }); }, setGradientCheckpointing: (gradientCheckpointing) => set({ gradientCheckpointing }), @@ -673,6 +771,19 @@ export const useTrainingConfigStore = create()( return s as unknown as TrainingConfigStore; }, partialize: partializePersistedState, + onRehydrateStorage: () => (state) => { + // datasetStreaming is persisted, but constraint-coupled fields like + // trainOnCompletions / maxSteps / evalSteps are NON_PERSISTED and + // rehydrate to defaults. That can resurrect an invalid combo (e.g. + // streaming=true with a default trainOnCompletions) that the backend + // rejects with 422. Reconcile immediately on load instead of relying + // on a post-mount effect. + if (!state) return; + const patch = streamingCompatiblePatch(state); + if (Object.keys(patch).length > 0) { + useTrainingConfigStore.setState(patch); + } + }, }, ), ); From 143aff844b52dd34e5a20e54f5086244c3c9c9fe Mon Sep 17 00:00:00 2001 From: Etherll <61019402+Etherll@users.noreply.github.com> Date: Fri, 19 Jun 2026 22:47:46 +0300 Subject: [PATCH 08/14] studio: fix review-team findings for streaming + main merge BLOCKER: streaming + raw-text/CPT crashed on len(IterableDataset). Guard it in the start route (reject format_type=="raw" or training_type=="Continued Pretraining") and in isStreamingSupported (datasetFormat !== "raw"). Also: - models/training.py: validate hf_dataset/subset/split (charset+length, block ..//); cap dataset slice indices (le=1e9); note validator ordering - chat_templates.py: guard _apply_custom_mapping .map() for streaming - trainer.py: warn when packing+streaming - training-config-store.ts: persist-migration bump to v11 (standalone datasetStreaming backfill); add isVisionModel to NON_PERSISTED; toast on silent streamingCompatiblePatch mutations in the 4 indirect setters - tests: route rejections (max_steps, raw/cpt), slice cap, unsafe hf_dataset --- studio/backend/core/training/trainer.py | 6 ++ studio/backend/models/training.py | 55 +++++++++++ studio/backend/routes/training.py | 9 ++ .../backend/tests/test_training_streaming.py | 92 +++++++++++++++++++ .../backend/utils/datasets/chat_templates.py | 7 +- .../studio/sections/dataset-section.tsx | 8 ++ .../training/stores/training-config-store.ts | 72 +++++++++++---- 7 files changed, 230 insertions(+), 19 deletions(-) diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index 75069b4fe14..c0791ca1791 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -3383,6 +3383,12 @@ def audio_vlm_collate_fn(examples): # Only add packing for text models (not DeepSeek OCR which is VLM) if not is_deepseek_ocr: packing_enabled = training_args.get("packing", False) + if packing_enabled and training_args.get("dataset_streaming", False): + logger.warning( + "Sequence packing is enabled with dataset streaming: " + "max_steps governs training length and packed-sample " + "counts are approximate since the stream length is unknown.\n" + ) config_args["packing"] = packing_enabled logger.info( f"Sequence packing: {'enabled' if packing_enabled else 'disabled'}\n" diff --git a/studio/backend/models/training.py b/studio/backend/models/training.py index a8e9322e979..0f7839d3437 100644 --- a/studio/backend/models/training.py +++ b/studio/backend/models/training.py @@ -28,6 +28,10 @@ _MIN_VISION_IMAGE_SIZE = 256 # 2048 was the most I could get most llms to work at without getting unstable _MAX_VISION_IMAGE_SIZE = 2048 +# Upper bound for dataset slice indices. Caps `.skip(n)` on streaming datasets so +# an absurd index can't make the loader iterate effectively forever (DoS guard). +# 1e9 is far beyond any realistic fine-tuning dataset row count. +_MAX_DATASET_SLICE_INDEX = 1_000_000_000 def _parse_lr(v: Any) -> float: @@ -103,11 +107,13 @@ class TrainingStartRequest(BaseModel): dataset_slice_start: Optional[int] = Field( None, ge = 0, + le = _MAX_DATASET_SLICE_INDEX, description = "Inclusive start row index for dataset slicing", ) dataset_slice_end: Optional[int] = Field( None, ge = 0, + le = _MAX_DATASET_SLICE_INDEX, description = "Inclusive end row index for dataset slicing", ) @@ -119,6 +125,9 @@ def _compat_split(cls, values: Any) -> Any: values.setdefault("train_split", values.pop("split")) return values + # NOTE: pydantic runs all `mode="after"` validators in definition order. A + # second one, `_check_steps_or_epochs`, is defined lower in this class; keep + # these cross-field checks order-independent so the two stay decoupled. @model_validator(mode = "after") def _validate_dataset_slice(self) -> "TrainingStartRequest": # Only the ordering is validated here. No upper bound is enforced on the @@ -136,6 +145,52 @@ def _validate_dataset_slice(self) -> "TrainingStartRequest": ) return self + @field_validator("hf_dataset") + @classmethod + def _check_hf_dataset(cls, v: Optional[str]) -> Optional[str]: + # Constrain the HF dataset id to a safe charset + length to shrink the + # path-traversal / SSRF surface of `load_dataset(, ...)`. + if v is None: + return v + v = v.strip() + if not v: + return None + if len(v) > 256: + raise ValueError("hf_dataset is too long (max 256 chars)") + if ".." in v: + raise ValueError("hf_dataset must not contain '..'") + if not re.fullmatch(r"[A-Za-z0-9._\-/]+", v): + raise ValueError( + "hf_dataset may only contain letters, digits, '_', '-', '.', '/'" + ) + return v + + @field_validator("subset") + @classmethod + def _check_subset(cls, v: Optional[str]) -> Optional[str]: + if v is None: + return v + if len(v) > 128: + raise ValueError("subset is too long (max 128 chars)") + if not re.fullmatch(r"[A-Za-z0-9._\-]*", v): + raise ValueError("subset may only contain letters, digits, '_', '-', '.'") + return v + + @field_validator("train_split", "eval_split") + @classmethod + def _check_split_name(cls, v: Optional[str]) -> Optional[str]: + # Split names feed HF slice syntax (e.g. "train[:80%]"), so allow that + # charset but cap length and block path-traversal / NUL bytes. + if v is None: + return v + if len(v) > 128: + raise ValueError("split name is too long (max 128 chars)") + if "\x00" in v or ".." in v or "/" in v or "\\" in v: + raise ValueError("split name contains invalid characters") + if not re.fullmatch(r"[A-Za-z0-9_\-\[\]:%.+ ]*", v): + raise ValueError("split name contains invalid characters") + return v + @field_validator("learning_rate", mode = "before") @classmethod def _check_learning_rate(cls, v): diff --git a/studio/backend/routes/training.py b/studio/backend/routes/training.py index 439acd0e8a9..7b02c4d09a9 100644 --- a/studio/backend/routes/training.py +++ b/studio/backend/routes/training.py @@ -217,6 +217,15 @@ async def start_training( status_code = 422, detail = "dataset_streaming with evaluation requires a separate eval_split.", ) + if request.format_type == "raw" or request.training_type == "Continued Pretraining": + raise HTTPException( + status_code = 422, + detail = ( + "dataset_streaming is not supported for raw-text or " + "continued-pretraining mode yet (raw-text preprocessing " + "requires a finite dataset length)." + ), + ) # Convert request to kwargs for backend training_kwargs = { diff --git a/studio/backend/tests/test_training_streaming.py b/studio/backend/tests/test_training_streaming.py index 1652ca42d42..3c704ba4bf4 100644 --- a/studio/backend/tests/test_training_streaming.py +++ b/studio/backend/tests/test_training_streaming.py @@ -195,6 +195,73 @@ def test_streaming_start_requires_separate_eval_split(eval_split): assert "separate eval_split" in exc_info.value.detail +def test_streaming_start_rejects_missing_max_steps(): + training_route = _load_route_module( + "training_route_module_for_streaming_max_steps_test", + "routes/training.py", + ) + request = TrainingStartRequest( + model_name = "unsloth/test", + training_type = "LoRA/QLoRA", + hf_dataset = "org/dataset", + format_type = "chatml", + dataset_streaming = True, + max_steps = 0, + ) + + backend = SimpleNamespace( + current_job_id = None, + is_training_active = lambda: False, + start_training = lambda **kwargs: pytest.fail("backend should not start"), + ) + + with patch.object(training_route, "get_training_backend", return_value = backend): + with pytest.raises(HTTPException) as exc_info: + asyncio.run( + training_route.start_training(request, current_subject = "test-user") + ) + + assert exc_info.value.status_code == 422 + assert "max_steps" in exc_info.value.detail + + +@pytest.mark.parametrize( + "training_type, format_type", + [ + ("LoRA/QLoRA", "raw"), # raw-text format alone + ("Continued Pretraining", "chatml"), # CPT training_type alone + ], +) +def test_streaming_start_rejects_raw_text_and_cpt(training_type, format_type): + training_route = _load_route_module( + "training_route_module_for_streaming_raw_cpt_test", + "routes/training.py", + ) + request = TrainingStartRequest( + model_name = "unsloth/test", + training_type = training_type, + hf_dataset = "org/dataset", + format_type = format_type, + dataset_streaming = True, + max_steps = 10, + ) + + backend = SimpleNamespace( + current_job_id = None, + is_training_active = lambda: False, + start_training = lambda **kwargs: pytest.fail("backend should not start"), + ) + + with patch.object(training_route, "get_training_backend", return_value = backend): + with pytest.raises(HTTPException) as exc_info: + asyncio.run( + training_route.start_training(request, current_subject = "test-user") + ) + + assert exc_info.value.status_code == 422 + assert "raw-text or continued-pretraining" in exc_info.value.detail + + def test_dataset_slice_bounds_are_non_negative(): with pytest.raises(ValidationError): TrainingStartRequest( @@ -236,6 +303,31 @@ def test_dataset_slice_accepts_equal_and_ordered_bounds(): assert ordered.dataset_slice_end == 9 +def test_dataset_slice_rejects_above_max_index(): + # Upper bound guards streaming `.skip(n)` against pathological indices. + with pytest.raises(ValidationError): + TrainingStartRequest( + model_name = "unsloth/test", + training_type = "LoRA/QLoRA", + format_type = "alpaca", + dataset_slice_start = 2_000_000_000, + ) + + +@pytest.mark.parametrize( + "bad_hf_dataset", + ["../../etc/passwd", "org/../../secret", "a" * 257], +) +def test_hf_dataset_rejects_unsafe_values(bad_hf_dataset): + with pytest.raises(ValidationError): + TrainingStartRequest( + model_name = "unsloth/test", + training_type = "LoRA/QLoRA", + format_type = "alpaca", + hf_dataset = bad_hf_dataset, + ) + + def test_is_streaming_dataset_detects_hf_iterable(): assert is_streaming_dataset(_iterable_dataset([{"a": 1}])) is True diff --git a/studio/backend/utils/datasets/chat_templates.py b/studio/backend/utils/datasets/chat_templates.py index dbbeadf64b0..5f566aa41ed 100644 --- a/studio/backend/utils/datasets/chat_templates.py +++ b/studio/backend/utils/datasets/chat_templates.py @@ -252,7 +252,12 @@ def _apply_custom_mapping(examples): return result try: - dataset = dataset.map(_apply_custom_mapping, batched = True, batch_size = batch_size) + # Mirror the other call sites: omit eager-only kwargs (num_proc/desc) + # for streaming IterableDatasets, whose .map() rejects them. + custom_map_kwargs = {"batched": True, "batch_size": batch_size} + if not is_streaming_dataset(dataset): + custom_map_kwargs["desc"] = "Applying custom ChatML mapping" + dataset = dataset.map(_apply_custom_mapping, **custom_map_kwargs) # Update to use conversations format final_format = "chatml_conversations" chat_column = "conversations" diff --git a/studio/frontend/src/features/studio/sections/dataset-section.tsx b/studio/frontend/src/features/studio/sections/dataset-section.tsx index ac6bbca2bb4..117b359beed 100644 --- a/studio/frontend/src/features/studio/sections/dataset-section.tsx +++ b/studio/frontend/src/features/studio/sections/dataset-section.tsx @@ -46,6 +46,8 @@ import { useTrainingConfigStore, type LocalDatasetInfo, } from "@/features/training"; +// Imported directly from the store module rather than the "@/features/training" +// barrel to avoid an import cycle (the barrel re-exports this section's siblings). import { hasSeparateStreamingEvalSplit } from "@/features/training/stores/training-config-store"; import { useNavigate } from "@tanstack/react-router"; import { @@ -225,6 +227,12 @@ export function DatasetSection() { maxSteps > 0 && !trainOnCompletions && hasSeparateStreamingEvalSplit({ evalSteps, datasetSplit, datasetEvalSplit }) && + // Raw-text / CPT mode preprocesses via a finite-length code path + // (prepare_raw_text_dataset -> len(dataset)) that crashes on an + // IterableDataset, so streaming is unsupported there. CPT always forces + // datasetFormat="raw", so this single check covers both cases; the backend + // also rejects format_type=="raw" / Continued Pretraining defensively. + datasetFormat !== "raw" && !isVisionModel && !isAudioModel && !isDatasetImage && 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 50fbe138668..8f7fbd58354 100644 --- a/studio/frontend/src/features/training/stores/training-config-store.ts +++ b/studio/frontend/src/features/training/stores/training-config-store.ts @@ -130,6 +130,7 @@ const NON_PERSISTED_STATE_KEYS: ReadonlySet = new Set "isDatasetAudio", "trainOnCompletions", "maxPositionEmbeddings", + "isVisionModel", ]); function partializePersistedState( @@ -208,6 +209,25 @@ function streamingCompatiblePatch( return patch; } +// streamingCompatiblePatch can silently flip streaming-coupled fields. Surface a +// toast when it does, so the indirect setters (split / eval-split / max-steps / +// eval-steps) match setDatasetStreaming's "tell the user what changed" behavior. +function notifyStreamingCompat(patch: Partial): void { + if (patch.datasetStreaming === false) { + toast.info("Streaming turned off: streaming needs a fixed Max Steps > 0."); + return; + } + const disabled = [ + patch.trainOnCompletions === false && "assistant-completions-only", + patch.evalSteps === 0 && "evaluation (needs a separate eval split)", + ].filter(Boolean); + if (disabled.length > 0) { + toast.info( + `Adjusted for streaming. Disabled incompatible options: ${disabled.join(", ")}.`, + ); + } +} + type TrainingMethodStatePatch = Partial< Pick< TrainingConfigState, @@ -678,14 +698,16 @@ export const useTrainingConfigStore = create()( setDatasetSplit: (datasetSplit) => { const state = get(); const nextState = { ...state, datasetSplit }; + const streamingPatch = streamingCompatiblePatch(nextState); set({ datasetSplit, datasetManualMapping: emptyManualMapping(), isDatasetImage: null, isDatasetAudio: false, isCheckingDataset: false, - ...streamingCompatiblePatch(nextState), + ...streamingPatch, }); + notifyStreamingCompat(streamingPatch); const datasetName = state.datasetSource === "huggingface" @@ -712,11 +734,17 @@ export const useTrainingConfigStore = create()( setDatasetEvalSplit: (datasetEvalSplit) => { const state = get(); const evalSteps = datasetEvalSplit ? 0.1 : 0; + const streamingPatch = streamingCompatiblePatch({ + ...state, + datasetEvalSplit, + evalSteps, + }); set({ datasetEvalSplit, evalSteps, - ...streamingCompatiblePatch({ ...state, datasetEvalSplit, evalSteps }), + ...streamingPatch, }); + notifyStreamingCompat(streamingPatch); }, setDatasetStreaming: (datasetStreaming) => { if (!datasetStreaming) { @@ -816,19 +844,24 @@ export const useTrainingConfigStore = create()( setWarmupSteps: (warmupSteps) => set({ warmupSteps }), setMaxSteps: (maxSteps) => { const state = get(); + // streamingCompatiblePatch already turns streaming off when maxSteps<=0, + // so no separate datasetStreaming reset is needed here. + const streamingPatch = streamingCompatiblePatch({ ...state, maxSteps }); set({ maxSteps, - ...(maxSteps > 0 ? {} : { datasetStreaming: false }), - ...streamingCompatiblePatch({ ...state, maxSteps }), + ...streamingPatch, }); + notifyStreamingCompat(streamingPatch); }, setSaveSteps: (saveSteps) => set({ saveSteps }), setEvalSteps: (evalSteps) => { const state = get(); + const streamingPatch = streamingCompatiblePatch({ ...state, evalSteps }); set({ evalSteps, - ...streamingCompatiblePatch({ ...state, evalSteps }), + ...streamingPatch, }); + notifyStreamingCompat(streamingPatch); }, setPacking: (packing) => set({ packing }), setTrainOnCompletions: (trainOnCompletions) => { @@ -886,7 +919,7 @@ export const useTrainingConfigStore = create()( }, { name: "unsloth_training_config_v1", - version: 10, + version: 11, migrate: (persisted, version) => { const s = persisted as Record; if (version < 2 && s.datasetSubset == null && s.datasetConfig != null) { @@ -922,20 +955,23 @@ export const useTrainingConfigStore = create()( s.weightDecay = DEFAULT_HYPERPARAMS.weightDecay; } } - if (version < 10) { - s.datasetStreaming ??= false; - if (s.trainingMethod === "cpt") { - // Backfill CPT defaults for state persisted before they existed. - s.loraRank = 128; - s.loraAlpha = 32; - s.loraVariant = "rslora"; - s.targetModules = CPT_TARGET_MODULES; - s.datasetFormat = "raw"; - if (s.learningRate == null || s.learningRate === LR_DEFAULT_LORA) { - s.learningRate = LR_DEFAULT_CPT; - } + if (version < 10 && s.trainingMethod === "cpt") { + // Backfill CPT defaults for state persisted before they existed. + s.loraRank = 128; + s.loraAlpha = 32; + s.loraVariant = "rslora"; + s.targetModules = CPT_TARGET_MODULES; + s.datasetFormat = "raw"; + if (s.learningRate == null || s.learningRate === LR_DEFAULT_LORA) { + s.learningRate = LR_DEFAULT_CPT; } } + if (version < 11) { + // Standalone bump: users already on main's v10 (CPT) skipped the + // streaming backfill when it was nested under v<10, so give it its + // own version guard. + s.datasetStreaming ??= false; + } return s as unknown as TrainingConfigStore; }, partialize: partializePersistedState, From f1acb51fdf4c4d0c2a05458a100380ec02f32f7d Mon Sep 17 00:00:00 2001 From: Etherll <61019402+Etherll@users.noreply.github.com> Date: Fri, 19 Jun 2026 22:47:46 +0300 Subject: [PATCH 09/14] studio: enable raw-text/CPT dataset streaming + streaming UX polish - raw_text: keep the lazy filter but skip len()-based row counting for IterableDatasets so raw-text / CPT can stream; guard the eval-size log - routes/trainer: drop the raw/CPT streaming block; add a defensive not-streaming guard on the eval auto-split (train_test_split) - dataset-section: streaming toggle is visible-but-disabled and lists the exact unmet requirement(s) in its tooltip; block embedding models - training-start-overlay: show "streaming (no full download)" instead of a stuck download bar for streaming runs - trim the streaming test suite to the high-value cases --- studio/backend/core/training/trainer.py | 12 +- studio/backend/routes/training.py | 9 - .../backend/tests/test_training_streaming.py | 262 +++++++----------- studio/backend/utils/datasets/raw_text.py | 17 ++ .../studio/sections/dataset-section.tsx | 144 ++++++---- .../studio/training-start-overlay.tsx | 10 +- studio/frontend/src/i18n/locales/en.ts | 1 + 7 files changed, 220 insertions(+), 235 deletions(-) diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index 7fd78f4f430..805afe73bdc 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -2619,11 +2619,19 @@ def _apply_raw_text_prep(ds: Dataset, split_name: str) -> Dataset: } if has_separate_eval_source and eval_dataset is not None: + eval_rows = ( + f"{len(eval_dataset):,} rows" + if hasattr(eval_dataset, "__len__") + else "streaming" + ) logger.info( f"{_raw_mode_label().capitalize()}: eval dataset " - f"({len(eval_dataset)} rows) kept as raw text\n" + f"({eval_rows}) kept as raw text\n" ) - elif eval_enabled and not has_separate_eval_source: + elif eval_enabled and not has_separate_eval_source and not dataset_streaming: + # _resolve_eval_split_from_dataset does a train_test_split (needs + # len/random access). Streaming always provides a separate eval + # split (route-enforced), so this auto-split is non-streaming only. split_result = self._resolve_eval_split_from_dataset(dataset) if split_result is not None: train_portion, eval_dataset = split_result diff --git a/studio/backend/routes/training.py b/studio/backend/routes/training.py index f43f65ba035..8aaefa3ed6e 100644 --- a/studio/backend/routes/training.py +++ b/studio/backend/routes/training.py @@ -215,15 +215,6 @@ async def start_training( status_code = 422, detail = "dataset_streaming with evaluation requires a separate eval_split.", ) - if request.format_type == "raw" or request.training_type == "Continued Pretraining": - raise HTTPException( - status_code = 422, - detail = ( - "dataset_streaming is not supported for raw-text or " - "continued-pretraining mode yet (raw-text preprocessing " - "requires a finite dataset length)." - ), - ) # Convert request to backend kwargs. training_kwargs = { diff --git a/studio/backend/tests/test_training_streaming.py b/studio/backend/tests/test_training_streaming.py index 3c704ba4bf4..e726282a790 100644 --- a/studio/backend/tests/test_training_streaming.py +++ b/studio/backend/tests/test_training_streaming.py @@ -5,7 +5,6 @@ import asyncio import importlib.util -import sys from pathlib import Path from types import SimpleNamespace from unittest.mock import patch @@ -16,11 +15,7 @@ from models.training import TrainingStartRequest from utils.datasets.chat_templates import apply_chat_template_to_dataset -from utils.datasets.format_conversion import ( - convert_alpaca_to_chatml, - convert_chatml_to_alpaca, - standardize_chat_format, -) +from utils.datasets.format_conversion import convert_chatml_to_alpaca from utils.datasets.iterable import is_streaming_dataset datasets = pytest.importorskip("datasets") @@ -57,6 +52,10 @@ def _iterable_dataset(rows): return datasets.IterableDataset.from_generator(lambda: iter(rows)) +# --- Streaming keeps dataset.map() lazy: eager-only kwargs (num_proc/desc) are +# omitted for IterableDatasets, which reject them. One per module. --- + + def test_chat_template_mapping_omits_eager_kwargs_for_streaming(monkeypatch): seen_kwargs = [] original_map = datasets.IterableDataset.map @@ -130,6 +129,77 @@ def spy_map(self, *args, **kwargs): assert all("desc" not in kwargs for kwargs in seen_kwargs) +# --- Streaming detection --- + + +def test_is_streaming_dataset_detects_hf_iterable(): + assert is_streaming_dataset(_iterable_dataset([{"a": 1}])) is True + + +def test_is_streaming_dataset_false_for_plain_list(): + assert is_streaming_dataset([{"a": 1}]) is False + + +# --- Raw-text / CPT streaming: keep the lazy filter, skip the len()-based +# counting that would TypeError on an IterableDataset (the BLOCKER fix). --- + + +def test_drop_invalid_text_rows_streaming_keeps_filter_skips_len(): + from utils.datasets.raw_text import _drop_invalid_text_rows + + stream = datasets.Dataset.from_list( + [{"text": "keep1"}, {"text": None}, {"text": "keep2"}] + ).to_iterable_dataset() + assert not hasattr(stream, "__len__") + + filtered, notices = _drop_invalid_text_rows( + stream, mode_title = "Raw text", split_scope = "this dataset" + ) + + # Result still streams; only string-'text' rows survive. + assert [row["text"] for row in filtered] == ["keep1", "keep2"] + assert any(n.level == "info" for n in notices) + + +# --- Request validation --- + + +def test_dataset_slice_bounds_are_non_negative(): + with pytest.raises(ValidationError): + TrainingStartRequest( + model_name = "unsloth/test", + training_type = "LoRA/QLoRA", + format_type = "alpaca", + dataset_slice_start = -1, + ) + + with pytest.raises(ValidationError): + TrainingStartRequest( + model_name = "unsloth/test", + training_type = "LoRA/QLoRA", + format_type = "alpaca", + dataset_slice_start = 5, + dataset_slice_end = 4, + ) + + +@pytest.mark.parametrize( + "bad_hf_dataset", + ["../../etc/passwd", "org/../../secret", "a" * 257], +) +def test_hf_dataset_rejects_unsafe_values(bad_hf_dataset): + with pytest.raises(ValidationError): + TrainingStartRequest( + model_name = "unsloth/test", + training_type = "LoRA/QLoRA", + format_type = "alpaca", + hf_dataset = bad_hf_dataset, + ) + + +# --- Start-route streaming compatibility guards --- + + def test_streaming_start_rejects_train_on_completions_before_backend_start(): training_route = _load_route_module( "training_route_module_for_streaming_completion_test", @@ -228,13 +298,15 @@ def test_streaming_start_rejects_missing_max_steps(): @pytest.mark.parametrize( "training_type, format_type", [ - ("LoRA/QLoRA", "raw"), # raw-text format alone - ("Continued Pretraining", "chatml"), # CPT training_type alone + ("LoRA/QLoRA", "raw"), # raw-text format + ("Continued Pretraining", "chatml"), # CPT ], ) -def test_streaming_start_rejects_raw_text_and_cpt(training_type, format_type): +def test_streaming_start_accepts_raw_text_and_cpt(training_type, format_type): + # Streaming + raw-text / CPT is supported: _drop_invalid_text_rows skips its + # len()-based checks for IterableDatasets, so the start route must NOT reject. training_route = _load_route_module( - "training_route_module_for_streaming_raw_cpt_test", + "training_route_module_for_streaming_raw_cpt_accept_test", "routes/training.py", ) request = TrainingStartRequest( @@ -246,171 +318,27 @@ def test_streaming_start_rejects_raw_text_and_cpt(training_type, format_type): max_steps = 10, ) + captured = {} + + def _start_training(**kwargs): + captured.update(kwargs) + return True + backend = SimpleNamespace( - current_job_id = None, + current_job_id = "job_test", is_training_active = lambda: False, - start_training = lambda **kwargs: pytest.fail("backend should not start"), + start_training = _start_training, ) with patch.object(training_route, "get_training_backend", return_value = backend): - with pytest.raises(HTTPException) as exc_info: - asyncio.run( + with patch.object(training_route, "load_model_defaults", return_value = {}): + response = asyncio.run( training_route.start_training(request, current_subject = "test-user") ) - assert exc_info.value.status_code == 422 - assert "raw-text or continued-pretraining" in exc_info.value.detail - - -def test_dataset_slice_bounds_are_non_negative(): - with pytest.raises(ValidationError): - TrainingStartRequest( - model_name = "unsloth/test", - training_type = "LoRA/QLoRA", - format_type = "alpaca", - dataset_slice_start = -1, - ) - - with pytest.raises(ValidationError): - TrainingStartRequest( - model_name = "unsloth/test", - training_type = "LoRA/QLoRA", - format_type = "alpaca", - dataset_slice_start = 5, - dataset_slice_end = 4, - ) - - -def test_dataset_slice_accepts_equal_and_ordered_bounds(): - # start == end is intentionally allowed (single-row slice); start < end too. - equal = TrainingStartRequest( - model_name = "unsloth/test", - training_type = "LoRA/QLoRA", - format_type = "alpaca", - dataset_slice_start = 5, - dataset_slice_end = 5, - ) - assert equal.dataset_slice_start == 5 - assert equal.dataset_slice_end == 5 - - ordered = TrainingStartRequest( - model_name = "unsloth/test", - training_type = "LoRA/QLoRA", - format_type = "alpaca", - dataset_slice_start = 2, - dataset_slice_end = 9, - ) - assert ordered.dataset_slice_end == 9 - - -def test_dataset_slice_rejects_above_max_index(): - # Upper bound guards streaming `.skip(n)` against pathological indices. - with pytest.raises(ValidationError): - TrainingStartRequest( - model_name = "unsloth/test", - training_type = "LoRA/QLoRA", - format_type = "alpaca", - dataset_slice_start = 2_000_000_000, - ) - - -@pytest.mark.parametrize( - "bad_hf_dataset", - ["../../etc/passwd", "org/../../secret", "a" * 257], -) -def test_hf_dataset_rejects_unsafe_values(bad_hf_dataset): - with pytest.raises(ValidationError): - TrainingStartRequest( - model_name = "unsloth/test", - training_type = "LoRA/QLoRA", - format_type = "alpaca", - hf_dataset = bad_hf_dataset, - ) - - -def test_is_streaming_dataset_detects_hf_iterable(): - assert is_streaming_dataset(_iterable_dataset([{"a": 1}])) is True - - -def test_is_streaming_dataset_false_for_plain_list(): - assert is_streaming_dataset([{"a": 1}]) is False - - -def test_is_streaming_dataset_torch_branch_when_datasets_unavailable(): - torch = pytest.importorskip("torch") - - class _TorchIterable(torch.utils.data.IterableDataset): - def __iter__(self): - return iter([1, 2, 3]) - - # Force the `from datasets import IterableDataset` import to fail so the - # torch detection branch is exercised. - with patch.dict(sys.modules, {"datasets": None}): - assert is_streaming_dataset(_TorchIterable()) is True - - -def test_is_streaming_dataset_false_when_both_backends_unavailable(): - # Both `datasets` and `torch.utils.data` imports fail -> graceful False. - with patch.dict(sys.modules, {"datasets": None, "torch.utils.data": None}): - assert is_streaming_dataset(object()) is False - - -def test_standardize_chat_format_omits_eager_kwargs_for_streaming(monkeypatch): - seen_kwargs = [] - original_map = datasets.IterableDataset.map - - def spy_map(self, *args, **kwargs): - seen_kwargs.append(dict(kwargs)) - return original_map(self, *args, **kwargs) - - monkeypatch.setattr(datasets.IterableDataset, "map", spy_map) - - dataset = _iterable_dataset( - [ - { - "conversations": [ - {"from": "human", "value": "Hi"}, - {"from": "gpt", "value": "Hello"}, - ] - } - ] - ) - result = standardize_chat_format( - dataset, - tokenizer = _Tokenizer(), - batch_size = 1, - num_proc = 2, - ) - - next(iter(result)) - assert seen_kwargs - assert all("num_proc" not in kwargs for kwargs in seen_kwargs) - assert all("desc" not in kwargs for kwargs in seen_kwargs) - - -def test_convert_alpaca_to_chatml_omits_eager_kwargs_for_streaming(monkeypatch): - seen_kwargs = [] - original_map = datasets.IterableDataset.map - - def spy_map(self, *args, **kwargs): - seen_kwargs.append(dict(kwargs)) - return original_map(self, *args, **kwargs) - - monkeypatch.setattr(datasets.IterableDataset, "map", spy_map) - - converted = convert_alpaca_to_chatml( - _iterable_dataset( - [{"instruction": "Question", "input": "", "output": "Answer"}] - ), - batch_size = 1, - num_proc = 2, - ) - - row = next(iter(converted)) - assert "conversations" in row - assert seen_kwargs - assert all("num_proc" not in kwargs for kwargs in seen_kwargs) - assert all("desc" not in kwargs for kwargs in seen_kwargs) + assert response.status == "queued" + assert captured["dataset_streaming"] is True + assert captured["format_type"] == format_type def test_streaming_start_happy_path_reaches_backend(): diff --git a/studio/backend/utils/datasets/raw_text.py b/studio/backend/utils/datasets/raw_text.py index 03315fb2870..1a1c09c93d3 100644 --- a/studio/backend/utils/datasets/raw_text.py +++ b/studio/backend/utils/datasets/raw_text.py @@ -40,7 +40,24 @@ def _split_scope(split_name: str | None) -> str: def _drop_invalid_text_rows( dataset: Dataset, *, mode_title: str, split_scope: str ) -> tuple[Dataset, list[RawTextNotice]]: + # Lazy filter — drops rows whose 'text' is null/non-string before they reach + # the tokenizer. Works on both Dataset and streaming IterableDataset. filtered_dataset = dataset.filter(lambda ex: isinstance(ex["text"], str)) + + # Streaming datasets (IterableDataset) have no __len__, so we can't count the + # dropped rows or verify the result is non-empty without consuming the whole + # stream. Keep the filter, skip only the len()-based diagnostics. + if not hasattr(dataset, "__len__"): + return filtered_dataset, [ + RawTextNotice( + message = ( + f"{mode_title}: streaming dataset — rows with null or " + f"non-string 'text' in {split_scope} are dropped on the fly." + ), + level = "info", + ) + ] + dropped_rows = len(dataset) - len(filtered_dataset) if not dropped_rows: return filtered_dataset, [] diff --git a/studio/frontend/src/features/studio/sections/dataset-section.tsx b/studio/frontend/src/features/studio/sections/dataset-section.tsx index 65ec61ebc6d..aa48d00a5a7 100644 --- a/studio/frontend/src/features/studio/sections/dataset-section.tsx +++ b/studio/frontend/src/features/studio/sections/dataset-section.tsx @@ -3,6 +3,7 @@ import { SectionCard } from "@/components/section-card"; import { Button } from "@/components/ui/button"; +import { Checkbox } from "@/components/ui/checkbox"; import { Collapsible, CollapsibleContent, @@ -175,6 +176,7 @@ export function DatasetSection() { evalSteps, isVisionModel, isAudioModel, + isEmbeddingModel, isDatasetImage, isDatasetAudio, uploadedFile, @@ -208,6 +210,7 @@ export function DatasetSection() { evalSteps: s.evalSteps, isVisionModel: s.isVisionModel, isAudioModel: s.isAudioModel, + isEmbeddingModel: s.isEmbeddingModel, isDatasetImage: s.isDatasetImage, isDatasetAudio: s.isDatasetAudio, uploadedFile: s.uploadedFile, @@ -222,25 +225,41 @@ export function DatasetSection() { })), ); - // Streaming is only supported for Hugging Face text datasets. - // Hide the toggle for vision/audio models or datasets detected as multimodal, - // since downstream preprocessing (convert_to_vlm_format, audio collators) - // requires random access and would crash on an IterableDataset. - const isStreamingSupported = - datasetSource === "huggingface" && - maxSteps > 0 && - !trainOnCompletions && - hasSeparateStreamingEvalSplit({ evalSteps, datasetSplit, datasetEvalSplit }) && - // Raw-text / CPT mode preprocesses via a finite-length code path - // (prepare_raw_text_dataset -> len(dataset)) that crashes on an - // IterableDataset, so streaming is unsupported there. CPT always forces - // datasetFormat="raw", so this single check covers both cases; the backend - // also rejects format_type=="raw" / Continued Pretraining defensively. - datasetFormat !== "raw" && - !isVisionModel && - !isAudioModel && - !isDatasetImage && - !isDatasetAudio; + // Streaming is only supported for Hugging Face text datasets. Rather than + // hiding the toggle when a constraint isn't met, keep it visible but disabled + // and list the exact unmet requirement(s) in its tooltip — a control that + // silently disappears is confusing. Downstream preprocessing + // (convert_to_vlm_format, audio collators) needs random access and would + // crash on an IterableDataset, hence the constraints below. + const streamingBlockers: string[] = []; + if (datasetSource !== "huggingface") + streamingBlockers.push( + "Use a Hugging Face dataset (not a local upload or S3 source).", + ); + if (maxSteps <= 0) + streamingBlockers.push( + "Set Max Steps > 0 — streaming datasets have no known length.", + ); + if (trainOnCompletions) + streamingBlockers.push('Turn off "Assistant completions only".'); + if (!hasSeparateStreamingEvalSplit({ evalSteps, datasetSplit, datasetEvalSplit })) + streamingBlockers.push( + "Pick a separate eval split — evaluation is on but no distinct eval split is set.", + ); + if (isVisionModel) + streamingBlockers.push("Vision models don't support streaming."); + if (isAudioModel) + streamingBlockers.push("Audio models don't support streaming."); + if (isEmbeddingModel) + streamingBlockers.push( + "Embedding models don't support streaming (training needs the full dataset).", + ); + if (isDatasetImage) + streamingBlockers.push("This dataset looks like images, which can't stream."); + if (isDatasetAudio) + streamingBlockers.push("This dataset looks like audio, which can't stream."); + + const isStreamingSupported = streamingBlockers.length === 0; // If streaming was previously enabled but the config became incompatible // (model switched to vision, dataset detected as image, etc.), clear it so @@ -1162,43 +1181,56 @@ export function DatasetSection() {
- {isStreamingSupported && ( -
- - Streaming Mode - - - - - - Load Hugging Face datasets using streaming mode. Requires max_steps > 0 and a separate eval split when evaluation is enabled. - - - - - - -

- Only applies to Hugging Face text datasets. -

-
- )} +
+ setDatasetStreaming(!!v)} + /> + + + + + + + {isStreamingSupported ? ( + + Stream Hugging Face text datasets instead of + downloading them. + + ) : ( +
+

+ Streaming unavailable. To enable: +

+
    + {streamingBlockers.map((reason) => ( +
  • {reason}
  • + ))} +
+
+ )} +
+
+
diff --git a/studio/frontend/src/features/studio/training-start-overlay.tsx b/studio/frontend/src/features/studio/training-start-overlay.tsx index 3b0cce658b1..c69b25577c5 100644 --- a/studio/frontend/src/features/studio/training-start-overlay.tsx +++ b/studio/frontend/src/features/studio/training-start-overlay.tsx @@ -263,6 +263,10 @@ export function TrainingStartOverlay({ const configuredModel = useTrainingConfigStore((s) => s.selectedModel); const datasetSource = useTrainingConfigStore((s) => s.datasetSource); const dataset = useTrainingConfigStore((s) => s.dataset); + // Streaming runs never fully download the dataset (only small metadata lands + // in the HF cache), so the cache-watching download bar would sit near 0% + // forever and read as "stuck downloading". Show a streaming note instead. + const datasetStreaming = useTrainingConfigStore((s) => s.datasetStreaming); // Only HF datasets have a download phase to track; uploaded files are already // on disk by the time the overlay shows up. const hfDatasetName = datasetSource === "huggingface" ? dataset : null; @@ -380,7 +384,11 @@ export function TrainingStartOverlay({ step: currentStep, })} - {datasetDownload.downloadedBytes > 0 || datasetDownload.cachePath ? ( + {datasetStreaming ? ( + + {t("studio.trainingStart.datasetStreaming")} + + ) : datasetDownload.downloadedBytes > 0 || datasetDownload.cachePath ? ( Date: Fri, 19 Jun 2026 19:49:49 +0000 Subject: [PATCH 10/14] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- studio/backend/core/training/trainer.py | 46 +++++-------------- studio/backend/models/training.py | 12 ++--- .../backend/tests/test_training_streaming.py | 16 ++----- .../utils/datasets/format_conversion.py | 1 + studio/backend/utils/datasets/iterable.py | 2 - 5 files changed, 19 insertions(+), 58 deletions(-) diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index 805afe73bdc..62c70402e6d 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -2395,9 +2395,7 @@ def _apply_raw_text_prep(ds: Dataset, split_name: str) -> Dataset: load_kwargs["name"] = subset if dataset_streaming: - self._update_progress( - status_message = f"Streaming dataset: {dataset_source}..." - ) + self._update_progress(status_message = f"Streaming dataset: {dataset_source}...") dataset = load_dataset(**load_kwargs, streaming = True) # Optional iterable slicing @@ -2494,9 +2492,7 @@ def _apply_raw_text_prep(ds: Dataset, split_name: str) -> Dataset: if subset: probe_kwargs["config_name"] = subset try: - available_splits = get_dataset_split_names( - **probe_kwargs - ) + available_splits = get_dataset_split_names(**probe_kwargs) except Exception as probe_err: raise ValueError( f"Could not list splits for '{dataset_source}' " @@ -2508,9 +2504,7 @@ def _apply_raw_text_prep(ds: Dataset, split_name: str) -> Dataset: f"dataset '{dataset_source}'. Available splits: " f"{available_splits}" ) - eval_dataset = load_dataset( - **eval_load_kwargs, streaming = True - ) + eval_dataset = load_dataset(**eval_load_kwargs, streaming = True) else: eval_dataset = load_dataset(**eval_load_kwargs) @@ -2520,9 +2514,7 @@ def _apply_raw_text_prep(ds: Dataset, split_name: str) -> Dataset: f"Loaded eval split '{eval_split}' with {len(eval_dataset)} rows\n" ) else: - logger.info( - f"Loaded eval split '{eval_split}' in streaming mode\n" - ) + logger.info(f"Loaded eval split '{eval_split}' in streaming mode\n") elif eval_split and eval_split == effective_train: if dataset_streaming: raise ValueError( @@ -2708,9 +2700,7 @@ def _apply_raw_text_prep(ds: Dataset, split_name: str) -> Dataset: ) eval_dataset = eval_info["dataset"] logger.info("Eval dataset formatted successfully\n") - elif ( - eval_enabled and not has_separate_eval_source and not dataset_streaming - ): + elif eval_enabled and not has_separate_eval_source and not dataset_streaming: # No separate eval source — split the already-formatted dataset formatted_dataset = dataset_info["dataset"] split_result = self._resolve_eval_split_from_dataset(formatted_dataset) @@ -3402,9 +3392,7 @@ def audio_vlm_collate_fn(examples): # Audio VLM (e.g. Gemma 3N + audio): raw Dataset from _format_audio_vlm_dataset # Notebook uses processing_class=processor.tokenizer (text tokenizer only) # Raw-text runs are routed to the text path below. - train_dataset = ( - dataset["dataset"] if isinstance(dataset, dict) else dataset - ) + train_dataset = dataset["dataset"] if isinstance(dataset, dict) else dataset processing_class = ( self.tokenizer.tokenizer if hasattr(self.tokenizer, "tokenizer") @@ -3445,9 +3433,7 @@ def audio_vlm_collate_fn(examples): if isinstance(self.tokenizer, ProcessorMixin) and hasattr( self.tokenizer, "tokenizer" ): - logger.info( - "Unwrapping Processor → raw tokenizer for text-only SFTTrainer" - ) + logger.info("Unwrapping Processor → raw tokenizer for text-only SFTTrainer") sft_tokenizer = self.tokenizer.tokenizer if is_cpt: @@ -3581,9 +3567,7 @@ def audio_vlm_collate_fn(examples): # every sample becomes all -100, and Unsloth drops them, leaving # 0 usable samples. Skip this len()-based check for streaming. if detect_streaming_dataset(self.trainer.train_dataset): - logger.info( - "Skipping post-filter length check for streaming dataset\n" - ) + logger.info("Skipping post-filter length check for streaming dataset\n") else: filtered_len = len(self.trainer.train_dataset) original_dataset_obj = ( @@ -3591,11 +3575,7 @@ def audio_vlm_collate_fn(examples): ) original_len = len(original_dataset_obj) dropped = original_len - filtered_len - drop_pct = ( - round(100 * dropped / original_len, 1) - if original_len > 0 - else 0 - ) + drop_pct = round(100 * dropped / original_len, 1) if original_len > 0 else 0 if filtered_len == 0 or drop_pct > 30: max_seq = training_args.get("max_seq_length", 2048) @@ -3618,9 +3598,7 @@ def audio_vlm_collate_fn(examples): f"({drop_pct}%) were dropped (all labels " f"masked). {filtered_len} samples remain.\n" ) - logger.info( - f"Post-filter dataset size: {filtered_len} samples\n" - ) + logger.info(f"Post-filter dataset size: {filtered_len} samples\n") except Exception as e: logger.warning(f"Failed to apply train on responses only: {e}") @@ -3634,9 +3612,7 @@ def audio_vlm_collate_fn(examples): # ========== PROGRESS TRACKING ========== self.trainer.add_callback(self._create_progress_callback()) - train_dataset_obj = ( - dataset["dataset"] if isinstance(dataset, dict) else dataset - ) + train_dataset_obj = dataset["dataset"] if isinstance(dataset, dict) else dataset is_streaming_dataset = detect_streaming_dataset(train_dataset_obj) max_steps_value = training_args.get("max_steps") diff --git a/studio/backend/models/training.py b/studio/backend/models/training.py index 6e721340397..c821112f56b 100644 --- a/studio/backend/models/training.py +++ b/studio/backend/models/training.py @@ -128,16 +128,12 @@ class TrainingStartRequest(BaseModel): format_type: str = Field(..., description = "Dataset format type") subset: Optional[str] = None train_split: Optional[str] = Field("train", description = "Training split name") - eval_split: Optional[str] = Field( - None, description = "Eval split name. None = auto-detect" - ) + eval_split: Optional[str] = Field(None, description = "Eval split name. None = auto-detect") dataset_streaming: bool = Field( False, description = "Whether to load the Hugging Face dataset in streaming mode", ) - eval_steps: float = Field( - 0.00, description = "Fraction of total steps between evals (0-1)" - ) + eval_steps: float = Field(0.00, description = "Fraction of total steps between evals (0-1)") dataset_slice_start: Optional[int] = Field( None, ge = 0, @@ -194,9 +190,7 @@ def _check_hf_dataset(cls, v: Optional[str]) -> Optional[str]: if ".." in v: raise ValueError("hf_dataset must not contain '..'") if not re.fullmatch(r"[A-Za-z0-9._\-/]+", v): - raise ValueError( - "hf_dataset may only contain letters, digits, '_', '-', '.', '/'" - ) + raise ValueError("hf_dataset may only contain letters, digits, '_', '-', '.', '/'") return v @field_validator("subset") diff --git a/studio/backend/tests/test_training_streaming.py b/studio/backend/tests/test_training_streaming.py index e726282a790..a0d90089464 100644 --- a/studio/backend/tests/test_training_streaming.py +++ b/studio/backend/tests/test_training_streaming.py @@ -43,9 +43,7 @@ def apply_chat_template( ): assert tokenize is False assert add_generation_prompt is False - return "\n".join( - f"{message['role']}: {message['content']}" for message in conversation - ) + return "\n".join(f"{message['role']}: {message['content']}" for message in conversation) def _iterable_dataset(rows): @@ -223,9 +221,7 @@ def test_streaming_start_rejects_train_on_completions_before_backend_start(): with patch.object(training_route, "get_training_backend", return_value = backend): with pytest.raises(HTTPException) as exc_info: - asyncio.run( - training_route.start_training(request, current_subject = "test-user") - ) + asyncio.run(training_route.start_training(request, current_subject = "test-user")) assert exc_info.value.status_code == 422 assert "train_on_completions" in exc_info.value.detail @@ -257,9 +253,7 @@ def test_streaming_start_requires_separate_eval_split(eval_split): with patch.object(training_route, "get_training_backend", return_value = backend): with pytest.raises(HTTPException) as exc_info: - asyncio.run( - training_route.start_training(request, current_subject = "test-user") - ) + asyncio.run(training_route.start_training(request, current_subject = "test-user")) assert exc_info.value.status_code == 422 assert "separate eval_split" in exc_info.value.detail @@ -287,9 +281,7 @@ def test_streaming_start_rejects_missing_max_steps(): with patch.object(training_route, "get_training_backend", return_value = backend): with pytest.raises(HTTPException) as exc_info: - asyncio.run( - training_route.start_training(request, current_subject = "test-user") - ) + asyncio.run(training_route.start_training(request, current_subject = "test-user")) assert exc_info.value.status_code == 422 assert "max_steps" in exc_info.value.detail diff --git a/studio/backend/utils/datasets/format_conversion.py b/studio/backend/utils/datasets/format_conversion.py index 708eded6811..0f57e6ae5e0 100644 --- a/studio/backend/utils/datasets/format_conversion.py +++ b/studio/backend/utils/datasets/format_conversion.py @@ -37,6 +37,7 @@ def standardize_chat_format( """ import collections import itertools + # Check if vision tokenizer is used is_vlm = False if tokenizer is not None: diff --git a/studio/backend/utils/datasets/iterable.py b/studio/backend/utils/datasets/iterable.py index d2c5eed8996..8408ef75f30 100644 --- a/studio/backend/utils/datasets/iterable.py +++ b/studio/backend/utils/datasets/iterable.py @@ -8,7 +8,6 @@ def is_streaming_dataset(dataset) -> bool: """Return True for iterable datasets that do not support eager map kwargs.""" try: from datasets import IterableDataset as HfIterableDataset - if isinstance(dataset, HfIterableDataset): return True except ImportError: @@ -16,7 +15,6 @@ def is_streaming_dataset(dataset) -> bool: try: from torch.utils.data import IterableDataset as TorchIterableDataset - return isinstance(dataset, TorchIterableDataset) except ImportError: return False From 583f8fa7bf932fd96cae25aca223486103389fc6 Mon Sep 17 00:00:00 2001 From: Etherll <61019402+Etherll@users.noreply.github.com> Date: Fri, 19 Jun 2026 23:16:59 +0300 Subject: [PATCH 11/14] studio: address streaming review (MLX/embedding guards, sliced eval split, rehydrate timing) - routes: reject dataset_streaming for embedding training and on Apple Silicon (MLX); both loaders materialize the full dataset instead of streaming - trainer: validate the base eval split name so streaming eval accepts HF slice syntax such as "validation[:1000]" - training-config-store: defer the onRehydrateStorage setState to a microtask so it doesn't hit the store's TDZ during synchronous hydration - test: streaming start rejects embedding models --- studio/backend/core/training/trainer.py | 6 +++- studio/backend/routes/training.py | 12 +++++++ .../backend/tests/test_training_streaming.py | 34 +++++++++++++++++++ .../training/stores/training-config-store.ts | 5 ++- 4 files changed, 55 insertions(+), 2 deletions(-) diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index 62c70402e6d..9a67222fd1f 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -2498,7 +2498,11 @@ def _apply_raw_text_prep(ds: Dataset, split_name: str) -> Dataset: f"Could not list splits for '{dataset_source}' " f"to validate eval_split='{eval_split}': {probe_err}" ) - if eval_split not in available_splits: + # HF split slicing (e.g. "validation[:1000]") is a + # valid split expression; validate the base split name, + # not the whole slice expression. + base_eval_split = eval_split.split("[", 1)[0] + if base_eval_split not in available_splits: raise ValueError( f"Requested eval split '{eval_split}' not found in " f"dataset '{dataset_source}'. Available splits: " diff --git a/studio/backend/routes/training.py b/studio/backend/routes/training.py index 8aaefa3ed6e..8e57c6b6d7c 100644 --- a/studio/backend/routes/training.py +++ b/studio/backend/routes/training.py @@ -198,6 +198,18 @@ async def start_training( status_code = 400, detail = "dataset_streaming is not supported for vision or audio datasets.", ) + if request.is_embedding: + raise HTTPException( + status_code = 400, + detail = "dataset_streaming is not supported for embedding training; the embedding loader needs the full dataset.", + ) + from utils.hardware import hardware as _hw + + if _hw.DEVICE == _hw.DeviceType.MLX: + raise HTTPException( + status_code = 400, + detail = "dataset_streaming is not yet supported on Apple Silicon (MLX); the MLX loader materializes the full dataset.", + ) if request.max_steps is None or request.max_steps <= 0: raise HTTPException( status_code = 422, diff --git a/studio/backend/tests/test_training_streaming.py b/studio/backend/tests/test_training_streaming.py index a0d90089464..94f8160f62a 100644 --- a/studio/backend/tests/test_training_streaming.py +++ b/studio/backend/tests/test_training_streaming.py @@ -287,6 +287,40 @@ def test_streaming_start_rejects_missing_max_steps(): assert "max_steps" in exc_info.value.detail +def test_streaming_start_rejects_embedding_models(): + # The embedding training path loads the full dataset (no streaming) and uses + # len/select, so the route must reject streaming for embedding runs even on a + # direct API call (the UI blocker doesn't cover that). + training_route = _load_route_module( + "training_route_module_for_streaming_embedding_test", + "routes/training.py", + ) + request = TrainingStartRequest( + model_name = "unsloth/test", + training_type = "LoRA/QLoRA", + hf_dataset = "org/dataset", + format_type = "chatml", + dataset_streaming = True, + is_embedding = True, + max_steps = 10, + ) + + backend = SimpleNamespace( + current_job_id = None, + is_training_active = lambda: False, + start_training = lambda **kwargs: pytest.fail("backend should not start"), + ) + + with patch.object(training_route, "get_training_backend", return_value = backend): + with pytest.raises(HTTPException) as exc_info: + asyncio.run( + training_route.start_training(request, current_subject = "test-user") + ) + + assert exc_info.value.status_code == 400 + assert "embedding" in exc_info.value.detail + + @pytest.mark.parametrize( "training_type, format_type", [ 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 62766510338..85ce25aa97f 100644 --- a/studio/frontend/src/features/training/stores/training-config-store.ts +++ b/studio/frontend/src/features/training/stores/training-config-store.ts @@ -1004,7 +1004,10 @@ export const useTrainingConfigStore = create()( if (!state) return; const patch = streamingCompatiblePatch(state); if (Object.keys(patch).length > 0) { - useTrainingConfigStore.setState(patch); + // Sync localStorage hydration runs inside create(), before + // useTrainingConfigStore is assigned (TDZ). Defer to a microtask so the + // store exists when we reconcile the persisted streaming combo. + queueMicrotask(() => useTrainingConfigStore.setState(patch)); } }, }, From 1c59866e81c83ef13f7e965ff2317cd88005c38f Mon Sep 17 00:00:00 2001 From: Etherll <61019402+Etherll@users.noreply.github.com> Date: Sat, 20 Jun 2026 09:25:11 +0300 Subject: [PATCH 12/14] studio: harden HF dataset streaming (column_names, split slicing, empty/eval bounds, gating) Address a deeper streaming review: - raw_text: resolve_column_names() guards IterableDataset.column_names=None (from_generator / unresolved features) so raw-text and CPT streaming no longer raise TypeError before training - models/routes: reject HF slice syntax in train_split/eval_split when streaming (load_dataset(streaming=True) raises "Bad split"); reject mixed sources (local/S3) and embedding/MLX streaming at the API, not just in the UI - trainer: an empty post-slice/filter stream fails preflight with a clear message; streaming eval is capped (STREAMING_EVAL_MAX_SAMPLES) so each eval terminates; the manual-slice shortcut falls back to a regular load when train_split is sliced - format_conversion: streaming conversions preflight the first mapped row so format errors surface before training, not mid-iteration - frontend: block streaming on Apple Silicon; clear datasetStreaming when a dataset is detected as image/audio at start --- studio/backend/core/training/trainer.py | 48 ++++-- studio/backend/models/training.py | 18 ++ studio/backend/routes/training.py | 19 +++ .../backend/tests/test_training_streaming.py | 155 ++++++++++++++++++ .../utils/datasets/format_conversion.py | 45 ++++- studio/backend/utils/datasets/raw_text.py | 33 +++- .../studio/sections/dataset-section.tsx | 7 + .../training/hooks/use-training-actions.ts | 4 + 8 files changed, 313 insertions(+), 16 deletions(-) diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index 9a67222fd1f..3b39f1559c3 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -123,7 +123,7 @@ def _ensure_real_packages(*names: str) -> None: from utils.datasets import format_and_template_dataset from utils.datasets import MODEL_TO_TEMPLATE_MAPPER, TEMPLATE_TO_RESPONSES_MAPPER from utils.datasets.iterable import is_streaming_dataset as detect_streaming_dataset -from utils.datasets.raw_text import prepare_raw_text_dataset +from utils.datasets.raw_text import prepare_raw_text_dataset, resolve_column_names from utils.paths import ( ensure_dir, resolve_dataset_path, @@ -139,6 +139,11 @@ def _ensure_real_packages(*names: str) -> None: logger = get_logger(__name__) +# A streaming eval dataset has no __len__, so a streaming evaluation would +# iterate the entire (potentially unbounded) source on every eval step. Cap it +# to a fixed sample count so each evaluation terminates predictably. +STREAMING_EVAL_MAX_SAMPLES = 500 + def _build_report_targets(training_args) -> list[str] | str: report_to: list[str] = [] @@ -2435,8 +2440,14 @@ def _apply_raw_text_prep(ds: Dataset, split_name: str) -> Dataset: # rows and materialize them (avoids downloading the whole dataset); # the eager [start, end] trim happens further below. _slice_start = dataset_slice_start or 0 + # streaming=True rejects HF slice syntax (e.g. "train[:50%]") + # with "Bad split", so the streaming shortcut is unusable when + # train_split already carries a slice expression, so fall back to + # the regular download path, which handles HF slice syntax. + _split_has_slice = (train_split or "").find("[") != -1 if ( - dataset_slice_end is not None + not _split_has_slice + and dataset_slice_end is not None and dataset_slice_end >= 0 and dataset_slice_end >= _slice_start ): @@ -2498,17 +2509,26 @@ def _apply_raw_text_prep(ds: Dataset, split_name: str) -> Dataset: f"Could not list splits for '{dataset_source}' " f"to validate eval_split='{eval_split}': {probe_err}" ) - # HF split slicing (e.g. "validation[:1000]") is a - # valid split expression; validate the base split name, - # not the whole slice expression. - base_eval_split = eval_split.split("[", 1)[0] - if base_eval_split not in available_splits: + # Streaming rejects HF slice syntax, and the request + # validator already blocks bracketed streaming splits, + # so eval_split here is always a bare split name. + if eval_split not in available_splits: raise ValueError( f"Requested eval split '{eval_split}' not found in " f"dataset '{dataset_source}'. Available splits: " f"{available_splits}" ) eval_dataset = load_dataset(**eval_load_kwargs, streaming = True) + # A streaming eval dataset has no __len__; bound it so + # each evaluation terminates instead of consuming the + # whole stream. .take() stays lazy and survives the + # later format/raw-text .map() passes. + if not hasattr(eval_dataset, "__len__"): + eval_dataset = eval_dataset.take(STREAMING_EVAL_MAX_SAMPLES) + logger.info( + f"Streaming eval split capped to " + f"{STREAMING_EVAL_MAX_SAMPLES} samples\n" + ) else: eval_dataset = load_dataset(**eval_load_kwargs) @@ -2641,9 +2661,13 @@ def _apply_raw_text_prep(ds: Dataset, split_name: str) -> Dataset: ) logger.info(f"Raw-text dataset ready ({n_display} samples)\n") - if "text" not in train_dataset.column_names: + # Streaming datasets can report column_names as None, which would + # make "text" not in None raise TypeError; resolve_column_names + # falls back to features/first-row probing. + train_columns = resolve_column_names(train_dataset) + if "text" not in train_columns: raise ValueError( - f"Raw-text dataset missing 'text' column: {train_dataset.column_names}" + f"Raw-text dataset missing 'text' column: {train_columns}" ) return (dataset_info, eval_dataset) @@ -2905,7 +2929,11 @@ def _preflight_first_batch(self) -> Optional[str]: loader = self.trainer.get_train_dataloader() batch = next(iter(loader)) except StopIteration: - return None + return ( + "Cannot start training: the dataset produced no training rows. " + "This usually means a split/slice or streaming filter removed every " + "row. Check your train split, slice range, and dataset filters." + ) except Exception as e: model = self.model_name or "this model" return ( diff --git a/studio/backend/models/training.py b/studio/backend/models/training.py index c821112f56b..670d5f911de 100644 --- a/studio/backend/models/training.py +++ b/studio/backend/models/training.py @@ -493,6 +493,24 @@ def _check_lora_dropout(cls, v: float) -> float: description = "S3 bucket configuration for loading datasets from AWS S3. Requires boto3 to be installed.", ) + @model_validator(mode = "after") + def _validate_streaming_splits(self) -> "TrainingStartRequest": + # Streaming load_dataset does not accept HF slice syntax (e.g. "train[:50%]" + # or "train[:20]"). Probe-confirmed: raises ValueError: Bad split. Reject + # early with a clear message so the user knows to use a plain split name. + if self.dataset_streaming: + for field_name, split_val in ( + ("train_split", self.train_split), + ("eval_split", self.eval_split), + ): + if split_val is not None and "[" in split_val: + raise ValueError( + f"dataset_streaming does not support HF slice syntax in {field_name} " + f"(got {split_val!r}); streaming load_dataset raises 'Bad split' on " + "bracket expressions. Use a plain split name (e.g. 'train', 'validation')." + ) + return self + @model_validator(mode = "after") def _check_steps_or_epochs(self) -> "TrainingStartRequest": # Each accepts 0 as "use the other"; both 0 means nothing to train. diff --git a/studio/backend/routes/training.py b/studio/backend/routes/training.py index 8e57c6b6d7c..6d12a54f082 100644 --- a/studio/backend/routes/training.py +++ b/studio/backend/routes/training.py @@ -227,6 +227,25 @@ async def start_training( status_code = 422, detail = "dataset_streaming with evaluation requires a separate eval_split.", ) + # Streaming is HF-only: reject when the request also carries a local + # dataset path or an S3 config; those sources cannot be streamed via + # HF's streaming loader. + if request.local_datasets: + raise HTTPException( + status_code = 400, + detail = ( + "dataset_streaming is HF-only; remove local_datasets / S3 source. " + "Streaming is not supported with local file paths." + ), + ) + if request.s3_config is not None: + raise HTTPException( + status_code = 400, + detail = ( + "dataset_streaming is HF-only; remove local_datasets / S3 source. " + "Streaming is not supported with S3 datasets." + ), + ) # Convert request to backend kwargs. training_kwargs = { diff --git a/studio/backend/tests/test_training_streaming.py b/studio/backend/tests/test_training_streaming.py index 94f8160f62a..f73b592f464 100644 --- a/studio/backend/tests/test_training_streaming.py +++ b/studio/backend/tests/test_training_streaming.py @@ -406,3 +406,158 @@ def _start_training(**kwargs): assert captured["dataset_streaming"] is True assert captured["max_steps"] == 10 assert captured["eval_split"] == "validation" + + +# streaming rejects HF slice syntax in train_split / eval_split + + +@pytest.mark.parametrize( + "field, value", + [ + ("train_split", "train[:50%]"), + ("train_split", "train[:20]"), + ("eval_split", "validation[:1000]"), + ], +) +def test_streaming_rejects_bracketed_split_syntax(field, value): + # The model_validator _validate_streaming_splits raises ValidationError when + # dataset_streaming=True and a split contains "[" (HF slice syntax). + kwargs = { + "model_name": "unsloth/test", + "training_type": "LoRA/QLoRA", + "hf_dataset": "org/dataset", + "format_type": "chatml", + "dataset_streaming": True, + "max_steps": 10, + field: value, + } + with pytest.raises(ValidationError) as exc_info: + TrainingStartRequest(**kwargs) + detail = str(exc_info.value) + assert "slice" in detail.lower() or "bracket" in detail.lower() or "[" in detail + + +# streaming rejects mixed sources (local_datasets) + + +def test_streaming_start_rejects_local_datasets(): + # dataset_streaming + local_datasets -> 400, 'local' in detail + training_route = _load_route_module( + "training_route_module_for_streaming_local_datasets_test", + "routes/training.py", + ) + request = TrainingStartRequest( + model_name = "unsloth/test", + training_type = "LoRA/QLoRA", + hf_dataset = "org/dataset", + format_type = "chatml", + dataset_streaming = True, + max_steps = 10, + ) + # Bypass Pydantic's local-path validation by injecting directly after construction. + object.__setattr__(request, "local_datasets", ["/some/local/file.jsonl"]) + + backend = SimpleNamespace( + current_job_id = None, + is_training_active = lambda: False, + start_training = lambda **kwargs: pytest.fail("backend should not start"), + ) + + with patch.object(training_route, "get_training_backend", return_value = backend): + with pytest.raises(HTTPException) as exc_info: + asyncio.run(training_route.start_training(request, current_subject = "test-user")) + + assert exc_info.value.status_code == 400 + assert "local" in exc_info.value.detail.lower() or "hf-only" in exc_info.value.detail.lower() + + +# _drop_invalid_text_rows handles from_generator with column_names=None + + +def test_drop_invalid_text_rows_from_generator_none_column_names(): + # from_generator IterableDatasets have column_names=None; resolve_column_names + # must fall back to first-row probe. _drop_invalid_text_rows must not raise + # TypeError and must filter correctly. + from utils.datasets.raw_text import _drop_invalid_text_rows + + def _gen(): + yield {"text": "valid row"} + yield {"text": None} # invalid, should be dropped + yield {"text": "another row"} + + stream = datasets.IterableDataset.from_generator(_gen) + # Precondition: column_names is None on a raw from_generator dataset. + assert stream.column_names is None, ( + "precondition failed: expected column_names=None for from_generator dataset" + ) + + filtered, notices = _drop_invalid_text_rows( + stream, mode_title = "Raw text", split_scope = "test split" + ) + + rows = list(filtered) + assert [r["text"] for r in rows] == ["valid row", "another row"] + # At least one info/warning notice about dropped rows. + assert len(notices) >= 1 + + +# _preflight_first_batch returns error string on empty dataloader + + +def test_preflight_first_batch_returns_error_on_empty_stream(): + # StopIteration from an empty dataloader must return a clear + # error string (not None). Test via a minimal stub, no real model needed. + import types + import sys + + # Minimal stub trainer whose get_train_dataloader() yields nothing. + class _EmptyLoader: + def __iter__(self): + return iter([]) + + class _StubTrainer: + def get_train_dataloader(self): + return _EmptyLoader() + + # Load UnslothTrainer class from trainer.py via importlib to avoid heavy imports. + trainer_path = ( + _BACKEND_ROOT / "core" / "training" / "trainer.py" + ) + spec = importlib.util.spec_from_file_location("trainer_module", trainer_path) + trainer_mod = importlib.util.module_from_spec(spec) + # Provide a minimal sys.modules shim so top-level imports in trainer.py don't + # crash when optional heavy deps (torch, unsloth) are absent. + _orig_import = __builtins__.__import__ if hasattr(__builtins__, "__import__") else __import__ + + try: + spec.loader.exec_module(trainer_mod) + except Exception: + # trainer.py has optional heavy imports; access _preflight_first_batch directly. + pass + + # If we successfully loaded the module, find the trainer class. + trainer_cls = None + for name, obj in vars(trainer_mod).items() if "trainer_mod" in dir() else []: + if hasattr(obj, "_preflight_first_batch"): + trainer_cls = obj + break + + if trainer_cls is None: + pytest.skip( + "Could not load trainer module (missing optional deps: torch/unsloth)." + ) + + # Build a bare instance without calling __init__ (avoids needing real deps). + instance = object.__new__(trainer_cls) + instance.trainer = _StubTrainer() + instance.model_name = "stub-model" + + result = instance._preflight_first_batch() + + assert result is not None, ( + "_preflight_first_batch must return an error string (not None) when the " + "training dataloader is empty." + ) + assert isinstance(result, str) + # The message should indicate there are no training rows / empty dataset. + assert any(kw in result.lower() for kw in ("empty", "no training", "no rows", "stream")) diff --git a/studio/backend/utils/datasets/format_conversion.py b/studio/backend/utils/datasets/format_conversion.py index 0f57e6ae5e0..95c9a005345 100644 --- a/studio/backend/utils/datasets/format_conversion.py +++ b/studio/backend/utils/datasets/format_conversion.py @@ -161,7 +161,20 @@ def _standardize_dataset(examples): dataset_map_kwargs["num_proc"] = num_proc dataset_map_kwargs["desc"] = "Standardizing chat format" - return dataset.map(_standardize_dataset, **dataset_map_kwargs) + result = dataset.map(_standardize_dataset, **dataset_map_kwargs) + + # For streaming, force the first mapped row through now so any + # column/format errors surface before training begins (not mid-iteration). + # IterableDataset re-iterates from the generator source, so this is safe. + if is_streaming_dataset(dataset): + try: + next(iter(result)) + except Exception as exc: + raise ValueError( + f"Streaming chat-format standardization failed on the first row: {exc}" + ) from exc + + return result def convert_chatml_to_alpaca( @@ -232,7 +245,20 @@ def _convert(examples): dataset_map_kwargs["num_proc"] = num_proc dataset_map_kwargs["desc"] = "Converting ChatML to Alpaca format" - return dataset.map(_convert, **dataset_map_kwargs) + result = dataset.map(_convert, **dataset_map_kwargs) + + # For streaming, force the first mapped row through now so any + # column/format errors surface before training begins (not mid-iteration). + # IterableDataset re-iterates from the generator source, so this is safe. + if is_iterable: + try: + next(iter(result)) + except Exception as exc: + raise ValueError( + f"Streaming ChatML-to-Alpaca conversion failed on the first row: {exc}" + ) from exc + + return result def convert_alpaca_to_chatml( @@ -285,7 +311,20 @@ def _convert(examples): dataset_map_kwargs["num_proc"] = num_proc dataset_map_kwargs["desc"] = "Converting Alpaca to ChatML format" - return dataset.map(_convert, **dataset_map_kwargs) + result = dataset.map(_convert, **dataset_map_kwargs) + + # For streaming, force the first mapped row through now so any + # column/format errors surface before training begins (not mid-iteration). + # IterableDataset re-iterates from the generator source, so this is safe. + if is_iterable: + try: + next(iter(result)) + except Exception as exc: + raise ValueError( + f"Streaming Alpaca-to-ChatML conversion failed on the first row: {exc}" + ) from exc + + return result def _format_eta(seconds): diff --git a/studio/backend/utils/datasets/raw_text.py b/studio/backend/utils/datasets/raw_text.py index 1a1c09c93d3..112528fdd0b 100644 --- a/studio/backend/utils/datasets/raw_text.py +++ b/studio/backend/utils/datasets/raw_text.py @@ -22,10 +22,36 @@ class RawTextPreparationResult: notices: list[RawTextNotice] +def resolve_column_names(dataset) -> list[str]: + """Return the column names for *dataset*, guarding against None. + + IterableDataset.column_names is None until HF datasets>=X materialises + it from the first batch; .map() also keeps it None. Resolution order: + 1. dataset.column_names if truthy (regular Dataset or HF>=4.4) + 2. keys of dataset.features if available + 3. bounded first-row probe, consumes one element, safe on IterableDataset + because HF re-iterates from the generator on the next pass + 4. [] as a last resort so callers never see None + """ + col_names = getattr(dataset, "column_names", None) + if col_names: + return list(col_names) + + features = getattr(dataset, "features", None) + if features: + return list(features.keys()) + + try: + first_row = next(iter(dataset)) + return list(first_row.keys()) + except Exception: + return [] + + def _string_columns(dataset: Dataset) -> list[str]: feature_map = getattr(dataset, "features", {}) or {} string_cols: list[str] = [] - for col in dataset.column_names: + for col in resolve_column_names(dataset): feature = feature_map.get(col) dtype = str(getattr(feature, "dtype", "")) if dtype in {"string", "large_string"}: @@ -92,12 +118,13 @@ def prepare_raw_text_dataset( mode_title = mode_label.capitalize() split_scope = _split_scope(split_name) - if "text" not in dataset.column_names: + col_names = resolve_column_names(dataset) + if "text" not in col_names: string_cols = _string_columns(dataset) if not string_cols: raise ValueError( f"{mode_title} training requires a string 'text' column but none " - f"was found in {split_scope} (columns: {dataset.column_names})." + f"was found in {split_scope} (columns: {col_names})." ) renamed_col = string_cols[0] diff --git a/studio/frontend/src/features/studio/sections/dataset-section.tsx b/studio/frontend/src/features/studio/sections/dataset-section.tsx index aa48d00a5a7..75ca84dfd71 100644 --- a/studio/frontend/src/features/studio/sections/dataset-section.tsx +++ b/studio/frontend/src/features/studio/sections/dataset-section.tsx @@ -84,6 +84,7 @@ import { useShallow } from "zustand/react/shallow"; import { DocumentUploadRedirectDialog } from "./document-upload-redirect-dialog"; import { translate, useT } from "@/i18n"; import { S3ConfigForm } from "./s3-config-form"; +import { usePlatformStore } from "@/config/env"; const TRAINING_UPLOAD_EXTENSIONS = [ ".csv", @@ -225,6 +226,8 @@ export function DatasetSection() { })), ); + const platformDeviceType = usePlatformStore((s) => s.deviceType); + // Streaming is only supported for Hugging Face text datasets. Rather than // hiding the toggle when a constraint isn't met, keep it visible but disabled // and list the exact unmet requirement(s) in its tooltip — a control that @@ -258,6 +261,10 @@ export function DatasetSection() { streamingBlockers.push("This dataset looks like images, which can't stream."); if (isDatasetAudio) streamingBlockers.push("This dataset looks like audio, which can't stream."); + if (platformDeviceType === "mac") + streamingBlockers.push( + "Streaming isn't supported on Apple Silicon (MLX) yet.", + ); const isStreamingSupported = streamingBlockers.length === 0; diff --git a/studio/frontend/src/features/training/hooks/use-training-actions.ts b/studio/frontend/src/features/training/hooks/use-training-actions.ts index 3dec8d3a32d..8c2b3e9e771 100644 --- a/studio/frontend/src/features/training/hooks/use-training-actions.ts +++ b/studio/frontend/src/features/training/hooks/use-training-actions.ts @@ -88,6 +88,10 @@ export function useTrainingActions() { useTrainingConfigStore.setState({ isDatasetImage: isImage, isDatasetAudio: isAudio, + // Streaming is unsupported for image/audio datasets; clear the flag + // so buildTrainingStartPayload never ships dataset_streaming=true + // for a modality the backend would reject with a 422. + ...(isImage || isAudio ? { datasetStreaming: false } : {}), }); } From 14b7d30bc307a4247d068441c93d01be8fa9c0b9 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Sat, 20 Jun 2026 06:26:19 +0000 Subject: [PATCH 13/14] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- studio/backend/core/training/trainer.py | 4 +--- .../backend/tests/test_training_streaming.py | 20 +++++++------------ 2 files changed, 8 insertions(+), 16 deletions(-) diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index 3b39f1559c3..bb25a836932 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -2666,9 +2666,7 @@ def _apply_raw_text_prep(ds: Dataset, split_name: str) -> Dataset: # falls back to features/first-row probing. train_columns = resolve_column_names(train_dataset) if "text" not in train_columns: - raise ValueError( - f"Raw-text dataset missing 'text' column: {train_columns}" - ) + raise ValueError(f"Raw-text dataset missing 'text' column: {train_columns}") return (dataset_info, eval_dataset) elif self.is_audio_vlm: diff --git a/studio/backend/tests/test_training_streaming.py b/studio/backend/tests/test_training_streaming.py index f73b592f464..1fc59d792dd 100644 --- a/studio/backend/tests/test_training_streaming.py +++ b/studio/backend/tests/test_training_streaming.py @@ -313,9 +313,7 @@ def test_streaming_start_rejects_embedding_models(): with patch.object(training_route, "get_training_backend", return_value = backend): with pytest.raises(HTTPException) as exc_info: - asyncio.run( - training_route.start_training(request, current_subject = "test-user") - ) + asyncio.run(training_route.start_training(request, current_subject = "test-user")) assert exc_info.value.status_code == 400 assert "embedding" in exc_info.value.detail @@ -482,14 +480,14 @@ def test_drop_invalid_text_rows_from_generator_none_column_names(): def _gen(): yield {"text": "valid row"} - yield {"text": None} # invalid, should be dropped + yield {"text": None} # invalid, should be dropped yield {"text": "another row"} stream = datasets.IterableDataset.from_generator(_gen) # Precondition: column_names is None on a raw from_generator dataset. - assert stream.column_names is None, ( - "precondition failed: expected column_names=None for from_generator dataset" - ) + assert ( + stream.column_names is None + ), "precondition failed: expected column_names=None for from_generator dataset" filtered, notices = _drop_invalid_text_rows( stream, mode_title = "Raw text", split_scope = "test split" @@ -520,9 +518,7 @@ def get_train_dataloader(self): return _EmptyLoader() # Load UnslothTrainer class from trainer.py via importlib to avoid heavy imports. - trainer_path = ( - _BACKEND_ROOT / "core" / "training" / "trainer.py" - ) + trainer_path = _BACKEND_ROOT / "core" / "training" / "trainer.py" spec = importlib.util.spec_from_file_location("trainer_module", trainer_path) trainer_mod = importlib.util.module_from_spec(spec) # Provide a minimal sys.modules shim so top-level imports in trainer.py don't @@ -543,9 +539,7 @@ def get_train_dataloader(self): break if trainer_cls is None: - pytest.skip( - "Could not load trainer module (missing optional deps: torch/unsloth)." - ) + pytest.skip("Could not load trainer module (missing optional deps: torch/unsloth).") # Build a bare instance without calling __init__ (avoids needing real deps). instance = object.__new__(trainer_cls) From 1d80936f35cf7fc3130744f63551aeb81bad7cfd Mon Sep 17 00:00:00 2001 From: Etherll <61019402+Etherll@users.noreply.github.com> Date: Mon, 22 Jun 2026 15:48:05 +0300 Subject: [PATCH 14/14] studio: fix CI for streaming PR (lint blocker + no-torch sandbox + preflight test) - trainer.py: drop unused `IterableDataset` import (hoist safety-net blocker). - test_training_streaming.py: only select real classes (isinstance type) when locating the trainer class, so a MagicMock-stubbed global is never passed to object.__new__ (fixes TypeError on the Python 3.10-3.13 jobs). - no-torch import sandboxes (test_e2e_no_torch_sandbox.py, test_studio_import_no_torch.py): teach the chat_templates/format_conversion exec stubs and the full-import-chain copy list about the new `.iterable` module so the AFTER/runtime cases import without torch again. --- studio/backend/core/training/trainer.py | 2 +- .../backend/tests/test_training_streaming.py | 5 ++++- tests/python/test_e2e_no_torch_sandbox.py | 14 +++++++++++++ tests/python/test_studio_import_no_torch.py | 20 +++++++++++++++++++ 4 files changed, 39 insertions(+), 2 deletions(-) diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index bb25a836932..6041817851c 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -114,7 +114,7 @@ def _ensure_real_packages(*names: str) -> None: from typing import Any, Dict, List, Optional, Callable from dataclasses import dataclass import pandas as pd -from datasets import Dataset, IterableDataset +from datasets import Dataset from utils.datasets.cache_safe import load_dataset_cache_safe as load_dataset from core.inference.llama_cpp import _hf_offline_if_dns_dead diff --git a/studio/backend/tests/test_training_streaming.py b/studio/backend/tests/test_training_streaming.py index 1fc59d792dd..8ff016d3bfb 100644 --- a/studio/backend/tests/test_training_streaming.py +++ b/studio/backend/tests/test_training_streaming.py @@ -534,7 +534,10 @@ def get_train_dataloader(self): # If we successfully loaded the module, find the trainer class. trainer_cls = None for name, obj in vars(trainer_mod).items() if "trainer_mod" in dir() else []: - if hasattr(obj, "_preflight_first_batch"): + # Only real classes — when heavy deps are stubbed with MagicMock, + # hasattr() is always True on a mock, so guard on isinstance(obj, type) + # to avoid picking a mock instance (object.__new__ would then reject it). + if isinstance(obj, type) and hasattr(obj, "_preflight_first_batch"): trainer_cls = obj break diff --git a/tests/python/test_e2e_no_torch_sandbox.py b/tests/python/test_e2e_no_torch_sandbox.py index 7da7cf0eaeb..23d4a2c2b3f 100644 --- a/tests/python/test_e2e_no_torch_sandbox.py +++ b/tests/python/test_e2e_no_torch_sandbox.py @@ -27,6 +27,7 @@ FORMAT_DETECTION = DATASETS_DIR / "format_detection.py" MODEL_MAPPINGS = DATASETS_DIR / "model_mappings.py" VLM_PROCESSING = DATASETS_DIR / "vlm_processing.py" +ITERABLE = DATASETS_DIR / "iterable.py" HARDWARE_PY = HARDWARE_DIR / "hardware.py" # Studio venv for server tests @@ -280,9 +281,13 @@ def test_after_chat_templates_imports(self, no_torch_venv): mm = types.ModuleType('model_mappings') mm.MODEL_TO_TEMPLATE_MAPPER = {{}} sys.modules['model_mappings'] = mm + it = types.ModuleType('iterable') + it.is_streaming_dataset = lambda *a, **k: False + sys.modules['iterable'] = it source = open({str(CHAT_TEMPLATES)!r}).read() source = source.replace('from .format_detection import', 'from format_detection import') source = source.replace('from .model_mappings import', 'from model_mappings import') + source = source.replace('from .iterable import', 'from iterable import') exec(source) print("OK") """) @@ -323,6 +328,7 @@ def test_after_full_import_chain_imports(self, no_torch_venv, sandbox_dir): VLM_PROCESSING, DATA_COLLATORS, CHAT_TEMPLATES, + ITERABLE, ]: if src.exists(): shutil.copy2(src, pkg_dir / src.name) @@ -431,10 +437,14 @@ def test_alpaca_template_accessible(self, no_torch_venv): mm = types.ModuleType('model_mappings') mm.MODEL_TO_TEMPLATE_MAPPER = {{}} sys.modules['model_mappings'] = mm + it = types.ModuleType('iterable') + it.is_streaming_dataset = lambda *a, **k: False + sys.modules['iterable'] = it ns = {{}} source = open({str(CHAT_TEMPLATES)!r}).read() source = source.replace('from .format_detection import', 'from format_detection import') source = source.replace('from .model_mappings import', 'from model_mappings import') + source = source.replace('from .iterable import', 'from iterable import') exec(source, ns) assert 'Instruction' in ns['DEFAULT_ALPACA_TEMPLATE'] print("OK") @@ -544,11 +554,15 @@ def test_lazy_torch_fails_at_call_time_not_import_time(self, no_torch_venv, sand mm = types.ModuleType('model_mappings') mm.MODEL_TO_TEMPLATE_MAPPER = {{}} sys.modules['model_mappings'] = mm + it = types.ModuleType('iterable') + it.is_streaming_dataset = lambda *a, **k: False + sys.modules['iterable'] = it ns = {{}} source = open({str(CHAT_TEMPLATES)!r}).read() source = source.replace('from .format_detection import', 'from format_detection import') source = source.replace('from .model_mappings import', 'from model_mappings import') + source = source.replace('from .iterable import', 'from iterable import') exec(source, ns) # Import succeeds -- this is the fix diff --git a/tests/python/test_studio_import_no_torch.py b/tests/python/test_studio_import_no_torch.py index 86dc8581aba..c4efbc8cea8 100644 --- a/tests/python/test_studio_import_no_torch.py +++ b/tests/python/test_studio_import_no_torch.py @@ -254,10 +254,15 @@ def test_exec_with_stubs(self, no_torch_venv): model_mappings.MODEL_TO_TEMPLATE_MAPPER = {{}} sys.modules['model_mappings'] = model_mappings + iterable = types.ModuleType('iterable') + iterable.is_streaming_dataset = lambda *a, **k: False + sys.modules['iterable'] = iterable + # Read and transform the source: replace relative imports with absolute source = open({str(CHAT_TEMPLATES)!r}).read() source = source.replace('from .format_detection import', 'from format_detection import') source = source.replace('from .model_mappings import', 'from model_mappings import') + source = source.replace('from .iterable import', 'from iterable import') exec(source) @@ -295,10 +300,15 @@ def test_default_alpaca_template_defined(self, no_torch_venv): model_mappings.MODEL_TO_TEMPLATE_MAPPER = {{}} sys.modules['model_mappings'] = model_mappings + iterable = types.ModuleType('iterable') + iterable.is_streaming_dataset = lambda *a, **k: False + sys.modules['iterable'] = iterable + ns = {{}} source = open({str(CHAT_TEMPLATES)!r}).read() source = source.replace('from .format_detection import', 'from format_detection import') source = source.replace('from .model_mappings import', 'from model_mappings import') + source = source.replace('from .iterable import', 'from iterable import') exec(source, ns) assert 'DEFAULT_ALPACA_TEMPLATE' in ns, "DEFAULT_ALPACA_TEMPLATE not defined" @@ -379,6 +389,10 @@ def test_convert_chatml_to_alpaca_no_torch(self, no_torch_venv): datasets_mod.IterableDataset = type('IterableDataset', (), {{}}) sys.modules['datasets'] = datasets_mod + iterable_mod = types.ModuleType('iterable') + iterable_mod.is_streaming_dataset = lambda *a, **k: False + sys.modules['iterable'] = iterable_mod + # Stub utils.hardware utils_mod = types.ModuleType('utils') hardware_mod = types.ModuleType('utils.hardware') @@ -390,6 +404,7 @@ def test_convert_chatml_to_alpaca_no_torch(self, no_torch_venv): # Read and exec format_conversion.py source = open({str(FORMAT_CONVERSION)!r}).read() source = source.replace('from .format_detection import', 'from format_detection import') + source = source.replace('from .iterable import', 'from iterable import') ns = {{'__name__': '__test__'}} exec(source, ns) @@ -437,6 +452,10 @@ def test_convert_alpaca_to_chatml_no_torch(self, no_torch_venv): datasets_mod.IterableDataset = type('IterableDataset', (), {{}}) sys.modules['datasets'] = datasets_mod + iterable_mod = types.ModuleType('iterable') + iterable_mod.is_streaming_dataset = lambda *a, **k: False + sys.modules['iterable'] = iterable_mod + utils_mod = types.ModuleType('utils') hardware_mod = types.ModuleType('utils.hardware') hardware_mod.dataset_map_num_proc = lambda n=None: 1 @@ -446,6 +465,7 @@ def test_convert_alpaca_to_chatml_no_torch(self, no_torch_venv): source = open({str(FORMAT_CONVERSION)!r}).read() source = source.replace('from .format_detection import', 'from format_detection import') + source = source.replace('from .iterable import', 'from iterable import') ns = {{'__name__': '__test__'}} exec(source, ns)