From e1f82b4446998a715a80a872ef377d172da50e91 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Thu, 2 Jul 2026 09:56:49 +0000 Subject: [PATCH 1/2] Refuse non-SDXL base models at diffusion training start The trainer only supports the SDXL U-Net, but a FLUX / Qwen-Image / Z-Image repo or a GGUF filename passed as base_model was accepted and then failed minutes later inside StableDiffusionXLPipeline.from_pretrained with an unrelated-looking error. Add a name-based guard in normalized() so known DiT-family names and .gguf checkpoints are rejected up front, which the API start route surfaces as an immediate 400 with a message that says exactly which bases are trainable. Unrecognisable names still pass through so custom local SDXL checkpoints keep working. --- .../core/training/diffusion_lora_trainer.py | 40 +++++++++++++++++++ .../tests/test_diffusion_lora_trainer.py | 31 ++++++++++++++ 2 files changed, 71 insertions(+) diff --git a/studio/backend/core/training/diffusion_lora_trainer.py b/studio/backend/core/training/diffusion_lora_trainer.py index 0e20b6ede5..b213626b54 100644 --- a/studio/backend/core/training/diffusion_lora_trainer.py +++ b/studio/backend/core/training/diffusion_lora_trainer.py @@ -28,6 +28,7 @@ import json import math import os import random +import re import time from dataclasses import dataclass, field, replace from pathlib import Path @@ -42,6 +43,18 @@ _CAPTION_EXTS = (".txt", ".caption") # diffusers' canonical single-file LoRA name, so load_lora_weights(dir) finds it. DEFAULT_LORA_FILENAME = "pytorch_lora_weights.safetensors" +# Families Studio can LOAD but not train (DiT architectures). A base-model name that +# clearly belongs to one is refused in normalized(), so a wrong pick fails at start +# (an instant HTTP 400 through the API) instead of minutes later inside +# StableDiffusionXLPipeline.from_pretrained. Single tokens match on word boundaries; +# hyphenated markers match as substrings of the hyphen-condensed name. +_NON_SDXL_TOKENS = frozenset({"flux", "sd3", "kontext", "pixart", "sana", "lumina", "cogview"}) +_NON_SDXL_PHRASES = ("qwen-image", "z-image", "stable-diffusion-3", "hunyuan-dit") +_ONLY_SDXL_HINT = ( + "Only SDXL bases can be trained right now (e.g. stabilityai/stable-diffusion-xl-base-1.0 " + "or stabilityai/sdxl-turbo). Other families can load LoRAs but not train them yet." +) + EventCb = Callable[[dict[str, Any]], None] # Returns a falsy value to keep training, or a truthy stop signal: bare True, or a dict # that may carry ``save=False`` to cancel without saving a partial adapter. @@ -89,6 +102,7 @@ class DiffusionLoraConfig: Also coerces values that arrive as strings/blanks through the Studio config path (``learning_rate`` is preserved as a string there; ``hf_token`` defaults to "").""" + assert_trainable_base_model(self.base_model) if self.train_steps < 1: raise ValueError("train_steps must be >= 1") if self.train_batch_size < 1: @@ -276,6 +290,32 @@ def _encode_sdxl_prompts( return prompt_embeds, pooled +def assert_trainable_base_model(base_model: str) -> None: + """Refuse base models that are recognisably not SDXL, before anything is downloaded. + + Purely name-based: a GGUF filename or a known DiT-family name (FLUX / Qwen-Image / + Z-Image / SD3 / ...) can never train on the SDXL U-Net trainer, so failing here turns + a confusing mid-run crash into an immediate, actionable error. Names this cannot + classify pass through; from_pretrained still fails cleanly on a genuinely wrong pick.""" + name = str(base_model or "").strip().lower() + if name.endswith(".gguf"): + raise ValueError( + f"'{base_model}' is a GGUF checkpoint, which can't be trained. {_ONLY_SDXL_HINT}" + ) + condensed = re.sub(r"[^a-z0-9]+", "-", name) + hit = next( + (p for p in _NON_SDXL_PHRASES if p in condensed), + None, + ) or next( + (t for t in condensed.split("-") if t in _NON_SDXL_TOKENS), + None, + ) + if hit: + raise ValueError( + f"'{base_model}' looks like a {hit} model, which isn't trainable. {_ONLY_SDXL_HINT}" + ) + + def _assert_trusted_base_model(base_model: str) -> None: """Gate the training base model the same way the inference backend gates non-GGUF loads: a local path or a trusted repo (``unsloth/*`` or an allowlisted official base). This runs diff --git a/studio/backend/tests/test_diffusion_lora_trainer.py b/studio/backend/tests/test_diffusion_lora_trainer.py index 2f83eee1d3..e238fdaff7 100644 --- a/studio/backend/tests/test_diffusion_lora_trainer.py +++ b/studio/backend/tests/test_diffusion_lora_trainer.py @@ -193,3 +193,34 @@ def test_config_rejects_nonpositive_learning_rate(): DiffusionLoraConfig( base_model = "b", data_dir = "d", output_dir = "o", learning_rate = 0 ).normalized() + + +def test_config_rejects_known_non_sdxl_base_models(): + # Known DiT families and GGUF checkpoints must fail at normalise time (an instant + # 400 via the API) instead of minutes later inside from_pretrained. + for bad in ( + "unsloth/FLUX.1-dev-GGUF", + "black-forest-labs/FLUX.1-schnell", + "unsloth/Qwen-Image-2512-unsloth-bnb-4bit", + "Tongyi-MAI/Z-Image-Turbo", + "stabilityai/stable-diffusion-3-medium", + "unsloth/FLUX.1-Kontext-dev", + "z-image-turbo-Q4_K_M.gguf", + ): + with pytest.raises(ValueError, match = "SDXL"): + DiffusionLoraConfig( + base_model = bad, data_dir = "d", output_dir = "o" + ).normalized() + + +def test_config_accepts_sdxl_and_unknown_base_models(): + # SDXL names and unclassifiable custom names/paths must pass the guard (a wrong + # custom pick still fails cleanly in from_pretrained). + for ok in ( + "stabilityai/stable-diffusion-xl-base-1.0", + "stabilityai/sdxl-turbo", + "/data/checkpoints/my-custom-sdxl", + "my-finetune", + ): + cfg = DiffusionLoraConfig(base_model = ok, data_dir = "d", output_dir = "o").normalized() + assert cfg.base_model == ok From cfde12451c45a177f0cd7be275c6688a6e14209b Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Thu, 2 Jul 2026 09:57:26 +0000 Subject: [PATCH 2/2] Add diffusion dataset upload and training info endpoints Training an image LoRA required knowing the Studio home layout and copying files onto the server by hand, which is the most confusing step of the whole flow. Two small endpoints fix that: - GET /api/train/diffusion/info reports the datasets and outputs roots plus every dataset folder that contains images (with image/caption counts), so the UI can offer a picker instead of a blind free-text path. - POST /api/train/diffusion/dataset uploads images and optional caption .txt / metadata.jsonl files into a named folder under the datasets root, creating it on first use and accumulating on repeat uploads so large sets can arrive in batches. Names are validated to a single path component and files stream to disk under the same per-upload size cap as LLM dataset uploads. The returned name is a valid data_dir for /diffusion/start. --- studio/backend/models/training.py | 31 ++++ studio/backend/routes/training.py | 141 +++++++++++++++++- .../backend/tests/test_diffusion_training.py | 82 ++++++++++ 3 files changed, 253 insertions(+), 1 deletion(-) diff --git a/studio/backend/models/training.py b/studio/backend/models/training.py index 876d60213a..ec8900c4a3 100644 --- a/studio/backend/models/training.py +++ b/studio/backend/models/training.py @@ -738,3 +738,34 @@ class DiffusionTrainingStatusResponse(BaseModel): lora_path: Optional[str] = None started_at: Optional[float] = None updated_at: Optional[float] = None + + +class DiffusionDatasetSummary(BaseModel): + """One image-dataset folder under the Studio datasets root.""" + + name: str + path: str + image_count: int + caption_count: int + + +class DiffusionTrainingInfoResponse(BaseModel): + """Where diffusion training reads/writes on this Studio, plus usable datasets. + + Lets the UI show real on-disk locations and offer existing dataset folders, + instead of asking users to know the Studio home layout.""" + + datasets_root: str + outputs_root: str + datasets: List[DiffusionDatasetSummary] + + +class DiffusionDatasetUploadResponse(BaseModel): + """Result of uploading images/captions into a named dataset folder. Counts are + for the whole folder after the upload, so repeat uploads show the running total.""" + + name: str + path: str + image_count: int + caption_count: int + uploaded: int diff --git a/studio/backend/routes/training.py b/studio/backend/routes/training.py index dfdef01411..63e3c74570 100644 --- a/studio/backend/routes/training.py +++ b/studio/backend/routes/training.py @@ -7,7 +7,7 @@ Training API routes import sys from pathlib import Path -from fastapi import APIRouter, Depends, HTTPException, Request +from fastapi import APIRouter, Depends, File, Form, HTTPException, Request, UploadFile from fastapi.responses import StreamingResponse from typing import Dict, Optional, Any import structlog @@ -58,6 +58,9 @@ from models import ( TrainingProgress, ) from models.training import ( + DiffusionDatasetSummary, + DiffusionDatasetUploadResponse, + DiffusionTrainingInfoResponse, DiffusionTrainingStartRequest, DiffusionTrainingStartResponse, DiffusionTrainingStatusResponse, @@ -1192,3 +1195,139 @@ async def diffusion_training_status(current_subject: str = Depends(get_current_s """Poll the current diffusion training job's status/progress (JSON).""" from core.training.diffusion_training_service import get_diffusion_training_service return DiffusionTrainingStatusResponse(**get_diffusion_training_service().status()) + + +# Extensions accepted into an image-training dataset folder: images the trainer reads, +# plus its caption sources (per-image sidecars and metadata/captions jsonl). +_DIFFUSION_DATASET_IMAGE_EXTS = {".png", ".jpg", ".jpeg", ".webp", ".bmp"} +_DIFFUSION_DATASET_TEXT_EXTS = {".txt", ".caption", ".jsonl"} + + +def _diffusion_dataset_summary(folder: Path) -> DiffusionDatasetSummary: + images = captions = 0 + for f in folder.iterdir(): + if not f.is_file(): + continue + ext = f.suffix.lower() + if ext in _DIFFUSION_DATASET_IMAGE_EXTS: + images += 1 + elif ext in (".txt", ".caption"): + captions += 1 + return DiffusionDatasetSummary( + name = folder.name, path = str(folder), image_count = images, caption_count = captions + ) + + +@router.get("/diffusion/info", response_model = DiffusionTrainingInfoResponse) +async def diffusion_training_info(current_subject: str = Depends(get_current_subject)): + """Describe where diffusion training reads/writes, and list usable dataset folders. + + A dataset folder is any direct child of the datasets root that contains at least one + image. The UI uses this to offer a picker instead of a blind free-text path.""" + from utils.paths import datasets_root, outputs_root + + def scan() -> DiffusionTrainingInfoResponse: + root = datasets_root() + found: list[DiffusionDatasetSummary] = [] + try: + children = sorted(p for p in root.iterdir() if p.is_dir()) + except OSError: + children = [] + for child in children: + try: + summary = _diffusion_dataset_summary(child) + except OSError: + continue + if summary.image_count > 0: + found.append(summary) + return DiffusionTrainingInfoResponse( + datasets_root = str(root), outputs_root = str(outputs_root()), datasets = found + ) + + return await asyncio.to_thread(scan) + + +_DATASET_NAME_RE = None # compiled lazily; module keeps its import block torch-free + + +def _clean_diffusion_dataset_name(name: str) -> str: + """Validate a dataset folder name: a single path component, no traversal, printable.""" + import re + global _DATASET_NAME_RE + if _DATASET_NAME_RE is None: + _DATASET_NAME_RE = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._ -]{0,127}$") + cleaned = (name or "").strip() + if not _DATASET_NAME_RE.fullmatch(cleaned) or ".." in cleaned: + raise HTTPException( + status_code = 400, + detail = ( + "Dataset name must be a plain folder name (letters, numbers, dots, " + "dashes, spaces; no slashes), e.g. 'my-style-photos'." + ), + ) + return cleaned + + +@router.post("/diffusion/dataset", response_model = DiffusionDatasetUploadResponse) +async def upload_diffusion_dataset( + name: str = Form(...), + files: list[UploadFile] = File(...), + current_subject: str = Depends(get_current_subject), +): + """Upload training images (and optional caption .txt / metadata.jsonl files) into a + named folder under the Studio datasets root, creating it if needed. Repeat uploads + into the same name accumulate, so large datasets can arrive in batches. The returned + name can be passed directly as ``data_dir`` to /diffusion/start.""" + from utils.paths import datasets_root + from utils.upload_limits import get_upload_limit_bytes, get_upload_limit_label + + cleaned = _clean_diffusion_dataset_name(name) + folder = datasets_root() / cleaned + folder.mkdir(parents = True, exist_ok = True) + + limit_bytes = get_upload_limit_bytes() + total_bytes = 0 + uploaded = 0 + allowed = _DIFFUSION_DATASET_IMAGE_EXTS | _DIFFUSION_DATASET_TEXT_EXTS + for f in files: + filename = Path(f.filename or "").name.strip().replace("\x00", "") + ext = Path(filename).suffix.lower() + if not filename or ext not in allowed: + exts = ", ".join(sorted(allowed)) + raise HTTPException( + status_code = 400, + detail = f"Unsupported file '{f.filename}'. Allowed: {exts}", + ) + dest = folder / filename + complete = False + try: + with open(dest, "wb") as out: + while chunk := await f.read(1024 * 1024): + total_bytes += len(chunk) + if total_bytes > limit_bytes: + raise HTTPException( + status_code = 413, + detail = ( + "Dataset upload too large. " + f"Maximum is {get_upload_limit_label()} per upload; " + "add the remaining images in another batch." + ), + ) + out.write(chunk) + complete = True + finally: + if not complete: + try: + dest.unlink(missing_ok = True) + except OSError: + pass + uploaded += 1 + + summary = _diffusion_dataset_summary(folder) + return DiffusionDatasetUploadResponse( + name = cleaned, + path = str(folder), + image_count = summary.image_count, + caption_count = summary.caption_count, + uploaded = uploaded, + ) diff --git a/studio/backend/tests/test_diffusion_training.py b/studio/backend/tests/test_diffusion_training.py index 35ff816046..6fdfe7c009 100644 --- a/studio/backend/tests/test_diffusion_training.py +++ b/studio/backend/tests/test_diffusion_training.py @@ -377,3 +377,85 @@ def test_stale_pump_events_cannot_corrupt_new_job(): # The current job's events still apply. svc._apply_event({"type": "progress", "step": 9}, proc = current) assert svc.status()["step"] == 9 + + +# ── /diffusion/info + /diffusion/dataset (dataset discovery + upload) ───────── +@pytest.fixture +def dataset_roots(client, monkeypatch, tmp_path): + # The endpoints import these lazily per-request, so patching the package attr works. + import utils.paths as up + + ds_root = tmp_path / "assets" / "datasets" + out_root = tmp_path / "outputs" + ds_root.mkdir(parents = True) + out_root.mkdir(parents = True) + monkeypatch.setattr(up, "datasets_root", lambda: ds_root) + monkeypatch.setattr(up, "outputs_root", lambda: out_root) + return ds_root, out_root + + +def test_diffusion_info_lists_image_dataset_folders(client, dataset_roots): + ds_root, out_root = dataset_roots + good = ds_root / "cat-photos" + good.mkdir() + (good / "a.png").write_bytes(b"x") + (good / "b.jpg").write_bytes(b"x") + (good / "a.txt").write_text("a cat") + (ds_root / "empty-dir").mkdir() # no images -> not a dataset + (ds_root / "stray.txt").write_text("not a folder") + + r = client.get("/api/train/diffusion/info") + assert r.status_code == 200, r.text + body = r.json() + assert body["datasets_root"] == str(ds_root) + assert body["outputs_root"] == str(out_root) + assert [d["name"] for d in body["datasets"]] == ["cat-photos"] + assert body["datasets"][0]["image_count"] == 2 + assert body["datasets"][0]["caption_count"] == 1 + + +def test_diffusion_dataset_upload_accumulates(client, dataset_roots): + ds_root, _ = dataset_roots + files = [ + ("files", ("a.png", b"png-bytes", "image/png")), + ("files", ("b.JPG", b"jpg-bytes", "image/jpeg")), + ("files", ("a.txt", b"a caption", "text/plain")), + ] + r = client.post("/api/train/diffusion/dataset", data = {"name": "my style"}, files = files) + assert r.status_code == 200, r.text + body = r.json() + assert body["name"] == "my style" + assert body["uploaded"] == 3 + assert body["image_count"] == 2 + assert body["caption_count"] == 1 + assert (ds_root / "my style" / "a.png").read_bytes() == b"png-bytes" + + # A second batch into the same name accumulates (large sets arrive in chunks). + r = client.post( + "/api/train/diffusion/dataset", + data = {"name": "my style"}, + files = [("files", ("c.webp", b"w", "image/webp"))], + ) + assert r.status_code == 200, r.text + assert r.json()["uploaded"] == 1 + assert r.json()["image_count"] == 3 + + +def test_diffusion_dataset_upload_rejects_traversal_names(client, dataset_roots): + for bad in ("../evil", "a/b", ".hidden", " "): + r = client.post( + "/api/train/diffusion/dataset", + data = {"name": bad}, + files = [("files", ("a.png", b"x", "image/png"))], + ) + assert r.status_code == 400, f"{bad!r}: {r.status_code}" + + +def test_diffusion_dataset_upload_rejects_unsupported_files(client, dataset_roots): + r = client.post( + "/api/train/diffusion/dataset", + data = {"name": "ok-name"}, + files = [("files", ("weights.exe", b"mz", "application/octet-stream"))], + ) + assert r.status_code == 400 + assert "Unsupported file" in r.json()["detail"]