Merge branch 'diffusion-krea2' into diffusion-train-perf2

This commit is contained in:
Daniel Han 2026-07-05 02:11:26 +00:00
commit 551c38bd4a
19 changed files with 262 additions and 27 deletions

View file

@ -59,18 +59,20 @@ from .diffusion_speed import (
SPEED_OFF,
apply_speed_optims,
compile_eligible,
normalize_speed_mode,
resolve_speed_mode,
restore_backend_flags,
snapshot_backend_flags,
)
from .diffusion_attention import (
apply_attention_backend,
normalize_attention_backend,
select_attention_backend,
)
from . import diffusion_compile_cache as compile_cache
from . import diffusion_gguf_compile as gguf_compile
from .diffusion_cache import apply_step_cache
from .diffusion_precision import quantize_text_encoders
from .diffusion_cache import apply_step_cache, normalize_transformer_cache
from .diffusion_precision import normalize_te_quant, quantize_text_encoders
from .diffusion_prequant import (
load_prequantized_transformer,
resolve_prequant_source,
@ -698,6 +700,15 @@ class DiffusionBackend:
if self._load_token != token:
return
logger.error("diffusion.load_failed: %s", exc)
# Free the debris of a failed construction (e.g. a load-time OOM): _state was
# never committed, and the next load's _unload_locked early-returns on a None
# state, so nothing else releases the reserved VRAM. Guarded: a sticky CUDA
# error makes synchronize() raise, which would skip stamping the REAL error
# below and leave the client polling forever.
try:
clear_gpu_cache()
except Exception: # noqa: BLE001
pass
# Redact native paths: this error is surfaced verbatim via the
# load-progress poll, and Studio can run as a shared server.
from utils.native_path_leases import redact_native_paths
@ -903,6 +914,14 @@ class DiffusionBackend:
model_kind = model_kind,
)
kind = resolve_model_kind(gguf_filename, model_kind)
# Validate every mode string that can raise NOW, before this load evicts the
# previous pipeline below: their first in-line uses all sit past _unload_locked,
# where a bad request would cost the user their working model.
transformer_quant = normalize_transformer_quant(transformer_quant)
normalize_speed_mode(speed_mode)
normalize_attention_backend(attention_backend)
normalize_transformer_cache(transformer_cache)
normalize_te_quant(text_encoder_quant)
# For a full pipeline the repo itself supplies every component, so it is its
# own base; the single-file kinds resolve the companion base diffusers repo.
base = (
@ -975,7 +994,7 @@ class DiffusionBackend:
transformer_quant_engaged = None
if (
kind == "gguf"
and normalize_transformer_quant(transformer_quant) is not None
and transformer_quant is not None # normalized above, pre-eviction
and dense_transformer_supported(target)
and plan.offload_policy == OFFLOAD_NONE
):
@ -1005,7 +1024,12 @@ class DiffusionBackend:
# clear_gpu_cache() could not otherwise reclaim that VRAM before the
# GGUF build (the OOM-fallback path this cleanup exists for).
del exc
clear_gpu_cache()
# Guarded: after an OOM/sticky CUDA error synchronize() can
# raise, and this fallback path must still reach the GGUF build.
try:
clear_gpu_cache()
except Exception: # noqa: BLE001
pass
if pipe is None:
if kind == "pipeline":
@ -1989,7 +2013,8 @@ class DiffusionBackend:
gen = _GenState(total_steps = steps)
def _on_step(pipe, step_index, timestep, callback_kwargs):
now = time.time()
# Monotonic: a wall-clock adjustment (NTP) mid-denoise would skew the ETA.
now = time.monotonic()
gen.step = step_index + 1
if gen.first_step_at == 0.0:
gen.first_step_at = now
@ -2066,9 +2091,8 @@ class DiffusionBackend:
self._cancel_event.set()
with self._lock:
# Abort an in-flight denoise too by setting ITS cancel event, so the step
# callback stops it. unload does NOT take _generate_lock — it must return
# promptly; the running generate keeps its own pipe reference, so freeing
# _state here can't crash it, and its VRAM is reclaimed when it returns
# callback stops it. The running generate keeps its own pipe reference, so
# freeing _state here can't crash it; its VRAM is reclaimed when it exits
# (within ~one step thanks to the cancel).
if self._active_generate_cancel is not None:
self._active_generate_cancel.set()
@ -2077,6 +2101,14 @@ class DiffusionBackend:
# committing) and drop the marker so the next load starts clean.
self._load_token += 1
self._loading = None
# Wait for the signalled denoise to actually exit before reporting unloaded:
# callers treat this return as "VRAM is free" (the GPU arbiter hands the GPU
# to chat next; the training routes size their run against it), and the
# denoise holds its pipe until the next step callback. generate() holds
# _generate_lock for its full body, so a bare acquire is the exit barrier
# (never while holding _lock -- generate takes _lock inside _generate_lock).
with self._generate_lock:
pass
return self.status()
def _unload_locked(self) -> None:
@ -2101,7 +2133,7 @@ class DiffusionBackend:
uninstall_patches()
uninstall_arch_patches()
# NOTE: we deliberately do NOT call state.pipe.unload_lora_weights() here. unload()
# sets the cancel event but does not take _generate_lock, so a LoRA-backed denoise
# only acquires _generate_lock AFTER this teardown, so a LoRA-backed denoise
# can still be running on this same pipe for up to one more callback; mutating its
# adapter layers now would race that in-flight generation. The whole pipe is dropped
# just below (self._state = None; del state; clear_gpu_cache()), so the adapter

View file

@ -129,6 +129,11 @@ def select_attention_backend(
backend = _ALIASES[alias]
if backend == "native":
return None
# Every explicit kernel here (cuDNN / flash* / sage) is CUDA+NVIDIA-only; on
# ROCm / MPS / CPU diffusers accepts the name at set time and the first
# generation crashes, so drop to the native default up front.
if not _is_cuda_nvidia(target):
return None
# An arch-gated kernel (flash3/flash4) on a card that can't run it would set fine
# then crash mid-generation, so drop it to the native default up front.
if not _backend_arch_supported(backend):

View file

@ -89,7 +89,10 @@ def apply_step_cache(
_warn(logger, mode, RuntimeError("transformer has no cache_context (not a CacheMixin)"))
return None
try:
from diffusers import FirstBlockCacheConfig
try:
from diffusers import FirstBlockCacheConfig
except ImportError: # older diffusers exports it only from diffusers.hooks
from diffusers.hooks import FirstBlockCacheConfig
config = FirstBlockCacheConfig(threshold = thr)
enable_cache(config)
@ -101,6 +104,12 @@ def apply_step_cache(
logger.info("diffusion.cache: %s engaged (threshold=%s)", mode, thr)
return mode
except Exception as exc: # noqa: BLE001 — incompatible model -> run uncached
# enable_cache can fail after hooking some blocks; drop any partial hooks so
# the reported-uncached model doesn't actually run half-cached.
try:
transformer.disable_cache()
except Exception: # noqa: BLE001
pass
_warn(logger, mode, exc)
return None

View file

@ -452,10 +452,19 @@ def resolve_base_repo(fam: DiffusionFamily, base_repo: Optional[str]) -> str:
_GENERATION_DEFAULTS: tuple[tuple[str, int, float], ...] = (
("z-image-turbo", 9, 0.0),
("flux.1-schnell", 4, 0.0),
# Kontext (editing) before the generic flux.1: ~28 steps, lower guidance (~2.5).
("kontext", 28, 2.5),
("flux.1", 28, 3.5),
("flux.2-klein", 4, 0.0),
# FLUX.2-dev is the full (non-distilled) model: more steps + real guidance.
("flux.2-dev", 28, 4.0),
("qwen-image", 20, 4.0),
("z-image", 20, 4.0),
# SDXL: Turbo is distilled (few steps, no CFG); base/full SDXL wants ~30 steps and
# real CFG (~7). "sdxl-turbo" must precede the generic "sdxl" substring match.
("sdxl-turbo", 3, 0.0),
("stable-diffusion-xl", 30, 7.0),
("sdxl", 30, 7.0),
)
# Unrecognised model: distilled few-step / no-CFG shape, matching the UI fallback.
_GENERATION_DEFAULT_FALLBACK = (9, 0.0)

View file

@ -269,7 +269,12 @@ def build_sd_cpp_command(
if params.seed is not None:
cmd += ["--seed", str(int(params.seed))]
if params.batch_count and params.batch_count != 1:
cmd += ["--batch-count", str(int(params.batch_count))]
# sd-cli names the extra batch images itself (output_2.png, ...) and the runner
# collects only the literal --output path, so a CLI batch would silently drop
# every image after the first. Batches go through the sdcpp server API instead.
raise ValueError(
"sd-cli runs are single-image; use the sdcpp server API for batch generation."
)
cmd += ["--output", output_path]
if threads is not None:

View file

@ -241,6 +241,9 @@ class _SdLoading:
repo_id: str
base_repo: str
# Companion asset repos (VAE / text encoders) this load fetches, so the
# delete-cached guard protects them for the whole download/finalize window.
asset_repos: tuple[str, ...] = ()
expected_bytes: int = 0
downloaded_bytes: int = 0
error: Optional[str] = None
@ -407,7 +410,17 @@ class SdCppDiffusionBackend:
self._load_token += 1
token = self._load_token
self._cancel_event.clear()
self._loading = _SdLoading(repo_id = repo_id, base_repo = base)
self._loading = _SdLoading(
repo_id = repo_id,
base_repo = base,
asset_repos = tuple(
dict.fromkeys(
r
for r, _f, kind in self._asset_specs(repo_id, gguf_filename, fam)
if kind != "diffusion_model"
)
),
)
threading.Thread(
target = self._run_load,
@ -678,12 +691,16 @@ class SdCppDiffusionBackend:
def loading_repo_ids(self) -> tuple[str, ...]:
"""Repo ids an in-flight background load is downloading (empty when idle).
Mirrors the diffusers backend so the delete-cached guard can query whichever
engine is active without caring which one it got."""
engine is active without caring which one it got. Includes the companion
VAE / text-encoder repos: deleting one of those mid-load would remove files
the committed SdCppModelFiles paths need."""
with self._lock:
loading = self._loading
if loading is None or loading.error is not None:
return ()
return tuple(r for r in (loading.repo_id, loading.base_repo) if r)
return tuple(
r for r in (loading.repo_id, loading.base_repo, *loading.asset_repos) if r
)
# ── Generate ───────────────────────────────────────────────────────────

View file

@ -364,6 +364,9 @@ class SdCppEngine:
def _prepare_out(output_path: str) -> Path:
out = Path(output_path)
out.parent.mkdir(parents = True, exist_ok = True)
# Drop a stale file at the target so the post-run is_file() check proves THIS
# run produced the image, not a leftover from an earlier run at the same path.
out.unlink(missing_ok = True)
return out
def _run(

View file

@ -400,6 +400,10 @@ class DiffusionLoraConfig:
f"base_precision={base_precision!r} trains in bf16 compute; set "
f"mixed_precision to bf16."
)
# A zero/negative gamma would zero out (or invert) the min-SNR weight and
# silently train on a degenerate loss; None is the documented disable.
if self.snr_gamma is not None and float(self.snr_gamma) <= 0:
raise ValueError("snr_gamma must be > 0, or null to disable min-SNR weighting")
# learning_rate can arrive as a string ("1e-4") from the Studio config path, which
# preserves it as a string after validation; coerce so AdamW receives a float.
try:

View file

@ -713,10 +713,14 @@ class DiffusionTrainingStartRequest(BaseModel):
default_factory = lambda: ["to_k", "to_q", "to_v", "to_out.0"],
description = "U-Net modules to attach LoRA to",
)
max_grad_norm: float = Field(1.0, gt = 0, description = "Gradient clipping max-norm")
max_grad_norm: float = Field(
1.0, ge = 0, description = "Gradient clipping max-norm; 0 disables clipping"
)
seed: int = Field(42)
mixed_precision: Literal["bf16", "fp16", "no"] = Field("bf16")
snr_gamma: Optional[float] = Field(5.0, description = "Min-SNR loss weighting; null disables")
snr_gamma: Optional[float] = Field(
5.0, gt = 0, description = "Min-SNR loss weighting; null disables"
)
gradient_checkpointing: bool = Field(True)
lr_scheduler: str = Field("constant")
lr_warmup_steps: int = Field(0, ge = 0)

View file

@ -12076,6 +12076,21 @@ async def openai_image_generations(
# isn't loaded; the global handler turns this into the OpenAI envelope.
raise HTTPException(status_code = 503, detail = _NO_IMAGE_MODEL_MSG)
# An edit-only model (Qwen-Image-Edit, FLUX Kontext) needs an input image this API
# cannot supply; refuse up front with a 400 instead of letting the backend's
# ValueError surface as a sanitized 500.
workflows = status.get("workflows") or []
if workflows and "txt2img" not in workflows:
raise HTTPException(
status_code = 400,
detail = openai_error_body(
"The loaded image model is edit-only (it requires an input image); "
"load a text-to-image model to use this endpoint.",
status = 400,
param = "model",
),
)
# Fall back to the resolved base repo so a local-path load (whose repo_id is a
# filesystem path) still gets the right per-model steps/guidance.
steps, guidance = default_generation_params(status.get("repo_id"), status.get("base_repo"))

View file

@ -1224,6 +1224,16 @@ async def start_diffusion_training(
except ValueError as e:
raise HTTPException(status_code = 400, detail = str(e))
# 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
# instead of tearing down the user's chat/Images model and 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 loaded chat/Images model, and never surfaces as a confusing mid-load 401.

View file

@ -75,7 +75,7 @@ def test_auto_stays_native_off_nvidia(monkeypatch):
def test_explicit_backend_honored_regardless_of_speed(monkeypatch):
monkeypatch.setattr(att, "_is_cuda_nvidia", lambda target: False)
monkeypatch.setattr(att, "_is_cuda_nvidia", lambda target: True)
# Pin a high capability so the arch-gated flash4 isn't dropped by the runtime check.
monkeypatch.setattr(att, "_cuda_capability", lambda: (10, 0))
assert select_attention_backend(_target(), "sage", speed_active = False) == "sage"
@ -83,6 +83,15 @@ def test_explicit_backend_honored_regardless_of_speed(monkeypatch):
assert select_attention_backend(_target(), "cudnn", speed_active = False) == "_native_cudnn"
def test_explicit_backend_dropped_off_nvidia_cuda(monkeypatch):
# Explicit cuDNN/flash/sage on ROCm / MPS / CPU passes diffusers' set-time check
# and crashes at the first generation, so selection drops to the native default.
monkeypatch.setattr(att, "_is_cuda_nvidia", lambda target: False)
monkeypatch.setattr(att, "_cuda_capability", lambda: (10, 0))
for alias in ("sage", "flash", "flash4", "cudnn"):
assert select_attention_backend(_target(device = "mps"), alias, speed_active = True) is None
def test_explicit_native_returns_none():
# native is the default -> nothing to set.
assert select_attention_backend(_target(), "native", speed_active = True) is None

View file

@ -1392,6 +1392,34 @@ def test_load_promotes_fp16_to_fp32_for_zimage_only(fake_runtime, monkeypatch, t
assert q["dtype"] == "float16" # fp16-compatible family keeps fp16 on pre-Ampere
def test_bad_mode_strings_fail_before_eviction(fake_runtime):
# Every mode normalizer that can raise runs BEFORE the load evicts the previous
# pipeline, so a bad request never costs the user their working model.
backend = DiffusionBackend()
fam = detect_family("unsloth/Z-Image-GGUF")
backend._state = _LoadState(
pipe = object(),
family = fam,
repo_id = "r",
base_repo = "b",
device = "cpu",
dtype = "float32",
cpu_offload = False,
)
for kwargs in (
{"transformer_quant": "int7"},
{"speed_mode": "warp"},
{"attention_backend": "bogus"},
{"transformer_cache": "bogus"},
{"text_encoder_quant": "fp3"},
):
with pytest.raises(ValueError):
backend.load_pipeline(
"unsloth/Z-Image-GGUF", gguf_filename = "m.gguf", **kwargs
)
assert backend._state is not None
# Lock split + mid-denoise cancellation
@ -1435,17 +1463,25 @@ def test_generate_lock_split_keeps_status_and_unload_responsive(fake_runtime):
assert backend.status()["loaded"] is True
assert backend.generate_progress()["active"] is True
# unload() must return promptly (it does not wait on _generate_lock) and signal
# THIS in-flight generation's cancel event.
cancel_ref = backend._active_generate_cancel
assert cancel_ref is not None
# unload() signals THIS generation's cancel event, then waits for the denoise to
# actually exit before returning: callers treat its return as "VRAM is free" (the
# GPU arbiter hands the GPU to chat on it). Release the pipe once the cancel
# lands, standing in for the step callback of a real pipeline.
releaser = threading.Thread(target = lambda: (cancel_ref.wait(5), release.set()))
releaser.start()
backend.unload()
assert backend._active_generate_cancel is not None
assert backend._active_generate_cancel.is_set()
releaser.join(5)
assert cancel_ref.is_set()
assert backend.status()["loaded"] is False
release.set()
t.join(5)
# The cancelled generation raised rather than returning a now-evicted image.
# The cancelled generation raised rather than returning a now-evicted image, and
# it had already exited (deregistering its cancel) before unload() returned.
assert "exc" in out and "cancelled" in str(out["exc"]).lower()
assert backend._active_generate_cancel is None
def test_callback_cancellation_interrupts_denoise(fake_runtime):

View file

@ -136,6 +136,29 @@ def test_incompatible_model_runs_uncached(monkeypatch):
assert apply_step_cache(_pipe(t), mode = "fbcache") is None
def test_enable_cache_failure_rolls_back_partial_hooks(monkeypatch):
# enable_cache can raise after hooking some blocks; the reported-uncached model
# must not actually run half-cached, so the failure path calls disable_cache.
_stub_diffusers(monkeypatch)
t = _MixinTransformer(fail = True)
t.disabled = False
t.disable_cache = lambda: setattr(t, "disabled", True)
assert apply_step_cache(_pipe(t), mode = "fbcache") is None
assert t.disabled is True
def test_config_import_falls_back_to_hooks_module(monkeypatch):
# Older diffusers exports FirstBlockCacheConfig only from diffusers.hooks.
diffusers = types.ModuleType("diffusers") # no FirstBlockCacheConfig attribute
monkeypatch.setitem(sys.modules, "diffusers", diffusers)
hooks = types.ModuleType("diffusers.hooks")
hooks.FirstBlockCacheConfig = _Config
monkeypatch.setitem(sys.modules, "diffusers.hooks", hooks)
t = _MixinTransformer()
assert apply_step_cache(_pipe(t), mode = "fbcache") == TC_FBCACHE
assert t.enabled_with.threshold == DEFAULT_FBCACHE_THRESHOLD
def test_missing_transformer_is_none(monkeypatch):
_stub_diffusers(monkeypatch)
pipe = types.SimpleNamespace(transformer = None)
@ -143,7 +166,10 @@ def test_missing_transformer_is_none(monkeypatch):
def test_diffusers_unavailable_runs_uncached(monkeypatch):
# no diffusers import -> best-effort returns None, load proceeds uncached.
# no diffusers import -> best-effort returns None, load proceeds uncached. Block the
# hooks module too: the config import falls back to diffusers.hooks, which a REAL
# earlier import in the test session may have left cached in sys.modules.
monkeypatch.setitem(sys.modules, "diffusers", None)
monkeypatch.setitem(sys.modules, "diffusers.hooks", None)
t = _MixinTransformer()
assert apply_step_cache(_pipe(t), mode = "fbcache") is None

View file

@ -235,6 +235,16 @@ def test_config_rejects_zero_lora_alpha():
DiffusionLoraConfig(base_model = "b", data_dir = "d", output_dir = "o", lora_alpha = 0).normalized()
def test_config_rejects_nonpositive_snr_gamma():
# gamma <= 0 zeroes/inverts the min-SNR weight; None is the documented disable.
with pytest.raises(ValueError, match = "snr_gamma"):
DiffusionLoraConfig(base_model = "b", data_dir = "d", output_dir = "o", snr_gamma = 0).normalized()
cfg = DiffusionLoraConfig(
base_model = "b", data_dir = "d", output_dir = "o", snr_gamma = None
).normalized()
assert cfg.snr_gamma is None
def test_config_coerces_string_learning_rate():
# The Studio config path preserves learning_rate as a string; normalize to float.
cfg = DiffusionLoraConfig(

View file

@ -434,6 +434,20 @@ def test_config_from_dict_epoch_mode_drops_max_steps_sentinel():
assert cfg_explicit.train_steps == 25
def test_route_start_accepts_zero_max_grad_norm(client):
# 0 is the documented "disable clipping" value (the trainer skips clip_grad_norm_);
# the request model must not reject it.
r = client.post("/api/train/diffusion/start", json = {**_BODY, "max_grad_norm": 0.0})
assert r.status_code == 200, r.text
assert client._fake.started_with["max_grad_norm"] == 0.0
def test_route_start_rejects_nonpositive_snr_gamma(client):
# gamma <= 0 zeroes/inverts the min-SNR loss weight; null is the disable value.
r = client.post("/api/train/diffusion/start", json = {**_BODY, "snr_gamma": 0})
assert r.status_code == 422
def test_route_start_rejects_uncontained_paths(client):
# An absolute path outside the Studio dataset roots is a 400, not silently accepted.
r = client.post("/api/train/diffusion/start", json = {**_BODY, "data_dir": "/etc"})

View file

@ -166,10 +166,18 @@ def test_build_appends_offload_and_extra_args_last():
def test_build_negative_prompt_and_batch():
files = SdCppModelFiles(diffusion_model = "/m/z.gguf")
params = SdCppGenParams(prompt = "x", negative_prompt = "blurry", batch_count = 3)
params = SdCppGenParams(prompt = "x", negative_prompt = "blurry")
cmd = build_sd_cpp_command("/bin/sd-cli", files, params, output_path = "/o.png")
assert _pair(cmd, "--negative-prompt") == "blurry"
assert _pair(cmd, "--batch-count") == "3"
# A CLI batch would silently drop every image after the first (the runner only
# collects the literal --output path), so the builder rejects it outright.
with pytest.raises(ValueError, match = "single-image"):
build_sd_cpp_command(
"/bin/sd-cli",
files,
SdCppGenParams(prompt = "x", batch_count = 3),
output_path = "/o.png",
)
def test_build_omits_unset_optional_params():

View file

@ -323,6 +323,22 @@ def test_generate_raises_when_no_output_despite_success(tmp_path, monkeypatch):
)
def test_generate_does_not_return_stale_preexisting_output(tmp_path, monkeypatch):
# A leftover file at the target path must not satisfy the post-run output check
# when the run itself produced nothing: the target is cleared before the run.
e = _engine(tmp_path)
out = tmp_path / "img.png"
out.write_bytes(b"stale")
_patch_popen(monkeypatch, lines = ["ok"], returncode = 0, out_file = out, write = False)
with pytest.raises(RuntimeError, match = "no image"):
e.generate(
SdCppModelFiles(diffusion_model = "/m/z.gguf"),
SdCppGenParams(prompt = "x"),
output_path = str(out),
)
assert not out.exists()
def test_generate_raises_when_binary_missing():
e = SdCppEngine(binary = None)
with pytest.raises(RuntimeError, match = "not found"):

View file

@ -1982,7 +1982,9 @@ export function HubModelPicker({
);
// Local ./models entries. Chat-only Studio runs GGUF (any host) and MLX (Mac
// only), so raw checkpoints there are hidden (mirrors the cached non-GGUF
// rule). An MLX build a Mac user dropped in ./models stays selectable.
// rule). An MLX build a Mac user dropped in ./models stays selectable. A
// task-scoped picker (Images) is exempt: the image backend loads local
// diffusers/safetensors pipelines even on chat-only (no-GPU, native) hosts.
const sortedLocalDir = useMemo(
() =>
sortLocalModels(
@ -1990,6 +1992,7 @@ export function HubModelPicker({
(m) =>
passesTaskGate(m.task, m.model_id ?? m.id, task) &&
(!chatOnly ||
task != null ||
localModelIsGguf(m) ||
(isMac && localModelIsMlx(m))) &&
localModelMatchesFormat(m, formatFilter) &&