refactor(zimage): migrate to generic native_loader

Replaces zimage's four family loaders with thin wrappers binding
native_loader's generics to z-image's prefix tuples and resolve_targets.

resolve_targets folds the legacy attention.qkv split and attention.out
alias rename into the path-resolution step. Fused qkv now emits three
ChunkSpec(idx, total=3) entries; attention.out / attention.out.0 /
attention.wo aliases collapse to attention.to_out.0.

BARE_DIFFUSERS_PREFIXES allows bare paths starting with layers. /
noise_refiner. / context_refiner. to pass through to the loader. This
matches real Z-Image LoRAs exported via
ZImageTransformer2DModel.save_lora_adapter().

LoHA on fused qkv now binds via NetworkModuleHadaChunk (added to shared
infra by the flux2 PR) instead of being skipped. Test renamed to
test_loha_legacy_fused_qkv_chunked.

parse_key returns (prefix_used, base, suffix) instead of the old
(network_key, suffix); parse test updated.
This commit is contained in:
CalamitousFelicitousness
2026-05-18 23:14:52 +01:00
parent c7e8e7a029
commit c3b7379fd3
2 changed files with 158 additions and 361 deletions
+129 -346
View File
@@ -1,379 +1,162 @@
"""Z-Image native adapter loader.
Runs when :func:`modules.lora.lora_overrides.get_method` returns ``'native'``
(``lora_force_diffusers`` off and ``zimage`` in ``allow_native``). Reads the
safetensors directly and writes into sdnext's existing
``network_layer_mapping``, returning a ``Network`` populated with
``NetworkModule*`` entries that ``network_activate`` will apply. If the
setting is on, the diffusers PEFT path handles the file instead.
(``lora_force_diffusers`` off and ``zimage`` in ``allow_native``).
Entry points, one per family:
Entry points, one per family: :func:`try_load_lora` (plus DoRA),
:func:`try_load_lokr`, :func:`try_load_loha`, :func:`try_load_oft`.
- LoRA (+ DoRA) via :func:`try_load_lora`
- LoKR via :func:`try_load_lokr`
- LoHA via :func:`try_load_loha` (not yet validated against a real Z-Image adapter)
- OFT via :func:`try_load_oft` (not yet validated against a real Z-Image adapter)
Recognized key prefixes for every family: ``diffusion_model.``,
``transformer.``, ``lora_unet_``, or bare. Diffusers-PEFT ``lora_A``/``lora_B``
are normalized to ``lora_down``/``lora_up``.
Recognized key prefixes: ``diffusion_model.``, ``transformer.``,
``lora_unet_``, or bare paths starting with the known block-level prefixes
(``layers.``, ``noise_refiner.``, ``context_refiner.``).
Pre-refactor Z-Image attention layouts (fused ``attention.qkv``, bare
``attention.out`` / ``attention.wo``) are rewritten to the current diffusers
``to_q``/``to_k``/``to_v`` and ``to_out.0``. For LoRA the fused qkv up-weight
is chunked along dim 0 at load time. For LoKR the split is deferred to apply
time via :class:`NetworkModuleLokrChunk`, which materializes ``kron(w1, w2)``
once per forward pass and returns the designated slice.
``attention.out`` / ``attention.wo``) are rewritten by :func:`resolve_targets`
to the current diffusers ``to_q``/``to_k``/``to_v`` and ``to_out.0``. For
LoRA the fused qkv up-weight is chunked at load time; for LoKR the split is
deferred to apply time via :class:`network_lokr.NetworkModuleLokrChunk`.
Fused ``attention.qkv`` for LoHA and OFT is skipped with a warning: no
``NetworkModuleHadaChunk`` exists, and an OFT block structure is tied to the
target module's ``out_features``, so a Q/K/V split is not a drop-in.
slice variant exists for LoHA's Hadamard product on Linear targets, and an
OFT block structure is tied to the target module's ``out_features`` so a
Q/K/V split is not a drop-in.
"""
import os
import time
import torch
from modules import shared, sd_models
from modules.logger import log
from modules.lora import network, network_lora, network_lokr, network_hada, network_oft, lora_convert
from modules.lora import lora_common as l
from modules.lora import native_loader
from modules.lora.native_loader import ChunkSpec
KNOWN_PREFIXES = ("diffusion_model.", "transformer.", "lora_unet_")
# === Arch-specific prefix configuration ===
# Every family also picks up the universal optional keys
# (alpha, scale, bias, dora_scale) via base NetworkModule.__init__.
LORA_SUFFIXES = (
".lora_down.weight", ".lora_up.weight",
".lora_A.weight", ".lora_B.weight",
".alpha", ".dora_scale", ".bias", ".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",
)
KNOWN_PREFIXES = native_loader.KNOWN_PREFIXES_DEFAULT
# Presence of any of these substrings anywhere in a key marks a file as belonging to that family.
LORA_MARKERS = (".lora_down.weight", ".lora_up.weight", ".lora_A.weight", ".lora_B.weight")
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")
SUFFIX_NORMALIZE = {
"lora_A.weight": "lora_down.weight",
"lora_B.weight": "lora_up.weight",
}
ATTENTION_OUT_ALIASES = ("attention_out", "attention_out_0", "attention_wo")
ATTENTION_OUT_TARGET = "attention_to_out_0"
ATTENTION_QKV_SUFFIX = "attention_qkv"
ATTENTION_QKV_TARGETS = ("attention_to_q", "attention_to_k", "attention_to_v")
BARE_DIFFUSERS_PREFIXES = ("layers.", "noise_refiner.", "context_refiner.")
def try_load_lora(name, network_on_disk, lora_scale):
"""Try loading a Z-Image LoRA (plus DoRA) as native modules."""
t0 = time.time()
state_dict = sd_models.read_state_dict(network_on_disk.filename, what='network')
if not has_marker(state_dict, LORA_MARKERS):
return None
# === Re-exports for test/back-compat ===
mapping = resolve_mapping()
net = new_network(name, network_on_disk)
LORA_SUFFIXES = native_loader.LORA_SUFFIXES
LOKR_SUFFIXES = native_loader.LOKR_SUFFIXES
LOHA_SUFFIXES = native_loader.LOHA_SUFFIXES
OFT_SUFFIXES = native_loader.OFT_SUFFIXES
groups = group_by_suffixes(state_dict, LORA_SUFFIXES)
groups = expand_legacy_attention_lora(groups)
LORA_MARKERS = native_loader.LORA_MARKERS
LOKR_MARKERS = native_loader.LOKR_MARKERS
LOHA_MARKERS = native_loader.LOHA_MARKERS
OFT_MARKERS = native_loader.OFT_MARKERS
unmapped = 0
shape_mismatch = 0
for network_key, w in groups.items():
if 'lora_down.weight' not in w or 'lora_up.weight' not in w:
continue
sd_module = mapping.get(network_key)
if sd_module is None:
unmapped += 1
continue
if not shapes_match(sd_module, w['lora_down.weight'], w['lora_up.weight']):
log.warning(f'Network load: type=LoRA name="{name}" key={network_key} shape mismatch')
shape_mismatch += 1
continue
nw = network.NetworkWeights(network_key=network_key, sd_key=network_key, w=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=shape_mismatch)
def try_load_lokr(name, network_on_disk, lora_scale):
"""Try loading a Z-Image LoKR as native modules."""
t0 = time.time()
state_dict = sd_models.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)
groups, chunk_info = expand_legacy_attention_lokr(groups)
unmapped = 0
for network_key, 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
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)
chunked = chunk_info.get(network_key)
if chunked is not None:
idx, num = chunked
net.modules[network_key] = network_lokr.NetworkModuleLokrChunk(net, nw, idx, num)
else:
net.modules[network_key] = network_lokr.NetworkModuleLokr(net, nw)
return finalize_network(net, name, 'LoKR', lora_scale, t0, unmapped=unmapped)
def try_load_loha(name, network_on_disk, lora_scale):
"""Try loading a Z-Image LoHA as native modules. Fused attention.qkv groups are skipped."""
t0 = time.time()
state_dict = sd_models.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)
groups = rename_attention_out(groups)
unmapped = 0
skipped_qkv = 0
for network_key, w in groups.items():
if network_key.endswith("_" + ATTENTION_QKV_SUFFIX):
log.warning(f'Network load: type=LoHA name="{name}" key={network_key} fused qkv skipped (unsupported)')
skipped_qkv += 1
continue
if not all(k in w for k in ("hada_w1_a", "hada_w1_b", "hada_w2_a", "hada_w2_b")):
continue
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_hada.NetworkModuleHada(net, nw)
return finalize_network(net, name, 'LoHA', lora_scale, t0, unmapped=unmapped, skipped=skipped_qkv)
def try_load_oft(name, network_on_disk, lora_scale):
"""Try loading a Z-Image OFT adapter as native modules. Fused attention.qkv groups are skipped."""
t0 = time.time()
state_dict = sd_models.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)
groups = rename_attention_out(groups)
unmapped = 0
skipped_qkv = 0
for network_key, w in groups.items():
if network_key.endswith("_" + ATTENTION_QKV_SUFFIX):
log.warning(f'Network load: type=OFT name="{name}" key={network_key} fused qkv skipped (unsupported)')
skipped_qkv += 1
continue
if not ("oft_blocks" in w or "oft_diag" in w):
continue
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_oft.NetworkModuleOFT(net, nw)
return finalize_network(net, name, 'OFT', lora_scale, t0, unmapped=unmapped, skipped=skipped_qkv)
def has_marker(state_dict, markers):
return any(any(m in k for m in markers) for k in state_dict)
def resolve_mapping():
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):
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):
if len(net.modules) == 0:
if unmapped or mismatch or skipped:
log.debug(
f'Network load: type={family} name="{name}" native no-match'
f' unmapped={unmapped} mismatch={mismatch} 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:
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 group_by_suffixes(state_dict, suffixes):
"""Group state_dict entries by target module.
Returns ``{network_key: {suffix: tensor, ...}}`` where ``network_key`` follows
the sdnext convention ``lora_transformer_<path_with_underscores>``. Only keys
whose suffix appears in ``suffixes`` are kept; ``lora_A``/``lora_B`` are
normalized to ``lora_down``/``lora_up``.
"""
groups: dict[str, dict[str, torch.Tensor]] = {}
for key, value in state_dict.items():
parsed = parse_key(key, suffixes)
if parsed is None:
continue
network_key, suffix = parsed
slot = groups.get(network_key)
if slot is None:
slot = {}
groups[network_key] = slot
slot[suffix] = value
return groups
SUFFIX_NORMALIZE = native_loader.SUFFIX_NORMALIZE
BARE_DIFFUSERS_PREFIX_USED = native_loader.BARE_DIFFUSERS_PREFIX_USED
has_marker = native_loader.has_marker
def parse_key(key, suffixes):
stripped = key
for p in KNOWN_PREFIXES:
if key.startswith(p):
stripped = key[len(p):]
break
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)
network_key = 'lora_transformer_' + base.replace('.', '_')
return network_key, suffix
"""Z-Image-bound :func:`native_loader.parse_key`."""
return native_loader.parse_key(
key, suffixes,
prefixes=KNOWN_PREFIXES,
bare_diffusers_prefixes=BARE_DIFFUSERS_PREFIXES,
)
def rename_attention_out(groups):
"""Rename legacy ``attention.out`` / ``attention.wo`` keys to ``attention.to_out.0``."""
out: dict[str, dict[str, torch.Tensor]] = {}
for key, w in groups.items():
new_key = None
for alias in ATTENTION_OUT_ALIASES:
if key.endswith("_" + alias):
new_key = key[: -len(alias)] + ATTENTION_OUT_TARGET
break
out[new_key or key] = w
return out
def group_by_suffixes(state_dict, suffixes):
"""Z-Image-bound :func:`native_loader.group_by_suffixes`."""
return native_loader.group_by_suffixes(
state_dict, suffixes,
prefixes=KNOWN_PREFIXES,
bare_diffusers_prefixes=BARE_DIFFUSERS_PREFIXES,
)
def expand_legacy_attention_lora(groups):
"""Rename attention.out, then split fused attention.qkv into three per-projection groups.
# === Target resolution (arch-specific) ===
The down-weight is shared across Q/K/V and the up-weight is chunked along
dim 0 (the concatenated output dim in the fused layout).
def resolve_targets(prefix_used, base):
"""Return ``[(diffusers_path, ChunkSpec | None), ...]`` for a parsed group key.
Handles two legacy Z-Image attention layouts:
- Fused ``attention.qkv`` is split into three Q/K/V targets, each carrying
a :class:`ChunkSpec` for the loader to apply.
- ``attention.out`` / ``attention.out.0`` / ``attention.wo`` are aliased to
the current diffusers ``attention.to_out.0`` path.
Everything else (modern split-attention paths, MLP, norms, embedders) is
returned verbatim.
"""
groups = rename_attention_out(groups)
out: dict[str, dict[str, torch.Tensor]] = {}
for key, w in groups.items():
if not key.endswith("_" + ATTENTION_QKV_SUFFIX):
out[key] = w
continue
stem = key[: -len(ATTENTION_QKV_SUFFIX)]
down = w.get("lora_down.weight")
up = w.get("lora_up.weight")
if down is None or up is None or up.shape[0] % 3 != 0:
out[key] = w
continue
chunks = torch.chunk(up, 3, dim=0)
alpha = w.get("alpha")
dora = w.get("dora_scale")
for target, chunk in zip(ATTENTION_QKV_TARGETS, chunks):
split_key = stem + target
split = {
"lora_down.weight": down,
"lora_up.weight": chunk.contiguous(),
}
if alpha is not None:
split["alpha"] = alpha
if dora is not None:
split["dora_scale"] = dora
out[split_key] = split
return out
if prefix_used == "lora_unet_":
return _underscore_to_diffusers_targets(base)
if prefix_used == "transformer.":
return [(base, None)]
if prefix_used == BARE_DIFFUSERS_PREFIX_USED:
return [(base, None)]
if prefix_used in (None, "diffusion_model."):
return _dotted_to_diffusers_targets(base)
return []
def expand_legacy_attention_lokr(groups):
"""Rename attention.out, then split fused attention.qkv for LoKR.
def _dotted_to_diffusers_targets(base):
"""For BFL / bare-BFL keys like ``layers.0.attention.qkv``."""
if base.endswith(".attention.qkv"):
stem = base[:-len(".attention.qkv")]
return [
(f"{stem}.attention.to_q", ChunkSpec(idx=0, total=3)),
(f"{stem}.attention.to_k", ChunkSpec(idx=1, total=3)),
(f"{stem}.attention.to_v", ChunkSpec(idx=2, total=3)),
]
for alias in (".attention.out.0", ".attention.out", ".attention.wo"):
if base.endswith(alias):
stem = base[:-len(alias)]
return [(f"{stem}.attention.to_out.0", None)]
return [(base, None)]
For LoKR the split is deferred to apply time via ``NetworkModuleLokrChunk``
(returned via a parallel ``chunk_info`` dict keyed by the split network key).
The three resulting groups share the same tensor dict by shallow copy; the
chunk module computes ``kron(w1, w2)`` once and slices the designated row
range.
Returns ``(groups, chunk_info)`` where ``chunk_info[network_key] == (index, num)``.
"""
groups = rename_attention_out(groups)
out: dict[str, dict[str, torch.Tensor]] = {}
chunk_info: dict[str, tuple[int, int]] = {}
for key, w in groups.items():
if not key.endswith("_" + ATTENTION_QKV_SUFFIX):
out[key] = w
continue
stem = key[: -len(ATTENTION_QKV_SUFFIX)]
for i, target in enumerate(ATTENTION_QKV_TARGETS):
split_key = stem + target
out[split_key] = dict(w)
chunk_info[split_key] = (i, 3)
return out, chunk_info
def _underscore_to_diffusers_targets(base):
"""For kohya flat-underscore keys like ``layers_0_attention_qkv``."""
if base.endswith("_attention_qkv"):
stem = base[:-len("_attention_qkv")]
return [
(f"{stem}_attention_to_q", ChunkSpec(idx=0, total=3)),
(f"{stem}_attention_to_k", ChunkSpec(idx=1, total=3)),
(f"{stem}_attention_to_v", ChunkSpec(idx=2, total=3)),
]
for alias in ("_attention_out_0", "_attention_out", "_attention_wo"):
if base.endswith(alias):
stem = base[:-len(alias)]
return [(f"{stem}_attention_to_out_0", None)]
return [(base, None)]
# === Native loaders (thin wrappers over native_loader generics) ===
_BIND_KWARGS = dict(
resolve_targets=resolve_targets,
prefixes=KNOWN_PREFIXES,
bare_diffusers_prefixes=BARE_DIFFUSERS_PREFIXES,
arch_name="zimage",
)
def try_load_lora(name, network_on_disk, lora_scale):
return native_loader.try_load_lora(name, network_on_disk, lora_scale, **_BIND_KWARGS)
def try_load_lokr(name, network_on_disk, lora_scale):
return native_loader.try_load_lokr(name, network_on_disk, lora_scale, **_BIND_KWARGS)
def try_load_loha(name, network_on_disk, lora_scale):
return native_loader.try_load_loha(name, network_on_disk, lora_scale, **_BIND_KWARGS)
def try_load_oft(name, network_on_disk, lora_scale):
return native_loader.try_load_oft(name, network_on_disk, lora_scale, **_BIND_KWARGS)
def try_load(name, network_on_disk, lora_scale):
"""Run every Z-Image family loader, merge any that match."""
return native_loader.try_load_chain(
name, network_on_disk, lora_scale,
family_loaders=(try_load_lora, try_load_lokr, try_load_loha, try_load_oft),
)
+29 -15
View File
@@ -414,25 +414,27 @@ CAT_PARSE = category('parse')
def test_parse_key_all_prefixes():
"""parse_key recognizes BFL, PEFT, kohya, and bare keys with the right base+suffix."""
"""parse_key recognizes BFL, PEFT, kohya, and bare-diffusers keys.
Returns (prefix_used, base, suffix) - prefix_used is the matched
KNOWN_PREFIXES element, BARE_DIFFUSERS_PREFIX_USED for bare paths
matching BARE_DIFFUSERS_PREFIXES, or None when no prefix is recognized.
"""
bd = Z.BARE_DIFFUSERS_PREFIX_USED
cases = [
# BFL prefix -> dotted base, lora_A normalized to lora_down
('diffusion_model.layers.0.attention.to_q.lora_A.weight',
Z.LORA_SUFFIXES,
('lora_transformer_layers_0_attention_to_q', 'lora_down.weight')),
# PEFT prefix
('diffusion_model.', 'layers.0.attention.to_q', 'lora_down.weight')),
('transformer.layers.0.attention.to_v.lora_B.weight',
Z.LORA_SUFFIXES,
('lora_transformer_layers_0_attention_to_v', 'lora_up.weight')),
# kohya flat underscore form
('transformer.', 'layers.0.attention.to_v', 'lora_up.weight')),
('lora_unet_layers_0_attention_to_k.lora_down.weight',
Z.LORA_SUFFIXES,
('lora_transformer_layers_0_attention_to_k', 'lora_down.weight')),
# Bare dotted (no recognized prefix)
('lora_unet_', 'layers_0_attention_to_k', 'lora_down.weight')),
# Bare path starting with a known block prefix
('layers.0.attention.to_out.0.lora_A.weight',
Z.LORA_SUFFIXES,
('lora_transformer_layers_0_attention_to_out_0', 'lora_down.weight')),
# Unrelated keys reject cleanly
(bd, 'layers.0.attention.to_out.0', 'lora_down.weight')),
('random.unrelated.key', Z.LORA_SUFFIXES, None),
]
for key, suffixes, expected in cases:
@@ -580,11 +582,23 @@ def test_loha_bfl_proj():
return True
def test_loha_legacy_fused_qkv_skipped():
"""LoHA on legacy fused attention.qkv is skipped with a warning (no chunk variant)."""
def test_loha_legacy_fused_qkv_chunked():
"""LoHA on legacy fused attention.qkv emits 3 HadaChunk modules.
Post-migration the generic LoHA loader uses NetworkModuleHadaChunk for
equal-chunks dispatch (the chunk class was added in the flux2 PR and is
now shared infrastructure).
"""
net = _load_via(Z.try_load_loha, sd_loha_legacy_fused_qkv_skipped())
# Either no net returned (nothing matched) or net with zero modules
assert net is None or len(net.modules) == 0, f'expected no modules, got {net.modules if net else None}'
assert net is not None and len(net.modules) == 3, f'got {net.modules if net else None}'
expected = {
'lora_transformer_layers_0_attention_to_q',
'lora_transformer_layers_0_attention_to_k',
'lora_transformer_layers_0_attention_to_v',
}
assert set(net.modules) == expected, f'got {set(net.modules)}'
for mod in net.modules.values():
assert isinstance(mod, network_hada.NetworkModuleHadaChunk)
return True
@@ -689,7 +703,7 @@ def run_tests():
test_lokr_bfl_adaln,
test_lokr_legacy_fused_qkv_chunked,
test_loha_bfl_proj,
test_loha_legacy_fused_qkv_skipped,
test_loha_legacy_fused_qkv_chunked,
test_oft_lycoris_no_npe,
test_oft_legacy_fused_qkv_skipped,
]: