Spaces:
Sleeping
Sleeping
fix(lora): accept XL SFT PEFT-style LoRA keys in validator
Browse filesXL SFT uses LLM-style naming (layers.X.cross_attn.k_proj) not
diffusion-style (transformer.blocks.X.attn.to_k). LoRAs trained
with PEFT wrap keys as base_model.model.layers.*. Add both prefix
and module-suffix sets so community XL SFT LoRAs pass validation.
- lora_stack.py +17 -9
lora_stack.py
CHANGED
|
@@ -27,9 +27,12 @@ from pathlib import Path
|
|
| 27 |
|
| 28 |
_log = logging.getLogger("ams.lora")
|
| 29 |
|
| 30 |
-
#
|
| 31 |
-
|
| 32 |
-
|
|
|
|
|
|
|
|
|
|
| 33 |
_MAX_FILE_BYTES = 500 * 1024 * 1024 # 500 MB cap
|
| 34 |
_MAX_RANK = 256
|
| 35 |
|
|
@@ -91,12 +94,16 @@ def sniff(path: Path | str) -> LoRAInfo:
|
|
| 91 |
continue
|
| 92 |
if not isinstance(v, dict) or "shape" not in v:
|
| 93 |
continue
|
| 94 |
-
#
|
| 95 |
-
#
|
| 96 |
-
#
|
| 97 |
-
if
|
|
|
|
|
|
|
|
|
|
|
|
|
| 98 |
has_ace_prefix = True
|
| 99 |
-
# Extract module suffix from
|
| 100 |
for suffix in _EXPECTED_MODULES:
|
| 101 |
if f".{suffix}.lora_A.weight" in k or f".{suffix}.lora_B.weight" in k:
|
| 102 |
target_modules.add(suffix)
|
|
@@ -109,7 +116,8 @@ def sniff(path: Path | str) -> LoRAInfo:
|
|
| 109 |
"OK"
|
| 110 |
if compatible
|
| 111 |
else (
|
| 112 |
-
f"Expected ACE-Step DiT modules ({sorted(
|
|
|
|
| 113 |
f"{sorted(set(header.keys()) - {'__metadata__'})[:3]}…"
|
| 114 |
)
|
| 115 |
)
|
|
|
|
| 27 |
|
| 28 |
_log = logging.getLogger("ams.lora")
|
| 29 |
|
| 30 |
+
# ACE-Step v1 / diffusers-style DiT: transformer.blocks.X.attn.to_q …
|
| 31 |
+
_EXPECTED_MODULES_V1 = {"to_q", "to_k", "to_v", "to_out.0", "ff.net.0.proj", "ff.net.2"}
|
| 32 |
+
# ACE-Step v1.5 XL SFT: LLM-style naming — layers.X.cross_attn.k_proj …
|
| 33 |
+
# LoRAs trained with PEFT wrap these as base_model.model.layers.X.cross_attn.k_proj
|
| 34 |
+
_EXPECTED_MODULES_XL = {"q_proj", "k_proj", "v_proj", "o_proj"}
|
| 35 |
+
_EXPECTED_MODULES = _EXPECTED_MODULES_V1 | _EXPECTED_MODULES_XL
|
| 36 |
_MAX_FILE_BYTES = 500 * 1024 * 1024 # 500 MB cap
|
| 37 |
_MAX_RANK = 256
|
| 38 |
|
|
|
|
| 94 |
continue
|
| 95 |
if not isinstance(v, dict) or "shape" not in v:
|
| 96 |
continue
|
| 97 |
+
# v1 / diffusers DiT: keys start with "transformer." or "transformer_blocks."
|
| 98 |
+
# v1.5 XL SFT + PEFT: keys start with "base_model.model.layers."
|
| 99 |
+
# Reject SDXL UNet ("unet.") and LLM-only LoRAs even if they share suffix names.
|
| 100 |
+
if (
|
| 101 |
+
k.startswith("transformer.")
|
| 102 |
+
or k.startswith("transformer_blocks.")
|
| 103 |
+
or k.startswith("base_model.model.layers.")
|
| 104 |
+
):
|
| 105 |
has_ace_prefix = True
|
| 106 |
+
# Extract module suffix from both naming conventions.
|
| 107 |
for suffix in _EXPECTED_MODULES:
|
| 108 |
if f".{suffix}.lora_A.weight" in k or f".{suffix}.lora_B.weight" in k:
|
| 109 |
target_modules.add(suffix)
|
|
|
|
| 116 |
"OK"
|
| 117 |
if compatible
|
| 118 |
else (
|
| 119 |
+
f"Expected ACE-Step DiT modules (v1: {sorted(_EXPECTED_MODULES_V1)} or "
|
| 120 |
+
f"XL SFT: {sorted(_EXPECTED_MODULES_XL)}), got modules in: "
|
| 121 |
f"{sorted(set(header.keys()) - {'__metadata__'})[:3]}…"
|
| 122 |
)
|
| 123 |
)
|