Merge remote-tracking branch 'origin/diffusion-lora-ux' into diffusion-lora-training

This commit is contained in:
Daniel Han 2026-07-02 01:23:23 +00:00
commit e0db8f4612
9 changed files with 512 additions and 101 deletions

View file

@ -191,10 +191,13 @@ def _snap_to_multiple(img: Any, multiple: int = 16) -> Any:
# no unsloth-hosted GGUF, so without this its curated catalog entry could not load.
# Exact-match, lowercased, so it cannot be widened by a typo-squat. Extend
# deliberately, and never add a repo that carries pickled weights or remote code.
# The SDXL refiner is intentionally NOT here: it is an img2img-only refiner pipeline
# (StableDiffusionXLImg2ImgPipeline), but this backend loads every ``sdxl`` repo as the
# base txt2img StableDiffusionXLPipeline and advertises txt2img, so allowlisting the
# refiner would surface the wrong workflow and call it without its required input image.
_TRUSTED_NON_GGUF_REPOS = frozenset(
{
"stabilityai/stable-diffusion-xl-base-1.0",
"stabilityai/stable-diffusion-xl-refiner-1.0",
"stabilityai/sdxl-turbo",
}
)
@ -208,9 +211,18 @@ def _is_trusted_diffusion_repo(repo_id: str) -> bool:
non-GGUF paths are gated to the ``unsloth/*`` org (the curated safetensors models),
a short allowlist of official safetensors-only base repos (``_TRUSTED_NON_GGUF_REPOS``,
e.g. the SDXL base), and local paths the user explicitly pointed at (already on their
disk). The GGUF path is unchanged and stays open to any repo, as before."""
if Path(repo_id).expanduser().exists():
return True
disk). The GGUF path is unchanged and stays open to any repo, as before.
A bare ``owner/name`` HF id is never a real filesystem path, and an id with invalid
characters makes ``Path.exists()`` raise OSError; treat any such failure as "not a
local path" so the trust decision falls through to the org/allowlist checks (the
loader's validate_load_request raises the clear FileNotFoundError for a genuinely
missing local pick)."""
try:
if Path(repo_id).expanduser().exists():
return True
except OSError:
pass
rid = repo_id.strip().lower()
return rid.startswith("unsloth/") or rid in _TRUSTED_NON_GGUF_REPOS
@ -434,6 +446,16 @@ class DiffusionBackend:
f"pass family_override with that family name. (Video models and image models "
f"whose diffusers transformer has no single-file loader are not supported.)"
)
# A GGUF load builds a transformer-only file via the generic GGUF branch
# (UNet2DConditionModel.from_single_file(subfolder="transformer", GGUFQuantizationConfig)).
# Families whose single file IS the whole pipeline (SDXL) have no transformer-only
# GGUF path, so reject GGUF here -- before the route evicts the current model and
# the background load fails deep in from_single_file.
if kind == "gguf" and fam.single_file_is_pipeline:
raise ValueError(
f"'{fam.name}' checkpoints are whole-pipeline single files and have no GGUF "
f"transformer variant; load the .safetensors pipeline instead of a GGUF."
)
# Non-GGUF loads (a single-file safetensors transformer, or a full pipeline)
# are gated to the unsloth org or a local path -- they fetch + deserialise
# weights, so an arbitrary remote repo is rejected here, before any work.
@ -568,6 +590,7 @@ class DiffusionBackend:
base,
kwargs.get("hf_token"),
kind = kind,
single_file_is_pipeline = bool(fam and fam.single_file_is_pipeline),
)
with self._lock:
# Stamp progress only if this load is still current; a superseding
@ -634,6 +657,7 @@ class DiffusionBackend:
hf_token: Optional[str],
*,
kind: str = "gguf",
single_file_is_pipeline: bool = False,
) -> tuple[int, list[str]]:
"""Total download size for the progress bar, plus the base-repo files to
fetch (the prefetch reuses this list, so the base is listed only once).
@ -641,7 +665,9 @@ class DiffusionBackend:
For a ``pipeline`` load the whole repo IS the pipeline (``base_repo`` is the
repo itself), so the transformer/ subfolder is INCLUDED -- unlike the GGUF /
single-file paths, where the transformer is the single file and the base repo
supplies only the companions."""
supplies only the companions. For a ``single_file_is_pipeline`` family (SDXL) the
single file is the WHOLE pipeline, so the base repo supplies only config/tokenizer
(no weights) and its weight files are skipped."""
from huggingface_hub import HfApi
api = HfApi()
@ -671,9 +697,16 @@ class DiffusionBackend:
if gguf_filename and not Path(repo_id).expanduser().exists():
info = api.model_info(repo_id, files_metadata = True, token = hf_token)
total += sum(s.size or 0 for s in info.siblings if s.rfilename == gguf_filename)
# A whole-pipeline single file (SDXL) needs only the base repo's config/tokenizer,
# not its (unused, multi-GB) weight files.
base_filter = (
_base_config_file_downloaded
if (kind == "single_file" and single_file_is_pipeline)
else _base_file_downloaded
)
base_info = api.model_info(base_repo, files_metadata = True, token = hf_token)
for s in base_info.siblings:
if _base_file_downloaded(s.rfilename):
if base_filter(s.rfilename):
base_files.append(s.rfilename)
total += s.size or 0
except Exception as exc: # noqa: BLE001 — estimate is best-effort
@ -719,6 +752,13 @@ class DiffusionBackend:
model_kind: Optional[str] = None,
_load_token: Optional[int] = None,
) -> dict[str, Any]:
# A blank / whitespace-only token must degrade to anonymous access, not be passed
# as an explicit credential (from_single_file / from_pretrained / the Hub client
# can error on a malformed token instead of falling back). Normalize once here so
# every load branch and the size estimate below use a real token or None.
hf_token = hf_token.strip() if isinstance(hf_token, str) else hf_token
hf_token = hf_token or None
# Validate first (cheap, no torch/diffusers) so a direct call with a bad
# family fails with ValueError even in a no-diffusers runtime.
fam = self.validate_load_request(
@ -1253,11 +1293,30 @@ class DiffusionBackend:
if cn_model is None:
if cancel.is_set():
raise RuntimeError(DIFFUSION_CANCELLED_MSG)
cn_model = (
getattr(diffusers, model_cls_name)
.from_pretrained(resolved_cn.path, torch_dtype = state.dtype, token = state.hf_token)
.to(state.device)
import torch
# state.dtype is the display string saved at load ("bfloat16"), NOT a
# torch.dtype; pass the real dtype so diffusers loads the ControlNet at the
# base compute dtype instead of silently defaulting to float32 (extra VRAM).
cn_dtype = getattr(torch, str(state.dtype).replace("torch.", ""), None)
cn_model = getattr(diffusers, model_cls_name).from_pretrained(
resolved_cn.path,
torch_dtype = cn_dtype,
# An empty / malformed token means anonymous access; the HF client can
# raise on a blank credential instead of falling back, so coerce to None.
token = state.hf_token or None,
)
# Placement must follow the base model's offload policy. A resident base moves
# the ControlNet resident too; an offloaded (low-VRAM) base streams it through
# the device with group offloading instead of forcing the whole module onto the
# GPU, which would defeat the offload and risk an OOM. Best-effort: any failure
# falls back to the resident placement (the prior behaviour).
if getattr(state, "offload_policy", OFFLOAD_NONE) != OFFLOAD_NONE and (
_offload_controlnet_module(cn_model, state.device, logger)
):
pass
else:
cn_model = cn_model.to(state.device)
if cancel.is_set():
# An unload raced the blocking download above and already cleared the
# ControlNet caches; caching now would pin the module past the unload.
@ -1288,7 +1347,10 @@ class DiffusionBackend:
if denoiser is None or vae is None:
return
try:
target_dtype = denoiser.dtype
# Read the dtype from a parameter (not denoiser.dtype): a plain nn.Module has no
# .dtype, and a torch.compile'd/ wrapped denoiser can obscure it; this also
# matches how the VAE dtype is read on the next line.
target_dtype = next(denoiser.parameters()).dtype
if next(vae.parameters()).dtype != target_dtype:
vae.to(dtype = target_dtype)
except (StopIteration, AttributeError, RuntimeError):
@ -1332,6 +1394,15 @@ class DiffusionBackend:
)
resolved = diffusion_lora.resolve_specs(specs, hf_token = state.hf_token, cancel_event = cancel)
# The shared catalog scans both .safetensors and .gguf, but diffusers'
# load_lora_weights only takes safetensors; a .gguf adapter would otherwise fail
# deep in generation. Reject it here as a clean 400 before touching the pipe.
bad = [r.id for r in resolved if r.fmt != "safetensors"]
if bad:
raise ValueError(
"GGUF LoRA adapters are not supported on the diffusers engine "
f"({', '.join(bad)}); use a .safetensors adapter, or the native engine."
)
# Unique adapter names (diffusers requires distinct names; sanitized stems can collide).
uniq: list[tuple[str, str, float]] = []
seen: set[str] = set()
@ -1439,8 +1510,24 @@ class DiffusionBackend:
pipe = state.pipe
init_pil = mask_pil = None
control_pil = None
cn_scale = cn_gstart = cn_gend = None
cn_scale = cn_gstart = cn_gend = cn_mode = None
ref_extra: list = []
# Validate parameter dependencies up front: mask / upscale / reference all
# need an input image, and reference conditioning needs a family that
# supports it. Without these guards an unsupported combination would be
# silently ignored and quietly fall back to txt2img / img2img.
if init_image is None:
if mask_image is not None:
raise ValueError("mask_image requires an input image (init_image).")
if upscale is not None and upscale > 1.0:
raise ValueError("upscale requires an input image (init_image).")
if reference_images:
raise ValueError("reference_images require an input image (init_image).")
if reference_images and not getattr(state.family, "reference", False):
raise ValueError(
f"Reference images are not supported for the '{state.family.name}' "
"model family."
)
if getattr(state.family, "edit", False):
# Instruction editing: the loaded pipe is the edit pipeline. It always
# needs an input image; the prompt is the edit instruction. No mask, no
@ -1509,41 +1596,55 @@ class DiffusionBackend:
# and passes a control map. v1 conditions txt2img only (not img2img/inpaint/edit).
if controlnet is not None:
from core.inference import diffusion_controlnet
if workflow != "txt2img":
raise ValueError(
"ControlNet currently combines with plain text-to-image only, not the "
f"{workflow} workflow."
)
if not diffusion_controlnet.supports_controlnet(
engine = "diffusers",
family = state.family.name,
has_controlnet_pipeline = bool(
getattr(state.family, "controlnet_pipeline_class", None)
),
model_kind = state.kind,
transformer_quant = state.transformer_quant,
):
raise ValueError(
"ControlNet is not supported for this model/quantisation on the "
"diffusers engine (needs a bf16 or bnb-4bit load of a family with a "
"ControlNet pipeline; not GGUF-via-diffusers or torchao fp8/int8)."
)
cn_id, cn_image_b64, cn_type, cn_strength, cn_gs, cn_ge = controlnet
resolved_cn = diffusion_controlnet.resolve_controlnet(
cn_id,
family = state.family.name,
hf_token = state.hf_token,
cancel_event = cancel,
)
pipe = self._controlnet_pipe(state, resolved_cn, cancel)
workflow = "controlnet"
src = _decode_b64_image(cn_image_b64, mode = "RGB")
# Control map at the OUTPUT size so it aligns with the generated latents.
control_pil = diffusion_controlnet.preprocess_control(src, cn_type).resize(
(width, height), Image.LANCZOS
)
cn_scale, cn_gstart, cn_gend = cn_strength, cn_gs, cn_ge
# strength 0 disables ControlNet (documented on the request model, and the
# frontend slider allows it): skip the whole path so a no-op selection never
# pays the multi-GB ControlNet download / VRAM cost.
if cn_strength in (None, 0, 0.0):
controlnet = None
else:
if workflow != "txt2img":
raise ValueError(
"ControlNet currently combines with plain text-to-image only, not "
f"the {workflow} workflow."
)
if not diffusion_controlnet.supports_controlnet(
engine = "diffusers",
family = state.family.name,
has_controlnet_pipeline = bool(
getattr(state.family, "controlnet_pipeline_class", None)
),
model_kind = state.kind,
transformer_quant = state.transformer_quant,
):
raise ValueError(
"ControlNet is not supported for this model/quantisation on the "
"diffusers engine (needs a bf16 or bnb-4bit load of a family with a "
"ControlNet pipeline; not GGUF-via-diffusers or torchao fp8/int8)."
)
# Decode + preprocess the control image FIRST so a malformed / unsupported
# image fails as a clean 400 BEFORE any ControlNet download or pipe build,
# rather than after paying that cost. Control map at the OUTPUT size so it
# aligns with the generated latents.
src = _decode_b64_image(cn_image_b64, mode = "RGB")
control_pil = diffusion_controlnet.preprocess_control(src, cn_type).resize(
(width, height), Image.LANCZOS
)
try:
resolved_cn = diffusion_controlnet.resolve_controlnet(
cn_id, family = state.family.name
)
except FileNotFoundError as exc:
# An unknown / missing ControlNet id is a bad selection -> 400, not a
# generic 500 (the route maps ValueError, not FileNotFoundError).
raise ValueError(str(exc)) from exc
pipe = self._controlnet_pipe(state, resolved_cn, cancel)
workflow = "controlnet"
cn_scale, cn_gstart, cn_gend = cn_strength, cn_gs, cn_ge
# Flux Union ControlNet selects the active mode by an integer
# ``control_mode`` (canny/depth/pose/...); map the chosen control type so
# the union model applies the right head instead of a default/wrong one.
cn_mode = diffusion_controlnet.union_control_mode(cn_id, cn_type)
# Auto-resize odd-sized inputs to a multiple of 16 for the workflows whose
# OUTPUT size is taken from the input image (img2img / inpaint / extend / edit),
# so an upload like 186px tall no longer fails the pipeline's divisibility check.
@ -1556,9 +1657,8 @@ class DiffusionBackend:
mask_pil = mask_pil.resize(init_pil.size, _PILImage.NEAREST)
if init_pil is not None:
# Keep the VAE encode dtype consistent with the input image.
self._align_vae_dtype(
pipe, getattr(state.family, "denoiser_attr", "transformer")
)
# state.family is always a DiffusionFamily, which defines denoiser_attr.
self._align_vae_dtype(pipe, state.family.denoiser_attr)
# Pipelines vary in which kwargs they accept (img2img derives size from the
# input image and may reject width/height; a distilled pipe may take no
@ -1620,6 +1720,10 @@ class DiffusionBackend:
kwargs["control_guidance_start"] = cn_gstart
if "control_guidance_end" in call_params and cn_gend is not None:
kwargs["control_guidance_end"] = cn_gend
# Union ControlNet mode index (Flux); only when the pipe accepts it and the
# selected control type maps to a known mode.
if "control_mode" in call_params and cn_mode is not None:
kwargs["control_mode"] = cn_mode
gen = _GenState(total_steps = steps)
@ -1724,13 +1828,13 @@ class DiffusionBackend:
if state.eager_patched:
uninstall_patches()
uninstall_arch_patches()
# Drop any LoRA adapters applied to the pipe so a later reference-path load is
# bit-identical and the freed transformer carries no adapter layers. Idempotent.
try:
if getattr(state.pipe, "_unsloth_loras", ()):
state.pipe.unload_lora_weights()
except Exception: # noqa: BLE001 -- best-effort cleanup on teardown
pass
# 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
# 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
# tensors are freed with it -- no explicit unload is needed for memory or for a
# later load (which builds a fresh pipe).
# Drop the workflow pipes built around this load's modules so they don't pin the
# freed pipeline (they only re-wire its components, but holding the wrappers
# would keep the modules alive past unload).
@ -1861,6 +1965,34 @@ def _hf_base_model(repo_id: str, hf_token: Optional[str]) -> Optional[str]:
return base if isinstance(base, str) and base.strip() else None
def _offload_controlnet_module(cn_model: Any, device: str, logger: Any) -> bool:
"""Stream a ControlNet module through ``device`` via diffusers group offloading.
Used when the base model was loaded with an offload policy: forcing the ControlNet
fully resident with ``.to(device)`` would defeat that low-VRAM placement and can OOM.
Group offloading is applied to this single module (it does not touch the base pipe's
existing hooks), so it is isolated and reversible. Returns True on success; on any
failure the caller falls back to a resident placement, so this never blocks a load."""
try:
import torch
from diffusers.hooks import apply_group_offloading
onload = torch.device(device)
apply_group_offloading(
cn_model,
onload_device = onload,
offload_device = torch.device("cpu"),
offload_type = "block_level",
num_blocks_per_group = 1,
use_stream = onload.type == "cuda",
)
return True
except Exception as exc: # noqa: BLE001 — offload is best-effort; resident is the fallback
if logger is not None:
logger.warning("diffusion.controlnet: group offload failed (%s); loading resident", exc)
return False
def _base_file_downloaded(rfilename: str) -> bool:
"""True for base-repo files ``from_pretrained`` actually fetches.
@ -1876,6 +2008,33 @@ def _base_file_downloaded(rfilename: str) -> bool:
return not rfilename.startswith("assets/")
# Weight file extensions the base repo need NOT supply when the single file is the whole
# pipeline (SDXL): from_single_file(config=base) reads only the base repo's structure
# (config/tokenizer/scheduler) and takes the weights from the single file.
_BASE_WEIGHT_EXTS = (
".safetensors",
".bin",
".ckpt",
".pt",
".pth",
".gguf",
".onnx",
".onnx_data",
".msgpack",
".h5",
".pb",
)
def _base_config_file_downloaded(rfilename: str) -> bool:
"""True for base-repo files needed to BUILD a pipeline structure around a whole-pipeline
single file WITHOUT its weights: config / tokenizer / scheduler JSON, but no weight
tensors (the single file supplies those). Used for ``single_file_is_pipeline`` families."""
if not _base_file_downloaded(rfilename):
return False
return not rfilename.lower().endswith(_BASE_WEIGHT_EXTS)
def _pipeline_file_downloaded(rfilename: str) -> bool:
"""True for files a full-pipeline ``from_pretrained`` fetches.

View file

@ -19,7 +19,6 @@ backend read an arbitrary location.
from __future__ import annotations
import re
import threading
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Optional
@ -138,21 +137,28 @@ def _catalog_by_id() -> dict[str, ControlNetCatalogEntry]:
return {e.id: e for e in (list(_CURATED) + _scan_local())}
def resolve_controlnet(
spec_id: str,
*,
family: Optional[str] = None,
hf_token: Optional[str] = None,
cancel_event: Optional[threading.Event] = None,
) -> ResolvedControlNet:
def resolve_controlnet(spec_id: str, *, family: Optional[str] = None) -> ResolvedControlNet:
"""Resolve a ControlNet id to a loadable repo id / local dir.
Accepts a catalog/local id, or a bare public HF repo id (``owner/name``). The backend
loads the result with ``ControlNetModelClass.from_pretrained(path)`` (download + cache
handled there, like the base pipeline). Raises on an unknown id -> the caller maps to 400.
``family`` (the loaded base family) enforces catalog compatibility: a ControlNet is
architecture-specific, so a catalog entry tagged for another family is rejected here
with a clear error rather than being loaded through the wrong pipeline class later.
"""
entry = _catalog_by_id().get(spec_id)
if entry is not None:
# A curated/local entry may declare the families it is built for. A client that
# bypasses the UI filter (direct API call) could send an entry for another family;
# reject it before any download so it never reaches the wrong ControlNet pipeline.
fam = (family or "").strip().lower()
if entry.families and fam and fam not in {f.lower() for f in entry.families}:
raise ValueError(
f"ControlNet '{spec_id}' is for {', '.join(entry.families)}, not the loaded "
f"'{family}' model; pick a ControlNet built for this family."
)
if entry.source == "local":
path = entry.local_path or ""
if not path or not Path(path).is_dir():
@ -174,6 +180,33 @@ def resolve_controlnet(
)
# Union ControlNet mode indices. A single "union" model covers several control modes and
# selects the active one via an integer ``control_mode`` argument; these are the standard
# indices used by the FLUX.1 / Qwen-Image union ControlNets. "passthrough" (an already-made
# map) carries no intrinsic mode, so it maps to nothing (the caller omits control_mode).
_UNION_CONTROL_MODES: dict[str, int] = {
"canny": 0,
"tile": 1,
"depth": 2,
"blur": 3,
"pose": 4,
"gray": 5,
"lq": 6,
}
def union_control_mode(spec_id: str, control_type: str) -> Optional[int]:
"""The integer ``control_mode`` for a union ControlNet, or None.
Returns a mode only for a curated *union* catalog entry AND a control type that maps to a
known index; otherwise None so the caller omits the kwarg (a non-union ControlNet has a
single fixed mode, and 'passthrough' does not name one). Pure lookup, no network."""
entry = _catalog_by_id().get(spec_id)
if entry is None or not entry.is_union:
return None
return _UNION_CONTROL_MODES.get((control_type or "").strip().lower())
def preprocess_control(image: Any, control_type: str) -> Any:
"""Turn a source image into a control map.

View file

@ -95,26 +95,31 @@ def sanitize_alias(raw: str) -> str:
def _scan_local() -> list[LoraCatalogEntry]:
entries: list[LoraCatalogEntry] = []
root = loras_dir()
try:
children = sorted(root.iterdir())
except OSError:
return entries
for p in children:
if not p.is_file():
continue
return []
files = [p for p in children if p.is_file() and p.suffix.lower() in _ALL_EXTS]
# Two files that share a stem but differ in extension (foo.safetensors + foo.gguf)
# would collide on id (== stem), so the frontend select value and resolve_one's
# id->entry lookup could only ever address one of them. Disambiguate a colliding
# stem by keeping the full filename as the id; a unique stem stays the clean stem.
stem_counts: dict[str, int] = {}
for p in files:
stem_counts[p.stem] = stem_counts.get(p.stem, 0) + 1
entries: list[LoraCatalogEntry] = []
for p in files:
ext = p.suffix.lower()
if ext not in _ALL_EXTS:
continue
try:
size = p.stat().st_size
except OSError:
size = 0
entry_id = p.name if stem_counts.get(p.stem, 0) > 1 else p.stem
entries.append(
LoraCatalogEntry(
id = p.stem,
display_name = p.stem,
id = entry_id,
display_name = entry_id,
source = "local",
fmt = "gguf" if ext == ".gguf" else "safetensors",
local_path = str(p),
@ -157,6 +162,9 @@ def resolve_one(
shared xet-fallback helper. Raises FileNotFoundError/ValueError on an unresolvable or
unsupported id -- the caller maps that to a clear 400.
"""
# An empty / whitespace token sent verbatim to HfApi triggers an auth error instead
# of falling back to anonymous access; normalise it to None.
hf_token = hf_token.strip() if hf_token and hf_token.strip() else None
entry = _catalog_by_id().get(spec_id)
if entry is not None:
if entry.source == "local":
@ -176,6 +184,17 @@ def resolve_one(
if "/" in spec_id:
repo_id, _, weight_name = spec_id.partition(":")
weight_name = weight_name or None
if weight_name is not None:
# A client-supplied weight file must stay a plain filename inside the repo:
# reject traversal / absolute paths so it can never resolve outside the HF
# cache dir once handed to the downloader.
if (
".." in weight_name
or weight_name.startswith(("/", "\\", "~"))
or "\\" in weight_name
or os.path.isabs(weight_name)
):
raise ValueError(f"invalid LoRA weight file path '{weight_name}'")
if weight_name is None:
weight_name = _pick_repo_weight_file(repo_id, hf_token)
ext = os.path.splitext(weight_name)[1].lower()
@ -218,12 +237,19 @@ def resolve_specs(
hf_token: Optional[str] = None,
cancel_event: Optional[threading.Event] = None,
) -> list[ResolvedLora]:
"""Resolve request (id, weight) pairs, dropping zero-weight entries."""
"""Resolve request (id, weight) pairs, dropping zero-weight entries.
A stale / unknown id raises FileNotFoundError inside resolve_one; convert it to
ValueError so the route (which maps only ValueError to a 400) reports bad client
input instead of a generic 500."""
out: list[ResolvedLora] = []
for spec_id, weight in specs:
if weight == 0:
continue
out.append(resolve_one(spec_id, weight, hf_token = hf_token, cancel_event = cancel_event))
try:
for spec_id, weight in specs:
if weight == 0:
continue
out.append(resolve_one(spec_id, weight, hf_token = hf_token, cancel_event = cancel_event))
except FileNotFoundError as exc:
raise ValueError(str(exc)) from exc
return out
@ -265,20 +291,27 @@ _TAG_RE = re.compile(r"<lora:([^:>]+):([^>]+)>")
def inject_prompt_tags(prompt: str, resolved: list[ResolvedLora]) -> str:
"""Append `<lora:ALIAS:WEIGHT>` tags to the prompt, skipping any the user already typed.
"""Append `<lora:ALIAS:WEIGHT>` tags for the selected adapters, using the backend-
validated weights.
sd-cli strips these tags before they reach the model, so appending them is safe and
deterministic. Duplicate protection: if the prompt already contains a tag for the same
alias, we don't add a second one.
deterministic. A selected adapter's weight is validated (0-2) and recorded in the
request/gallery, so the injected tag must WIN over any `<lora:ALIAS:...>` the user
typed for that same alias: strip a user tag whose alias matches a selected adapter,
then append the validated one. Tags for aliases the user typed that are NOT selected
are left untouched (free-form use).
"""
existing = {m.group(1) for m in _TAG_RE.finditer(prompt)}
tags = [
f"<lora:{r.alias}:{_fmt_weight(r.weight)}>" for r in resolved if r.alias not in existing
]
selected = {r.alias for r in resolved}
# Drop any user-typed tag whose alias is one of the selected adapters, so the typed
# weight can't override the validated weight (or slip outside the 0-2 bounds).
cleaned = _TAG_RE.sub(lambda m: "" if m.group(1) in selected else m.group(0), prompt)
# Collapse whitespace left by stripped tags without disturbing the user's text.
cleaned = re.sub(r"[ \t]{2,}", " ", cleaned).strip()
tags = [f"<lora:{r.alias}:{_fmt_weight(r.weight)}>" for r in resolved]
if not tags:
return prompt
sep = "" if not prompt or prompt.endswith(" ") else " "
return f"{prompt}{sep}{' '.join(tags)}"
return cleaned
sep = "" if not cleaned or cleaned.endswith(" ") else " "
return f"{cleaned}{sep}{' '.join(tags)}"
def _fmt_weight(w: float) -> str:

View file

@ -1856,6 +1856,14 @@ class ControlNetSpec(BaseModel):
1.0, ge = 0.0, le = 1.0, description = "Fraction of steps at which ControlNet ends"
)
@model_validator(mode = "after")
def _check_guidance_range(self) -> "ControlNetSpec":
# An inverted range (start > end) means "act over no steps"; reject it as a clean
# 422 instead of letting the diffusers pipeline raise a 500 deep in the denoise.
if self.guidance_start > self.guidance_end:
raise ValueError("guidance_start must be <= guidance_end")
return self
class DiffusionGenerateRequest(BaseModel):
"""Request to generate one image from the loaded diffusion model."""

View file

@ -483,6 +483,38 @@ def test_generate_img2img_unsupported_family_raises(fake_runtime, tmp_path, monk
backend.generate(prompt = "x", steps = 4, init_image = _tiny_png_b64())
def test_generate_rejects_conditioning_without_init_image(fake_runtime, tmp_path):
"""mask / upscale / reference all need an input image; without one they must raise a
clear ValueError rather than silently degrading to txt2img."""
(tmp_path / "model.gguf").write_bytes(b"x")
backend = DiffusionBackend()
backend.load_pipeline(
str(tmp_path), gguf_filename = "model.gguf", base_repo = "base/repo", family_override = "z-image"
)
with pytest.raises(ValueError, match = "mask_image requires"):
backend.generate(prompt = "x", steps = 4, mask_image = _mask_b64(64))
with pytest.raises(ValueError, match = "upscale requires"):
backend.generate(prompt = "x", steps = 4, upscale = 2.0)
with pytest.raises(ValueError, match = "reference_images require"):
backend.generate(prompt = "x", steps = 4, reference_images = [_tiny_png_b64()])
def test_generate_rejects_reference_on_unsupported_family(fake_runtime, tmp_path):
"""A non-reference family rejects reference_images instead of silently dropping them."""
(tmp_path / "model.gguf").write_bytes(b"x")
backend = DiffusionBackend()
backend.load_pipeline(
str(tmp_path), gguf_filename = "model.gguf", base_repo = "base/repo", family_override = "z-image"
)
with pytest.raises(ValueError, match = "Reference images are not supported"):
backend.generate(
prompt = "x",
steps = 4,
init_image = _tiny_png_b64(),
reference_images = [_tiny_png_b64()],
)
def test_generate_upscale_enlarges_and_low_strength(fake_runtime, tmp_path):
"""An init_image + upscale factor routes generate() through the family's img2img
pipeline (hires fix): the source is enlarged to size*factor (rounded to /16) before the

View file

@ -45,6 +45,26 @@ def test_resolve_controlnet_rejects_filesystem_like_ids():
dc.resolve_controlnet(bad)
def test_resolve_controlnet_enforces_family_match():
# A curated entry tagged for another family must be rejected before download so it
# never reaches the wrong ControlNet pipeline class.
with pytest.raises(ValueError, match = "not the"):
dc.resolve_controlnet("qwen-union", family = "flux.1")
# The matching family resolves fine, and no family (unfiltered) is permissive.
assert dc.resolve_controlnet("qwen-union", family = "qwen-image").path
assert dc.resolve_controlnet("qwen-union").path
def test_union_control_mode_maps_only_union_entries():
# Union entries map a known control type to its integer mode; passthrough / unknown
# types and non-union ids return None so the caller omits control_mode.
assert dc.union_control_mode("flux-union-pro", "canny") == 0
assert dc.union_control_mode("flux-union-pro", "depth") == 2
assert dc.union_control_mode("flux-union-pro", "pose") == 4
assert dc.union_control_mode("flux-union-pro", "passthrough") is None
assert dc.union_control_mode("some/bare-repo", "canny") is None
def test_resolve_controlnet_local(tmp_path, monkeypatch):
d = tmp_path / "controlnets"
d.mkdir()

View file

@ -37,10 +37,19 @@ def test_inject_prompt_tags_appends_with_spacing():
assert dl.inject_prompt_tags("x", [r1]) == "x <lora:s:1>"
def test_inject_prompt_tags_dedupes_user_typed_tag():
def test_inject_prompt_tags_validated_weight_overrides_user_typed():
r = dl.ResolvedLora("id", "style", "/p", "safetensors", 0.8)
# user already wrote a tag for the same alias -> not duplicated
assert dl.inject_prompt_tags("a cat <lora:style:1>", [r]) == "a cat <lora:style:1>"
# A user-typed tag for a SELECTED adapter is replaced by the backend-validated weight
# (so the recorded/validated 0-2 weight wins over whatever was typed), not duplicated.
assert dl.inject_prompt_tags("a cat <lora:style:1>", [r]) == "a cat <lora:style:0.8>"
def test_inject_prompt_tags_keeps_unselected_user_tags():
r = dl.ResolvedLora("id", "style", "/p", "safetensors", 0.8)
# A user tag for an alias that is NOT one of the selected adapters is left untouched.
out = dl.inject_prompt_tags("a cat <lora:other:0.5>", [r])
assert "<lora:other:0.5>" in out
assert "<lora:style:0.8>" in out
def test_inject_prompt_tags_empty_returns_prompt():
@ -129,6 +138,41 @@ def test_resolve_specs_drops_zero_weight(tmp_path, monkeypatch):
assert len(out) == 1 and out[0].weight == 1.0
def test_resolve_specs_maps_unknown_id_to_valueerror(tmp_path, monkeypatch):
# An unknown / stale id raises FileNotFoundError in resolve_one; resolve_specs must
# surface it as ValueError so the route returns 400, not a generic 500.
d = tmp_path / "loras"
d.mkdir()
monkeypatch.setattr(dl, "loras_dir", lambda: d)
with pytest.raises(ValueError):
dl.resolve_specs([("nope", 1.0)])
def test_scan_local_disambiguates_identical_stems(tmp_path, monkeypatch):
# foo.safetensors and foo.gguf must get distinct ids so each is addressable; a
# unique stem keeps its clean stem id.
d = tmp_path / "loras"
d.mkdir()
(d / "foo.safetensors").write_bytes(b"x")
(d / "foo.gguf").write_bytes(b"y")
(d / "solo.safetensors").write_bytes(b"z")
monkeypatch.setattr(dl, "loras_dir", lambda: d)
by_id = {e.id: e for e in dl.list_loras()}
assert "foo.safetensors" in by_id and "foo.gguf" in by_id
assert by_id["foo.safetensors"].fmt == "safetensors"
assert by_id["foo.gguf"].fmt == "gguf"
assert "solo" in by_id # unique stem is untouched
def test_resolve_one_rejects_traversal_weight_name(tmp_path, monkeypatch):
# A client-supplied weight file with traversal / absolute path is rejected before it
# can reach the downloader (it must stay a plain filename inside the repo).
monkeypatch.setattr(dl, "loras_dir", lambda: tmp_path)
for bad in ("owner/name:../secret.safetensors", "owner/name:/etc/x.safetensors"):
with pytest.raises(ValueError):
dl.resolve_one(bad, 1.0)
# ── Request-model validation ────────────────────────────────────────────────
@ -264,3 +308,21 @@ def test_diffusers_apply_rejects_unsupported_quant():
[("styleA", 1.0)],
threading.Event(),
)
def test_diffusers_apply_rejects_gguf_adapter(monkeypatch):
# A .gguf adapter (discoverable in the shared catalog) cannot load on the diffusers
# engine; it must be rejected as a clean 400 before touching the pipe.
import threading
monkeypatch.setattr(
dl,
"resolve_specs",
lambda specs, **_: [
dl.ResolvedLora(i, dl.sanitize_alias(i), f"/{i}.gguf", "gguf", w) for i, w in specs
],
)
pipe = _FakePipe()
with pytest.raises(ValueError, match = "GGUF LoRA"):
_backend()._apply_loras(_fake_state(pipe), [("styleA", 1.0)], threading.Event())
assert pipe.loaded == [] # never touched the pipe

View file

@ -15,6 +15,8 @@ from __future__ import annotations
import types
import pytest
from core.inference import diffusion_lora
from core.inference.diffusion import (
DiffusionBackend,
@ -70,7 +72,8 @@ def test_sdxl_base_repos_are_trusted_non_gguf():
# Official safetensors-only base repos are allowlisted so their catalog entries load.
assert _is_trusted_diffusion_repo("stabilityai/stable-diffusion-xl-base-1.0")
assert _is_trusted_diffusion_repo("stabilityai/sdxl-turbo")
assert _is_trusted_diffusion_repo("stabilityai/stable-diffusion-xl-refiner-1.0")
# The refiner is img2img-only and is intentionally NOT allowlisted (see
# test_sdxl_refiner_not_trusted).
# Case-insensitive match.
assert _is_trusted_diffusion_repo("StabilityAI/SDXL-Turbo")
# A random repo (even one that detects as SDXL) is NOT trusted for a non-GGUF load.
@ -100,9 +103,11 @@ class _FakeVae:
def test_align_vae_dtype_uses_unet_denoiser():
# For SDXL the denoiser lives at pipe.unet; _align_vae_dtype must read it (a pipe
# with only .unet and no .transformer) and cast the VAE to the U-Net's dtype.
# with only .unet and no .transformer) and cast the VAE to the U-Net's dtype. The
# dtype is read from a parameter (denoiser has no .dtype), so use a _FakeVae denoiser.
vae = _FakeVae(dtype = "float32")
pipe = types.SimpleNamespace(unet = types.SimpleNamespace(dtype = "bfloat16"), vae = vae)
unet = _FakeVae(dtype = "bfloat16")
pipe = types.SimpleNamespace(unet = unet, vae = vae)
DiffusionBackend._align_vae_dtype(pipe, "unet")
assert vae.moved_to == "bfloat16"
@ -110,7 +115,8 @@ def test_align_vae_dtype_uses_unet_denoiser():
def test_align_vae_dtype_transformer_default_unchanged():
# DiT default: reads pipe.transformer; a pipe with no transformer is a safe no-op.
vae = _FakeVae(dtype = "float32")
pipe = types.SimpleNamespace(transformer = types.SimpleNamespace(dtype = "bfloat16"), vae = vae)
transformer = _FakeVae(dtype = "bfloat16")
pipe = types.SimpleNamespace(transformer = transformer, vae = vae)
DiffusionBackend._align_vae_dtype(pipe)
assert vae.moved_to == "bfloat16"
# No denoiser attribute -> no-op (does not raise, does not move the VAE).
@ -147,3 +153,39 @@ def test_pipeline_prefetch_skips_non_torch_artifacts():
assert not keep("unet/flax_model.msgpack")
assert not keep("vae_decoder/model.onnx_data")
assert not keep("assets/preview.png")
def test_sdxl_refiner_not_trusted():
# The refiner is an img2img-only pipeline; the sdxl family loads every repo as the
# base txt2img pipeline, so the refiner must NOT be allowlisted for a non-GGUF load.
assert not _is_trusted_diffusion_repo("stabilityai/stable-diffusion-xl-refiner-1.0")
# The base and turbo remain trusted.
assert _is_trusted_diffusion_repo("stabilityai/stable-diffusion-xl-base-1.0")
assert _is_trusted_diffusion_repo("stabilityai/sdxl-turbo")
def test_sdxl_gguf_load_rejected_up_front():
# SDXL has no transformer-only GGUF variant (its single file is the whole pipeline),
# so a GGUF request must fail cheap validation before the GPU handoff.
backend = DiffusionBackend()
with pytest.raises(ValueError, match = "no GGUF"):
backend.validate_load_request(
"some-org/my-sdxl.gguf", gguf_filename = "my-sdxl.gguf", family_override = "sdxl"
)
def test_base_config_filter_skips_weights():
# For a whole-pipeline single file, the base repo supplies only config/tokenizer, not
# its (unused) weight tensors.
from core.inference.diffusion import _base_config_file_downloaded as keep
assert keep("model_index.json")
assert keep("text_encoder/config.json")
assert keep("tokenizer/vocab.json")
assert keep("scheduler/scheduler_config.json")
assert not keep("unet/diffusion_pytorch_model.safetensors")
assert not keep("vae/diffusion_pytorch_model.bin")
assert not keep("text_encoder/model.onnx")
# transformer/ and assets/ stay excluded (inherited from _base_file_downloaded).
assert not keep("transformer/config.json")
assert not keep("assets/x.png")

View file

@ -985,24 +985,42 @@ export function ImagesPage({ active = true }: { active?: boolean }) {
}, [images, hasMore, selectedId, quant]);
// Refresh the LoRA picker's suggestions when the loaded model (family) changes. A LoRA is
// trained for a specific base family, so a model swap invalidates the current selection --
// clear it (the user re-adds a suggestion or types a Hub repo id for the new family). We do
// trained for a specific base family, so a real model SWAP invalidates the current selection
// -- clear it then (the user re-adds a suggestion or types a Hub repo id for the new family).
// But do NOT clear on the first load or an unload: a user can restore a saved recipe (which
// sets loras) BEFORE the model finishes loading, and clearing on that load->capable
// transition would silently drop the restored adapters. We track the previously-loaded
// family in a ref and clear only when it changes to a different loaded family. We also do
// NOT filter the selection against the discovered catalog: a valid pick can be a free-text
// Hugging Face repo id that is not in the (often empty) curated list.
const loraCapable = Boolean(status?.loaded && status?.supports_lora);
const prevLoraFamilyRef = useRef<string | null | undefined>(undefined);
useEffect(() => {
if (!loraCapable) {
// Options are gone with the model, but keep the selection: it may have just been
// restored while the model is (re)loading. It is only SENT when loraCapable, and a
// real family swap below clears it.
setAvailableLoras([]);
setLoras([]);
return;
}
setLoras([]);
const fam = status?.family ?? null;
const prev = prevLoraFamilyRef.current;
if (prev != null && prev !== fam) {
setLoras([]);
}
prevLoraFamilyRef.current = fam;
let cancelled = false;
listDiffusionLoras(status?.family ?? undefined)
.then((list) => {
if (!cancelled) setAvailableLoras(list);
})
.catch(() => {
// Clear only the OPTIONS on a failed catalog refresh. Unlike the catalog-only
// picker below the stack, this free-text picker holds selections (bare HF repo
// ids) that are valid without being in the catalog; a transient refresh failure
// must not wipe them. Stale cross-family selections are already cleared by the
// family-swap check above, and hidden LoRAs are never sent (handleGenerate is
// gated on loraCapable).
if (!cancelled) setAvailableLoras([]);
});
return () => {
@ -1564,7 +1582,11 @@ export function ImagesPage({ active = true }: { active?: boolean }) {
reference_images: condRefImages,
// Drop empty (no id typed yet) and zero-weight rows, and trim hand-typed repo ids,
// so the recipe records only adapters that actually applied. Empty -> omit entirely.
// Gate on loraCapable: a restore can leave adapters in state while the loaded model
// does not support LoRA (picker hidden), and sending them would fail generation with
// no visible row to remove.
loras: (() => {
if (!loraCapable) return undefined;
const active = loras
.map((l) => ({ id: l.id.trim(), weight: l.weight }))
.filter((l) => l.id && l.weight > 0);
@ -2058,7 +2080,7 @@ export function ImagesPage({ active = true }: { active?: boolean }) {
)}
{loras.map((sel, i) => (
<div
key={i}
key={sel.id || i}
className="space-y-1.5 rounded-lg border border-border bg-muted/30 p-2"
>
<div className="flex items-center gap-2">
@ -2132,7 +2154,7 @@ export function ImagesPage({ active = true }: { active?: boolean }) {
hint="Condition the image on a control map (edges / depth / pose). Union models cover many types. Use 'Canny' to trace edges from your image, or 'Passthrough' if it is already a control map."
>
<div className="space-y-2 rounded-lg border border-border bg-muted/30 p-2">
<Select value={controlnetId} onValueChange={setControlnetId}>
<Select value={controlnetId || undefined} onValueChange={setControlnetId}>
<SelectTrigger className="h-8 w-full text-xs">
<SelectValue placeholder="Select a ControlNet" />
</SelectTrigger>