From 1fc8bf53c763598c0354fd6e1b81365c32c5b2f9 Mon Sep 17 00:00:00 2001 From: Sanat Bhargava Date: Mon, 22 Jun 2026 20:18:18 +0530 Subject: [PATCH] Add Hugging Face dataset streaming mode to Studio (#4946) * Add HF dataset streaming mode to Studio * Added default value for datasetStreaming in training-config-store.ts * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Handle None max_steps for streaming validation * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * 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: 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: 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: 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 * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * 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: 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 * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * 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. --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> Co-authored-by: Roland Tannous <115670425+rolandtannous@users.noreply.github.com> Co-authored-by: Roland Tannous Co-authored-by: Etherll <61019402+Etherll@users.noreply.github.com> --- studio/backend/core/training/trainer.py | 412 ++++++++----- studio/backend/core/training/training.py | 1 + studio/backend/core/training/worker.py | 1 + studio/backend/models/training.py | 100 +++- studio/backend/routes/training.py | 63 ++ .../backend/tests/test_training_streaming.py | 560 ++++++++++++++++++ .../backend/utils/datasets/chat_templates.py | 27 +- .../backend/utils/datasets/dataset_utils.py | 5 +- .../utils/datasets/format_conversion.py | 68 ++- studio/backend/utils/datasets/iterable.py | 20 + studio/backend/utils/datasets/raw_text.py | 50 +- .../studio/sections/dataset-section.tsx | 130 +++- .../studio/sections/params-section.tsx | 13 +- .../studio/training-start-overlay.tsx | 10 +- .../src/features/training/api/mappers.ts | 1 + .../training/hooks/use-training-actions.ts | 4 + .../training/stores/training-config-store.ts | 167 +++++- .../src/features/training/types/api.ts | 1 + .../src/features/training/types/config.ts | 2 + studio/frontend/src/i18n/locales/en.ts | 1 + tests/python/test_e2e_no_torch_sandbox.py | 14 + tests/python/test_studio_import_no_torch.py | 20 + 22 files changed, 1481 insertions(+), 189 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 bdf2a6f030..a13a75a72e 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -61,7 +61,7 @@ import structlog 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 @@ -72,7 +72,8 @@ from utils.models import is_vision_model, detect_audio_type from utils.models.model_config import _env_offline from utils.datasets import format_and_template_dataset from utils.datasets import MODEL_TO_TEMPLATE_MAPPER, TEMPLATE_TO_RESPONSES_MAPPER -from utils.datasets.raw_text import prepare_raw_text_dataset +from utils.datasets.iterable import is_streaming_dataset as detect_streaming_dataset +from utils.datasets.raw_text import prepare_raw_text_dataset, resolve_column_names from utils.paths import ( ensure_dir, resolve_dataset_path, @@ -88,6 +89,11 @@ from utils.subprocess_compat import ( 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] = [] @@ -2224,17 +2230,18 @@ class UnslothTrainer: def load_and_format_dataset( self, - dataset_source: str, + dataset_source: Optional[str], format_type: str = "auto", - local_datasets: list = None, - local_eval_datasets: list = None, - custom_format_mapping: dict = None, - subset: str = None, + 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: str = None, + eval_split: Optional[str] = None, + dataset_streaming: bool = False, eval_steps: float = 0.00, - dataset_slice_start: int = None, - dataset_slice_end: int = None, + dataset_slice_start: Optional[int] = None, + dataset_slice_end: Optional[int] = None, is_cpt: bool = False, s3_config: dict = None, ) -> Optional[tuple]: @@ -2342,47 +2349,89 @@ class UnslothTrainer: 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 needed rows, not the whole dataset. - rows_to_stream = dataset_slice_end + 1 + if dataset_streaming: + self._update_progress(status_message = f"Streaming dataset: {dataset_source}...") + 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) + # 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." + ) + 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"[dataset-slice] Manual slice specified " - f"(start={dataset_slice_start}, end={dataset_slice_end}), " - f"streaming {rows_to_stream} rows\n" - ) - stream = load_dataset(**load_kwargs, streaming = True) - dataset = Dataset.from_list(list(stream.take(rows_to_stream))) - logger.info( - f"[dataset-slice] Downloaded {len(dataset)} rows " - f"(requested {rows_to_stream})\n" - ) - self._update_progress( - status_message = f"Streamed {len(dataset)} rows from HuggingFace" + f"Loaded Hugging Face dataset in streaming mode: {dataset_source}\n" ) + self._update_progress(status_message = f"Streaming {dataset_source}") else: - self._update_progress( - status_message = f"Downloading dataset: {dataset_source}..." - ) - dataset = load_dataset(**load_kwargs) + # Non-streaming: if a slice end is given, stream only the needed + # 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 ( + not _split_has_slice + and dataset_slice_end is not None + and dataset_slice_end >= 0 + and dataset_slice_end >= _slice_start + ): + 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" + ) + stream = load_dataset(**load_kwargs, streaming = True) + dataset = Dataset.from_list(list(stream.take(rows_to_stream))) + logger.info( + f"[dataset-slice] Downloaded {len(dataset)} rows " + f"(requested {rows_to_stream})\n" + ) + else: + self._update_progress( + status_message = f"Downloading dataset: {dataset_source}..." + ) + 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) if eval_enabled: effective_train = train_split or "train" @@ -2392,17 +2441,69 @@ class UnslothTrainer: 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: + # 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}" + ) + # 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) + 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: + 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 set eval_steps to 0." + ) # Same split as training — split 80/20 after formatting logger.info( f"Eval split '{eval_split}' is the same as train split — will split 80/20\n" ) else: + if dataset_streaming: + raise ValueError( + "Streaming mode currently requires an explicit eval split when evaluation is enabled." + ) # Auto-detect eval split from HF (separate dataset or None) eval_dataset = self._auto_detect_eval_split_from_hf( dataset_source = dataset_source, @@ -2416,8 +2517,12 @@ class UnslothTrainer: if dataset is None: raise ValueError("No dataset provided") - # Apply index range slicing if requested (inclusive 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). + # Streaming already sliced lazily via skip()/take() above; the non-streaming + # manual-slice path fetched up to end+1 rows and is trimmed to [start, end] here. + 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 = dataset_slice_end if dataset_slice_end is not None else total_rows - 1 @@ -2480,11 +2585,19 @@ class UnslothTrainer: } 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 @@ -2498,10 +2611,12 @@ class UnslothTrainer: ) logger.info(f"Raw-text dataset ready ({n_display} samples)\n") - if "text" not in train_dataset.column_names: - raise ValueError( - f"Raw-text dataset missing 'text' column: {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_columns}") return (dataset_info, eval_dataset) elif self.is_audio_vlm: @@ -2540,13 +2655,16 @@ class UnslothTrainer: 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") # ========== 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, @@ -2557,8 +2675,8 @@ class UnslothTrainer: 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) @@ -2759,7 +2877,11 @@ class UnslothTrainer: 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 ( @@ -2799,8 +2921,15 @@ class UnslothTrainer: f"columns are mapped correctly for '{model}'." ) - def _train_worker(self, dataset: Dataset, **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 platforms, register compiled-cache dirs on sys.path/PYTHONPATH # before any dataset.map() so spawned workers can import compiled @@ -3129,7 +3258,7 @@ class UnslothTrainer: else: # Default 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) @@ -3137,7 +3266,7 @@ class UnslothTrainer: 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"] @@ -3159,7 +3288,10 @@ class UnslothTrainer: 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" @@ -3213,6 +3345,12 @@ class UnslothTrainer: # Packing for text models only (DeepSeek OCR 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" @@ -3231,10 +3369,10 @@ class UnslothTrainer: logger.info("Training configuration prepared\n") # ========== TRAINER INITIALIZATION ========== if self.is_audio_vlm and not raw_text_mode: - # Audio VLM (e.g. Gemma 3N): raw Dataset from - # _format_audio_vlm_dataset, processing_class=processor.tokenizer. - # Raw-text runs go to the text path below. - train_dataset = dataset if isinstance(dataset, Dataset) else dataset["dataset"] + # 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 processing_class = ( self.tokenizer.tokenizer if hasattr(self.tokenizer, "tokenizer") @@ -3275,9 +3413,7 @@ class UnslothTrainer: if isinstance(self.tokenizer, ProcessorMixin) and hasattr( self.tokenizer, "tokenizer" ): - logger.info( - f" ⚠️ 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: @@ -3405,58 +3541,44 @@ class UnslothTrainer: ) logger.info("Train on responses only configured successfully\n") - # Safety net: train_on_responses_only masks non-response - # tokens with -100. If max_seq_length is too short, the - # response is truncated away, every sample becomes all -100, - # and Unsloth drops them, leaving 0 usable 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'." + # ── Safety net: check if all samples were filtered out ── + # train_on_responses_only masks non-response tokens with -100; + # if max_seq_length is too short the response is truncated away, + # 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") + else: + filtered_len = len(self.trainer.train_dataset) + original_dataset_obj = ( + dataset["dataset"] if isinstance(dataset, dict) else dataset ) - logger.error(error_msg) - self._update_progress(error = error_msg, is_training = False) - return + 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 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") + 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 - # [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, - # ) + 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}") @@ -3470,27 +3592,41 @@ class UnslothTrainer: # ========== PROGRESS TRACKING ========== self.trainer.add_callback(self._create_progress_callback()) - num_samples = None - if hasattr(self.trainer, "train_dataset") and self.trainer.train_dataset is not None: - try: - num_samples = len(self.trainer.train_dataset) - except TypeError: - logger.debug( - "train_dataset does not support len(); falling back to " - "raw dataset size for step estimation." - ) + train_dataset_obj = dataset["dataset"] if isinstance(dataset, dict) else dataset + is_streaming_dataset = detect_streaming_dataset(train_dataset_obj) - if num_samples is None: - num_samples = len(dataset["dataset"] if isinstance(dataset, dict) else dataset) + max_steps_value = training_args.get("max_steps") + max_steps = 0 if max_steps_value is None else int(max_steps_value) - 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), - ) + 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 = max_steps + else: + # Prefer the trainer's processed dataset length (post + # train-on-responses filtering); fall back to the raw dataset + # if it has no len(). + num_samples = None + if getattr(self.trainer, "train_dataset", None) is not None: + try: + num_samples = len(self.trainer.train_dataset) + except TypeError: + num_samples = None + if num_samples is None: + 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), + max_steps, + ) + + self._update_progress(total_steps = total_steps) # ========== START TRAINING ========== # Fail fast on an invalid first batch (empty/float input_ids) vs a step-1 crash. preflight_error = self._preflight_first_batch() diff --git a/studio/backend/core/training/training.py b/studio/backend/core/training/training.py index 052826e579..34bb4331a7 100644 --- a/studio/backend/core/training/training.py +++ b/studio/backend/core/training/training.py @@ -281,6 +281,7 @@ class TrainingBackend: "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 cd243c6a21..70c3bc9a28 100644 --- a/studio/backend/core/training/worker.py +++ b/studio/backend/core/training/worker.py @@ -2801,6 +2801,7 @@ def run_training_process(*, event_queue: Any, stop_queue: Any, config: dict) -> 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 ae3061d943..670d5f911d 100644 --- a/studio/backend/models/training.py +++ b/studio/backend/models/training.py @@ -27,6 +27,10 @@ _MAX_LORA_ALPHA = 32_768 _MIN_VISION_IMAGE_SIZE = 256 # 2048 is the highest most llms stay stable at _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 class S3Config(BaseModel): @@ -125,12 +129,22 @@ class TrainingStartRequest(BaseModel): 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") + 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)") dataset_slice_start: Optional[int] = Field( - None, description = "Inclusive start row index for dataset slicing" + None, + ge = 0, + le = _MAX_DATASET_SLICE_INDEX, + 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, + le = _MAX_DATASET_SLICE_INDEX, + description = "Inclusive end row index for dataset slicing", ) @model_validator(mode = "before") @@ -141,6 +155,70 @@ class TrainingStartRequest(BaseModel): 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 + # 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 + + @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): @@ -415,6 +493,24 @@ class TrainingStartRequest(BaseModel): 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 0b0191ba12..8da0dc508f 100644 --- a/studio/backend/routes/training.py +++ b/studio/backend/routes/training.py @@ -185,6 +185,68 @@ async def start_training( ) request.resume_from_checkpoint = resume_checkpoint + # 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.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, + 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.", + ) + # 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 = { "model_name": request.model_name, @@ -199,6 +261,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/backend/tests/test_training_streaming.py b/studio/backend/tests/test_training_streaming.py new file mode 100644 index 0000000000..8ff016d3bf --- /dev/null +++ b/studio/backend/tests/test_training_streaming.py @@ -0,0 +1,560 @@ +# 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 +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_chatml_to_alpaca +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)) + + +# --- 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 + + 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) + + +# --- 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", + "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_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 + + +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", + [ + ("LoRA/QLoRA", "raw"), # raw-text format + ("Continued Pretraining", "chatml"), # CPT + ], +) +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_accept_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, + ) + + 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["format_type"] == format_type + + +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" + + +# 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 []: + # 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 + + 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/chat_templates.py b/studio/backend/utils/datasets/chat_templates.py index 82a30fd55b..1da0bb818f 100644 --- a/studio/backend/utils/datasets/chat_templates.py +++ b/studio/backend/utils/datasets/chat_templates.py @@ -7,6 +7,7 @@ Apply chat templates to datasets and generate dataset info summaries. """ 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__) @@ -238,7 +239,13 @@ def apply_chat_template_to_dataset( 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" is_standardized = True @@ -295,13 +302,9 @@ def apply_chat_template_to_dataset( '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() @@ -362,18 +365,14 @@ def apply_chat_template_to_dataset( 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() @@ -384,7 +383,7 @@ def apply_chat_template_to_dataset( # Monitor dataset.map() tqdm progress and relay it. _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/dataset_utils.py b/studio/backend/utils/datasets/dataset_utils.py index faa3deac70..f81a20ad6e 100644 --- a/studio/backend/utils/datasets/dataset_utils.py +++ b/studio/backend/utils/datasets/dataset_utils.py @@ -1203,7 +1203,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/backend/utils/datasets/format_conversion.py b/studio/backend/utils/datasets/format_conversion.py index cb24bd96ba..95c9a00534 100644 --- a/studio/backend/utils/datasets/format_conversion.py +++ b/studio/backend/utils/datasets/format_conversion.py @@ -5,7 +5,7 @@ import os -from datasets import IterableDataset +from .iterable import is_streaming_dataset from loggers import get_logger logger = get_logger(__name__) @@ -37,9 +37,8 @@ def standardize_chat_format( """ import collections import itertools - from datasets import IterableDataset - # Detect a vision tokenizer + # Check if vision tokenizer is used is_vlm = False if tokenizer is not None: if hasattr(tokenizer, "image_processor") or hasattr(tokenizer, "tokenizer"): @@ -151,7 +150,7 @@ def standardize_chat_format( "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: @@ -162,7 +161,20 @@ def standardize_chat_format( 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( @@ -178,11 +190,7 @@ def convert_chatml_to_alpaca( - "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): chatml_data = examples.get(chat_column) if chat_column else None @@ -226,7 +234,7 @@ def convert_chatml_to_alpaca( "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: @@ -237,7 +245,20 @@ def convert_chatml_to_alpaca( 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( @@ -250,11 +271,7 @@ def convert_alpaca_to_chatml( Output: 'conversations' column with standard 'role'/'content' dicts. """ - 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 = [] @@ -283,7 +300,7 @@ def convert_alpaca_to_chatml( "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: @@ -294,7 +311,20 @@ def convert_alpaca_to_chatml( 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/iterable.py b/studio/backend/utils/datasets/iterable.py new file mode 100644 index 0000000000..8408ef75f3 --- /dev/null +++ b/studio/backend/utils/datasets/iterable.py @@ -0,0 +1,20 @@ +# 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/backend/utils/datasets/raw_text.py b/studio/backend/utils/datasets/raw_text.py index 03315fb287..112528fdd0 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"}: @@ -40,7 +66,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, [] @@ -75,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 32f7990d07..285644c5d4 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, @@ -44,6 +45,9 @@ 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 { ArrowDown01Icon, @@ -77,6 +81,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", @@ -162,13 +167,21 @@ export function DatasetSection() { setDatasetSplit, datasetEvalSplit, setDatasetEvalSplit, + datasetStreaming, + setDatasetStreaming, + trainOnCompletions, + maxSteps, + evalSteps, + isVisionModel, + isAudioModel, + isEmbeddingModel, + isDatasetImage, + isDatasetAudio, uploadedFile, uploadedEvalFile, setUploadedEvalFile, hfToken, modelType, - isVisionModel, - isAudioModel, datasetSliceStart, setDatasetSliceStart, datasetSliceEnd, @@ -188,13 +201,21 @@ export function DatasetSection() { setDatasetSplit: s.setDatasetSplit, datasetEvalSplit: s.datasetEvalSplit, setDatasetEvalSplit: s.setDatasetEvalSplit, + datasetStreaming: s.datasetStreaming, + setDatasetStreaming: s.setDatasetStreaming, + trainOnCompletions: s.trainOnCompletions, + maxSteps: s.maxSteps, + evalSteps: s.evalSteps, + isVisionModel: s.isVisionModel, + isAudioModel: s.isAudioModel, + isEmbeddingModel: s.isEmbeddingModel, + isDatasetImage: s.isDatasetImage, + isDatasetAudio: s.isDatasetAudio, uploadedFile: s.uploadedFile, uploadedEvalFile: s.uploadedEvalFile, setUploadedEvalFile: s.setUploadedEvalFile, hfToken: s.hfToken, modelType: s.modelType, - isVisionModel: s.isVisionModel, - isAudioModel: s.isAudioModel, datasetSliceStart: s.datasetSliceStart, setDatasetSliceStart: s.setDatasetSliceStart, datasetSliceEnd: s.datasetSliceEnd, @@ -202,6 +223,57 @@ 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 + // 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."); + if (platformDeviceType === "mac") + streamingBlockers.push( + "Streaming isn't supported on Apple Silicon (MLX) yet.", + ); + + 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 + // 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">( @@ -1116,6 +1188,56 @@ export function DatasetSection() { +
+ 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/sections/params-section.tsx b/studio/frontend/src/features/studio/sections/params-section.tsx index f254921fa6..059d665122 100644 --- a/studio/frontend/src/features/studio/sections/params-section.tsx +++ b/studio/frontend/src/features/studio/sections/params-section.tsx @@ -1067,11 +1067,22 @@ export function ParamsSection(): ReactElement { store.setTrainOnCompletions(!!v)} /> diff --git a/studio/frontend/src/features/studio/training-start-overlay.tsx b/studio/frontend/src/features/studio/training-start-overlay.tsx index 3b0cce658b..c69b25577c 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 ? ( = new Set "isDatasetAudio", "trainOnCompletions", "maxPositionEmbeddings", + "isVisionModel", "s3Config", ]); @@ -165,6 +168,68 @@ 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; +} + +// 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, @@ -649,15 +714,19 @@ 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, + ...streamingPatch, }); + notifyStreamingCompat(streamingPatch); - const state = get(); const datasetName = state.datasetSource === "huggingface" ? state.dataset @@ -681,10 +750,53 @@ export const useTrainingConfigStore = create()( runDatasetCheck(datasetName, split); }, setDatasetEvalSplit: (datasetEvalSplit) => { + const state = get(); + const evalSteps = datasetEvalSplit ? 0.1 : 0; + const streamingPatch = streamingCompatiblePatch({ + ...state, + datasetEvalSplit, + evalSteps, + }); set({ datasetEvalSplit, - evalSteps: datasetEvalSplit ? 0.1 : 0, + evalSteps, + ...streamingPatch, }); + notifyStreamingCompat(streamingPatch); + }, + 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 }), @@ -748,13 +860,34 @@ export const useTrainingConfigStore = create()( set({ gradientAccumulation }), setWeightDecay: (weightDecay) => set({ weightDecay }), setWarmupSteps: (warmupSteps) => set({ warmupSteps }), - setMaxSteps: (maxSteps) => set({ maxSteps }), + 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, + ...streamingPatch, + }); + notifyStreamingCompat(streamingPatch); + }, setSaveSteps: (saveSteps) => set({ saveSteps }), - setEvalSteps: (evalSteps) => set({ evalSteps }), + setEvalSteps: (evalSteps) => { + const state = get(); + const streamingPatch = streamingCompatiblePatch({ ...state, evalSteps }); + set({ + evalSteps, + ...streamingPatch, + }); + notifyStreamingCompat(streamingPatch); + }, setPacking: (packing) => set({ packing }), setTrainOnCompletions: (trainOnCompletions) => { _trainOnCompletionsManuallySet = true; - set({ trainOnCompletions }); + set({ + trainOnCompletions, + ...(trainOnCompletions ? { datasetStreaming: false } : {}), + }); }, setGradientCheckpointing: (gradientCheckpointing) => set({ gradientCheckpointing }), @@ -805,7 +938,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) { @@ -852,9 +985,31 @@ export const useTrainingConfigStore = create()( 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, + 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) { + // 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)); + } + }, }, ), ); diff --git a/studio/frontend/src/features/training/types/api.ts b/studio/frontend/src/features/training/types/api.ts index 1d90a3647f..4f7a41bdea 100644 --- a/studio/frontend/src/features/training/types/api.ts +++ b/studio/frontend/src/features/training/types/api.ts @@ -18,6 +18,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 68ede77d6e..d24ce31a6a 100644 --- a/studio/frontend/src/features/training/types/config.ts +++ b/studio/frontend/src/features/training/types/config.ts @@ -29,6 +29,7 @@ export interface TrainingConfigState { datasetSubset: string | null; datasetSplit: string | null; datasetEvalSplit: string | null; + datasetStreaming: boolean; datasetManualMapping: DatasetManualMapping; datasetSystemPrompt: string; datasetUserTemplate: string; @@ -107,6 +108,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; diff --git a/studio/frontend/src/i18n/locales/en.ts b/studio/frontend/src/i18n/locales/en.ts index 7bc86b57ec..1ea352f081 100644 --- a/studio/frontend/src/i18n/locales/en.ts +++ b/studio/frontend/src/i18n/locales/en.ts @@ -845,6 +845,7 @@ export const en = { resumingTraining: "Resuming training...", startingTraining: "starting training...", dataset: "Dataset", + datasetStreaming: "Dataset: streaming (no full download)", modelWeights: "Model weights", }, tour: { diff --git a/tests/python/test_e2e_no_torch_sandbox.py b/tests/python/test_e2e_no_torch_sandbox.py index 7da7cf0eae..23d4a2c2b3 100644 --- a/tests/python/test_e2e_no_torch_sandbox.py +++ b/tests/python/test_e2e_no_torch_sandbox.py @@ -27,6 +27,7 @@ CHAT_TEMPLATES = DATASETS_DIR / "chat_templates.py" 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 @@ class TestBeforeAfterImportChain: 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 @@ class TestBeforeAfterImportChain: VLM_PROCESSING, DATA_COLLATORS, CHAT_TEMPLATES, + ITERABLE, ]: if src.exists(): shutil.copy2(src, pkg_dir / src.name) @@ -431,10 +437,14 @@ class TestDataclassInstantiation: 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 @@ class TestEdgeCasesBrokenTorch: 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 86dc8581ab..c4efbc8cea 100644 --- a/tests/python/test_studio_import_no_torch.py +++ b/tests/python/test_studio_import_no_torch.py @@ -254,10 +254,15 @@ class TestChatTemplatesNoTorchVenv: 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 @@ class TestChatTemplatesNoTorchVenv: 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 @@ class TestFormatConversionNoTorchVenv: 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 @@ class TestFormatConversionNoTorchVenv: # 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 @@ class TestFormatConversionNoTorchVenv: 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 @@ class TestFormatConversionNoTorchVenv: 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)