Diffusion LoRA: harden resolution, native tag precedence, and diffusers teardown
Address review findings on the LoRA path: - resolve_one: normalise a blank/whitespace hf_token to None (anonymous access) and reject a client-supplied weight file with traversal / absolute path. - resolve_specs: convert FileNotFoundError from an unknown/stale id to ValueError so the route returns 400 instead of a generic 500. - _scan_local: disambiguate local adapters that share a stem (foo.safetensors vs foo.gguf) so each is uniquely addressable. - inject_prompt_tags: the backend-validated weight now wins over a user-typed <lora:ALIAS:...> for a selected adapter; unselected user tags are left alone. - diffusers _apply_loras: reject a .gguf adapter with a clear error before touching the pipe (diffusers loads safetensors only). - _unload_locked: drop the explicit unload_lora_weights() on teardown; the pipe is dropped wholesale (freeing adapters), so the previous call could race an in-flight denoise on the same pipe. - Images page: use a stable LoRA key and clear the selection (not just the options) when the catalog refresh fails.
This commit is contained in:
parent
c2b1c8a5a1
commit
50a313f93d
4 changed files with 149 additions and 36 deletions
|
|
@ -1249,6 +1249,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()
|
||||
|
|
@ -1579,13 +1588,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).
|
||||
|
|
|
|||
|
|
@ -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,21 @@ 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 +293,29 @@ _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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -975,7 +975,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;
|
||||
|
|
@ -1968,7 +1973,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">
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue