Merge remote-tracking branch 'origin/diffusion-lora' into diffusion-controlnet

This commit is contained in:
Daniel Han 2026-07-02 01:18:40 +00:00
commit 582e2dcf39
5 changed files with 205 additions and 39 deletions

View file

@ -191,9 +191,18 @@ def _is_trusted_diffusion_repo(repo_id: str) -> bool:
on an arbitrary repo, which fetches and deserialises third-party weights. So the
non-GGUF paths are gated to the ``unsloth/*`` org (the curated safetensors models) and
to 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
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 unsloth/ check (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
return repo_id.strip().lower().startswith("unsloth/")
@ -1313,6 +1322,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()
@ -1422,6 +1440,22 @@ class DiffusionBackend:
control_pil = 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
@ -1721,13 +1755,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).

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

@ -468,6 +468,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

@ -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

@ -986,7 +986,12 @@ export function ImagesPage({ active = true }: { active?: boolean }) {
setLoras((prev) => prev.filter((s) => ids.has(s.id)));
})
.catch(() => {
if (!cancelled) setAvailableLoras([]);
if (cancelled) return;
// Clear the SELECTED adapters too, not just the options: leaving a stale `loras`
// selection in state (with the picker now hidden/empty) would still be posted by
// handleGenerate and could apply adapters from the previous model, or fail.
setAvailableLoras([]);
setLoras([]);
});
return () => {
cancelled = true;
@ -2015,7 +2020,7 @@ export function ImagesPage({ active = true }: { active?: boolean }) {
<div className="space-y-2">
{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">