Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
24 commits
Select commit Hold shift + click to select a range
741b1f2
Add HF dataset streaming mode to Studio
sanatb187 Apr 10, 2026
7482312
Added default value for datasetStreaming in training-config-store.ts
sanatb187 Apr 10, 2026
e1e7218
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Apr 10, 2026
169ab65
Handle None max_steps for streaming validation
sanatb187 Apr 10, 2026
682bc3b
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Apr 10, 2026
aed46be
Merge branch 'main' into feat/studio-dataset-streaming-mode
rolandtannous Apr 10, 2026
0139b8c
Merge branch 'main' into feat/studio-dataset-streaming-mode
rolandtannous Apr 14, 2026
8890be4
Merge branch 'main' into feat/studio-dataset-streaming-mode
rolandtannous Apr 14, 2026
69c85c0
studio: fast-fail streaming validation and guard incompatible modes
rolandtannous Apr 14, 2026
ebea1a6
studio: add streaming dataset tests, iterable helper, and streaming t…
Etherll Jun 19, 2026
d3ca08a
Merge branch 'main' (4c06c1dcc) into feat/studio-dataset-streaming-mode
Etherll Jun 19, 2026
143aff8
studio: fix review-team findings for streaming + main merge
Etherll Jun 19, 2026
249fb1e
Merge current main (42965df2e, +291 commits) into streaming branch
Etherll Jun 19, 2026
f1acb51
studio: enable raw-text/CPT dataset streaming + streaming UX polish
Etherll Jun 19, 2026
d6ea22f
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Jun 19, 2026
a589492
Merge branch 'main' into feat/studio-dataset-streaming-mode
Etherll Jun 19, 2026
583f8fa
studio: address streaming review (MLX/embedding guards, sliced eval s…
Etherll Jun 19, 2026
1c59866
studio: harden HF dataset streaming (column_names, split slicing, emp…
Etherll Jun 20, 2026
14b7d30
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Jun 20, 2026
8820ef1
Merge branch 'main' into feat/studio-dataset-streaming-mode
Etherll Jun 20, 2026
1d80936
studio: fix CI for streaming PR (lint blocker + no-torch sandbox + pr…
Etherll Jun 22, 2026
40cf6d4
Merge branch 'main' into feat/studio-dataset-streaming-mode
Etherll Jun 22, 2026
08bba7f
Merge branch 'main' into feat/studio-dataset-streaming-mode
Etherll Jun 22, 2026
7a4ddd0
Merge branch 'main' into feat/studio-dataset-streaming-mode
Etherll Jun 22, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
414 changes: 275 additions & 139 deletions studio/backend/core/training/trainer.py

Large diffs are not rendered by default.

1 change: 1 addition & 0 deletions studio/backend/core/training/training.py
Original file line number Diff line number Diff line change
Expand Up @@ -281,6 +281,7 @@ def start_training(
"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"),
Expand Down
1 change: 1 addition & 0 deletions studio/backend/core/training/worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -2801,6 +2801,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),
Comment thread
sanatb187 marked this conversation as resolved.
eval_steps = config.get("eval_steps", 0.00),
dataset_slice_start = config.get("dataset_slice_start"),
dataset_slice_end = config.get("dataset_slice_end"),
Expand Down
100 changes: 98 additions & 2 deletions studio/backend/models/training.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,10 @@
_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):
Expand Down Expand Up @@ -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")
Expand All @@ -141,6 +155,70 @@ 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
# 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(<id>, ...)`.
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):
Expand Down Expand Up @@ -415,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.
Expand Down
63 changes: 63 additions & 0 deletions studio/backend/routes/training.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Comment thread
Etherll marked this conversation as resolved.
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:
Comment thread
Etherll marked this conversation as resolved.
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,
Expand All @@ -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,
Expand Down
Loading