techfreakworm commited on
Commit
158626f
·
unverified ·
1 Parent(s): f9a1a8e

fix(lora): accept XL SFT PEFT-style LoRA keys in validator

Browse files

XL 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.

Files changed (1) hide show
  1. 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
- # Expected DiT module suffixes for ACE-Step 1.5 XL SFT.
31
- # Match against `*.to_q.lora_A.weight`, etc.
32
- _EXPECTED_MODULES = {"to_q", "to_k", "to_v", "to_out.0", "ff.net.0.proj", "ff.net.2"}
 
 
 
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
- # ACE-Step DiT keys start with "transformer." (the diffusers DiT prefix).
95
- # SDXL UNet LoRAs start with "unet." — reject those even though the
96
- # inner attention layer names overlap (`.to_q.lora_A.weight`).
97
- if k.startswith("transformer.") or k.startswith("transformer_blocks."):
 
 
 
 
98
  has_ace_prefix = True
99
- # Extract module suffix from things like "transformer.blocks.0.attn.to_q.lora_A.weight"
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(_EXPECTED_MODULES)}), got modules in: "
 
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
  )