diff --git a/modules/lora/native_adapter.py b/modules/lora/native_adapter.py index 40cd29e97..002c57d82 100644 --- a/modules/lora/native_adapter.py +++ b/modules/lora/native_adapter.py @@ -92,6 +92,9 @@ def _resolve_prefix(network_prefix, prefix_used): SUFFIX_NORMALIZE = { "lora_A.weight": "lora_down.weight", "lora_B.weight": "lora_up.weight", + # bare parameter names, saved by wrappers that hold the factors as nn.Parameter (alibaba-pai PDD files) + "lora_down": "lora_down.weight", + "lora_up": "lora_up.weight", } @@ -103,6 +106,7 @@ SUFFIX_NORMALIZE = { LORA_SUFFIXES = ( ".lora_down.weight", ".lora_up.weight", ".lora_mid.weight", ".lora_A.weight", ".lora_B.weight", + ".lora_down", ".lora_up", # diff_b: bias delta some saves pair with the weight LoRA, applied as ex_bias. # magnitude / lora_magnitude_vector: DoRA row norms (ai-toolkit / PEFT key # names); converted onto the dora_scale path by try_load_lora. @@ -154,7 +158,7 @@ FULL_SUFFIXES = ( # on accidental overlaps with other families. LORA_MARKERS = ( - ".lora_down.weight", ".lora_up.weight", + ".lora_down", ".lora_up", # bare and .weight forms alike ".lora_A.weight", ".lora_B.weight", # PEFT named-adapter saves embed the slot name as ``.lora_A..weight``; # the trailing-dot forms catch every variant. diff --git a/pipelines/minimax/minimax_lora.py b/pipelines/minimax/minimax_lora.py index 5709772f4..17079f74d 100644 --- a/pipelines/minimax/minimax_lora.py +++ b/pipelines/minimax/minimax_lora.py @@ -85,20 +85,6 @@ BARE_DIFFUSERS_PREFIX_USED = native_adapter.BARE_DIFFUSERS_PREFIX_USED has_marker = native_adapter.has_marker -MINIMAX_EXTRA_SUFFIXES = (".lora_down", ".lora_up", ".lora_A", ".lora_B") -MINIMAX_LORA_SUFFIXES = native_adapter.LORA_SUFFIXES + MINIMAX_EXTRA_SUFFIXES -MINIMAX_SUFFIX_NORMALIZE = { - "lora_down": "lora_down.weight", - "lora_up": "lora_up.weight", - "lora_A": "lora_down.weight", - "lora_B": "lora_up.weight", -} - - -def normalize_mini_max_suffix(suffix: str) -> str: - return MINIMAX_SUFFIX_NORMALIZE.get(suffix, suffix) - - def _flattened(dotted): return re.escape(dotted.replace(".", "_")) @@ -132,11 +118,7 @@ def parse_key(key, suffixes): key = native_adapter.unwrap_peft_wrapper(key) if key.startswith("dit."): key = "diffusion_model." + key[len("dit."):] - parsed = native_adapter.parse_key(key, suffixes, prefixes=KNOWN_PREFIXES) - if parsed is None: - return None - prefix_used, base, suffix = parsed - return prefix_used, base, normalize_mini_max_suffix(suffix) + return native_adapter.parse_key(key, suffixes, prefixes=KNOWN_PREFIXES) def group_by_suffixes(state_dict, suffixes, *, prefixes=None): # pylint: disable=unused-argument