mirror of
https://github.com/vladmandic/automatic
synced 2026-08-25 22:20:46 +02:00
25b7961e4e
A delta that does not fit its target module cannot apply, and applying only the layers that do fit leaves the model in a state nothing was trained for, so try_load_chain drops the whole file when any family reports a mismatch. Bias deltas were never checked against the target bias and could only surface at apply time; a module with no bias stays a non-mismatch, since whole architectures are built bias=False. - check bias deltas against the module bias in the lora, norm and full loaders - carry the mismatch count on the network so the chain can refuse the file - record refused writes in the infotext so a partial apply is not read as clean - point the krea2 full-diff test at a module that has a bias
1134 lines
48 KiB
Python
1134 lines
48 KiB
Python
"""Shared scaffolding for native adapter loaders.
|
|
|
|
Each per-arch native adapter loader implements the same algorithm:
|
|
|
|
1. Read the safetensors state dict
|
|
2. Test for family-specific markers; bail out if absent
|
|
3. Resolve the diffusers ``network_layer_mapping``
|
|
4. Group state-dict entries by ``(prefix, base)``
|
|
5. For each group, ask the arch to resolve targets (diffusers paths + chunk specs)
|
|
6. Instantiate a :class:`network.NetworkModule*` per resolved target
|
|
7. Return a populated :class:`network.Network` (or ``None`` if no matches)
|
|
|
|
This module holds the parts of that algorithm that don't vary between
|
|
architectures. Constants (suffix and marker tables), helpers (``has_marker``,
|
|
``resolve_mapping``, ``new_network``, ``finalize_network``, ``shapes_match``),
|
|
the parameterized parsing primitives (``parse_key``, ``group_by_suffixes``),
|
|
and the two cross-arch key normalizations (``unwrap_peft_wrapper`` and
|
|
``strip_peft_adapter_name``) live here.
|
|
|
|
The variance that does remain is captured by ``ChunkSpec`` (how a row range
|
|
of a fused weight is described) and the per-arch ``resolve_targets`` callable
|
|
each loader passes in (how a parsed ``(prefix, base)`` maps to one or more
|
|
diffusers paths plus optional chunk descriptors).
|
|
|
|
Per-arch loader modules import this module and pass their own ``prefixes``,
|
|
``bare_prefixes``, ``bare_diffusers_prefixes``, and ``resolve_targets`` to the
|
|
generic helpers.
|
|
"""
|
|
|
|
import os
|
|
import time
|
|
from dataclasses import dataclass
|
|
|
|
import torch
|
|
|
|
from modules import shared, sd_models, sd_models_utils
|
|
from modules.logger import log
|
|
from modules.lora import (
|
|
lora_convert, network, network_boft, network_full, network_glora,
|
|
network_hada, network_ia3, network_lokr, network_lora, network_norm,
|
|
network_oft,
|
|
)
|
|
from modules.lora import lora_common as l
|
|
|
|
|
|
# Universal prefix list shared by every native arch loader. Per-arch loaders
|
|
# extend this with arch-specific entries when their files use additional
|
|
# vendor-specific naming conventions. ``lora_transformer_`` is not a vendor
|
|
# format but sdnext's own internal transformer namespace; files already saved
|
|
# in it (e.g. OneTrainer) pass through verbatim, see :func:`resolve_group_targets`.
|
|
# ``lycoris_`` is the LyCORIS-standalone save format: the wrapped diffusers
|
|
# module path with dots rendered as underscores. It is arch-independent by
|
|
# construction (the wrapped layout equals the target layout), so it lives here
|
|
# rather than in any arch's prefix list.
|
|
KNOWN_PREFIXES_DEFAULT = ("diffusion_model.", "transformer.", "lora_unet_", "lora_transformer_", "lycoris_")
|
|
|
|
|
|
# Sentinel ``prefix_used`` value emitted by :func:`parse_key` when a bare path
|
|
# starting with a member of ``bare_diffusers_prefixes`` matches. A loader
|
|
# ``resolve_targets`` may dispatch on this string to rewrite the base path;
|
|
# when it declines, :func:`resolve_group_targets` binds the path verbatim.
|
|
BARE_DIFFUSERS_PREFIX_USED = "bare_diffusers"
|
|
|
|
|
|
# Default network-key prefix. Single-component arches keep this default;
|
|
# multi-component arches (those with separate text-encoder or adapter
|
|
# components alongside the transformer) pass a callable that picks per
|
|
# ``prefix_used``.
|
|
NETWORK_PREFIX_DEFAULT = "lora_transformer_"
|
|
|
|
|
|
# Prefixes whose parsed ``base`` is normally already a network-key tail
|
|
# (``arch_prefix + base.replace(".", "_")`` matches the stamped module name).
|
|
# :func:`resolve_group_targets` binds them verbatim unless the arch's
|
|
# ``resolve_targets`` claims the group first (needed when the arch's module
|
|
# tree diverges from the names these formats carry). ``lycoris_`` bases are
|
|
# already underscored, so the loader's ``.replace(".", "_")`` is a no-op on them.
|
|
PASSTHROUGH_PREFIXES_DEFAULT = ("transformer.", BARE_DIFFUSERS_PREFIX_USED, "lora_transformer_", "lycoris_")
|
|
|
|
|
|
def _resolve_prefix(network_prefix, prefix_used):
|
|
"""Return the network-key prefix for one parsed group.
|
|
|
|
``network_prefix`` is either a literal string (single-component arches) or
|
|
a ``Callable[[str | None], str]`` that picks based on which arch prefix
|
|
was matched (multi-component arches route to ``lora_te_`` / ``lora_llm_adapter_``
|
|
/ ``lora_transformer_`` etc.).
|
|
"""
|
|
return network_prefix(prefix_used) if callable(network_prefix) else network_prefix
|
|
|
|
|
|
SUFFIX_NORMALIZE = {
|
|
"lora_A.weight": "lora_down.weight",
|
|
"lora_B.weight": "lora_up.weight",
|
|
}
|
|
|
|
|
|
# === Family suffix tables ===
|
|
# Alpha / scale / bias / dora_scale flow into ``weights.w`` via the base
|
|
# ``network.NetworkModule.__init__`` and are listed here so they survive the
|
|
# suffix-filter pass in :func:`parse_key`.
|
|
|
|
LORA_SUFFIXES = (
|
|
".lora_down.weight", ".lora_up.weight", ".lora_mid.weight",
|
|
".lora_A.weight", ".lora_B.weight",
|
|
# 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.
|
|
# bias_indices/values/size: LyCORIS extraction sparse residual triplet,
|
|
# reconstructed into the dense-bias path by network.NetworkModule.
|
|
".alpha", ".dora_scale", ".magnitude", ".lora_magnitude_vector",
|
|
".bias", ".bias_indices", ".bias_values", ".bias_size", ".diff_b", ".scale",
|
|
)
|
|
LOKR_SUFFIXES = (
|
|
".lokr_w1", ".lokr_w2",
|
|
".lokr_w1_a", ".lokr_w1_b",
|
|
".lokr_w2_a", ".lokr_w2_b",
|
|
".lokr_t2",
|
|
".alpha", ".dora_scale", ".bias", ".scale",
|
|
)
|
|
LOHA_SUFFIXES = (
|
|
".hada_w1_a", ".hada_w1_b",
|
|
".hada_w2_a", ".hada_w2_b",
|
|
".hada_t1", ".hada_t2",
|
|
".alpha", ".dora_scale", ".bias", ".scale",
|
|
)
|
|
OFT_SUFFIXES = (
|
|
".oft_blocks", ".oft_diag",
|
|
".alpha", ".dora_scale", ".bias", ".scale",
|
|
)
|
|
IA3_SUFFIXES = (
|
|
".weight", ".on_input",
|
|
".alpha", ".scale",
|
|
)
|
|
GLORA_SUFFIXES = (
|
|
".a1.weight", ".a2.weight",
|
|
".b1.weight", ".b2.weight",
|
|
".alpha", ".dora_scale", ".scale",
|
|
)
|
|
NORM_SUFFIXES = (
|
|
".w_norm", ".b_norm",
|
|
".alpha", ".scale",
|
|
)
|
|
FULL_SUFFIXES = (
|
|
".diff", ".diff_b",
|
|
".alpha", ".scale",
|
|
)
|
|
|
|
|
|
# === Family marker tables ===
|
|
# Presence of any marker substring anywhere in a key triggers a try-load attempt
|
|
# for that family. Markers are deliberately narrower than suffixes (e.g. IA3's
|
|
# ``.on_input`` rather than ``.weight``) so :func:`has_marker` does not light up
|
|
# on accidental overlaps with other families.
|
|
|
|
LORA_MARKERS = (
|
|
".lora_down.weight", ".lora_up.weight",
|
|
".lora_A.weight", ".lora_B.weight",
|
|
# PEFT named-adapter saves embed the slot name as ``.lora_A.<name>.weight``;
|
|
# the trailing-dot forms catch every variant.
|
|
".lora_A.", ".lora_B.",
|
|
)
|
|
LOKR_MARKERS = (".lokr_w1", ".lokr_w2")
|
|
LOHA_MARKERS = (".hada_w1_a", ".hada_w1_b", ".hada_w2_a", ".hada_w2_b")
|
|
OFT_MARKERS = (".oft_blocks", ".oft_diag")
|
|
IA3_MARKERS = (".on_input",) # NOT .weight - too generic, overlaps every other family
|
|
GLORA_MARKERS = (".a1.weight", ".a2.weight", ".b1.weight", ".b2.weight")
|
|
NORM_MARKERS = (".w_norm",)
|
|
FULL_MARKERS = (".diff",)
|
|
|
|
|
|
# === Chunk descriptor ===
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ChunkSpec:
|
|
"""How to slice a fused weight along dim 0 for one target module.
|
|
|
|
Two forms supported:
|
|
|
|
- Equal chunks (``idx`` + ``total``): fused QKV split into Q/K/V via
|
|
``torch.chunk(up, total, dim=0)[idx]``. Used by flux2 / z-image where
|
|
Q, K and V have the same ``out_features``.
|
|
- Row range (``start`` + ``end``): asymmetric partition via
|
|
``up[start:end]``. Used by chroma's single-block ``linear1`` which
|
|
fuses Q / K / V / proj_mlp at unequal sizes
|
|
(``[3072, 3072, 3072, 12288]``).
|
|
|
|
Generic loaders check :attr:`is_equal_chunks` to decide between the two
|
|
forms and select the appropriate ``NetworkModule*Chunk`` /
|
|
``NetworkModule*SliceChunk`` variant.
|
|
"""
|
|
idx: int | None = None
|
|
total: int | None = None
|
|
start: int | None = None
|
|
end: int | None = None
|
|
|
|
@property
|
|
def is_equal_chunks(self) -> bool:
|
|
return self.idx is not None and self.total is not None
|
|
|
|
|
|
# === Key normalizations (applied universally by parse_key) ===
|
|
|
|
|
|
def unwrap_peft_wrapper(key):
|
|
"""Strip the ``base_model.model.`` prefix added by ``peft.save_pretrained``.
|
|
|
|
PeftModel.save_pretrained prepends this wrapper to every adapter key. The
|
|
content underneath can be any of the standard prefixes (BFL, diffusers PEFT,
|
|
kohya, bare); a single strip lets the rest of :func:`parse_key` handle the
|
|
unwrapped key normally. Mirrors the diffusers ``Flux2LoraLoaderMixin``
|
|
behavior of renaming ``base_model.model.`` to ``diffusion_model.`` before
|
|
feeding the converter.
|
|
"""
|
|
if key.startswith("base_model.model."):
|
|
return key[len("base_model.model."):]
|
|
return key
|
|
|
|
|
|
def strip_peft_adapter_name(key):
|
|
"""Normalize ``.lora_[AB].<adapter_name>.weight`` to ``.lora_[AB].weight``.
|
|
|
|
``peft.PeftModel`` and the diffusers ``save_lora_adapter`` exporter embed
|
|
the adapter slot name into the saved key (``"default"`` when not explicitly
|
|
set). Strip a single non-dotted name segment so the suffix table matches
|
|
without having to list every plausible adapter name.
|
|
"""
|
|
for inner in (".lora_A.", ".lora_B."):
|
|
idx = key.find(inner)
|
|
if idx == -1:
|
|
continue
|
|
rest = key[idx + len(inner):]
|
|
if rest == "weight" or not rest.endswith(".weight"):
|
|
continue
|
|
adapter_name = rest[:-len(".weight")]
|
|
if adapter_name and "." not in adapter_name:
|
|
return key[:idx] + inner + "weight"
|
|
return key
|
|
|
|
|
|
# === Core helpers ===
|
|
|
|
|
|
def has_marker(state_dict, markers):
|
|
"""Substring scan: does any key in ``state_dict`` contain any marker?"""
|
|
return any(any(m in k for m in markers) for k in state_dict)
|
|
|
|
|
|
def resolve_mapping():
|
|
"""Ensure ``network_layer_mapping`` is populated, return it (or empty dict)."""
|
|
sd_model = getattr(shared.sd_model, "pipe", shared.sd_model)
|
|
lora_convert.assign_network_names_to_compvis_modules(sd_model)
|
|
return getattr(shared.sd_model, "network_layer_mapping", {}) or {}
|
|
|
|
|
|
def new_network(name, network_on_disk):
|
|
"""Construct an empty :class:`network.Network` with the file's mtime stamped."""
|
|
net = network.Network(name, network_on_disk)
|
|
net.mtime = os.path.getmtime(network_on_disk.filename)
|
|
return net
|
|
|
|
|
|
def finalize_network(net, name, family, lora_scale, t0, unmapped=0, mismatch=0, skipped=0):
|
|
"""Emit the standard debug log line and return the populated network (or ``None``).
|
|
|
|
Returns ``None`` when no modules were bound. Records ``mismatch`` on the
|
|
network so :func:`try_load_chain` can refuse the file as a whole.
|
|
"""
|
|
net.mismatch = mismatch
|
|
if len(net.modules) == 0:
|
|
if mismatch:
|
|
log.error(
|
|
f'Network load: type={family} name="{name}" native no-match'
|
|
f' unmapped={unmapped} mismatch={mismatch} skipped={skipped}'
|
|
)
|
|
elif unmapped or skipped:
|
|
log.debug(
|
|
f'Network load: type={family} name="{name}" native no-match'
|
|
f' unmapped={unmapped} skipped={skipped}'
|
|
)
|
|
return None
|
|
log.debug(
|
|
f'Network load: type={family} name="{name}" native modules={len(net.modules)}'
|
|
f' unmapped={unmapped} mismatch={mismatch} skipped={skipped} scale={lora_scale}'
|
|
)
|
|
l.timer.activate += time.time() - t0
|
|
return net
|
|
|
|
|
|
def shapes_match(sd_module, down_w: torch.Tensor, up_w: torch.Tensor) -> bool:
|
|
"""LoRA-style rank-and-dim sanity check against the live module weight.
|
|
|
|
Honors SDNQ-quantized modules by reading the original shape from the
|
|
dequantizer rather than the packed weight tensor.
|
|
"""
|
|
if not hasattr(sd_module, "weight"):
|
|
return False
|
|
if hasattr(sd_module, "sdnq_dequantizer"):
|
|
mod_shape = sd_module.sdnq_dequantizer.original_shape
|
|
else:
|
|
mod_shape = sd_module.weight.shape
|
|
if len(mod_shape) < 2 or len(down_w.shape) < 2 or len(up_w.shape) < 2:
|
|
return False
|
|
return down_w.shape[1] == mod_shape[1] and up_w.shape[0] == mod_shape[0]
|
|
|
|
|
|
def lokr_kron_shape(w):
|
|
"""Return the ``(out, in_flat)`` shape of ``kron(w1, w2)`` from stored factors.
|
|
|
|
Each factor is either full (``lokr_w1``/``lokr_w2``) or rank-decomposed
|
|
(``lokr_w1_a @ lokr_w1_b``); a Tucker-rebuilt w2 (conv-only) stores its
|
|
parts as ``(rank, part)`` with the kernel dims carried by ``lokr_t2``.
|
|
Kernel dims are folded into ``in_flat``, matching how the delta reshapes
|
|
onto a conv weight.
|
|
"""
|
|
w1 = w.get("lokr_w1")
|
|
r1 = w1.shape[0] if w1 is not None else w["lokr_w1_a"].shape[0]
|
|
c1 = w1.shape[1] if w1 is not None else w["lokr_w1_b"].shape[1]
|
|
w2 = w.get("lokr_w2")
|
|
t2 = w.get("lokr_t2")
|
|
if w2 is not None:
|
|
r2 = w2.shape[0]
|
|
c2_flat = w2.numel() // r2
|
|
elif t2 is not None:
|
|
r2 = w["lokr_w2_a"].shape[1]
|
|
c2_flat = w["lokr_w2_b"].shape[1]
|
|
for d in t2.shape[2:]:
|
|
c2_flat *= d
|
|
else:
|
|
r2 = w["lokr_w2_a"].shape[0]
|
|
w2b = w["lokr_w2_b"]
|
|
c2_flat = w2b.numel() // w2b.shape[0]
|
|
return r1 * r2, c1 * c2_flat
|
|
|
|
|
|
def bias_delta_fits(sd_module, bias_w: torch.Tensor) -> bool:
|
|
"""Bias-delta sanity check against the live module bias.
|
|
|
|
A module with no bias is not a mismatch: whole architectures are built
|
|
``bias=False`` (flux2 has not one biased module), so a stray delta there is
|
|
one unappliable key and the apply pass counts it refused.
|
|
"""
|
|
bias = getattr(sd_module, "bias", None)
|
|
if bias is None:
|
|
return True
|
|
return tuple(bias_w.shape) == tuple(bias.shape)
|
|
|
|
|
|
def lokr_shapes_match(sd_module, kron_shape, chunk: ChunkSpec | None) -> bool:
|
|
"""Kron-vs-module dim check, honoring SDNQ original shapes and chunk rows.
|
|
|
|
The input dim is never chunked; the output dim must cover the full fused
|
|
weight for chunked targets (``total * out`` for equal chunks, ``end <=
|
|
kron_out`` with an exact row-range for slices).
|
|
"""
|
|
if not hasattr(sd_module, "weight"):
|
|
return False
|
|
if hasattr(sd_module, "sdnq_dequantizer"):
|
|
mod_shape = sd_module.sdnq_dequantizer.original_shape
|
|
else:
|
|
mod_shape = sd_module.weight.shape
|
|
if len(mod_shape) < 2:
|
|
return False
|
|
mod_in_flat = 1
|
|
for d in mod_shape[1:]:
|
|
mod_in_flat *= d
|
|
kron_out, kron_in_flat = kron_shape
|
|
if kron_in_flat != mod_in_flat:
|
|
return False
|
|
if chunk is None:
|
|
return kron_out == mod_shape[0]
|
|
if chunk.is_equal_chunks:
|
|
return kron_out == mod_shape[0] * chunk.total
|
|
return (chunk.end - chunk.start) == mod_shape[0] and kron_out >= chunk.end
|
|
|
|
|
|
# === Parsing primitives ===
|
|
|
|
|
|
def parse_key(key, suffixes, *, prefixes=KNOWN_PREFIXES_DEFAULT, bare_prefixes=(), bare_diffusers_prefixes=()):
|
|
"""Return ``(prefix_used, base, suffix_normalized)`` or ``None``.
|
|
|
|
``prefix_used`` is the matched element of ``prefixes``, ``BARE_DIFFUSERS_PREFIX_USED``
|
|
if a member of ``bare_diffusers_prefixes`` matched, or ``None`` for a key
|
|
that matched a member of ``bare_prefixes``. ``base`` is the path with prefix
|
|
and suffix removed. ``suffix_normalized`` is the suffix (without the leading
|
|
dot) after applying :data:`SUFFIX_NORMALIZE` (e.g. ``lora_A.weight`` becomes
|
|
``lora_down.weight``).
|
|
|
|
Always applies :func:`unwrap_peft_wrapper` and :func:`strip_peft_adapter_name`
|
|
to the raw key before format detection so callers do not have to opt in.
|
|
"""
|
|
key = unwrap_peft_wrapper(key)
|
|
key = strip_peft_adapter_name(key)
|
|
prefix_used = None
|
|
stripped = key
|
|
for p in prefixes:
|
|
if key.startswith(p):
|
|
prefix_used = p
|
|
stripped = key[len(p):]
|
|
break
|
|
if prefix_used is None:
|
|
if any(key.startswith(p) for p in bare_diffusers_prefixes):
|
|
prefix_used = BARE_DIFFUSERS_PREFIX_USED
|
|
elif not any(key.startswith(p) for p in bare_prefixes):
|
|
return None
|
|
|
|
matched_suffix = None
|
|
split_at = -1
|
|
for marker in suffixes:
|
|
if stripped.endswith(marker):
|
|
split_at = len(stripped) - len(marker)
|
|
matched_suffix = marker.lstrip(".")
|
|
break
|
|
if split_at < 0:
|
|
return None
|
|
|
|
base = stripped[:split_at]
|
|
if not base:
|
|
return None
|
|
|
|
suffix = SUFFIX_NORMALIZE.get(matched_suffix, matched_suffix)
|
|
return prefix_used, base, suffix
|
|
|
|
|
|
def group_by_suffixes(state_dict, suffixes, *, prefixes=KNOWN_PREFIXES_DEFAULT, bare_prefixes=(), bare_diffusers_prefixes=()):
|
|
"""Group state-dict entries by ``(prefix_used, base)``.
|
|
|
|
Returns ``{(prefix_used, base): {suffix: tensor, ...}}`` where each suffix
|
|
is the normalized form produced by :func:`parse_key`. Per-family loaders
|
|
apply their own key-presence gates on each group (e.g. LoRA requires both
|
|
``lora_down.weight`` and ``lora_up.weight``).
|
|
"""
|
|
groups: dict[tuple, dict[str, torch.Tensor]] = {}
|
|
for key, value in state_dict.items():
|
|
parsed = parse_key(
|
|
key, suffixes,
|
|
prefixes=prefixes,
|
|
bare_prefixes=bare_prefixes,
|
|
bare_diffusers_prefixes=bare_diffusers_prefixes,
|
|
)
|
|
if parsed is None:
|
|
continue
|
|
prefix_used, base, suffix = parsed
|
|
slot = groups.get((prefix_used, base))
|
|
if slot is None:
|
|
slot = {}
|
|
groups[(prefix_used, base)] = slot
|
|
slot[suffix] = value
|
|
return groups
|
|
|
|
|
|
# Surface ``sd_models.read_state_dict`` here so loader modules don't have to
|
|
# import ``sd_models`` directly; keeps the per-arch wrapper imports compact.
|
|
read_state_dict = sd_models.read_state_dict
|
|
|
|
|
|
def resolve_group_targets(resolve_targets, prefix_used, base):
|
|
"""Map a parsed ``(prefix_used, base)`` group to ``[(diffusers_path, chunk), ...]``.
|
|
|
|
The arch's ``resolve_targets`` is consulted first for every group.
|
|
Passthrough prefixes (:data:`PASSTHROUGH_PREFIXES_DEFAULT`) fall back to
|
|
verbatim binding when the resolver declines (returns empty), so arches
|
|
whose module tree matches the stamped names need no rewrite branch, while
|
|
arches with a divergent tree (e.g. krea2, whose transformer keeps checkpoint
|
|
names that differ from the ecosystem's diffusers names) can rewrite them.
|
|
"""
|
|
if prefix_used in PASSTHROUGH_PREFIXES_DEFAULT:
|
|
return resolve_targets(prefix_used, base) or [(base, None)]
|
|
return resolve_targets(prefix_used, base)
|
|
|
|
|
|
# === Generic family loaders ===
|
|
#
|
|
# Each loader takes a per-arch ``resolve_targets`` callable returning a list of
|
|
# ``(diffusers_path, ChunkSpec | None)`` tuples for each parsed group. The
|
|
# loader builds the network key as ``network_prefix + path.replace(".", "_")``
|
|
# where ``network_prefix`` defaults to ``"lora_transformer_"`` and may be
|
|
# either a literal string (single-component arches) or a callable that picks
|
|
# per ``prefix_used`` (multi-component arches such as Anima that route to
|
|
# ``lora_te_`` or ``lora_llm_adapter_`` based on which arch prefix matched).
|
|
# The loader then instantiates the appropriate ``network.NetworkModule*`` subclass.
|
|
#
|
|
# Fused-target handling per family:
|
|
#
|
|
# - LoRA: chunk at load time (lora_up is sliced along dim 0). Both equal and
|
|
# unequal ChunkSpec shapes are supported.
|
|
# - LoKR: defer to NetworkModuleLokrChunk (equal) or NetworkModuleLokrSliceChunk
|
|
# (unequal) which materialize the Kronecker product once and return the
|
|
# designated slice.
|
|
# - LoHA: only equal ChunkSpec is supported via NetworkModuleHadaChunk. Unequal
|
|
# slices and Tucker-decomposed LoHAs on fused targets are skipped with a
|
|
# warning (no slice variant exists, and Tucker keys cannot arise on Linear
|
|
# layers per LyCORIS upstream — see network_hada.NetworkModuleHadaChunk).
|
|
#
|
|
# DoRA on fused targets (LoRA/LoKR/LoHA): per-output dora_scale rows are
|
|
# sliced with the chunk via slice_dora_scale (exact, row norms are
|
|
# independent); per-input dora_scale couples the chunks through shared
|
|
# column norms and the group is skipped with a warning.
|
|
#
|
|
# Bias companions on fused targets: per-output diff_b deltas slice with the
|
|
# chunk (LoRA only; no other family's save format pairs them). The legacy
|
|
# weight-shaped "bias" key has no defined partition and skips the group.
|
|
# - OFT/BOFT: no chunk variant exists; fused targets are skipped with a
|
|
# warning. Discrimination is by ``oft_blocks.ndim`` (3-D OFT, 4-D BOFT),
|
|
# mirroring upstream LyCORIS ``algo_check``.
|
|
|
|
|
|
def slice_chunk_rows(t, chunk: ChunkSpec):
|
|
"""Slice dim 0 of ``t`` per ``chunk``.
|
|
|
|
Equal-chunks form uses ``torch.chunk`` (faster for the symmetric case);
|
|
row-range form uses tensor slicing for arbitrary partitions.
|
|
"""
|
|
if chunk.is_equal_chunks:
|
|
return torch.chunk(t, chunk.total, dim=0)[chunk.idx].contiguous()
|
|
return t[chunk.start:chunk.end].contiguous()
|
|
|
|
|
|
def _slice_lora_chunk(w, chunk: ChunkSpec):
|
|
"""Return a shallow copy of ``w`` with ``lora_up.weight`` sliced per ``chunk``."""
|
|
out = dict(w)
|
|
out["lora_up.weight"] = slice_chunk_rows(w["lora_up.weight"], chunk)
|
|
return out
|
|
|
|
|
|
def slice_dora_scale(w, chunk: ChunkSpec, fused_out):
|
|
"""Slice a per-output ``dora_scale`` to the chunk rows; ``None`` if unsliceable.
|
|
|
|
LyCORIS ``wd_on_out=True`` (the LoKr/LoHA default) stores per-output
|
|
magnitudes of shape ``(out, 1)`` (1-D ``(out,)`` also seen); the rows
|
|
partition exactly with the fused weight, so the chunk slice preserves the
|
|
DoRA math (row norms are independent across rows). ``wd_on_out=False``
|
|
stores per-input magnitudes whose column norms span every fused row,
|
|
coupling the chunks; no exact per-chunk equivalent exists and the caller
|
|
must skip the group.
|
|
"""
|
|
ds = w.get("dora_scale")
|
|
if ds is None:
|
|
return w
|
|
if ds.ndim >= 1 and ds.shape[0] == fused_out:
|
|
out = dict(w)
|
|
out["dora_scale"] = slice_chunk_rows(ds, chunk)
|
|
return out
|
|
return None
|
|
|
|
|
|
def slice_bias_delta(w, chunk: ChunkSpec, fused_out):
|
|
"""Slice a per-output ``diff_b`` bias delta to the chunk rows; ``None`` if unsliceable.
|
|
|
|
``diff_b`` stores one bias value per output feature, so it partitions with
|
|
the fused rows exactly like the up-weight.
|
|
"""
|
|
db = w.get("diff_b")
|
|
if db is None:
|
|
return w
|
|
if db.ndim >= 1 and db.shape[0] == fused_out:
|
|
out = dict(w)
|
|
out["diff_b"] = slice_chunk_rows(db, chunk)
|
|
return out
|
|
return None
|
|
|
|
|
|
def try_load_lora(name, network_on_disk, lora_scale, *,
|
|
resolve_targets, prefixes=KNOWN_PREFIXES_DEFAULT,
|
|
bare_prefixes=(), bare_diffusers_prefixes=(),
|
|
network_prefix=NETWORK_PREFIX_DEFAULT,
|
|
arch_name="generic"):
|
|
"""Generic LoRA loader (handles DoRA via the universal ``finalize_updown`` hook).
|
|
|
|
Fused targets are chunked at load time by slicing ``lora_up`` along dim 0;
|
|
the down-side is shared across the resolved targets.
|
|
"""
|
|
t0 = time.time()
|
|
state_dict = read_state_dict(network_on_disk.filename, what="network")
|
|
if not has_marker(state_dict, LORA_MARKERS):
|
|
return None
|
|
|
|
mapping = resolve_mapping()
|
|
net = new_network(name, network_on_disk)
|
|
groups = group_by_suffixes(
|
|
state_dict, LORA_SUFFIXES,
|
|
prefixes=prefixes,
|
|
bare_prefixes=bare_prefixes,
|
|
bare_diffusers_prefixes=bare_diffusers_prefixes,
|
|
)
|
|
|
|
unmapped = 0
|
|
mismatch = 0
|
|
skipped = 0
|
|
for (prefix, base), w in groups.items():
|
|
if "lora_down.weight" not in w or "lora_up.weight" not in w:
|
|
continue
|
|
# DoRA magnitude vectors: ai-toolkit saves `magnitude`, PEFT/diffusers
|
|
# `lora_magnitude_vector`. Both are 1-D per-output row norms with
|
|
# dora_scale semantics; reshape to (out, 1) so the apply-time
|
|
# orientation detection cannot misread square layers as per-input.
|
|
for mag_key in ("magnitude", "lora_magnitude_vector"):
|
|
mag = w.get(mag_key)
|
|
if mag is not None and "dora_scale" not in w:
|
|
w = dict(w)
|
|
w.pop(mag_key)
|
|
w["dora_scale"] = mag.reshape(-1, 1) if mag.ndim == 1 else mag
|
|
arch_prefix = _resolve_prefix(network_prefix, prefix)
|
|
for diffusers_path, chunk in resolve_group_targets(resolve_targets, prefix, base):
|
|
network_key = arch_prefix + diffusers_path.replace(".", "_")
|
|
sd_module = mapping.get(network_key)
|
|
if sd_module is None:
|
|
unmapped += 1
|
|
continue
|
|
|
|
target_w = w
|
|
if chunk is not None:
|
|
if "bias" in w or "bias_indices" in w:
|
|
# Weight-shaped bias residuals (dense or LyCORIS sparse
|
|
# triplet) are not partitioned onto fused targets.
|
|
log.warning(f'Network load: type=LoRA name="{name}" arch={arch_name} key={network_key} weight-shaped bias on fused target skipped (unsupported)')
|
|
skipped += 1
|
|
continue
|
|
fused_out = w["lora_up.weight"].shape[0]
|
|
target_w = _slice_lora_chunk(w, chunk)
|
|
target_w = slice_dora_scale(target_w, chunk, fused_out)
|
|
if target_w is None:
|
|
log.warning(f'Network load: type=LoRA name="{name}" arch={arch_name} key={network_key} per-input DoRA on fused target skipped (unsupported)')
|
|
skipped += 1
|
|
continue
|
|
target_w = slice_bias_delta(target_w, chunk, fused_out)
|
|
if target_w is None:
|
|
log.warning(f'Network load: type=LoRA name="{name}" arch={arch_name} key={network_key} non-per-output diff_b on fused target skipped (unsupported)')
|
|
skipped += 1
|
|
continue
|
|
|
|
if not shapes_match(sd_module, target_w["lora_down.weight"], target_w["lora_up.weight"]):
|
|
log.warning(
|
|
f'Network load: type=LoRA name="{name}" arch={arch_name} key={network_key}'
|
|
f' lora={target_w["lora_down.weight"].shape[1]}x{target_w["lora_up.weight"].shape[0]}'
|
|
f' module={getattr(sd_module, "weight", None).shape if hasattr(sd_module, "weight") else "?"}'
|
|
f' shape mismatch'
|
|
)
|
|
mismatch += 1
|
|
continue
|
|
|
|
if "diff_b" in target_w and not bias_delta_fits(sd_module, target_w["diff_b"]):
|
|
log.warning(
|
|
f'Network load: type=LoRA name="{name}" arch={arch_name} key={network_key}'
|
|
f' bias={tuple(target_w["diff_b"].shape)} module={tuple(sd_module.bias.shape)}'
|
|
f' bias shape mismatch'
|
|
)
|
|
mismatch += 1
|
|
continue
|
|
|
|
nw = network.NetworkWeights(network_key=network_key, sd_key=network_key, w=target_w, sd_module=sd_module)
|
|
net.modules[network_key] = network_lora.NetworkModuleLora(net, nw)
|
|
|
|
return finalize_network(net, name, "LoRA", lora_scale, t0, unmapped=unmapped, mismatch=mismatch, skipped=skipped)
|
|
|
|
|
|
def try_load_lokr(name, network_on_disk, lora_scale, *,
|
|
resolve_targets, prefixes=KNOWN_PREFIXES_DEFAULT,
|
|
bare_prefixes=(), bare_diffusers_prefixes=(),
|
|
network_prefix=NETWORK_PREFIX_DEFAULT,
|
|
arch_name="generic"):
|
|
"""Generic LoKR loader.
|
|
|
|
Stores only the compact LoKR factors and dispatches to
|
|
:class:`network_lokr.NetworkModuleLokrChunk` (equal chunks) or
|
|
:class:`network_lokr.NetworkModuleLokrSliceChunk` (row range) at apply
|
|
time. Both materialize ``kron(w1, w2)`` once per forward pass and return
|
|
the designated slice; full materialization happens lazily inside the
|
|
module rather than at load.
|
|
"""
|
|
t0 = time.time()
|
|
state_dict = read_state_dict(network_on_disk.filename, what="network")
|
|
if not has_marker(state_dict, LOKR_MARKERS):
|
|
return None
|
|
|
|
mapping = resolve_mapping()
|
|
net = new_network(name, network_on_disk)
|
|
groups = group_by_suffixes(
|
|
state_dict, LOKR_SUFFIXES,
|
|
prefixes=prefixes,
|
|
bare_prefixes=bare_prefixes,
|
|
bare_diffusers_prefixes=bare_diffusers_prefixes,
|
|
)
|
|
|
|
unmapped = 0
|
|
mismatch = 0
|
|
skipped = 0
|
|
for (prefix, base), w in groups.items():
|
|
has_1 = "lokr_w1" in w or ("lokr_w1_a" in w and "lokr_w1_b" in w)
|
|
has_2 = "lokr_w2" in w or ("lokr_w2_a" in w and "lokr_w2_b" in w)
|
|
if not (has_1 and has_2):
|
|
continue
|
|
arch_prefix = _resolve_prefix(network_prefix, prefix)
|
|
kron_shape = lokr_kron_shape(w)
|
|
for diffusers_path, chunk in resolve_group_targets(resolve_targets, prefix, base):
|
|
network_key = arch_prefix + diffusers_path.replace(".", "_")
|
|
sd_module = mapping.get(network_key)
|
|
if sd_module is None:
|
|
unmapped += 1
|
|
continue
|
|
if not lokr_shapes_match(sd_module, kron_shape, chunk):
|
|
log.warning(
|
|
f'Network load: type=LoKR name="{name}" arch={arch_name} key={network_key}'
|
|
f' kron={kron_shape[0]}x{kron_shape[1]}'
|
|
f' module={getattr(sd_module, "weight", None).shape if hasattr(sd_module, "weight") else "?"}'
|
|
f' shape mismatch'
|
|
)
|
|
mismatch += 1
|
|
continue
|
|
target_w = w
|
|
if chunk is not None:
|
|
if "bias" in w:
|
|
log.warning(f'Network load: type=LoKR name="{name}" arch={arch_name} key={network_key} weight-shaped bias on fused target skipped (unsupported)')
|
|
skipped += 1
|
|
continue
|
|
# Kron rows = w1 rows * w2 rows; the tucker w2_a orientation
|
|
# differs but tucker is conv-only and chunks are Linear-only.
|
|
w1 = w.get("lokr_w1")
|
|
w2 = w.get("lokr_w2")
|
|
w1_rows = w1.shape[0] if w1 is not None else w["lokr_w1_a"].shape[0]
|
|
w2_rows = w2.shape[0] if w2 is not None else w["lokr_w2_a"].shape[0]
|
|
target_w = slice_dora_scale(w, chunk, fused_out=w1_rows * w2_rows)
|
|
if target_w is None:
|
|
log.warning(f'Network load: type=LoKR name="{name}" arch={arch_name} key={network_key} per-input DoRA on fused target skipped (unsupported)')
|
|
skipped += 1
|
|
continue
|
|
nw = network.NetworkWeights(network_key=network_key, sd_key=network_key, w=target_w, sd_module=sd_module)
|
|
if chunk is None:
|
|
net.modules[network_key] = network_lokr.NetworkModuleLokr(net, nw)
|
|
elif chunk.is_equal_chunks:
|
|
net.modules[network_key] = network_lokr.NetworkModuleLokrChunk(net, nw, chunk.idx, chunk.total)
|
|
else:
|
|
net.modules[network_key] = network_lokr.NetworkModuleLokrSliceChunk(net, nw, chunk.start, chunk.end)
|
|
|
|
return finalize_network(net, name, "LoKR", lora_scale, t0, unmapped=unmapped, mismatch=mismatch, skipped=skipped)
|
|
|
|
|
|
def try_load_loha(name, network_on_disk, lora_scale, *,
|
|
resolve_targets, prefixes=KNOWN_PREFIXES_DEFAULT,
|
|
bare_prefixes=(), bare_diffusers_prefixes=(),
|
|
network_prefix=NETWORK_PREFIX_DEFAULT,
|
|
arch_name="generic"):
|
|
"""Generic LoHA (Hadamard product) loader.
|
|
|
|
Standard non-Tucker LoHA on fused targets uses
|
|
:class:`network_hada.NetworkModuleHadaChunk` for equal-chunks dispatch.
|
|
Tucker-decomposed LoHAs on fused targets and any unequal-chunks dispatch
|
|
are skipped with a warning: no slice variant exists, and Tucker keys
|
|
cannot arise on Linear layers per LyCORIS upstream.
|
|
"""
|
|
t0 = time.time()
|
|
state_dict = read_state_dict(network_on_disk.filename, what="network")
|
|
if not has_marker(state_dict, LOHA_MARKERS):
|
|
return None
|
|
|
|
mapping = resolve_mapping()
|
|
net = new_network(name, network_on_disk)
|
|
groups = group_by_suffixes(
|
|
state_dict, LOHA_SUFFIXES,
|
|
prefixes=prefixes,
|
|
bare_prefixes=bare_prefixes,
|
|
bare_diffusers_prefixes=bare_diffusers_prefixes,
|
|
)
|
|
|
|
unmapped = 0
|
|
skipped = 0
|
|
for (prefix, base), w in groups.items():
|
|
if not all(k in w for k in ("hada_w1_a", "hada_w1_b", "hada_w2_a", "hada_w2_b")):
|
|
continue
|
|
is_tucker = "hada_t1" in w or "hada_t2" in w
|
|
targets = resolve_group_targets(resolve_targets, prefix, base)
|
|
is_fused = any(t[1] is not None for t in targets)
|
|
if is_fused and is_tucker:
|
|
log.warning(f'Network load: type=LoHA name="{name}" arch={arch_name} key={base} Tucker fused QKV skipped (unsupported)')
|
|
skipped += 1
|
|
continue
|
|
arch_prefix = _resolve_prefix(network_prefix, prefix)
|
|
for diffusers_path, chunk in targets:
|
|
network_key = arch_prefix + diffusers_path.replace(".", "_")
|
|
sd_module = mapping.get(network_key)
|
|
if sd_module is None:
|
|
unmapped += 1
|
|
continue
|
|
target_w = w
|
|
if chunk is not None:
|
|
if "bias" in w:
|
|
log.warning(f'Network load: type=LoHA name="{name}" arch={arch_name} key={network_key} weight-shaped bias on fused target skipped (unsupported)')
|
|
skipped += 1
|
|
continue
|
|
target_w = slice_dora_scale(w, chunk, fused_out=w["hada_w1_a"].shape[0])
|
|
if target_w is None:
|
|
log.warning(f'Network load: type=LoHA name="{name}" arch={arch_name} key={network_key} per-input DoRA on fused target skipped (unsupported)')
|
|
skipped += 1
|
|
continue
|
|
nw = network.NetworkWeights(network_key=network_key, sd_key=network_key, w=target_w, sd_module=sd_module)
|
|
if chunk is None:
|
|
net.modules[network_key] = network_hada.NetworkModuleHada(net, nw)
|
|
elif chunk.is_equal_chunks:
|
|
net.modules[network_key] = network_hada.NetworkModuleHadaChunk(net, nw, chunk.idx, chunk.total)
|
|
else:
|
|
log.warning(f'Network load: type=LoHA name="{name}" arch={arch_name} key={network_key} unequal fused chunks unsupported')
|
|
skipped += 1
|
|
|
|
return finalize_network(net, name, "LoHA", lora_scale, t0, unmapped=unmapped, skipped=skipped)
|
|
|
|
|
|
def try_load_oft(name, network_on_disk, lora_scale, *,
|
|
resolve_targets, prefixes=KNOWN_PREFIXES_DEFAULT,
|
|
bare_prefixes=(), bare_diffusers_prefixes=(),
|
|
network_prefix=NETWORK_PREFIX_DEFAULT,
|
|
arch_name="generic"):
|
|
"""Generic OFT/BOFT loader.
|
|
|
|
OFT and BOFT share the ``oft_blocks`` save key and are discriminated by
|
|
tensor dimensionality (3-D OFT, 4-D BOFT), mirroring upstream
|
|
``algo_check``. Both kohya (``oft_blocks`` + alpha-as-constraint) and
|
|
LyCORIS (``oft_diag``) OFT layouts route through
|
|
:class:`network_oft.NetworkModuleOFT`; BOFT routes through
|
|
:class:`network_boft.NetworkModuleBOFT`.
|
|
|
|
Fused targets are skipped with a warning for both algorithms: OFT block
|
|
structure (and BOFT's per-stage block partition) is tied to the target
|
|
module's ``out_features``, so a per-Q/K/V split would require re-deriving
|
|
the rotations per chunk.
|
|
"""
|
|
t0 = time.time()
|
|
state_dict = read_state_dict(network_on_disk.filename, what="network")
|
|
if not has_marker(state_dict, OFT_MARKERS):
|
|
return None
|
|
|
|
mapping = resolve_mapping()
|
|
net = new_network(name, network_on_disk)
|
|
groups = group_by_suffixes(
|
|
state_dict, OFT_SUFFIXES,
|
|
prefixes=prefixes,
|
|
bare_prefixes=bare_prefixes,
|
|
bare_diffusers_prefixes=bare_diffusers_prefixes,
|
|
)
|
|
|
|
unmapped = 0
|
|
skipped = 0
|
|
for (prefix, base), w in groups.items():
|
|
if not ("oft_blocks" in w or "oft_diag" in w):
|
|
continue
|
|
is_boft = "oft_blocks" in w and w["oft_blocks"].ndim == 4
|
|
targets = resolve_group_targets(resolve_targets, prefix, base)
|
|
if any(t[1] is not None for t in targets):
|
|
log.warning(f'Network load: type={"BOFT" if is_boft else "OFT"} name="{name}" arch={arch_name} key={base} fused QKV skipped (unsupported)')
|
|
skipped += 1
|
|
continue
|
|
arch_prefix = _resolve_prefix(network_prefix, prefix)
|
|
for diffusers_path, _ in targets:
|
|
network_key = arch_prefix + diffusers_path.replace(".", "_")
|
|
sd_module = mapping.get(network_key)
|
|
if sd_module is None:
|
|
unmapped += 1
|
|
continue
|
|
nw = network.NetworkWeights(network_key=network_key, sd_key=network_key, w=w, sd_module=sd_module)
|
|
if is_boft:
|
|
net.modules[network_key] = network_boft.NetworkModuleBOFT(net, nw)
|
|
else:
|
|
net.modules[network_key] = network_oft.NetworkModuleOFT(net, nw)
|
|
|
|
return finalize_network(net, name, "OFT", lora_scale, t0, unmapped=unmapped, skipped=skipped)
|
|
|
|
|
|
def try_load_ia3(name, network_on_disk, lora_scale, *,
|
|
resolve_targets, prefixes=KNOWN_PREFIXES_DEFAULT,
|
|
bare_prefixes=(), bare_diffusers_prefixes=(),
|
|
network_prefix=NETWORK_PREFIX_DEFAULT,
|
|
arch_name="generic"):
|
|
"""Generic IA3 loader.
|
|
|
|
IA3 stores a per-row or per-column scale vector keyed under ``.weight``
|
|
plus an ``.on_input`` flag selecting which axis. ``.weight`` alone is too
|
|
generic for the marker scan (it overlaps every other family's
|
|
``.lora_down.weight`` / ``.hada_w*`` keys), so the marker gate insists on
|
|
``.on_input`` while the suffix table includes both.
|
|
|
|
Fused targets are skipped with a warning. There is no real-world IA3-on-DiT
|
|
prevalence to justify the asymmetry between ``on_input=True`` (which would
|
|
replicate cleanly across Q/K/V) and ``on_input=False`` (which would need
|
|
output-axis slicing).
|
|
"""
|
|
t0 = time.time()
|
|
state_dict = read_state_dict(network_on_disk.filename, what="network")
|
|
if not has_marker(state_dict, IA3_MARKERS):
|
|
return None
|
|
|
|
mapping = resolve_mapping()
|
|
net = new_network(name, network_on_disk)
|
|
groups = group_by_suffixes(
|
|
state_dict, IA3_SUFFIXES,
|
|
prefixes=prefixes,
|
|
bare_prefixes=bare_prefixes,
|
|
bare_diffusers_prefixes=bare_diffusers_prefixes,
|
|
)
|
|
|
|
unmapped = 0
|
|
skipped = 0
|
|
for (prefix, base), w in groups.items():
|
|
if not ("weight" in w and "on_input" in w):
|
|
continue
|
|
targets = resolve_group_targets(resolve_targets, prefix, base)
|
|
if any(t[1] is not None for t in targets):
|
|
log.warning(f'Network load: type=IA3 name="{name}" arch={arch_name} key={base} fused QKV skipped (unsupported)')
|
|
skipped += 1
|
|
continue
|
|
arch_prefix = _resolve_prefix(network_prefix, prefix)
|
|
for diffusers_path, _ in targets:
|
|
network_key = arch_prefix + diffusers_path.replace(".", "_")
|
|
sd_module = mapping.get(network_key)
|
|
if sd_module is None:
|
|
unmapped += 1
|
|
continue
|
|
nw = network.NetworkWeights(network_key=network_key, sd_key=network_key, w=w, sd_module=sd_module)
|
|
net.modules[network_key] = network_ia3.NetworkModuleIa3(net, nw)
|
|
|
|
return finalize_network(net, name, "IA3", lora_scale, t0, unmapped=unmapped, skipped=skipped)
|
|
|
|
|
|
def try_load_glora(name, network_on_disk, lora_scale, *,
|
|
resolve_targets, prefixes=KNOWN_PREFIXES_DEFAULT,
|
|
bare_prefixes=(), bare_diffusers_prefixes=(),
|
|
network_prefix=NETWORK_PREFIX_DEFAULT,
|
|
arch_name="generic"):
|
|
"""Generic GLoRA loader.
|
|
|
|
GLoRA stores four low-rank components (``a1`` / ``a2`` / ``b1`` / ``b2``)
|
|
and computes ``W_delta = w2b @ w1b + (target @ w2a) @ w1a``; the second
|
|
term is target-dependent. Fused targets are skipped: the target-dependent
|
|
term doesn't slice cleanly across projections without redirecting
|
|
calc_updown to a fused proxy weight, and the file pattern is vanishingly
|
|
rare on DiT architectures.
|
|
"""
|
|
t0 = time.time()
|
|
state_dict = read_state_dict(network_on_disk.filename, what="network")
|
|
if not has_marker(state_dict, GLORA_MARKERS):
|
|
return None
|
|
|
|
mapping = resolve_mapping()
|
|
net = new_network(name, network_on_disk)
|
|
groups = group_by_suffixes(
|
|
state_dict, GLORA_SUFFIXES,
|
|
prefixes=prefixes,
|
|
bare_prefixes=bare_prefixes,
|
|
bare_diffusers_prefixes=bare_diffusers_prefixes,
|
|
)
|
|
|
|
unmapped = 0
|
|
skipped = 0
|
|
for (prefix, base), w in groups.items():
|
|
if not all(k in w for k in ("a1.weight", "a2.weight", "b1.weight", "b2.weight")):
|
|
continue
|
|
targets = resolve_group_targets(resolve_targets, prefix, base)
|
|
if any(t[1] is not None for t in targets):
|
|
log.warning(f'Network load: type=GLoRA name="{name}" arch={arch_name} key={base} fused QKV skipped (unsupported)')
|
|
skipped += 1
|
|
continue
|
|
arch_prefix = _resolve_prefix(network_prefix, prefix)
|
|
for diffusers_path, _ in targets:
|
|
network_key = arch_prefix + diffusers_path.replace(".", "_")
|
|
sd_module = mapping.get(network_key)
|
|
if sd_module is None:
|
|
unmapped += 1
|
|
continue
|
|
nw = network.NetworkWeights(network_key=network_key, sd_key=network_key, w=w, sd_module=sd_module)
|
|
net.modules[network_key] = network_glora.NetworkModuleGLora(net, nw)
|
|
|
|
return finalize_network(net, name, "GLoRA", lora_scale, t0, unmapped=unmapped, skipped=skipped)
|
|
|
|
|
|
def try_load_norm(name, network_on_disk, lora_scale, *,
|
|
resolve_targets, prefixes=KNOWN_PREFIXES_DEFAULT,
|
|
bare_prefixes=(), bare_diffusers_prefixes=(),
|
|
network_prefix=NETWORK_PREFIX_DEFAULT,
|
|
arch_name="generic"): # pylint: disable=unused-argument
|
|
"""Generic Norm (LayerNorm / RMSNorm weight + bias delta) loader.
|
|
|
|
Norm targets are never fused, so the chunk dispatch is dropped.
|
|
|
|
Loader-local stamping: ``lora_convert.assign_network_names_to_compvis_modules``
|
|
deliberately skips setting ``module.network_layer_name`` for transformer
|
|
norm modules (except SD3) because of legacy CompVis UNet collisions. This
|
|
loader bypasses the guard locally: for each target it actually binds, it
|
|
sets ``network_layer_name`` directly on the host module so
|
|
``network_activate`` will apply the delta. Stamping is idempotent and only
|
|
touches modules a Norm adapter explicitly targets.
|
|
"""
|
|
t0 = time.time()
|
|
state_dict = read_state_dict(network_on_disk.filename, what="network")
|
|
if not has_marker(state_dict, NORM_MARKERS):
|
|
return None
|
|
|
|
mapping = resolve_mapping()
|
|
net = new_network(name, network_on_disk)
|
|
groups = group_by_suffixes(
|
|
state_dict, NORM_SUFFIXES,
|
|
prefixes=prefixes,
|
|
bare_prefixes=bare_prefixes,
|
|
bare_diffusers_prefixes=bare_diffusers_prefixes,
|
|
)
|
|
|
|
unmapped = 0
|
|
mismatch = 0
|
|
for (prefix, base), w in groups.items():
|
|
if "w_norm" not in w:
|
|
continue
|
|
targets = resolve_group_targets(resolve_targets, prefix, base)
|
|
if not targets:
|
|
unmapped += 1
|
|
continue
|
|
arch_prefix = _resolve_prefix(network_prefix, prefix)
|
|
for diffusers_path, chunk in targets:
|
|
if chunk is not None:
|
|
continue # norm targets are not fused
|
|
network_key = arch_prefix + diffusers_path.replace(".", "_")
|
|
sd_module = mapping.get(network_key)
|
|
if sd_module is None:
|
|
unmapped += 1
|
|
continue
|
|
if "b_norm" in w and not bias_delta_fits(sd_module, w["b_norm"]):
|
|
log.warning(
|
|
f'Network load: type=Norm name="{name}" arch={arch_name} key={network_key}'
|
|
f' bias={tuple(w["b_norm"].shape)} module={tuple(sd_module.bias.shape)}'
|
|
f' bias shape mismatch'
|
|
)
|
|
mismatch += 1
|
|
continue
|
|
if not getattr(sd_module, "network_layer_name", None):
|
|
sd_module.network_layer_name = network_key
|
|
nw = network.NetworkWeights(network_key=network_key, sd_key=network_key, w=w, sd_module=sd_module)
|
|
net.modules[network_key] = network_norm.NetworkModuleNorm(net, nw)
|
|
|
|
return finalize_network(net, name, "Norm", lora_scale, t0, unmapped=unmapped, mismatch=mismatch)
|
|
|
|
|
|
def try_load_full(name, network_on_disk, lora_scale, *,
|
|
resolve_targets, prefixes=KNOWN_PREFIXES_DEFAULT,
|
|
bare_prefixes=(), bare_diffusers_prefixes=(),
|
|
network_prefix=NETWORK_PREFIX_DEFAULT,
|
|
arch_name="generic"):
|
|
"""Generic Full (full-rank weight delta) loader.
|
|
|
|
Full adapters carry a complete weight delta (``diff``, same shape as the
|
|
host weight) and an optional bias delta (``diff_b``). Fused targets are
|
|
skipped with a warning: ``diff`` has the host weight's full shape and
|
|
row-slicing across three projections is well-defined arithmetically, but
|
|
no chunk class exists.
|
|
"""
|
|
t0 = time.time()
|
|
state_dict = read_state_dict(network_on_disk.filename, what="network")
|
|
if not has_marker(state_dict, FULL_MARKERS):
|
|
return None
|
|
|
|
mapping = resolve_mapping()
|
|
net = new_network(name, network_on_disk)
|
|
groups = group_by_suffixes(
|
|
state_dict, FULL_SUFFIXES,
|
|
prefixes=prefixes,
|
|
bare_prefixes=bare_prefixes,
|
|
bare_diffusers_prefixes=bare_diffusers_prefixes,
|
|
)
|
|
|
|
unmapped = 0
|
|
skipped = 0
|
|
mismatch = 0
|
|
for (prefix, base), w in groups.items():
|
|
if "diff" not in w:
|
|
continue
|
|
targets = resolve_group_targets(resolve_targets, prefix, base)
|
|
if any(t[1] is not None for t in targets):
|
|
log.warning(f'Network load: type=Full name="{name}" arch={arch_name} key={base} fused QKV skipped (unsupported)')
|
|
skipped += 1
|
|
continue
|
|
arch_prefix = _resolve_prefix(network_prefix, prefix)
|
|
for diffusers_path, _ in targets:
|
|
network_key = arch_prefix + diffusers_path.replace(".", "_")
|
|
sd_module = mapping.get(network_key)
|
|
if sd_module is None:
|
|
unmapped += 1
|
|
continue
|
|
if "diff_b" in w and not bias_delta_fits(sd_module, w["diff_b"]):
|
|
log.warning(
|
|
f'Network load: type=Full name="{name}" arch={arch_name} key={network_key}'
|
|
f' bias={tuple(w["diff_b"].shape)} module={tuple(sd_module.bias.shape)}'
|
|
f' bias shape mismatch'
|
|
)
|
|
mismatch += 1
|
|
continue
|
|
# Loader-local stamping, same as try_load_norm: a full-weight extraction carries the
|
|
# norm weights too, and assign_network_names_to_compvis_modules puts norms in the
|
|
# mapping but never stamps network_layer_name on them, which is what the apply pass
|
|
# keys off. Without this they bind and then silently never apply.
|
|
if not getattr(sd_module, "network_layer_name", None):
|
|
sd_module.network_layer_name = network_key
|
|
nw = network.NetworkWeights(network_key=network_key, sd_key=network_key, w=w, sd_module=sd_module)
|
|
net.modules[network_key] = network_full.NetworkModuleFull(net, nw)
|
|
|
|
return finalize_network(net, name, "Full", lora_scale, t0, unmapped=unmapped, mismatch=mismatch, skipped=skipped)
|
|
|
|
|
|
# === Per-arch umbrella ===
|
|
|
|
|
|
def try_load_chain(name, network_on_disk, lora_scale, family_loaders):
|
|
"""Run each family loader in order and merge any non-None results.
|
|
|
|
Per-arch loader modules expose a ``try_load(name, nod, scale)`` entry point
|
|
that the dispatcher in ``modules.lora.lora_load`` calls. That entry point
|
|
is a thin wrapper around this helper: it passes ``family_loaders`` as a
|
|
tuple of partial-applied generic loaders, each already bound to the arch's
|
|
``resolve_targets`` and prefix tuples.
|
|
"""
|
|
sd_models_utils.state_dict_cache.enable()
|
|
net = None
|
|
mismatch = 0
|
|
for try_fn in family_loaders:
|
|
sub = try_fn(name, network_on_disk, lora_scale)
|
|
if sub is None:
|
|
continue
|
|
mismatch += getattr(sub, 'mismatch', 0)
|
|
if net is None:
|
|
net = sub
|
|
else:
|
|
net.modules.update(sub.modules)
|
|
sd_models_utils.state_dict_cache.disable()
|
|
if net is not None and mismatch > 0: # applying only the layers that fit leaves the model in a state nothing was trained for
|
|
log.error(f'Network load: type=LoRA name="{name}" modules={len(net.modules)} mismatch={mismatch} shapes do not match the loaded model')
|
|
return None
|
|
return net
|