unsloth/studio/backend/routes/training.py
2026-07-26 21:39:57 +00:00

2476 lines
111 KiB
Python

# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""
Training API routes
"""
import sys
from pathlib import Path
from fastapi import APIRouter, Depends, File, Form, HTTPException, Request, UploadFile
from fastapi.responses import StreamingResponse
from typing import Dict, Optional, Any
import structlog
from loggers import get_logger
import asyncio
from datetime import datetime
import uuid as _uuid
# Add backend directory to path.
backend_path = Path(__file__).parent.parent.parent
if str(backend_path) not in sys.path:
sys.path.insert(0, str(backend_path))
try:
from core.training import get_training_backend
from core.training.resume import (
can_resume_run,
get_resume_checkpoint_path,
normalize_resume_output_dir,
)
from storage.studio_db import get_resumable_run_by_output_dir
from utils.models.model_config import load_model_defaults
from utils.paths import resolve_dataset_path
except ImportError:
# Fallback: parent directory.
parent_backend = backend_path.parent / "backend"
if str(parent_backend) not in sys.path:
sys.path.insert(0, str(parent_backend))
from core.training import get_training_backend
from core.training.resume import (
can_resume_run,
get_resume_checkpoint_path,
normalize_resume_output_dir,
)
from storage.studio_db import get_resumable_run_by_output_dir
from utils.models.model_config import load_model_defaults
from utils.paths import resolve_dataset_path
# Auth
from auth.authentication import authenticated_via_api_key, get_current_subject
from utils.utils import log_and_http_error
from models import (
TrainingStartRequest,
TrainingJobResponse,
TrainingStatus,
TrainingProgress,
)
from models.training import (
DiffusionCaptionUpdateRequest,
DiffusionDatasetExample,
DiffusionDatasetExamplesResponse,
DiffusionDatasetImageRecord,
DiffusionDatasetImagesResponse,
DiffusionDatasetImportRequest,
DiffusionDatasetImportResponse,
DiffusionDatasetSummary,
DiffusionDatasetUploadResponse,
DiffusionMetricHistory,
DiffusionTrainableFamily,
DiffusionTrainingInfoResponse,
DiffusionTrainingRunDetail,
DiffusionTrainingRunsResponse,
DiffusionTrainingRunSummary,
DiffusionTrainingStartRequest,
DiffusionTrainingStartResponse,
DiffusionTrainingStatusResponse,
DiffusionTrainingStopRequest,
)
from models.responses import TrainingStopResponse, TrainingMetricsResponse
from pydantic import BaseModel as PydanticBaseModel
from pydantic import ValidationError
class TrainingStopRequest(PydanticBaseModel):
save: bool = True
router = APIRouter()
logger = get_logger(__name__)
# Consecutive 1s polls without a step update that count as a stall. Applied only
# once stepping: the pre-first-step phase (model load + tokenization) can take far
# longer, and timing out there made a healthy long-prep run look frozen.
_PROGRESS_STALL_TIMEOUT_POLLS = 1800 # ~30 min at 1 poll/sec
def _validate_local_dataset_paths(paths: list[str], label: str = "Local dataset") -> list[str]:
"""Resolve and validate a list of local dataset paths. Returns validated absolute paths."""
validated = []
missing = []
for dataset_path in paths:
dataset_file = resolve_dataset_path(dataset_path)
if not dataset_file.exists():
missing.append(f"{dataset_path} (resolved: {dataset_file})")
continue
logger.info(f"Found {label.lower()} file: {dataset_file}")
validated.append(str(dataset_file))
if missing:
missing_detail = "; ".join(missing[:3])
raise HTTPException(
status_code = 400,
detail = f"{label} not found: {missing_detail}",
)
return validated
@router.get("/hardware")
async def get_hardware_utilization(current_subject: str = Depends(get_current_subject)):
"""
Live snapshot of GPU hardware utilization for the active backend.
Polled by the frontend during training.
"""
from utils.hardware import get_gpu_utilization
return get_gpu_utilization()
@router.get("/hardware/visible")
async def get_visible_hardware_utilization(current_subject: str = Depends(get_current_subject)):
from utils.hardware import get_visible_gpu_utilization
# Off the event loop: the ROCm fallbacks shell out (Windows perf counters, sysfs) and the System view polls this route.
return await asyncio.to_thread(get_visible_gpu_utilization)
def _background_video_generation_active() -> bool:
"""Whether a video clip is generating on the video backend's worker thread.
POST /video/generate returns at once and generates in the background, so an
in-flight clip is invisible to the keep-warm in-flight request count the
API-key training guards consult; ask the backend directly. Best-effort: a
probe failure must never block a training start."""
try:
from core.inference.video import get_video_backend
return bool(get_video_backend().generate_progress().get("active"))
except Exception as e: # noqa: BLE001
logger.warning("Could not check video generation state for training guard: %s", e)
return False
@router.post("/start")
async def start_training(
request: TrainingStartRequest,
current_subject: str = Depends(get_current_subject),
via_api_key: bool = Depends(authenticated_via_api_key),
):
"""
Start a training job.
Initiates training in the background and returns immediately. Use /status
to check progress.
"""
try:
logger.info(f"Starting training job with model: {request.model_name}")
# When Unsloth is driven as an inference API (API-key auth), refuse to start
# training while a request is in flight: training frees VRAM by unloading
# the chat model, which would kill the stream. The Unsloth UI (session auth)
# still starts training and coexists/frees VRAM as before. (A mixed UI+API
# session is not yet special-cased.)
if via_api_key is True:
from core.inference.llama_keepwarm import other_inference_request_count
if (
other_inference_request_count(current_request_counted = False) > 0
or _background_video_generation_active()
):
raise HTTPException(
status_code = 409,
detail = (
"Cannot start training over the API while an inference request is in "
"progress. Wait for it to finish, or start training from the Unsloth UI."
),
)
# No in-process ensure_transformers_version(): the subprocess
# (worker.py) activates the correct version before importing ML libs.
# A consented latest-transformers install stage-and-swaps .venv_t5_latest;
# a worker spawned mid-swap could activate a half-replaced sidecar.
from utils.transformers_latest import is_install_in_progress
if is_install_in_progress():
raise HTTPException(
status_code = 409,
detail = ("A transformers installation is in progress. Retry when it completes."),
)
backend = get_training_backend()
# S3 dataset loading needs the optional boto3 dependency. Reject early
# with a clear message so credentials are never accepted and then
# silently dropped on a host without boto3 installed.
if request.s3_config is not None:
from core.training.s3_dataset import boto3_available
if not boto3_available():
raise HTTPException(
status_code = 501,
detail = "S3 dataset loading requires boto3. Install it with: pip install boto3",
)
# Check before mutating state.
if backend.is_training_active():
existing_job_id: Optional[str] = getattr(backend, "current_job_id", "")
return TrainingJobResponse(
job_id = existing_job_id or "",
status = "error",
message = (
"Training is already in progress. "
"Stop current training before starting a new one."
),
error = "Training already active",
)
# A diffusion (SDXL) LoRA job runs in its own subprocess on the same GPU, so an LLM start must
# refuse while one is active or the two trainers contend for VRAM. Symmetric with the check in
# start_diffusion_training.
if _diffusion_training_active():
return TrainingJobResponse(
job_id = "",
status = "error",
message = (
"A diffusion (Images) LoRA training job is already running. "
"Stop it before starting an LLM training run."
),
error = "Diffusion training already active",
)
# Job ID; start_training() sets it on the backend only after the old
# pump thread is dead.
job_id = f"job_{datetime.now().strftime('%Y%m%d_%H%M%S')}_{_uuid.uuid4().hex[:8]}"
# Validate dataset paths if provided.
if request.local_datasets:
request.local_datasets = _validate_local_dataset_paths(
request.local_datasets, "Local dataset"
)
if request.local_eval_datasets and request.eval_steps > 0:
request.local_eval_datasets = _validate_local_dataset_paths(
request.local_eval_datasets, "Local eval dataset"
)
resume_output_dir: Optional[str] = None
resume_run: Optional[dict] = None
if request.resume_from_checkpoint:
try:
resume_output_dir = normalize_resume_output_dir(request.resume_from_checkpoint)
except ValueError as e:
# Deliberate user-facing validation message.
validation_message = str(e)
raise HTTPException(status_code = 400, detail = validation_message)
resume_run = get_resumable_run_by_output_dir(resume_output_dir)
if not resume_run or not can_resume_run(resume_run):
raise HTTPException(
status_code = 400,
detail = "Resume checkpoint must belong to a stopped or errored run with complete saved trainer state.",
)
resume_checkpoint = get_resume_checkpoint_path(resume_output_dir)
if not resume_checkpoint:
raise HTTPException(
status_code = 400,
detail = "Resume checkpoint must include saved trainer state.",
)
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,
"project_name": request.project_name,
"training_type": request.training_type,
"hf_token": request.hf_token or "",
"load_in_4bit": request.load_in_4bit,
"max_seq_length": request.max_seq_length,
"vision_image_size": request.vision_image_size,
"hf_dataset": request.hf_dataset or "",
"local_datasets": request.local_datasets,
"local_eval_datasets": request.local_eval_datasets,
"format_type": request.format_type,
"subset": request.subset,
"train_split": request.train_split,
"dataset_streaming": request.dataset_streaming,
"eval_split": request.eval_split,
"eval_steps": request.eval_steps,
"dataset_slice_start": request.dataset_slice_start,
"dataset_slice_end": request.dataset_slice_end,
"custom_format_mapping": request.custom_format_mapping,
"num_epochs": request.num_epochs,
"learning_rate": request.learning_rate,
"embedding_learning_rate": request.embedding_learning_rate,
"batch_size": request.batch_size,
"gradient_accumulation_steps": request.gradient_accumulation_steps,
"warmup_steps": request.warmup_steps,
"warmup_ratio": request.warmup_ratio,
"max_steps": request.max_steps,
"save_steps": request.save_steps,
"weight_decay": request.weight_decay,
"max_grad_norm": request.max_grad_norm,
"max_grad_value": request.max_grad_value,
"max_grad_leaf_norm": request.max_grad_leaf_norm,
"cast_norm_output_to_input_dtype": request.cast_norm_output_to_input_dtype,
"random_seed": request.random_seed,
"packing": request.packing,
"optim": request.optim,
"lr_scheduler_type": request.lr_scheduler_type,
"use_lora": request.use_lora,
"lora_r": request.lora_r,
"lora_alpha": request.lora_alpha,
"lora_dropout": request.lora_dropout,
"target_modules": request.target_modules if request.target_modules else None,
"gradient_checkpointing": request.gradient_checkpointing.strip()
if request.gradient_checkpointing and request.gradient_checkpointing.strip()
else "unsloth",
"use_rslora": request.use_rslora,
"use_loftq": request.use_loftq,
"use_dora": request.use_dora,
"train_on_completions": request.train_on_completions,
"finetune_vision_layers": request.finetune_vision_layers,
"finetune_language_layers": request.finetune_language_layers,
"finetune_attention_modules": request.finetune_attention_modules,
"finetune_mlp_modules": request.finetune_mlp_modules,
"is_dataset_image": request.is_dataset_image,
"is_dataset_audio": request.is_dataset_audio,
"is_embedding": request.is_embedding,
"enable_wandb": request.enable_wandb,
"wandb_token": request.wandb_token or "",
"wandb_project": request.wandb_project or "",
"enable_tensorboard": request.enable_tensorboard,
"tensorboard_dir": request.tensorboard_dir or "",
"output_dir": resume_output_dir,
"resume_from_checkpoint": request.resume_from_checkpoint,
"trust_remote_code": request.trust_remote_code,
"approved_remote_code_fingerprint": request.approved_remote_code_fingerprint,
"subject": current_subject,
"gpu_ids": request.gpu_ids,
"s3_config": request.s3_config.model_dump() if request.s3_config else None,
}
# Latest-sidecar models size and train 16-bit (same flip as chat load):
# 4-bit is disabled for brand-new architectures, so VRAM coexistence
# checks must not underestimate against a load the worker will refuse.
if training_kwargs["load_in_4bit"]:
from utils.transformers_version import latest_tier_active_for
if await asyncio.to_thread(
latest_tier_active_for,
training_kwargs["model_name"],
training_kwargs["hf_token"] or None,
):
training_kwargs["load_in_4bit"] = False
logger.info(
"Latest-transformers sidecar active for %s - sizing and "
"training in 16-bit (4-bit is disabled for brand-new "
"architectures)",
training_kwargs["model_name"],
)
# Training page has no trust_remote_code toggle, so honor the YAML default
# -- but only for genuine first-party (unsloth/nvidia) Hub repos, never a
# local path or a name merely starting with "unsloth/".
if not training_kwargs["trust_remote_code"]:
from utils.security.trusted_org import is_trusted_org_repo
model_defaults = load_model_defaults(request.model_name)
yaml_trust = model_defaults.get("training", {}).get("trust_remote_code", False)
if yaml_trust and is_trusted_org_repo(
request.model_name, hf_token = request.hf_token or None
):
logger.info(f"YAML config sets trust_remote_code=True for {request.model_name}")
training_kwargs["trust_remote_code"] = True
elif yaml_trust:
logger.warning(
"YAML sets trust_remote_code=True for %s but it is not a trusted "
"first-party repo; leaving disabled (user can opt in explicitly).",
request.model_name,
)
# Free VRAM for training: stop export, unload chat unless it can coexist.
# A before_spawn hook -> runs only after start_training's guards pass, so
# we never tear down chat/export VRAM for a start that is then refused.
def _free_vram_for_training() -> None:
try:
from core.export import get_export_backend
exp_backend = get_export_backend()
# Tear down the export subprocess whenever an export is in flight,
# not just once a checkpoint is loaded: during the load phase
# current_checkpoint is still unset while the worker is already
# allocating GPU memory, so gate on is_export_active() too.
if exp_backend.current_checkpoint or exp_backend.is_export_active():
logger.info("Shutting down export subprocess to free GPU memory for training")
exp_backend._shutdown_subprocess()
exp_backend.current_checkpoint = None
exp_backend.is_vision = False
exp_backend.is_peft = False
except Exception as e:
logger.warning("Could not shut down export subprocess: %s", e)
try:
# A resident or in-flight Images pipeline also holds GPU memory the run needs and can't be cheaply
# sized, so tear it down unconditionally like the export subprocess above (the chat block below
# fit-checks; diffusion can't). unload() no-ops when nothing is loaded and preempts an in-flight
# load; release the arbiter so it doesn't think the gone pipeline owns the GPU. Must precede the
# chat block, which early-returns.
from core.inference import gpu_arbiter
from core.inference.diffusion_engine_router import (
get_active_diffusion_engine,
)
# The ACTIVE engine, not the diffusers singleton: on a native (sd_cpp) selection the diffusers
# backend reports unloaded while the native engine still holds model state / a live generation.
diffusion = get_active_diffusion_engine()
if diffusion.is_loaded:
logger.info(
"Unloading diffusion (Images) model to free GPU memory for training"
)
diffusion.unload()
gpu_arbiter.release(gpu_arbiter.DIFFUSION)
except Exception as e:
logger.warning("Could not unload diffusion model for training: %s", e)
try:
# A resident or in-flight Video pipeline holds GPU memory the run needs too, and loads under the
# VIDEO arbiter owner the diffusion teardown above never touches. Tear it down the same way and
# release VIDEO, so a resident video session can't OOM the run. Must precede the chat block.
from core.inference import gpu_arbiter
from core.inference.video import get_video_backend
video = get_video_backend()
if video.status().get("loaded"):
logger.info("Unloading Video model to free GPU memory for training")
video.unload()
gpu_arbiter.release(gpu_arbiter.VIDEO)
except Exception as e:
logger.warning("Could not unload video model for training: %s", e)
try:
from routes.training_vram import (
can_keep_chat_during_training,
coordinate_models_for_training,
)
def _can_keep_resident_models():
return can_keep_chat_during_training(
model_name = training_kwargs["model_name"],
hf_token = training_kwargs["hf_token"],
training_type = training_kwargs["training_type"],
load_in_4bit = training_kwargs["load_in_4bit"],
batch_size = training_kwargs["batch_size"],
max_seq_length = training_kwargs["max_seq_length"],
lora_rank = training_kwargs["lora_r"],
target_modules = training_kwargs["target_modules"],
gradient_checkpointing = training_kwargs["gradient_checkpointing"],
optimizer = training_kwargs["optim"],
gpu_ids = training_kwargs["gpu_ids"],
)
freed = coordinate_models_for_training(_can_keep_resident_models)
if freed:
logger.info("Freed models for training: %s", freed)
except Exception as e:
logger.warning("Inference/training memory coordination failed; proceeding: %s", e)
# The hook runs only once start guards pass -> VRAM freed iff training starts.
from utils.transformers_version import SidecarSwapInProgress
try:
# Offloaded to a worker thread: the hook's diffusion/video unload() waits on the engines'
# generation locks until an in-flight denoise step hits its cancel callback (and the export
# subprocess teardown can take seconds), which would otherwise freeze every concurrent
# status/cancel/UI request. Overlapping starts are serialized by the backend's own guard.
success = await asyncio.to_thread(
backend.start_training,
job_id = job_id,
before_spawn = _free_vram_for_training,
resume_source_run_id = resume_run["id"] if resume_run else None,
**training_kwargs,
)
except SidecarSwapInProgress as exc:
# Expected loss of the race against a sidecar install: a retryable
# 409 matching the route-entry guard, not an internal error.
raise HTTPException(status_code = 409, detail = str(exc))
if not success:
progress_error = backend.trainer.training_progress.error
return TrainingJobResponse(
job_id = backend.current_job_id or "",
status = "error",
message = progress_error or "Failed to start training subprocess",
error = progress_error or "subprocess_start_failed",
)
return TrainingJobResponse(
job_id = job_id,
status = "queued",
message = "Training job queued and starting in subprocess",
error = None,
)
except HTTPException:
# Deliberate rejections (S3 not implemented, resume validation) must
# reach the client with their original status, not a generic 500.
raise
except ValueError as e:
logger.warning("Rejected training GPU selection: %s", e)
# Deliberate user-facing GPU-selection validation message.
validation_message = str(e)
raise HTTPException(status_code = 400, detail = validation_message)
except Exception as e:
raise log_and_http_error(
e,
500,
"Failed to start training",
event = "training.start_failed",
log = logger,
)
@router.post("/stop", response_model = TrainingStopResponse)
async def stop_training(
body: TrainingStopRequest = TrainingStopRequest(),
current_subject: str = Depends(get_current_subject),
):
"""
Stop the currently running training job.
Body:
save (bool): If True (default), save the model at the current checkpoint.
"""
try:
backend = get_training_backend()
is_active = backend.is_training_active()
logger.info("Stop requested: save=%s is_active=%s", body.save, is_active)
if not is_active:
return TrainingStopResponse(
status = "idle", message = "No training job is currently running"
)
if not backend.stop_training(save = body.save):
return TrainingStopResponse(
status = "idle", message = "No training job is currently running"
)
return TrainingStopResponse(
status = "stopped",
message = "Stop requested. Training will stop at the next safe step.",
)
except Exception as e:
raise log_and_http_error(
e,
500,
"Failed to stop training",
event = "training.stop_failed",
log = logger,
)
@router.post("/reset")
async def reset_training(current_subject: str = Depends(get_current_subject)):
"""Reset training state so the user can return to configuration."""
try:
backend = get_training_backend()
is_active = backend.is_training_active()
if is_active:
if backend._cancel_requested:
# Cancel (save=False) requested — force-terminate to reset immediately.
logger.info("Force-terminating subprocess for immediate reset (cancel path)")
backend.force_terminate()
else:
logger.warning("Rejected reset while training active: is_active=%s", is_active)
raise HTTPException(
status_code = 409,
detail = "Training is still running. Stop training and wait for it to finish before resetting.",
)
logger.info("Reset training state: clearing runtime + metric history")
backend._should_stop = False # Clear stop flag so status returns to idle
backend.trainer._update_progress(
is_training = False,
is_completed = False,
error = None,
status_message = "Ready to train",
step = 0,
loss = None,
epoch = 0,
total_steps = 0,
)
backend.loss_history = []
backend.lr_history = []
backend.step_history = []
backend.grad_norm_history = []
backend.grad_norm_step_history = []
return {"status": "ok"}
except HTTPException:
raise
except Exception as e:
raise log_and_http_error(
e,
500,
"Failed to reset training",
event = "training.reset_failed",
log = logger,
)
@router.get("/status")
async def get_training_status(current_subject: str = Depends(get_current_subject)):
"""
Get the current training status.
"""
try:
backend = get_training_backend()
job_id: str = getattr(backend, "current_job_id", "") or ""
is_active = backend.is_training_active()
try:
progress = backend.trainer.get_training_progress()
except Exception:
progress = None
status_message = (
getattr(progress, "status_message", None) if progress else None
) or "Ready to train"
error_message = getattr(progress, "error", None) if progress else None
trainer_stopped = getattr(backend, "_should_stop", False)
# Derive high-level phase
if error_message:
phase = "error"
elif is_active:
msg_lower = status_message.lower()
if "loading" in msg_lower or "importing" in msg_lower:
phase = "loading_model"
elif any(k in msg_lower for k in ["preparing", "initializing", "configuring"]):
phase = "configuring"
else:
phase = "training"
elif trainer_stopped:
phase = "stopped"
elif progress and getattr(progress, "is_completed", False):
phase = "completed"
else:
phase = "idle"
details = None
if progress:
details = {
"epoch": getattr(progress, "epoch", 0),
"step": getattr(progress, "step", 0),
"total_steps": getattr(progress, "total_steps", 0),
"loss": getattr(progress, "loss", None),
"learning_rate": getattr(progress, "learning_rate", None),
}
# Always present: an explicit null tells the client to drop a cached
# path (stop without save clears the run's output_dir).
details["output_dir"] = getattr(backend, "_output_dir", None) or None
# Metric history for chart recovery after SSE reconnection.
metric_history = None
if backend.step_history:
metric_history = {
"steps": list(backend.step_history),
"loss": list(backend.loss_history),
"lr": list(backend.lr_history),
"grad_norm": list(getattr(backend, "grad_norm_history", [])),
"grad_norm_steps": list(getattr(backend, "grad_norm_step_history", [])),
"eval_loss": list(backend.eval_loss_history),
"eval_steps": list(backend.eval_step_history),
}
return TrainingStatus(
job_id = job_id,
phase = phase,
is_training_running = is_active,
eval_enabled = backend.eval_enabled,
message = status_message,
error = error_message,
details = details,
metric_history = metric_history,
)
except Exception as e:
raise log_and_http_error(
e,
500,
"Failed to get training status",
event = "training.status_failed",
log = logger,
)
@router.get("/metrics", response_model = TrainingMetricsResponse)
async def get_training_metrics(current_subject: str = Depends(get_current_subject)):
"""
Get training metrics (loss, learning rate, steps).
"""
try:
backend = get_training_backend()
loss_history = backend.loss_history
lr_history = backend.lr_history
step_history = backend.step_history
grad_norm_history = getattr(backend, "grad_norm_history", [])
grad_norm_step_history = getattr(backend, "grad_norm_step_history", [])
current_loss = loss_history[-1] if loss_history else None
current_lr = lr_history[-1] if lr_history else None
current_step = step_history[-1] if step_history else None
return TrainingMetricsResponse(
loss_history = loss_history,
lr_history = lr_history,
step_history = step_history,
grad_norm_history = grad_norm_history,
grad_norm_step_history = grad_norm_step_history,
current_loss = current_loss,
current_lr = current_lr,
current_step = current_step,
)
except Exception as e:
raise log_and_http_error(
e,
500,
"Failed to get training metrics",
event = "training.metrics_failed",
log = logger,
)
@router.get("/progress")
async def stream_training_progress(
request: Request, current_subject: str = Depends(get_current_subject)
):
"""
Stream training progress via Server-Sent Events (SSE).
Real-time progress with reconnection support per the SSE spec:
- `id:` per event so the browser tracks position.
- `retry:` to control reconnection interval.
- Named `event:` types (progress, heartbeat, complete, error).
- Reads `Last-Event-ID` on reconnect to replay missed steps.
"""
# Read Last-Event-ID header for reconnection resume.
last_event_id = request.headers.get("last-event-id")
resume_from_step: Optional[int] = None
if last_event_id is not None:
try:
resume_from_step = int(last_event_id)
# Fires on every reconnect (each tab switch); the meaningful signal is
# the "replayed N missed steps" line below, logged only when N > 0.
logger.debug(f"SSE reconnect: resuming from step {resume_from_step}")
except ValueError:
logger.warning(f"Invalid Last-Event-ID: {last_event_id}")
async def event_generator():
backend = get_training_backend()
job_id: str = getattr(backend, "current_job_id", "") or ""
# ── Helpers ──────────────────────────────────────────────
def build_progress(
step: int,
loss: Optional[float],
learning_rate: Optional[float],
total_steps: int,
epoch: Optional[float] = None,
progress: Optional[Any] = None,
grad_norm_override: Optional[float] = None,
eval_loss_override: Optional[float] = None,
) -> TrainingProgress:
total = max(total_steps, 0)
if step < 0 or total == 0:
progress_percent = 0.0
else:
progress_percent = float(step) / float(total) * 100.0 if total > 0 else 0.0
# Pull values from the progress object if available.
elapsed_seconds = getattr(progress, "elapsed_seconds", None) if progress else None
eta_seconds = getattr(progress, "eta_seconds", None) if progress else None
grad_norm = grad_norm_override
if grad_norm is None and progress:
grad_norm = getattr(progress, "grad_norm", None)
num_tokens = getattr(progress, "num_tokens", None) if progress else None
eval_loss = eval_loss_override
if eval_loss is None and progress:
eval_loss = getattr(progress, "eval_loss", None)
return TrainingProgress(
job_id = job_id,
step = step,
total_steps = total,
loss = loss,
learning_rate = learning_rate,
progress_percent = progress_percent,
epoch = epoch,
elapsed_seconds = elapsed_seconds,
eta_seconds = eta_seconds,
grad_norm = grad_norm,
num_tokens = num_tokens,
eval_loss = eval_loss,
)
def format_sse(
data: str,
event: str = "progress",
event_id: Optional[int] = None,
) -> str:
"""Format a single SSE message with id/event/data fields."""
lines = []
if event_id is not None:
lines.append(f"id: {event_id}")
lines.append(f"event: {event}")
lines.append(f"data: {data}")
lines.append("") # trailing blank line
lines.append("") # double newline terminates the event
return "\n".join(lines)
# ── Retry directive ──────────────────────────────────────
# Reconnect after 3 seconds if the connection drops.
yield "retry: 3000\n\n"
# ── Replay missed steps on reconnect ─────────────────────
if resume_from_step is not None and backend.step_history:
replayed = 0
grad_norm_by_step = {
step_val: grad_val
for step_val, grad_val in zip(
getattr(backend, "grad_norm_step_history", []),
getattr(backend, "grad_norm_history", []),
)
}
for i, step_val in enumerate(backend.step_history):
if step_val > resume_from_step:
loss_val = backend.loss_history[i] if i < len(backend.loss_history) else None
lr_val = backend.lr_history[i] if i < len(backend.lr_history) else None
tp_replay = getattr(
getattr(backend, "trainer", None), "training_progress", None
)
total_replay = (
getattr(tp_replay, "total_steps", step_val) if tp_replay else step_val
)
epoch_replay = getattr(tp_replay, "epoch", None) if tp_replay else None
payload = build_progress(
step_val,
loss_val,
lr_val,
total_replay,
epoch_replay,
progress = tp_replay,
grad_norm_override = grad_norm_by_step.get(step_val),
)
yield format_sse(payload.model_dump_json(), event = "progress", event_id = step_val)
replayed += 1
if replayed:
logger.info(f"SSE reconnect: replayed {replayed} missed steps")
# ── Initial status (only on fresh connections) ───────────
if resume_from_step is None:
is_active = backend.is_training_active()
tp = getattr(getattr(backend, "trainer", None), "training_progress", None)
initial_total_steps = getattr(tp, "total_steps", 0) if tp else 0
initial_epoch = getattr(tp, "epoch", None) if tp else None
initial_progress = build_progress(
step = 0,
loss = None,
learning_rate = None,
total_steps = initial_total_steps,
epoch = initial_epoch,
progress = tp,
)
yield format_sse(initial_progress.model_dump_json(), event = "progress", event_id = 0)
# If not active, send final state and exit
if not is_active:
_live = (getattr(tp, "step", 0) or 0) if tp else 0
if backend.step_history or _live > 0:
final_step = backend.step_history[-1] if backend.step_history else 0
final_loss = backend.loss_history[-1] if backend.loss_history else None
final_lr = backend.lr_history[-1] if backend.lr_history else None
# Histories skip non-finite steps; report the live step with
# loss=None instead of the last finite pair.
if _live > final_step:
final_step = _live
final_loss = getattr(tp, "loss", None)
final_lr = getattr(tp, "learning_rate", final_lr)
final_total_steps = getattr(tp, "total_steps", final_step) if tp else final_step
final_epoch = getattr(tp, "epoch", None) if tp else None
payload = build_progress(
final_step,
final_loss,
final_lr,
final_total_steps,
final_epoch,
progress = tp,
)
yield format_sse(
payload.model_dump_json(), event = "complete", event_id = final_step
)
else:
yield format_sse(
build_progress(-1, None, None, 0, progress = tp).model_dump_json(),
event = "complete",
event_id = 0,
)
return
# ── Live polling loop ────────────────────────────────────
last_step = resume_from_step if resume_from_step is not None else -1
no_update_count = 0
# The stall timeout applies only once the run is stepping (pre-step prep
# may legitimately emit no step for a long time). On reconnect to an
# already-stepping run, seed from the resume point / history, else a worker
# that hangs after step N never times out for a client that reconnects past it.
seen_live_step = (resume_from_step is not None and resume_from_step > 0) or bool(
backend.step_history
)
while backend.is_training_active():
# Client gone: end the generator without falling through to the final
# "complete" frame, which a buffered/proxy consumer could otherwise read
# as a finished run while training is still active.
if await request.is_disconnected():
return
try:
tp_inner = getattr(getattr(backend, "trainer", None), "training_progress", None)
live_step = (getattr(tp_inner, "step", 0) or 0) if tp_inner else 0
if backend.step_history or live_step > 0:
current_step = backend.step_history[-1] if backend.step_history else 0
current_loss = backend.loss_history[-1] if backend.loss_history else None
current_lr = backend.lr_history[-1] if backend.lr_history else None
# Histories skip non-finite steps; follow the live progress
# step and report its loss (None until it recovers).
if live_step > current_step:
current_step = live_step
current_loss = getattr(tp_inner, "loss", None)
current_lr = getattr(tp_inner, "learning_rate", current_lr)
current_total_steps = (
getattr(tp_inner, "total_steps", current_step) if tp_inner else current_step
)
current_epoch = getattr(tp_inner, "epoch", None) if tp_inner else None
# Only send if the step changed.
if current_step != last_step:
progress_payload = build_progress(
current_step,
current_loss,
current_lr,
current_total_steps,
current_epoch,
progress = tp_inner,
)
yield format_sse(
progress_payload.model_dump_json(),
event = "progress",
event_id = current_step,
)
last_step = current_step
no_update_count = 0
seen_live_step = True
else:
no_update_count += 1
# Heartbeat every 10 seconds.
if no_update_count % 10 == 0:
heartbeat_payload = build_progress(
current_step,
current_loss,
current_lr,
current_total_steps,
current_epoch,
progress = tp_inner,
)
yield format_sse(
heartbeat_payload.model_dump_json(),
event = "heartbeat",
event_id = current_step,
)
else:
# No steps yet, but training is active (model loading, etc.).
no_update_count += 1
if no_update_count % 5 == 0:
# Pull total_steps + status so the frontend can show
# "Tokenizing…" etc.
tp_prep = getattr(
getattr(backend, "trainer", None),
"training_progress",
None,
)
prep_total = getattr(tp_prep, "total_steps", 0) if tp_prep else 0
preparing_payload = build_progress(
0,
None,
None,
prep_total,
progress = tp_prep,
)
yield format_sse(
preparing_payload.model_dump_json(),
event = "heartbeat",
event_id = 0,
)
# Fires only once stepping: a long pre-first-step prep phase is not
# a stall, and ending the stream there made a healthy run look frozen.
if seen_live_step and no_update_count > _PROGRESS_STALL_TIMEOUT_POLLS:
logger.warning("Progress stream timeout - no updates received")
tp_timeout = getattr(
getattr(backend, "trainer", None), "training_progress", None
)
timeout_payload = build_progress(last_step, None, None, 0, progress = tp_timeout)
yield format_sse(
timeout_payload.model_dump_json(),
event = "error",
event_id = last_step if last_step >= 0 else 0,
)
break
await asyncio.sleep(1) # Poll every second
except Exception as e:
logger.error(f"Error in progress stream: {e}", exc_info = True)
tp_error = getattr(getattr(backend, "trainer", None), "training_progress", None)
error_payload = build_progress(0, None, None, 0, progress = tp_error)
yield format_sse(
error_payload.model_dump_json(),
event = "error",
event_id = last_step if last_step >= 0 else 0,
)
break
# ── Final "complete" event ───────────────────────────────
final_step = backend.step_history[-1] if backend.step_history else last_step
final_loss = backend.loss_history[-1] if backend.loss_history else None
final_lr = backend.lr_history[-1] if backend.lr_history else None
final_tp = getattr(getattr(backend, "trainer", None), "training_progress", None)
# If the run ended on a non-finite stretch, report the live step with
# loss=None instead of rolling back to the last finite pair.
_final_live_step = (getattr(final_tp, "step", 0) or 0) if final_tp else 0
if _final_live_step > (final_step if final_step is not None else -1):
final_step = _final_live_step
final_loss = getattr(final_tp, "loss", None)
final_lr = getattr(final_tp, "learning_rate", final_lr)
final_total_steps = getattr(final_tp, "total_steps", final_step) if final_tp else final_step
final_epoch = getattr(final_tp, "epoch", None) if final_tp else None
final_payload = build_progress(
final_step,
final_loss,
final_lr,
final_total_steps,
final_epoch,
progress = final_tp,
)
yield format_sse(
final_payload.model_dump_json(),
event = "complete",
event_id = final_step if final_step >= 0 else 0,
)
return StreamingResponse(
event_generator(),
media_type = "text/event-stream",
headers = {
"Cache-Control": "no-cache",
"Connection": "keep-alive",
"X-Accel-Buffering": "no",
},
)
# ── Diffusion (SDXL) LoRA training ────────────────────────────────────────────
# A separate, lightweight job path from the LLM endpoints above: diffusion runs are driven by
# DiffusionTrainingService (its own subprocess + event pump), not the LLM TrainingBackend, so the
# two never contend and diffusion never triggers LLM lifecycle (DB run rows, plots, transfer).
def _diffusion_training_active() -> bool:
"""Whether a diffusion (SDXL) LoRA job is currently running. Best-effort so the
interlock never blocks a start just because the service could not be imported."""
try:
from core.training.diffusion_training_service import get_diffusion_training_service
return get_diffusion_training_service().is_active()
except Exception: # noqa: BLE001
return False
def _require_diffusion_dataset_mutable() -> None:
"""Reject a dataset mutation while a diffusion run is active.
The trainer re-opens dataset images during the loop, so mutating underneath it makes the run
nondeterministic or raises a FileNotFoundError mid-step. Fails open (a service-import failure
never blocks a mutation on an unknowable state), matching the start interlock."""
if _diffusion_training_active():
raise HTTPException(
status_code = 409,
detail = (
"Training images cannot be changed while diffusion training is active. "
"Stop the run before uploading, importing, editing captions, or deleting images."
),
)
def diffusion_dataset_interlock():
"""Dependency holding the dataset interlock for a whole mutating request.
The check above only covers the instant it runs: every one of these endpoints then hands its
filesystem work to a thread, and a ``/diffusion/start`` reserving in that gap would move
captions or images underneath the preflight or the running trainer. As a yield dependency the
registration spans the endpoint, so ``reserve()`` sees it and refuses instead. Fails open on an
import error, like the check it replaces."""
try:
from core.training.diffusion_training_service import (
TrainingActiveError,
get_diffusion_training_service,
)
service = get_diffusion_training_service()
except Exception: # noqa: BLE001 -- unknowable state never blocks a mutation
yield
return
try:
with service.dataset_mutation():
yield
except TrainingActiveError as exc:
raise HTTPException(status_code = 409, detail = str(exc)) from exc
def _free_gpu_for_diffusion_training() -> None:
"""Free GPU residents before the diffusion trainer spawns its own SDXL pipeline.
The trainer subprocess loads a full SDXL pipeline; an export worker, a resident
Images pipeline, or loaded chat models would otherwise keep their VRAM allocated and
OOM the run. Mirrors the LLM start path's pre-spawn cleanup (export + diffusion
pipeline + chat). Best-effort: a failure to free one resident never blocks the start."""
try:
from core.export import get_export_backend
exp_backend = get_export_backend()
if exp_backend.current_checkpoint or exp_backend.is_export_active():
logger.info("Shutting down export subprocess to free GPU memory for diffusion training")
exp_backend._shutdown_subprocess()
exp_backend.current_checkpoint = None
exp_backend.is_vision = False
exp_backend.is_peft = False
except Exception as e: # noqa: BLE001
logger.warning("Could not shut down export subprocess: %s", e)
try:
from core.inference import gpu_arbiter
from core.inference.diffusion_engine_router import get_active_diffusion_engine
# The ACTIVE engine, not the diffusers singleton: on a native (sd_cpp) selection the diffusers
# backend reports unloaded while the resident sd-server still holds the GPU, so unloading only
# the singleton is a no-op. Mirrors the LLM training start path.
diffusion = get_active_diffusion_engine()
if diffusion.is_loaded:
logger.info("Unloading resident Images pipeline to free GPU memory for training")
diffusion.unload() # no-op when nothing is loaded; also preempts an in-flight load
gpu_arbiter.release(gpu_arbiter.DIFFUSION)
except Exception as e: # noqa: BLE001
logger.warning("Could not unload Images pipeline for diffusion training: %s", e)
try:
# A resident Video pipeline loads under the VIDEO arbiter owner the Images teardown above
# doesn't free; unload it too (no-op when nothing is loaded) and release VIDEO so a resident
# video session can't OOM the diffusion trainer.
from core.inference import gpu_arbiter
from core.inference.video import get_video_backend
video = get_video_backend()
if video.status().get("loaded"):
logger.info("Unloading resident Video pipeline to free GPU memory for training")
video.unload() # no-op when nothing is loaded; also preempts an in-flight load
gpu_arbiter.release(gpu_arbiter.VIDEO)
except Exception as e: # noqa: BLE001
logger.warning("Could not unload Video pipeline for diffusion training: %s", e)
try:
# The SDXL trainer's footprint can't be cheaply sized against a resident chat model, so free chat
# unconditionally (like the LLM path does for an in-flight load) rather than risk an OOM.
from routes.training_vram import free_chat_models_for_training, summarize_resident_chat
if summarize_resident_chat()["any"]:
freed = free_chat_models_for_training(reason = "diffusion training starting")
logger.info("Freed chat model(s) for diffusion training: %s", freed)
except Exception as e: # noqa: BLE001
logger.warning("Could not free chat models for diffusion training: %s", e)
def _preflight_gated_base(base_model: str, hf_token: Optional[str]) -> None:
"""HEAD a remote base repo's model_index.json with the caller's token; raise HTTP 400 on
401/403 (gated / unauthorized) with an actionable message. Best-effort: a local path,
a non-repo string, or a network hiccup passes through so the trainer can surface any real
load error itself. Runs before GPU teardown so a doomed start never evicts a loaded model."""
import urllib.error
import urllib.request
repo = (base_model or "").strip()
# Only remote 'org/name' repos are gated; skip local paths and single-file names.
if (
not repo
or repo.count("/") != 1
or repo.startswith((".", "/", "~"))
or repo.endswith(".gguf")
):
return
url = f"https://huggingface.co/{repo}/resolve/main/model_index.json"
headers = {"Authorization": f"Bearer {hf_token}"} if hf_token else {}
req = urllib.request.Request(url, method = "HEAD", headers = headers)
try:
urllib.request.urlopen(req, timeout = 5)
except urllib.error.HTTPError as e:
if e.code in (401, 403):
raise HTTPException(
status_code = 400,
detail = (
f"Access to '{repo}' is gated or unauthorized. Accept the model's license "
f"on its Hugging Face page and add your HF token in Studio settings, then "
f"try again."
),
)
# 404 (e.g. a repo without a root model_index.json) and other codes are not an access problem;
# let the trainer surface any genuine load error.
except Exception: # noqa: BLE001 -- network/DNS hiccup must not block a start
return
def _resolve_diffusion_data_dir(raw: str) -> Path:
"""Resolve a diffusion-training ``data_dir``. The upload/labeling routes create and
manage image datasets directly under ``datasets_root()`` and the UI passes the bare
folder name back as ``data_dir``, but the generic :func:`resolve_dataset_path`
searches the LLM uploads and recipe dataset roots FIRST -- so an unrelated upload
file or recipe folder sharing that name would shadow the just-uploaded image
dataset (preflight 400 "not a directory", or training the wrong data). Prefer the
image dataset root for a bare single-component name that exists there; everything
else (explicit "uploads/..." / "recipes/..." prefixes, absolute paths, missing
names) resolves exactly as before."""
from utils.paths import datasets_root
value = str(raw or "").strip()
if value and "\x00" not in value:
p = Path(value)
# A single component that is not "..", so joining under datasets_root() cannot escape it.
if not p.is_absolute() and len(p.parts) == 1 and p.parts[0] != "..":
direct = datasets_root() / value
# Route a bare name through the same protected resolver the CRUD routes use, so a name to
# external-directory symlink is rejected here too (is_dir() follows the link). A broken symlink
# is included so it is rejected, not passed to resolve_dataset_path.
if direct.is_dir() or direct.is_symlink():
return _resolve_dataset_folder(value)
return resolve_dataset_path(raw)
@router.post("/diffusion/start", response_model = DiffusionTrainingStartResponse)
async def start_diffusion_training(
body: DiffusionTrainingStartRequest,
current_subject: str = Depends(get_current_subject),
via_api_key: bool = Depends(authenticated_via_api_key),
):
"""Start an SDXL LoRA training job from an image + caption dataset."""
from core.training.diffusion_training_service import get_diffusion_training_service
# Under API-key auth, refuse to start training while a request is in flight:
# _free_gpu_for_diffusion_training() below unloads the chat backends, killing the stream.
# Mirrors start_training.
if via_api_key is True:
from core.inference.llama_keepwarm import other_inference_request_count
if (
other_inference_request_count(current_request_counted = False) > 0
or _background_video_generation_active()
):
raise HTTPException(
status_code = 409,
detail = (
"Cannot start diffusion (Images) training over the API while an inference "
"request is in progress. Wait for it to finish, or start training from the "
"Studio UI."
),
)
# Interlock: refuse while an LLM training run holds the GPU (symmetric with the diffusion check
# in start_training), so the two trainers never contend for VRAM.
try:
if get_training_backend().is_training_active():
raise HTTPException(
status_code = 409,
detail = (
"An LLM training job is already running. "
"Stop it before starting diffusion (Images) training."
),
)
except HTTPException:
raise
except Exception: # noqa: BLE001 -- backend import/health issue must not block a start
pass
# Resolve + contain the dataset and output paths BEFORE spawning, so Studio-relative names work
# and absolute paths stay under a Studio root -- the trainer subprocess otherwise resolves them
# relative to its own cwd.
config = body.model_dump()
try:
from utils.paths import resolve_output_dir
config["data_dir"] = str(_resolve_diffusion_data_dir(config["data_dir"]))
config["output_dir"] = str(resolve_output_dir(config["output_dir"]))
except ValueError as e:
raise HTTPException(status_code = 400, detail = str(e))
# Validate the config BEFORE freeing resident GPU workloads, so a start then refused (bad numbers,
# non-SDXL base) never tears down the user's chat/Images model. service.start() re-runs this
# cheaply before spawn.
from core.training.diffusion_lora_trainer import _config_from_dict
try:
normalized_cfg = _config_from_dict(config).normalized()
except ValueError as e:
raise HTTPException(status_code = 400, detail = str(e))
# Preflight the requested DiT precision BEFORE freeing GPU residents: the trainer's own checks
# (bf16-capable GPU required; explicit int8 needs a functional torchao) fire only in the child,
# AFTER _free_gpu_for_diffusion_training() evicted the user's model. Fail fast (400) so a
# pre-Ampere GPU or stub-torchao host never tears down residents for a run that cannot start.
from core.training.diffusion_train_common import training_precision_preflight_error
_precision_reason = training_precision_preflight_error(
normalized_cfg.resolved_family, normalized_cfg.base_precision
)
if _precision_reason:
raise HTTPException(status_code = 400, detail = _precision_reason)
# Run the trainers' trust gate here too (both assert the same predicate before from_pretrained),
# so an untrusted/typoed base 400s BEFORE freeing GPU residents rather than failing in the child.
from core.training.diffusion_train_common import _assert_trusted_base_model
try:
_assert_trusted_base_model(config.get("base_model", ""))
except ValueError as e:
raise HTTPException(status_code = 400, detail = str(e))
# Preflight access to a gated base repo with the user's token BEFORE freeing GPU residents, so a
# missing/insufficient token fails fast (400) without tearing down the user's model and never
# surfaces as a confusing mid-load 401. Offloaded to a worker thread: it does a blocking urlopen
# HEAD (5s timeout) that would otherwise stall the event loop.
await asyncio.to_thread(
_preflight_gated_base, config.get("base_model", ""), config.get("hf_token")
)
from core.training import diffusion_train_common as _dtc
service = get_diffusion_training_service()
# Reserve the training slot BEFORE the dataset preflight (not just before freeing residents):
# is_active() otherwise flips true only at service.start(), so during this scan -- which
# decode-probes every image and can take a while -- a concurrent upload/caption/delete would pass
# _require_diffusion_dataset_mutable() and mutate the dataset the trainer is about to read, and a
# concurrent /images/load or /video/load would double-allocate VRAM. reserve() is a
# compare-and-set, so a second overlapping start 409s before touching anything; unreserve() runs
# in the finally ONLY when THIS request reserved.
reserved = False
try:
service.reserve()
reserved = True
# Preflight the dataset: a missing/empty/uncaptionable data_dir otherwise fails inside the spawned
# trainer AFTER the user's model was evicted. Same discovery the trainer runs, so the two cannot
# disagree.
try:
await asyncio.to_thread(
_dtc.discover_image_caption_pairs,
config["data_dir"],
instance_prompt = config.get("instance_prompt") or None,
caption_column = config.get("caption_column") or "text",
# Decode-probe every image now (cheap PIL header check) so a corrupt/zero-byte upload 400s BEFORE
# _free_gpu_for_diffusion_training() tears down the user's models.
verify_images = True,
)
except (FileNotFoundError, ValueError) as e:
raise HTTPException(status_code = 400, detail = str(e))
# Free resident GPU workloads (export / Images pipeline / chat) before the trainer loads its own
# pipeline. Offload the blocking teardown (engine unload waits on generation locks; export
# subprocess join can take seconds) to a worker thread so the event loop stays responsive.
await asyncio.to_thread(_free_gpu_for_diffusion_training)
job_id = service.start(config)
except ValueError as e:
raise HTTPException(status_code = 400, detail = str(e))
except RuntimeError as e:
# A job is already running (or a start is already reserved), or a dataset mutation is open
# (DatasetMutationInFlight) -- the same interlock from the other side, so also a 409.
raise HTTPException(status_code = 409, detail = str(e))
except HTTPException:
raise
except Exception as e:
raise log_and_http_error(
e,
500,
"Failed to start diffusion training",
event = "diffusion_training.start_failed",
log = logger,
)
finally:
# On success the now-live proc keeps is_active() true; on failure this clears the reservation so
# training isn't left permanently "active". Only the request that reserved clears it.
if reserved:
service.unreserve()
return DiffusionTrainingStartResponse(job_id = job_id, status = "running")
@router.post("/diffusion/stop")
async def stop_diffusion_training(
body: Optional[DiffusionTrainingStopRequest] = None,
current_subject: str = Depends(get_current_subject),
):
"""Request a clean stop of the running diffusion training job. The optional body's
``save`` mirrors the LLM /stop: true (default, also for an empty POST) exports the
partial adapter, false cancels without saving one."""
from core.training.diffusion_training_service import get_diffusion_training_service
save = body.save if body is not None else True
stopped = get_diffusion_training_service().stop(save = save)
return {"status": "stopping" if stopped else "idle"}
@router.get("/diffusion/status", response_model = DiffusionTrainingStatusResponse)
async def diffusion_training_status(current_subject: str = Depends(get_current_subject)):
"""Poll the current diffusion training job's status/progress (JSON)."""
from core.training.diffusion_training_service import get_diffusion_training_service
snap = get_diffusion_training_service().status()
# Fold the service's flat history arrays into the nested metric_history the UI charts.
metric_history = DiffusionMetricHistory(
steps = snap.pop("metric_steps", []),
loss = snap.pop("metric_loss", []),
lr = snap.pop("metric_lr", []),
grad_norm = snap.pop("metric_grad_norm", []),
)
return DiffusionTrainingStatusResponse(**snap, metric_history = metric_history)
@router.get("/diffusion/runs", response_model = DiffusionTrainingRunsResponse)
async def list_diffusion_training_runs(
limit: int = 20, current_subject: str = Depends(get_current_subject)
):
"""Previous diffusion training runs (terminal), newest first, from the persisted
per-run records. Summaries only; fetch one run for its config + metric logs."""
from core.training.diffusion_training_service import list_diffusion_runs
summaries: list[DiffusionTrainingRunSummary] = []
for r in list_diffusion_runs(limit = limit):
# list_diffusion_runs already skips non-dict / missing-id records, but a wrong-typed field would
# still raise here; catch it per record so one bad file never breaks the whole Previous runs panel.
try:
summaries.append(DiffusionTrainingRunSummary(**r))
except ValidationError:
continue
return DiffusionTrainingRunsResponse(runs = summaries)
@router.get("/diffusion/runs/{job_id}", response_model = DiffusionTrainingRunDetail)
async def get_diffusion_training_run(
job_id: str, current_subject: str = Depends(get_current_subject)
):
"""One persisted diffusion run's full record: summary + scrubbed start config + the
step/loss/grad-norm logs (for re-plotting a past run's charts)."""
from core.training.diffusion_training_service import get_diffusion_run
rec = get_diffusion_run(job_id)
# A valid-JSON file that is not an object (a truncated / hand-edited [] record) makes
# DiffusionTrainingRunDetail(**rec) raise TypeError -- not the ValidationError caught below -- and
# 500 the endpoint. Treat any non-dict record as absent, like the list route.
if not isinstance(rec, dict):
raise HTTPException(status_code = 404, detail = "No such training run.")
try:
return DiffusionTrainingRunDetail(**rec)
except ValidationError:
# A malformed on-disk record (hand-edited / older shape) reads as absent rather than 500 the
# endpoint, like the list route skips bad records.
raise HTTPException(status_code = 404, detail = "No such training run.")
# 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 _resolve_dataset_caption(
folder: Path, image_path: Path, meta_captions: dict[str, str]
) -> Optional[str]:
"""Resolve an image's caption using the same sidecar > metadata precedence the trainer
applies in ``discover_image_caption_pairs``. A per-image .txt/.caption sidecar wins and
is stripped, so an empty (tombstone) sidecar shadows metadata and yields "" -- the
trainer then skips that image (``if caption:``), so it must not count as captioned."""
caption: Optional[str] = None
for ext in (".txt", ".caption"):
sidecar = image_path.with_suffix(ext)
if sidecar.is_file():
try:
caption = sidecar.read_text(encoding = "utf-8").strip()
except (OSError, UnicodeError):
# Unreadable/invalid UTF-8 sidecar: no caption, not a 500.
caption = None
break
if caption is None:
try:
rel = image_path.relative_to(folder).as_posix()
except ValueError:
rel = None
caption = meta_captions.get(image_path.name) or (
meta_captions.get(rel) if rel is not None else None
)
return caption
def _diffusion_dataset_summary(folder: Path) -> DiffusionDatasetSummary:
# Count an image as captioned only when it resolves to a NON-EMPTY caption via the same sidecar
# over metadata precedence the trainer uses: an empty tombstone sidecar shadows a metadata row and
# makes the trainer skip the image, so counting it would mislabel an uncaptioned dataset.
meta_captions = _load_metadata_captions(folder)
images = captions = 0
for f in folder.iterdir():
if not f.is_file() or f.suffix.lower() not in _DIFFUSION_DATASET_IMAGE_EXTS:
continue
images += 1
if _resolve_dataset_caption(folder, f, meta_captions):
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:
# Skip hidden dirs: never user datasets, and an in-progress example import stages into a
# dot-prefixed sibling that must not surface as a dataset.
children = sorted(
p
for p in root.iterdir()
# Skip symlinked dirs: the CRUD resolver rejects them, so discovery must not advertise one as
# selectable (the read/caption/delete routes would refuse it).
if p.is_dir() and not p.is_symlink() and not p.name.startswith(".")
)
except OSError:
children = []
for child in children:
try:
summary = _diffusion_dataset_summary(child)
except OSError:
continue
if summary.image_count > 0:
found.append(summary)
from core.training.diffusion_train_common import family_train_infos
families = [DiffusionTrainableFamily(**info) for info in family_train_infos()]
return DiffusionTrainingInfoResponse(
datasets_root = str(root),
outputs_root = str(outputs_root()),
datasets = found,
families = families,
)
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),
_interlock: None = Depends(diffusion_dataset_interlock),
):
"""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."""
import os
import tempfile
from utils.upload_limits import get_upload_limit_bytes, get_upload_limit_label
_require_diffusion_dataset_mutable()
cleaned = _clean_diffusion_dataset_name(name)
# Run the same symlink + root-containment check as the read/caption/delete endpoints before any
# write, so a name to external-directory symlink can't make the staged upload write outside root.
folder = _resolve_dataset_folder(name, must_exist = False)
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
# Validate every filename up front so a valid image ahead of a bad one isn't left on disk when the
# 400 fires; the upload is all-or-nothing.
names: list[str] = []
for f in files:
# Normalise to a safe basename. Path.name doesn't split on a backslash on POSIX, so a Windows
# client sending a backslash path in the multipart filename would be stored verbatim; fold
# backslashes first so the true basename is taken for both separators. The read/caption/delete
# endpoints run the stored name through _safe_dataset_image_path, so a name still holding ".."
# here would list an image the grid can never preview, caption, or delete.
filename = Path((f.filename or "").replace("\\", "/")).name.strip().replace("\x00", "")
ext = Path(filename).suffix.lower()
if not filename or ".." in filename or ext not in allowed:
exts = ", ".join(sorted(allowed))
raise HTTPException(
status_code = 400,
detail = f"Unsupported file '{f.filename}'. Allowed: {exts}",
)
# Reject an EXACT duplicate name within THIS batch (two cat.png from different folders, or an API
# client repeating a part). The same-name exemption below is for SEPARATE repeat uploads, a
# deliberate overwrite; inside one batch the two parts are distinct files staged to the same
# destination on EVERY filesystem, so the later replace would silently discard the earlier one.
# Exact match only: a case VARIANT pair stays exempt per the stem guard.
fname_cf = filename.casefold()
if filename in names:
raise HTTPException(
status_code = 400,
detail = (
f"Duplicate file '{filename}' appears more than once in this upload. "
"Files sharing a name would overwrite each other; rename one before "
"uploading."
),
)
# Reject a second IMAGE sharing this stem but differing by extension (sample.png vs sample.jpg):
# both resolve to the same <stem>.txt sidecar (the kohya/diffusers convention the reader, editor
# and delete paths use), so keeping both would silently share -- and corrupt -- one caption. Check
# files already on disk and earlier images in THIS batch. Re-uploading the exact same name stays
# an overwrite; caption/text files are exempt.
if ext in _DIFFUSION_DATASET_IMAGE_EXTS:
stem = Path(filename).stem
# Compare stems (and the same-name guard) case-insensitively: on case-insensitive filesystems two
# images whose stems differ only by case resolve to the SAME <stem>.txt sidecar, so a
# case-sensitive check would let both corrupt one caption. A same-name case variant is exempt ONLY
# when its stem also differs in case (one file on case-insensitive filesystems, separate sidecars
# on Linux). An EXTENSION-case variant (cat.PNG vs cat.png) has equal stems, so it is rejected.
stem_cf = stem.casefold()
def _shares_sidecar(other_name: str) -> bool:
other = Path(other_name)
if (
other_name == filename
or other.suffix.lower() not in _DIFFUSION_DATASET_IMAGE_EXTS
or other.stem.casefold() != stem_cf
):
return False
# A casefold-equal full name is exempt unless the stems match EXACTLY (extension-case variants
# collide on one sidecar on case-sensitive filesystems).
return other.stem == stem or other_name.casefold() != fname_cf
clash = next(
(p.name for p in folder.iterdir() if p.is_file() and _shares_sidecar(p.name)),
None,
)
if clash is None:
clash = next((n for n in names if _shares_sidecar(n)), None)
if clash is not None:
raise HTTPException(
status_code = 400,
detail = (
f"Duplicate image name '{stem}'. '{clash}' is already in this "
f"dataset; two images sharing a name would share one '{stem}.txt' "
f"caption. Rename one before uploading."
),
)
names.append(filename)
# Stage each file to a temp name and move it into place only once the whole batch is written, so a
# mid-batch failure (size limit, disk error, disconnect) leaves the dataset untouched, including
# any pre-existing same-name file a direct write would have truncated.
staged: list[tuple[Path, Path]] = [] # (temp, final)
committed = False
try:
for f, filename in zip(files, names):
dest = folder / filename
# A filename-independent temp name so a long (but valid) filename can't overflow NAME_MAX once the
# staging suffix is added.
tmp = folder / f".upload-{_uuid.uuid4().hex}.part"
staged.append((tmp, dest))
with open(tmp, "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)
# Reject a decompression bomb before commit: a small compressible PNG can pass the byte limit yet
# decode to huge pixels and OOM the trainer's latent cache, so bound each image's dimensions from
# the header (mirrors diffusion._decode_b64_image).
if Path(filename).suffix.lower() in _DIFFUSION_DATASET_IMAGE_EXTS:
_validate_uploaded_training_image(tmp, filename)
uploaded += 1
# Re-check the interlock immediately before the commit: the entry guard only saw the pre-upload
# state, so a /diffusion/start could have reserved the training slot while we were streaming.
# Committing now would move images/captions underneath the trainer; a 409 here leaves the staged
# temps to the finally below.
_require_diffusion_dataset_mutable()
# Commit every staged file as one transaction. A plain replace loop is not atomic across files: a
# mid-loop failure leaves earlier destinations already overwritten while the request errors. Back
# up each pre-existing destination first, then on any failure drop the versions this request
# installed and restore every displaced original.
backups: list[tuple[Path, Optional[Path]]] = [] # (dest, backup path or None)
installed: list[Path] = []
try:
for tmp, dest in staged:
backup: Optional[Path] = None
if dest.exists():
backup = folder / f".upload-backup-{_uuid.uuid4().hex}.part"
dest.replace(backup)
backups.append((dest, backup))
tmp.replace(dest) # atomic on the same filesystem
installed.append(dest)
committed = True
except BaseException:
# Roll back: drop every new version, then restore every displaced original.
for dest in reversed(installed):
try:
dest.unlink(missing_ok = True)
except OSError:
pass
for dest, backup in reversed(backups):
if backup is not None and backup.exists():
try:
backup.replace(dest)
except OSError:
pass
raise
else:
for _, backup in backups:
if backup is not None:
try:
backup.unlink(missing_ok = True)
except OSError:
pass
finally:
if not committed:
for tmp, _ in staged:
try:
tmp.unlink(missing_ok = True)
except OSError:
pass
summary = _diffusion_dataset_summary(folder)
return DiffusionDatasetUploadResponse(
name = cleaned,
path = str(folder),
image_count = summary.image_count,
caption_count = summary.caption_count,
uploaded = uploaded,
)
# ── Dataset labeling (per-image caption editing) + one-click example imports ──
# Thumbnails live in a hidden subdir so they never appear in dataset listings or the trainer's
# image discovery (both scan only top-level files).
_THUMBS_DIRNAME = ".thumbs"
_MAX_CAPTION_CHARS = 2000
def _resolve_dataset_folder(name: str, *, must_exist: bool = True) -> Path:
"""Validate ``name`` (single component, no traversal) and resolve it under the Studio
datasets root. 404 when a read target is missing."""
from utils.paths import datasets_root
cleaned = _clean_diffusion_dataset_name(name)
root = datasets_root().resolve()
folder = root / cleaned
# Reject a symlinked dataset directory and prove the resolved folder stays under root:
# _safe_dataset_image_path only checks each image path, so a folder symlinked to an external
# directory would let read / caption / delete operate on files outside Studio.
if folder.is_symlink():
raise HTTPException(
status_code = 400,
detail = f"Dataset '{cleaned}' must not be a symbolic link.",
)
if must_exist and not folder.is_dir():
raise HTTPException(status_code = 404, detail = f"Dataset '{cleaned}' not found.")
try:
folder.resolve(strict = must_exist).relative_to(root)
except (OSError, ValueError):
raise HTTPException(
status_code = 400,
detail = f"Dataset '{cleaned}' escapes the Studio datasets directory.",
)
return folder
# Per-side dimension bound for uploaded training images, matching diffusion._decode_b64_image's
# 4096px inference guard, so a compressible PNG can't smuggle huge pixels past the byte limit.
_MAX_TRAINING_IMAGE_SIDE = 4096
def _validate_uploaded_training_image(path: Path, original_name: str) -> None:
"""Reject an uploaded training image whose decoded dimensions exceed the per-side limit.
Reads only the header (never img.load()), so a small-payload / huge-dimension file is caught
before it spikes memory. Bytes PIL cannot identify are left as-is (the upload contract accepts
arbitrary bytes under an image extension), so only oversized real images change behaviour."""
from PIL import Image, UnidentifiedImageError
try:
with Image.open(path) as image:
width, height = image.size
except Image.DecompressionBombError:
# Past Pillow's own hard limit (~179 MP) Image.open() raises before .size can be read. That error
# derives straight from Exception (not OSError/ValueError), so letting it escape 500s the upload;
# it is exactly the oversized image this guard rejects.
raise HTTPException(
status_code = 400,
detail = (
f"Image '{original_name}' is too large; maximum is "
f"{_MAX_TRAINING_IMAGE_SIDE}px per side."
),
)
except (OSError, UnidentifiedImageError, ValueError):
return # not a decodable image -> not a bomb; leave the existing contract
if width > _MAX_TRAINING_IMAGE_SIDE or height > _MAX_TRAINING_IMAGE_SIDE:
raise HTTPException(
status_code = 400,
detail = (
f"Image '{original_name}' is too large ({width}x{height}); maximum is "
f"{_MAX_TRAINING_IMAGE_SIDE}px per side."
),
)
def _safe_dataset_image_path(folder: Path, filename: str) -> Path:
"""Resolve ``filename`` to an image path strictly inside ``folder``. Rejects any path
separators / traversal / null bytes and non-image extensions."""
raw = filename or ""
if "/" in raw or "\\" in raw or ".." in raw or "\x00" in raw or raw != Path(raw).name:
raise HTTPException(status_code = 400, detail = "Invalid image filename.")
if Path(raw).suffix.lower() not in _DIFFUSION_DATASET_IMAGE_EXTS:
exts = ", ".join(sorted(_DIFFUSION_DATASET_IMAGE_EXTS))
raise HTTPException(status_code = 400, detail = f"Not an image file. Allowed: {exts}")
path = folder / raw
# Defense in depth: the real path must stay under the dataset folder.
try:
path.resolve().relative_to(folder.resolve())
except ValueError:
raise HTTPException(status_code = 400, detail = "Invalid image filename.")
return path
def _load_metadata_captions(folder: Path) -> dict[str, str]:
"""Read metadata.jsonl / captions.jsonl into {file_name: caption}, mirroring the
trainer's discovery (keys file_name/image/file; caption in the ``text`` column)."""
import json
out: dict[str, str] = {}
for meta_name in ("metadata.jsonl", "captions.jsonl"):
meta_path = folder / meta_name
if not meta_path.is_file():
continue
# Tolerate a bad upload (invalid UTF-8, or a line of non-object JSON): skip the record so the
# info / labeling / caption / summary endpoints don't 500.
try:
lines = meta_path.read_text(encoding = "utf-8").splitlines()
except (OSError, UnicodeError):
continue
for line in lines:
line = line.strip()
if not line:
continue
try:
row = json.loads(line)
except (json.JSONDecodeError, TypeError):
continue
if not isinstance(row, dict):
continue
key = row.get("file_name") or row.get("image") or row.get("file")
value = row.get("text")
# A JSON null is "no caption", not the string "None".
if key and value is not None:
out[str(key)] = str(value)
return out
def _image_record(
folder: Path, image_path: Path, meta_captions: dict[str, str]
) -> DiffusionDatasetImageRecord:
"""Build one image record, resolving its caption with sidecar > metadata precedence
(the same order the trainer uses). A per-image .txt / .caption sidecar wins because
it is the user's explicit edit from the labeling grid, which must override a
metadata.jsonl / captions.jsonl row for the image."""
caption: Optional[str] = None
source = "none"
for ext in (".txt", ".caption"):
sidecar = image_path.with_suffix(ext)
if sidecar.is_file():
try:
caption = sidecar.read_text(encoding = "utf-8").strip()
source = "sidecar"
except (OSError, UnicodeError):
# Unreadable / invalid UTF-8 sidecar (uploads store text sidecars as raw bytes):
# UnicodeDecodeError is a ValueError, not an OSError, so an OSError-only guard let it 500 the
# whole labeling grid. Read it as no caption, like the info summary does.
caption = None
break
if caption is None:
# Basename first, then the relative path as written in the jsonl (as_posix so a Windows backslash
# path still matches forward-slash keys): discover_image_caption_pairs's order.
meta = meta_captions.get(image_path.name)
if meta is None:
try:
meta = meta_captions.get(image_path.relative_to(folder).as_posix())
except ValueError:
meta = None
if meta is not None:
caption = meta
source = "metadata"
try:
size_bytes = image_path.stat().st_size
except OSError:
size_bytes = 0
width = height = 0
try:
from PIL import Image
with Image.open(image_path) as im:
width, height = im.size
except Exception: # noqa: BLE001 -- an unreadable image still lists (0x0) rather than 500
pass
return DiffusionDatasetImageRecord(
filename = image_path.name,
caption = caption,
caption_source = source, # type: ignore[arg-type]
width = width,
height = height,
size_bytes = size_bytes,
)
@router.get("/diffusion/dataset/{name}/images", response_model = DiffusionDatasetImagesResponse)
async def list_diffusion_dataset_images(
name: str, current_subject: str = Depends(get_current_subject)
):
"""List every image in a dataset folder with its resolved caption (including
uncaptioned images), for the labeling grid."""
folder = _resolve_dataset_folder(name)
def scan() -> DiffusionDatasetImagesResponse:
meta = _load_metadata_captions(folder)
records: list[DiffusionDatasetImageRecord] = []
for p in sorted(folder.iterdir()):
if p.is_file() and p.suffix.lower() in _DIFFUSION_DATASET_IMAGE_EXTS:
records.append(_image_record(folder, p, meta))
return DiffusionDatasetImagesResponse(name = folder.name, path = str(folder), images = records)
return await asyncio.to_thread(scan)
@router.get("/diffusion/dataset/{name}/image/{filename}")
async def get_diffusion_dataset_image(
name: str,
filename: str,
thumb: Optional[int] = None,
current_subject: str = Depends(get_current_subject),
):
"""Serve a dataset image. ``?thumb=<px>`` returns a cached downscaled JPEG (regenerated
when the source is newer), used by the labeling grid to stay light."""
from fastapi.responses import FileResponse
folder = _resolve_dataset_folder(name)
image_path = _safe_dataset_image_path(folder, filename)
if not image_path.is_file():
raise HTTPException(status_code = 404, detail = "Image not found.")
if not thumb:
return FileResponse(str(image_path))
size = max(32, min(1024, int(thumb)))
def make_thumb() -> Path:
from PIL import Image
thumbs_dir = folder / _THUMBS_DIRNAME
thumbs_dir.mkdir(exist_ok = True)
# Key on the full filename (stem + extension), not the stem: two images sharing a stem but
# differing by extension would otherwise collide on one cache file, and an mtime-newer cache for
# the first would be served for the second.
thumb_path = thumbs_dir / f"{image_path.name}_{size}.jpg"
src_mtime = image_path.stat().st_mtime
if thumb_path.is_file() and thumb_path.stat().st_mtime >= src_mtime:
return thumb_path
with Image.open(image_path) as im:
im = im.convert("RGB")
im.thumbnail((size, size), Image.LANCZOS)
im.save(thumb_path, format = "JPEG", quality = 85)
return thumb_path
try:
thumb_path = await asyncio.to_thread(make_thumb)
except Exception as e: # noqa: BLE001 -- fall back to the original on any decode failure
logger.warning("Thumbnail generation failed for %s: %s", image_path, e)
return FileResponse(str(image_path))
return FileResponse(str(thumb_path), media_type = "image/jpeg")
@router.put(
"/diffusion/dataset/{name}/caption/{filename}",
response_model = DiffusionDatasetImageRecord,
)
async def set_diffusion_dataset_caption(
name: str,
filename: str,
body: DiffusionCaptionUpdateRequest,
current_subject: str = Depends(get_current_subject),
_interlock: None = Depends(diffusion_dataset_interlock),
):
"""Write (or, when blank, clear) an image's ``.txt`` caption sidecar. Returns the
updated image record."""
_require_diffusion_dataset_mutable()
folder = _resolve_dataset_folder(name)
image_path = _safe_dataset_image_path(folder, filename)
if not image_path.is_file():
raise HTTPException(status_code = 404, detail = "Image not found.")
caption = (body.caption or "").strip()
if len(caption) > _MAX_CAPTION_CHARS:
raise HTTPException(
status_code = 400,
detail = f"Caption too long (max {_MAX_CAPTION_CHARS} characters).",
)
def write() -> DiffusionDatasetImageRecord:
sidecar = image_path.with_suffix(".txt")
if caption:
sidecar.write_text(caption, encoding = "utf-8")
image_path.with_suffix(".caption").unlink(missing_ok = True)
return _image_record(folder, image_path, _load_metadata_captions(folder))
# Blank must actually clear. Unlinking alone would resurface this image's metadata.jsonl /
# captions.jsonl caption, so when one exists write an EMPTY sidecar instead: both the reader and
# the trainer's discovery treat an existing sidecar as authoritative even when empty, a tombstone.
# With no metadata caption it is a plain cleanup.
meta = _load_metadata_captions(folder)
try:
rel = image_path.relative_to(folder).as_posix()
except ValueError:
rel = image_path.name
if image_path.name in meta or rel in meta:
sidecar.write_text("", encoding = "utf-8")
else:
sidecar.unlink(missing_ok = True)
image_path.with_suffix(".caption").unlink(missing_ok = True)
return _image_record(folder, image_path, meta)
return await asyncio.to_thread(write)
@router.delete("/diffusion/dataset/{name}/image/{filename}")
async def delete_diffusion_dataset_image(
name: str,
filename: str,
current_subject: str = Depends(get_current_subject),
_interlock: None = Depends(diffusion_dataset_interlock),
):
"""Remove an image, its caption sidecars, and any cached thumbnails."""
_require_diffusion_dataset_mutable()
folder = _resolve_dataset_folder(name)
image_path = _safe_dataset_image_path(folder, filename)
if not image_path.is_file():
raise HTTPException(status_code = 404, detail = "Image not found.")
def remove() -> dict:
import glob as _glob
image_path.unlink(missing_ok = True)
for ext in (".txt", ".caption"):
image_path.with_suffix(ext).unlink(missing_ok = True)
thumbs_dir = folder / _THUMBS_DIRNAME
if thumbs_dir.is_dir():
# Thumbs are keyed on the full filename (stem + extension), so match that here too; a stem-only
# glob would strand this image's thumbs or delete a same-stem sibling's. Escape the name: a raw
# glob metacharacter would match siblings' thumbs while leaving its own behind.
for t in thumbs_dir.glob(f"{_glob.escape(image_path.name)}_*.jpg"):
t.unlink(missing_ok = True)
return {"deleted": image_path.name}
return await asyncio.to_thread(remove)
# Curated, license-labelled example datasets for one-click import. ``loader`` picks the
# materialization strategy: "hf_dataset" streams rows from datasets.load_dataset (image + optional
# caption column); "imagefolder_jsonl" snapshot-downloads a dataset repo whose captions live in a
# *.jsonl (file_name/text) not a standard metadata.jsonl.
_DATASET_EXAMPLES: list[dict] = [
{
"id": "dreambooth-dog",
"label": "Dog (DreamBooth subject)",
"repo": "diffusers/dog-example",
"description": "5 photos of one dog. Teach a subject, then call it with the trigger.",
"license": "Google, research and demos",
"image_cap": 10,
"suggested_trigger": "a photo of sks dog",
"loader": "hf_dataset",
"caption_column": None,
"no_checks": False,
},
{
"id": "tuxemon",
"label": "Tuxemon (captioned style set)",
"repo": "linoyts/Tuxemon",
"description": "Captioned cartoon monster art. A style set, no trigger needed.",
"license": "cc-by-sa-3.0",
"image_cap": 60,
"suggested_trigger": None,
"loader": "hf_dataset",
"caption_column": "prompt",
"no_checks": True,
},
{
"id": "tarot-1920",
"label": "1920 Tarot (public domain style set)",
"repo": "multimodalart/1920-raider-waite-tarot-public-domain",
"description": "Captioned 1920 Raider-Waite tarot art. A permissive style set.",
"license": "public domain",
"image_cap": 60,
"suggested_trigger": None,
"loader": "imagefolder_jsonl",
"caption_column": "text",
"no_checks": True,
},
{
"id": "smithsonian-butterflies",
"label": "Smithsonian Butterflies",
"repo": "huggan/smithsonian_butterflies_subset",
"description": "100 butterfly photos. No captions, so use the trigger prompt.",
"license": "CC0",
"image_cap": 100,
# The metadata columns are species names / boilerplate alt-text, not captions, so train it as a
# subject set with the trigger prompt instead.
"suggested_trigger": "a photo of a sks butterfly",
"loader": "hf_dataset",
"caption_column": None,
"no_checks": False,
},
{
"id": "pixel-nouns",
"label": "Nouns (pixel avatars)",
"repo": "m1guelpf/nouns",
"description": "100 captioned pixel-art avatars. A style set, no trigger needed.",
"license": "cc0-1.0",
"image_cap": 100,
"suggested_trigger": None,
"loader": "hf_dataset",
"caption_column": "text",
"no_checks": False,
},
]
def _example_by_id(example_id: str) -> dict:
for entry in _DATASET_EXAMPLES:
if entry["id"] == example_id:
return entry
raise HTTPException(status_code = 404, detail = f"Unknown example dataset '{example_id}'.")
@router.get("/diffusion/dataset-examples", response_model = DiffusionDatasetExamplesResponse)
async def list_diffusion_dataset_examples(current_subject: str = Depends(get_current_subject)):
"""List the curated example datasets available for one-click import."""
return DiffusionDatasetExamplesResponse(
examples = [
DiffusionDatasetExample(
id = e["id"],
label = e["label"],
repo = e["repo"],
description = e["description"],
license = e["license"],
image_cap = e["image_cap"],
suggested_trigger = e["suggested_trigger"],
)
for e in _DATASET_EXAMPLES
]
)
def _detect_image_column(features) -> Optional[str]:
"""Return the first datasets Image-feature column name, else None."""
try:
from datasets import Image as HFImage
except Exception: # noqa: BLE001
HFImage = None # type: ignore[assignment]
for col, feat in features.items():
if HFImage is not None and isinstance(feat, HFImage):
return col
if type(feat).__name__ == "Image":
return col
return None
def _detect_caption_column(entry: dict, columns: list[str]) -> Optional[str]:
"""Pick the caption column: the entry's declared one if present, else a common name."""
declared = entry.get("caption_column")
if declared and declared in columns:
return declared
for cand in ("text", "prompt", "caption", "captions"):
if cand in columns:
return cand
return None
def _materialize_hf_dataset(entry: dict, dest: Path, cap: int) -> int:
"""Stream rows from datasets.load_dataset into ``dest`` as numbered images + optional
.txt sidecars. Returns the number of images written."""
from datasets import load_dataset
kwargs = {"split": "train"}
if entry.get("no_checks"):
kwargs["verification_mode"] = "no_checks"
ds = load_dataset(entry["repo"], **kwargs)
image_col = _detect_image_column(ds.features)
if image_col is None:
raise HTTPException(
status_code = 502,
detail = f"'{entry['repo']}' has no image column to import.",
)
caption_col = _detect_caption_column(entry, list(ds.features.keys()))
written = 0
for row in ds:
if written >= cap:
break
img = row[image_col]
if img is None:
continue
img = img.convert("RGB")
stem = f"img_{written:04d}"
img.save(dest / f"{stem}.png", format = "PNG")
if caption_col:
cap_text = row.get(caption_col)
if cap_text:
(dest / f"{stem}.txt").write_text(str(cap_text).strip(), encoding = "utf-8")
written += 1
return written
def _materialize_imagefolder_jsonl(entry: dict, dest: Path, cap: int) -> int:
"""Snapshot-download a dataset repo whose captions live in *.jsonl (file_name/text),
then copy referenced images + write .txt sidecars. Returns images written."""
import json
import shutil
from huggingface_hub import snapshot_download
caption_col = entry.get("caption_column") or "text"
snap = Path(
snapshot_download(
entry["repo"],
repo_type = "dataset",
allow_patterns = [
"*.jsonl",
"*.jpg",
"*.jpeg",
"*.png",
"*.webp",
"*.bmp",
"**/*.jpg",
"**/*.jpeg",
"**/*.png",
"**/*.webp",
"**/*.bmp",
],
)
)
# Map basename -> caption from every jsonl carrying file_name + caption column.
captions: dict[str, str] = {}
for jf in sorted(snap.rglob("*.jsonl")):
for line in jf.read_text(encoding = "utf-8").splitlines():
line = line.strip()
if not line:
continue
try:
row = json.loads(line)
except json.JSONDecodeError:
continue
fn = row.get("file_name") or row.get("image") or row.get("file")
value = row.get(caption_col)
# A JSON null is "no caption", not the string "None".
if fn and value is not None:
# First writer wins over sorted manifests, for deterministic results.
captions.setdefault(Path(str(fn)).name, str(value))
# Copy images (those with a caption first, so a cap keeps captioned pairs).
images = sorted(
p
for p in snap.rglob("*")
if p.is_file() and p.suffix.lower() in _DIFFUSION_DATASET_IMAGE_EXTS
)
images.sort(key = lambda p: (p.name not in captions, p.name))
written = 0
for src in images:
if written >= cap:
break
stem = f"img_{written:04d}"
shutil.copyfile(src, dest / f"{stem}{src.suffix.lower()}")
cap_text = captions.get(src.name)
if cap_text:
(dest / f"{stem}.txt").write_text(cap_text.strip(), encoding = "utf-8")
written += 1
return written
@router.post("/diffusion/dataset/import-example", response_model = DiffusionDatasetImportResponse)
async def import_diffusion_dataset_example(
body: DiffusionDatasetImportRequest,
current_subject: str = Depends(get_current_subject),
_interlock: None = Depends(diffusion_dataset_interlock),
):
"""Materialize a curated example dataset into a Studio dataset folder (images + .txt
captions), ready to train. Idempotent: a folder that already holds images is returned
as-is rather than re-downloaded."""
_require_diffusion_dataset_mutable()
entry = _example_by_id(body.id)
folder = _resolve_dataset_folder(body.name or entry["id"], must_exist = False)
def do_import() -> DiffusionDatasetImportResponse:
import os
import shutil
import tempfile
folder.mkdir(parents = True, exist_ok = True)
existing = _diffusion_dataset_summary(folder)
imported = 0
if existing.image_count == 0:
cap = int(entry["image_cap"])
# Materialize into a private staging dir and promote into the dataset folder only after the whole
# import succeeds. A partial materialize then leaves only the staging dir, never a half-filled
# dataset -- otherwise the image_count>0 idempotency check above would treat that partial as
# complete on retry and strand a truncated dataset (there is no dataset-delete flow). Staged as a
# hidden same-filesystem sibling so promotion is an atomic rename.
staging = Path(tempfile.mkdtemp(dir = folder.parent, prefix = f".{folder.name}.import-"))
try:
try:
if entry["loader"] == "imagefolder_jsonl":
imported = _materialize_imagefolder_jsonl(entry, staging, cap)
else:
imported = _materialize_hf_dataset(entry, staging, cap)
except HTTPException:
raise
except Exception as e: # noqa: BLE001 -- surface a readable fetch/parse failure
raise HTTPException(
status_code = 502,
detail = f"Could not import '{entry['repo']}': {e}",
)
if imported == 0:
raise HTTPException(
status_code = 502,
detail = f"No images found in '{entry['repo']}'.",
)
# Promote the fully-materialized staging dir as a UNIT. A per-file move loop is not atomic: a hard
# process death mid-loop would leave SOME images, which the image_count>0 idempotency check would
# accept as complete on retry. The folder was created empty here, so a single same-filesystem
# rename is atomic. If it holds unrelated non-image files (rmdir refuses), fall back to a per-file
# move rather than abort.
try:
os.rmdir(folder)
except OSError:
for p in staging.iterdir():
shutil.move(str(p), str(folder / p.name))
else:
os.replace(str(staging), str(folder))
finally:
shutil.rmtree(staging, ignore_errors = True)
summary = _diffusion_dataset_summary(folder)
return DiffusionDatasetImportResponse(
name = folder.name,
path = str(folder),
image_count = summary.image_count,
caption_count = summary.caption_count,
imported = imported,
license = entry["license"],
source_repo = entry["repo"],
)
return await asyncio.to_thread(do_import)