fix(lora): match the diffusers minimax lora converter

The reference fc1 is a fused [gate; value] SwiGLU projection and the
diffusers port stores [value; gate]. The native mapping did not swap the
halves, so gate and value deltas landed on each other's rows. The
mapping now also renames the standalone projections, reads a
metadata-only alpha, and accepts the musubi, peft dit and diffusers-named
layouts.
This commit is contained in:
CalamitousFelicitousness
2026-09-05 22:54:28 +01:00
parent e0a23b0c6e
commit 4a5dc98cb2
3 changed files with 752 additions and 44 deletions
+9
View File
@@ -563,11 +563,15 @@ def try_load_lora(name, network_on_disk, lora_scale, *,
bare_prefixes=(), bare_diffusers_prefixes=(),
network_prefix=NETWORK_PREFIX_DEFAULT,
group_by_suffixes_fn=group_by_suffixes,
network_alpha=None,
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.
``network_alpha`` is a file-level alpha for files without alpha tensors;
a file carrying any alpha of its own keeps those and ignores it.
"""
t0 = time.time()
state_dict = read_state_dict(network_on_disk.filename, what="network")
@@ -582,6 +586,8 @@ def try_load_lora(name, network_on_disk, lora_scale, *,
bare_prefixes=bare_prefixes,
bare_diffusers_prefixes=bare_diffusers_prefixes,
)
if network_alpha is not None and any("alpha" in w for w in groups.values()):
network_alpha = None
unmapped = 0
mismatch = 0
@@ -589,6 +595,9 @@ def try_load_lora(name, network_on_disk, lora_scale, *,
for (prefix, base), w in groups.items():
if "lora_down.weight" not in w or "lora_up.weight" not in w:
continue
if network_alpha is not None:
w = dict(w)
w["alpha"] = torch.tensor(float(network_alpha))
# 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
+78 -44
View File
@@ -3,10 +3,17 @@
MiniMax H3 is a modular video pipeline with a transformer and a text encoder.
This loader routes transformer keys to ``lora_transformer_`` and text-encoder
keys to ``lora_te_`` when native LoRAs are applied.
Published LoRAs target the reference module names. The mapping onto
``diffusers.MiniMaxH3Transformer3DModel`` follows the diffusers LoRA converter:
fused ``attn.qkv_proj`` splits into ``to_q``/``to_k``/``to_v``, and ``mlp.fc1``
lands on the fused SwiGLU projection with its two output halves swapped.
"""
import re
import torch
from modules.lora import native_adapter
@@ -16,14 +23,8 @@ KNOWN_PREFIXES = (
"diffusion_model.transformer.",
"diffusion_model.blocks.",
"diffusion_model.transformer_blocks.",
"diffusion_model.token_refiner.",
"diffusion_model.token_refiner.refiner_blocks.",
"diffusion_model.final_layer.",
"diffusion_model.video_patch_proj.",
"diffusion_model.audio_patch_proj.",
"diffusion_model.condition_proj.",
"diffusion_model.time_embedder.proj_in.",
"diffusion_model.time_embedder.proj_out.",
"diffusion_model.token_refiner.",
"diffusion_model.model.language_model.",
"text_encoders.",
"text_encoder.",
@@ -31,16 +32,27 @@ KNOWN_PREFIXES = (
"transformer.",
"transformer_blocks.",
"blocks.",
"token_refiner.",
"token_refiner.refiner_blocks.",
"final_layer.",
"video_patch_proj.",
"audio_patch_proj.",
"condition_proj.",
"time_embedder.proj_in.",
"time_embedder.proj_out.",
"token_refiner.",
) + native_adapter.KNOWN_PREFIXES_DEFAULT
# Reference keys outside the block stacks carry no arch prefix in reference saves; the base is the whole module path.
BARE_PREFIXES = ("video_patch_proj.", "audio_patch_proj.", "condition_proj.", "time_embedder.", "final_layer.")
# Diffusers module names saved without a component prefix (peft dumps, kohya-suffixed exports) bind verbatim.
BARE_DIFFUSERS_PREFIXES = ("transformer_blocks.", "token_refiner.refiner_blocks.", "proj_in.", "audio_proj_in.", "context_embedder.", "time_embedder.linear_", "norm_out.", "proj_out.", "audio_proj_out.")
STANDALONE_RENAMES = {
"video_patch_proj": "proj_in",
"audio_patch_proj": "audio_proj_in",
"condition_proj": "context_embedder",
"time_embedder.proj_in": "time_embedder.linear_1",
"time_embedder.proj_out": "time_embedder.linear_2",
"final_layer.adaln_proj.linear": "norm_out.linear",
"final_layer.video_out": "proj_out",
"final_layer.audio_out": "audio_proj_out",
}
# Re-export for tests / compatibility
LORA_SUFFIXES = native_adapter.LORA_SUFFIXES
@@ -80,26 +92,31 @@ def normalize_mini_max_suffix(suffix: str) -> str:
return MINIMAX_SUFFIX_NORMALIZE.get(suffix, suffix)
# musubi-tuner flattens every "." to "_" under a lora_unet_ prefix; the reference module names carry
# underscores of their own, so the dotted path is recovered by matching the whole flattened name.
_FLATTENED_MODULES = [
(r"blocks_(\d+)_attn_(qkv|out)_proj", r"blocks.\1.attn.\2_proj"),
(r"blocks_(\d+)_mlp_fc([12])", r"blocks.\1.mlp.fc\2"),
(r"blocks_(\d+)_adaln_proj_linear", r"blocks.\1.adaln_proj.linear"),
(r"token_refiner_blocks_(\d+)_attn_(qkv|out)_proj", r"token_refiner.blocks.\1.attn.\2_proj"),
(r"token_refiner_blocks_(\d+)_mlp_fc([12])", r"token_refiner.blocks.\1.mlp.fc\2"),
(r"(video|audio)_patch_proj", r"\1_patch_proj"),
(r"condition_proj", "condition_proj"),
(r"time_embedder_proj_(in|out)", r"time_embedder.proj_\1"),
(r"final_layer_adaln_proj_linear", "final_layer.adaln_proj.linear"),
(r"final_layer_(video|audio)_out", r"final_layer.\1_out"),
]
def _unflatten_lora_unet_key(key: str) -> str | None:
if not key.startswith("lora_unet_"):
return None
module_key, _, suffix = key[len("lora_unet_"):].partition(".")
if not suffix:
return None
patterns = [
(r"blocks_(\d+)_attn_out_proj", r"blocks.\1.attn.out_proj"),
(r"blocks_(\d+)_attn_qkv_proj", r"blocks.\1.attn.qkv_proj"),
(r"blocks_(\d+)_mlp_fc1", r"blocks.\1.mlp.fc1"),
(r"blocks_(\d+)_mlp_fc2", r"blocks.\1.mlp.fc2"), (r"token_refiner_blocks_(\d+)_attn_out_proj", r"token_refiner.blocks.\1.attn.out_proj"),
(r"token_refiner_blocks_(\d+)_attn_qkv_proj", r"token_refiner.blocks.\1.attn.qkv_proj"),
(r"token_refiner_blocks_(\d+)_mlp_fc1", r"token_refiner.blocks.\1.mlp.fc1"),
(r"token_refiner_blocks_(\d+)_mlp_fc2", r"token_refiner.blocks.\1.mlp.fc2"), ]
for pattern, replacement in patterns:
for pattern, replacement in _FLATTENED_MODULES:
if re.fullmatch(pattern, module_key):
return f"{re.sub(pattern, replacement, module_key)}.{suffix}"
return None
@@ -108,13 +125,28 @@ def parse_key(key, suffixes):
unflattened = _unflatten_lora_unet_key(key)
if unflattened is not None:
key = unflattened
parsed = native_adapter.parse_key(key, suffixes, prefixes=KNOWN_PREFIXES)
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, bare_prefixes=BARE_PREFIXES, bare_diffusers_prefixes=BARE_DIFFUSERS_PREFIXES)
if parsed is None:
return None
prefix_used, base, suffix = parsed
return prefix_used, base, normalize_mini_max_suffix(suffix)
def swap_swiglu_halves(slot):
"""Reorder fc1 output rows from the reference ``[gate; value]`` to the diffusers SwiGLU ``[value; gate]``."""
up = slot.get("lora_up.weight")
if up is None:
return
for key in ("lora_up.weight", "bias", "diff_b", "dora_scale"):
t = slot.get(key)
if t is not None and t.ndim >= 1 and t.shape[0] == up.shape[0]:
gate, value = t.chunk(2, dim=0)
slot[key] = torch.cat([value, gate], dim=0)
def group_by_suffixes(state_dict, suffixes, *, prefixes=None, bare_prefixes=(), bare_diffusers_prefixes=()): # pylint: disable=unused-argument
"""MiniMax-bound :func:`native_adapter.group_by_suffixes`."""
groups: dict[tuple, dict[str, object]] = {}
@@ -128,6 +160,9 @@ def group_by_suffixes(state_dict, suffixes, *, prefixes=None, bare_prefixes=(),
slot = {}
groups[(prefix_used, base)] = slot
slot[suffix] = value
for (_prefix_used, base), slot in groups.items():
if base.endswith(".mlp.fc1"):
swap_swiglu_halves(slot)
return groups
@@ -158,14 +193,14 @@ def _transformer_block_targets(target_prefix, base):
def resolve_targets(prefix_used, base):
"""Return ``[(diffusers_path, ChunkSpec | None), ...]`` for MiniMax keys."""
if prefix_used == "diffusion_model.":
if prefix_used == "diffusion_model." or prefix_used is None:
if base.startswith("transformer."):
return [(base[len("transformer."):], None)]
if base.startswith("text_encoder."):
return [(base[len("text_encoder."):], None)]
if base.startswith("text_encoders."):
return [(base[len("text_encoders."):], None)]
return [(base, None)]
return [(STANDALONE_RENAMES.get(base, base), None)]
if prefix_used in ("diffusion_model.transformer.", "transformer.", "lora_transformer_"):
return [(base, None)]
if prefix_used in ("diffusion_model.transformer_blocks.", "transformer_blocks."):
@@ -178,20 +213,6 @@ def resolve_targets(prefix_used, base):
if base.startswith("blocks."):
return _transformer_block_targets("token_refiner.refiner_blocks", base[len("blocks."):])
return [(f"token_refiner.refiner_{base}", None)]
if prefix_used in ("diffusion_model.final_layer.", "final_layer."):
if base == "adaln_proj.linear":
return [("norm_out.linear", None)]
return [(base, None)]
if prefix_used in ("diffusion_model.video_patch_proj.", "video_patch_proj."):
return [("proj_in." + base.split(".", 1)[1], None)] if "." in base else [("proj_in", None)]
if prefix_used in ("diffusion_model.audio_patch_proj.", "audio_patch_proj."):
return [("audio_proj_in." + base.split(".", 1)[1], None)] if "." in base else [("audio_proj_in", None)]
if prefix_used in ("diffusion_model.condition_proj.", "condition_proj."):
return [("context_embedder." + base.split(".", 1)[1], None)] if "." in base else [("context_embedder", None)]
if prefix_used in ("diffusion_model.time_embedder.proj_in.", "time_embedder.proj_in."):
return [("time_embedder.linear_1." + base.split(".", 1)[1], None)] if "." in base else [("time_embedder.linear_1", None)]
if prefix_used in ("diffusion_model.time_embedder.proj_out.", "time_embedder.proj_out."):
return [("time_embedder.linear_2." + base.split(".", 1)[1], None)] if "." in base else [("time_embedder.linear_2", None)]
if prefix_used in (
"diffusion_model.text_encoder.",
"diffusion_model.text_encoders.",
@@ -219,9 +240,22 @@ def network_prefix_for(prefix_used):
return "lora_transformer_"
def file_alpha(network_on_disk):
"""The training alpha some trainers record in the safetensors metadata instead of per-key tensors, or None."""
alpha = (getattr(network_on_disk, "metadata", None) or {}).get("alpha")
if alpha is None:
return None
try:
return float(alpha)
except (TypeError, ValueError):
return None
_BIND_KWARGS = dict(
resolve_targets=resolve_targets,
prefixes=KNOWN_PREFIXES,
bare_prefixes=BARE_PREFIXES,
bare_diffusers_prefixes=BARE_DIFFUSERS_PREFIXES,
network_prefix=network_prefix_for,
group_by_suffixes_fn=group_by_suffixes,
arch_name="minimaxh3",
@@ -229,7 +263,7 @@ _BIND_KWARGS = dict(
def try_load_lora(name, network_on_disk, lora_scale):
return native_adapter.try_load_lora(name, network_on_disk, lora_scale, **_BIND_KWARGS)
return native_adapter.try_load_lora(name, network_on_disk, lora_scale, network_alpha=file_alpha(network_on_disk), **_BIND_KWARGS)
def try_load_lokr(name, network_on_disk, lora_scale):
+665
View File
@@ -0,0 +1,665 @@
#!/usr/bin/env python
"""
Offline unit tests for the MiniMax H3 native adapter loader.
Published MiniMax H3 LoRAs target the reference module names (fused
``attn.qkv_proj``, fused SwiGLU ``mlp.fc1``, ``token_refiner.blocks``), while
sdnext loads upstream ``diffusers.MiniMaxH3Transformer3DModel``. The reference
fc1 is ``[gate; value]`` and diffusers' SwiGLU is ``[value; gate]``, so the
native mapping has to permute fc1 output rows the same way the diffusers LoRA
converter does. These tests pin ``pipelines.minimax.minimax_lora`` against that
converter: same targets, same per-module deltas.
Save formats exercised, each seen in a published LoRA:
- comfy / ai-toolkit (``diffusion_model.blocks.0.attn.qkv_proj.lora_A.weight``):
the CivitAI ecosystem and the larryvrh turbo files.
- bare reference names (``blocks.0.attn.qkv_proj.lora_A.weight``): the
unpruned larryvrh saves.
- musubi-tuner (``lora_unet_blocks_0_mlp_fc1.lora_down.weight`` + ``alpha``).
- comfy with block-diagonal fused qkv and tripled alpha: lightx2v's ComfyUI
exports.
- peft dump (``transformer_blocks.0.attn.to_q.lora_A.default.weight``) with the
training alpha in the file metadata: lightx2v's diffusers exports.
- diffusers names with kohya suffixes: the alibaba-pai Acc LoRAs.
- peft wrapper around a ``dit`` attribute (``base_model.model.dit.blocks.0...``):
the mvp-lab RAVEN LoRA.
The reference module tree is the real ``MiniMaxH3Transformer3DModel`` at tiny
dims, so module names and target shapes are authoritative. Every LoRA file
under the MiniMax H3 LoRA folder is also mapped through both paths and compared
per target, exactly when the file carries no alpha and through random probes
when the converter folds an alpha into the weights.
No running server required.
Usage:
python test/test-minimax-native-adapters.py
"""
import glob
import os
import sys
import tempfile
import time
import torch
import torch.nn as nn
import safetensors
import safetensors.torch
script_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, script_dir)
os.chdir(script_dir)
os.environ['SD_INSTALL_QUIET'] = '1'
# Bootstrap cmd_args before any module that pulls in shared.py.
import modules.cmd_args # pylint: disable=wrong-import-position
import installer # pylint: disable=wrong-import-position
_orig_argv = sys.argv
sys.argv = [sys.argv[0]]
try:
modules.cmd_args.parse_args()
finally:
sys.argv = _orig_argv
installer.add_args(modules.cmd_args.parser)
modules.cmd_args.parsed, _ = modules.cmd_args.parser.parse_known_args([])
from modules.errors import log # pylint: disable=wrong-import-position
from modules import shared # pylint: disable=wrong-import-position
from modules.lora import native_adapter # pylint: disable=wrong-import-position
from pipelines.minimax import minimax_lora as M # pylint: disable=wrong-import-position
from diffusers import MiniMaxH3Transformer3DModel # pylint: disable=wrong-import-position
from diffusers.loaders.lora_conversion_utils import _convert_non_diffusers_minimax_h3_lora_to_diffusers as convert_diffusers # pylint: disable=wrong-import-position
# ============================================================
# Test infrastructure
# ============================================================
results: dict[str, dict] = {}
def category(name: str):
if name not in results:
results[name] = {'passed': 0, 'failed': 0, 'tests': []}
return name
def record(cat: str, passed: bool, name: str, detail: str = ''):
status = 'PASS' if passed else 'FAIL'
results[cat]['passed' if passed else 'failed'] += 1
results[cat]['tests'].append((status, name))
msg = f' {status}: {name}'
if detail:
msg += f' ({detail})'
if passed:
log.info(msg)
else:
log.error(msg)
def run_test(cat: str, fn):
name = fn.__name__
try:
ok = fn()
if ok is False:
record(cat, False, name)
else:
record(cat, True, name)
except AssertionError as e:
record(cat, False, name, str(e))
except Exception as e: # pylint: disable=broad-except
record(cat, False, name, f'exception: {e}')
import traceback
traceback.print_exc()
# ============================================================
# Reference MiniMax H3 transformer (real class, tiny dims)
# ============================================================
# Upstream: heads=56, head_dim=128, hidden=5376, layers=50, refiner=2, ffn=14336.
# The attention inner dim exceeds the hidden size upstream; the test dims keep that.
NUM_LAYERS = 2
NUM_REFINER = 1
REF = MiniMaxH3Transformer3DModel(
num_attention_heads=2, attention_head_dim=8, hidden_size=12, num_layers=NUM_LAYERS, num_refiner_layers=NUM_REFINER,
ffn_dim=16, in_channels=4, audio_in_channels=4, patch_size=(1, 2, 2), text_dim=10, freq_dim=8,
time_embed_hidden_dim=12, time_embed_dim=6, rope_freq_dim=1,
)
# {diffusers dotted name: (out, in)} for every Linear.
LINEAR_SHAPES = {name: tuple(m.weight.shape) for name, m in REF.named_modules() if isinstance(m, nn.Linear)}
# {network key: module} exactly as assign_network_names_to_compvis_modules stamps it.
LINEAR_NETKEYS = {'lora_transformer_' + name.replace('.', '_') for name in LINEAR_SHAPES}
# Reference module names and where they land, used only to size synthetic tensors;
# correctness is judged against the diffusers converter, not this table.
STANDALONE = {
'video_patch_proj': 'proj_in',
'audio_patch_proj': 'audio_proj_in',
'condition_proj': 'context_embedder',
'time_embedder.proj_in': 'time_embedder.linear_1',
'time_embedder.proj_out': 'time_embedder.linear_2',
'final_layer.adaln_proj.linear': 'norm_out.linear',
'final_layer.video_out': 'proj_out',
'final_layer.audio_out': 'audio_proj_out',
}
BLOCK_LEAVES = {
'attn.qkv_proj': 'attn.to_q',
'attn.out_proj': 'attn.to_out.0',
'mlp.fc1': 'ff.net.0.proj',
'mlp.fc2': 'ff.net.2',
'adaln_proj.linear': 'adaln_proj.linear',
}
def ref_shape(ref_path: str):
"""(out, in) of the reference Linear a LoRA key names, read from the diffusers module it maps to."""
if ref_path in STANDALONE:
return LINEAR_SHAPES[STANDALONE[ref_path]]
if ref_path.startswith('blocks.'):
_, idx, leaf = ref_path.split('.', 2)
stack = 'transformer_blocks'
else:
_, _, idx, leaf = ref_path.split('.', 3)
stack = 'token_refiner.refiner_blocks'
out, inp = LINEAR_SHAPES[f'{stack}.{idx}.{BLOCK_LEAVES[leaf]}']
if leaf == 'attn.qkv_proj':
out *= 3
return out, inp
def all_ref_paths():
paths = list(STANDALONE)
for i in range(NUM_LAYERS):
paths += [f'blocks.{i}.{leaf}' for leaf in BLOCK_LEAVES]
for i in range(NUM_REFINER):
paths += [f'token_refiner.blocks.{i}.{leaf}' for leaf in BLOCK_LEAVES if leaf != 'adaln_proj.linear']
return paths
# ============================================================
# Mock pipeline wrapping the real transformer
# ============================================================
class _MockPipeline:
def __init__(self, transformer):
self.transformer = transformer
self.text_encoder = None
class _MockSdModel:
def __init__(self, pipe):
self.pipe = pipe
self.network_layer_mapping = {}
self.embedding_db = None
self.__class__.__name__ = 'MiniMaxH3ModularPipeline'
def install_mock_pipe():
"""Point shared.sd_model at a mock exposing the reference transformer; re-installed per load so stamps do not leak."""
sd_model = _MockSdModel(_MockPipeline(REF))
from modules.modeldata import model_data
model_data.sd_model = sd_model
return sd_model
# ============================================================
# Synthesizers and helpers
# ============================================================
RANK = 4
PEFT = ('lora_A.weight', 'lora_B.weight')
KOHYA = ('lora_down.weight', 'lora_up.weight')
def lora_pair(key_base, shape, suffix=PEFT, alpha=None, rank=RANK):
out, inp = shape
sd = {
f'{key_base}.{suffix[0]}': torch.randn(rank, inp),
f'{key_base}.{suffix[1]}': torch.randn(out, rank),
}
if alpha is not None:
sd[f'{key_base}.alpha'] = torch.tensor(float(alpha))
return sd
def synth(prefix='diffusion_model.', suffix=PEFT, alpha=None, flatten=False):
"""One LoRA pair for every reference Linear, in the requested reference-name layout."""
sd = {}
for path in all_ref_paths():
key = f'lora_unet_{path.replace(".", "_")}' if flatten else f'{prefix}{path}'
sd.update(lora_pair(key, ref_shape(path), suffix, alpha))
return sd
def synth_diffusers(prefix='', suffix=PEFT, infix='', alpha=None):
"""One LoRA pair for every Linear, keyed by its diffusers name; ``infix`` inserts a peft adapter slot."""
sd = {}
for path, shape in LINEAR_SHAPES.items():
pair = lora_pair(f'{prefix}{path}', shape, suffix, alpha)
for key, value in pair.items():
for marker in ('lora_A', 'lora_B', 'lora_down', 'lora_up'):
key = key.replace(f'.{marker}.weight', f'.{marker}{infix}.weight')
sd[key] = value
return sd
class TempLora:
"""Context manager: writes a state dict to a temp safetensors file."""
def __init__(self, state_dict, name='test', metadata=None):
self.state_dict = state_dict
self.name = name
self.metadata = metadata
self.path = None
def __enter__(self):
sd = {k: v.contiguous() for k, v in self.state_dict.items()}
fd, self.path = tempfile.mkstemp(suffix='.safetensors', prefix=f'{self.name}_')
os.close(fd)
safetensors.torch.save_file(sd, self.path)
return _MockNetworkOnDisk(self.path, self.name, self.metadata)
def __exit__(self, exc_type, exc_val, exc_tb):
if self.path and os.path.exists(self.path):
os.unlink(self.path)
class _MockNetworkOnDisk:
def __init__(self, filename, name, metadata=None):
self.filename = filename
self.name = name
self.shorthash = ''
self.sd_version = 'unknown'
self.metadata = metadata or {}
def load_native(state_dict, name='test', metadata=None):
install_mock_pipe()
with TempLora(state_dict, name=name, metadata=metadata) as nod:
return M.try_load(name, nod, lora_scale=1.0)
def native_deltas(net):
"""{network key: full delta} as the apply pass would compute it at multiplier 1."""
out = {}
for key, module in net.modules.items():
module.network.te_multiplier = 1.0
module.network.unet_multiplier = 1.0
up = module.up_model.weight.float()
down = module.down_model.weight.float()
out[key] = (up @ down) * module.calc_scale()
return out
def diffusers_deltas(state_dict, rewrite=None):
"""{network key: full delta} from the diffusers converter's output."""
if rewrite is not None:
state_dict = {k.replace(rewrite[0], rewrite[1], 1) if k.startswith(rewrite[0]) else k: v for k, v in state_dict.items()}
converted = convert_diffusers(dict(state_dict))
out = {}
for key, value in converted.items():
if not key.endswith('.lora_A.weight'):
continue
path = key[len('transformer.'):-len('.lora_A.weight')]
down = value.float()
up = converted[f'transformer.{path}.lora_B.weight'].float()
out['lora_transformer_' + path.replace('.', '_')] = up @ down
return out
def identity_deltas(state_dict, scale=1.0):
"""{network key: full delta} for a file already on diffusers names, at a uniform scale."""
out = {}
for key, value in state_dict.items():
for marker in ('lora_A', 'lora_down'):
idx = key.find(f'.{marker}')
if idx == -1:
continue
path = key[:idx].removeprefix('transformer.')
up_key = key.replace(marker, 'lora_B' if marker == 'lora_A' else 'lora_up', 1)
out['lora_transformer_' + path.replace('.', '_')] = (state_dict[up_key].float() @ value.float()) * scale
return out
def assert_same_deltas(native, reference):
assert set(native) == set(reference), f'targets differ: native-only={sorted(set(native) - set(reference))} reference-only={sorted(set(reference) - set(native))}'
for key, ref in reference.items():
got = native[key]
assert got.shape == ref.shape, f'{key}: shape {tuple(got.shape)} vs {tuple(ref.shape)}'
assert torch.allclose(got, ref, rtol=1e-5, atol=1e-5), f'{key}: delta differs, max abs {float((got - ref).abs().max()):.3e}'
def native_mapping(state_dict, network_alpha=None):
"""{diffusers path: (down, up, scale)} the native loader would bind, without a model: parse, resolve, chunk, file alpha."""
groups = M.group_by_suffixes(state_dict, M.LORA_SUFFIXES)
if network_alpha is not None and any('alpha' in w for w in groups.values()):
network_alpha = None
out = {}
for (prefix, base), w in groups.items():
if 'lora_down.weight' not in w or 'lora_up.weight' not in w:
continue
for path, chunk in native_adapter.resolve_group_targets(M.resolve_targets, prefix, base):
target = native_adapter._slice_lora_chunk(w, chunk) if chunk is not None else w # pylint: disable=protected-access
alpha = network_alpha if 'alpha' not in target else float(target['alpha'])
scale = 1.0 if alpha is None else alpha / target['lora_down.weight'].shape[0]
out[path] = (target['lora_down.weight'], target['lora_up.weight'], scale)
return out
REFERENCE_PREFIXES = ('diffusion_model.', 'blocks.', 'token_refiner.blocks.', 'final_layer.', 'lora_unet_', 'base_model.model.dit.',
'video_patch_proj.', 'audio_patch_proj.', 'condition_proj.', 'time_embedder.proj_')
def oracle_mapping(state_dict, network_alpha=None):
"""{diffusers path: (down, up, scale)} plus the keys no LoRA pair claims, from the layout's reference loader.
Reference-name layouts go through the diffusers converter, which folds any
alpha into the weights. Files already on diffusers names bind verbatim, with
a per-key alpha or else the file-level one applied as ``alpha / rank``.
"""
if any(k.startswith(REFERENCE_PREFIXES) for k in state_dict):
sd = {k.replace('base_model.model.dit.', 'diffusion_model.', 1) if k.startswith('base_model.model.dit.') else k: v for k, v in state_dict.items()}
converted = convert_diffusers(sd)
out = {}
for key, value in converted.items():
if key.endswith('.lora_A.weight'):
path = key[len('transformer.'):-len('.lora_A.weight')]
out[path] = (value, converted[f'transformer.{path}.lora_B.weight'], 1.0)
return out, []
has_alpha = any(k.endswith('.alpha') for k in state_dict)
out, ignored = {}, []
for key, value in state_dict.items():
marker = next((m for m in ('lora_A', 'lora_down') if f'.{m}' in key), None)
if marker is None:
if not any(f'.{m}' in key for m in ('lora_B', 'lora_up', 'alpha')):
ignored.append(key)
continue
path = key[:key.find(f'.{marker}')].removeprefix('transformer.')
up = state_dict[key.replace(marker, 'lora_B' if marker == 'lora_A' else 'lora_up', 1)]
alpha = state_dict.get(f'{path}.alpha', state_dict.get(f'transformer.{path}.alpha'))
if alpha is not None:
scale = float(alpha) / value.shape[0]
elif network_alpha is not None and not has_alpha:
scale = network_alpha / value.shape[0]
else:
scale = 1.0
out[path] = (value, up, scale)
return out, ignored
def assert_same_factors(native, oracle, label):
"""Exact when neither side scales; probe-equal on the effective delta otherwise."""
assert set(native) == set(oracle), f'{label}: targets differ: native-only={sorted(set(native) - set(oracle))[:4]} oracle-only={sorted(set(oracle) - set(native))[:4]}'
probes = 0
for path, (down_o, up_o, scale_o) in oracle.items():
down_n, up_n, scale_n = native[path]
if scale_n == 1.0 and scale_o == 1.0:
assert torch.equal(down_n, down_o), f'{label}: {path} down differs'
assert torch.equal(up_n, up_o), f'{label}: {path} up differs'
continue
probes += 1
x = torch.randn(4, down_o.shape[1])
got = (x @ down_n.float().t() @ up_n.float().t()) * scale_n
ref = (x @ down_o.float().t() @ up_o.float().t()) * scale_o
err = float((got - ref).abs().max() / ref.abs().max().clamp(min=1e-12))
assert err < 1e-4, f'{label}: {path} effective delta differs, relative error {err:.2e}'
return probes
# ============================================================
# Tests - resolution
# ============================================================
CAT_RESOLVE = category('resolve')
CAT_LOADER = category('loader')
CAT_REAL = category('real-files')
def test_every_reference_module_resolves_to_a_real_linear():
"""Every reference key layout resolves onto a Linear that exists in the diffusers model."""
for layout in ({'prefix': 'diffusion_model.'}, {'prefix': ''}, {'flatten': True, 'suffix': KOHYA}, {'prefix': 'base_model.model.dit.'}):
sd = synth(**layout)
got = set(native_mapping(sd))
expected = set(oracle_mapping(sd)[0])
assert got == expected, f'{layout}: native={sorted(got - expected)} missing={sorted(expected - got)}'
assert got <= set(LINEAR_SHAPES), f'{layout}: not real modules: {sorted(got - set(LINEAR_SHAPES))}'
return True
def test_fc1_output_halves_swapped():
"""fc1 gate rows land on the second half of ff.net.0.proj and value rows on the first."""
sd = lora_pair('diffusion_model.blocks.0.mlp.fc1', ref_shape('blocks.0.mlp.fc1'))
up = sd['diffusion_model.blocks.0.mlp.fc1.lora_B.weight']
down, bound_up, _scale = native_mapping(sd)['transformer_blocks.0.ff.net.0.proj']
half = up.shape[0] // 2
assert torch.equal(bound_up[:half], up[half:]), 'value rows must lead'
assert torch.equal(bound_up[half:], up[:half]), 'gate rows must trail'
assert torch.equal(down, sd['diffusion_model.blocks.0.mlp.fc1.lora_A.weight']), 'down is untouched'
return True
def test_fc1_row_extras_follow_the_swap():
"""Per-output extras on fc1 (bias delta, DoRA magnitude) are permuted with the up rows."""
sd = lora_pair('diffusion_model.blocks.0.mlp.fc1', ref_shape('blocks.0.mlp.fc1'))
out = ref_shape('blocks.0.mlp.fc1')[0]
sd['diffusion_model.blocks.0.mlp.fc1.diff_b'] = torch.arange(out, dtype=torch.float32)
sd['diffusion_model.blocks.0.mlp.fc1.dora_scale'] = torch.arange(out, dtype=torch.float32).reshape(out, 1)
groups = M.group_by_suffixes(sd, M.LORA_SUFFIXES)
slot = groups[('diffusion_model.blocks.', '0.mlp.fc1')]
half = out // 2
assert torch.equal(slot['diff_b'][:half], torch.arange(half, out, dtype=torch.float32))
assert torch.equal(slot['dora_scale'][:half, 0], torch.arange(half, out, dtype=torch.float32))
return True
def test_qkv_split_order():
"""Fused qkv rows split as [q; k; v] onto to_q / to_k / to_v."""
sd = lora_pair('diffusion_model.blocks.1.attn.qkv_proj', ref_shape('blocks.1.attn.qkv_proj'))
up = sd['diffusion_model.blocks.1.attn.qkv_proj.lora_B.weight']
mapping = native_mapping(sd)
for i, proj in enumerate(('to_q', 'to_k', 'to_v')):
_down, bound_up, _scale = mapping[f'transformer_blocks.1.attn.{proj}']
assert torch.equal(bound_up, torch.chunk(up, 3, dim=0)[i]), f'{proj} rows'
return True
def test_diffusers_peft_keys_bind_verbatim():
"""A diffusers-PEFT save is already on diffusers names and gets no permutation."""
sd = lora_pair('transformer.transformer_blocks.0.ff.net.0.proj', ref_shape('blocks.0.mlp.fc1'))
down, up, _scale = native_mapping(sd)['transformer_blocks.0.ff.net.0.proj']
assert torch.equal(up, sd['transformer.transformer_blocks.0.ff.net.0.proj.lora_B.weight'])
assert torch.equal(down, sd['transformer.transformer_blocks.0.ff.net.0.proj.lora_A.weight'])
return True
# ============================================================
# Tests - loader against the reference loaders
# ============================================================
def _loader_matches_converter(layout, name, rewrite=None):
sd = synth(**layout)
net = load_native(sd, name=name)
assert net is not None, 'nothing bound'
assert net.mismatch == 0, f'mismatch={net.mismatch}'
reference = diffusers_deltas(sd, rewrite=rewrite)
assert set(net.modules) <= LINEAR_NETKEYS, f'bound to non-linear keys: {sorted(set(net.modules) - LINEAR_NETKEYS)}'
assert_same_deltas(native_deltas(net), reference)
assert len(net.modules) == len(reference), f'bound {len(net.modules)} modules, converter has {len(reference)}'
return True
def test_comfy_layout_matches_converter():
"""diffusion_model.* keys with PEFT suffixes: the published turbo LoRA layout."""
return _loader_matches_converter({'prefix': 'diffusion_model.'}, 'comfy')
def test_bare_reference_layout_matches_converter():
"""Bare reference names, as the reference generate.py saves them."""
return _loader_matches_converter({'prefix': ''}, 'bare')
def test_musubi_kohya_alpha_matches_converter():
"""Flattened lora_unet_ names with kohya suffixes and a non-trivial alpha."""
return _loader_matches_converter({'flatten': True, 'suffix': KOHYA, 'alpha': RANK / 2}, 'musubi')
def test_peft_wrapped_dit_layout_matches_converter():
"""A peft dump wrapping the reference model under a dit attribute maps like diffusion_model."""
return _loader_matches_converter({'prefix': 'base_model.model.dit.'}, 'dit', rewrite=('base_model.model.dit.', 'diffusion_model.'))
def test_block_diagonal_qkv_with_tripled_alpha_matches_converter():
"""A fused qkv stored as stacked A and block-diagonal B with alpha tripled applies each projection at alpha / rank."""
out_q, inp = LINEAR_SHAPES['transformer_blocks.0.attn.to_q']
alpha = RANK / 2
downs = [torch.randn(RANK, inp) for _ in range(3)]
ups = [torch.randn(out_q, RANK) for _ in range(3)]
sd = {
'diffusion_model.blocks.0.attn.qkv_proj.lora_A.weight': torch.cat(downs, dim=0),
'diffusion_model.blocks.0.attn.qkv_proj.lora_B.weight': torch.block_diag(*ups),
'diffusion_model.blocks.0.attn.qkv_proj.alpha': torch.tensor(3 * alpha),
}
net = load_native(sd, name='blockdiag')
got = native_deltas(net)
assert_same_deltas(got, diffusers_deltas(sd))
for i, proj in enumerate(('to_q', 'to_k', 'to_v')):
expected = (ups[i] @ downs[i]) * (alpha / RANK)
assert torch.allclose(got[f'lora_transformer_transformer_blocks_0_attn_{proj}'], expected, rtol=1e-5, atol=1e-5), f'{proj} is not its own projection at alpha / rank'
return True
def test_peft_dump_layout_binds_verbatim():
"""Diffusers names carrying peft's .default. slot and no component prefix bind one to one."""
sd = synth_diffusers(infix='.default')
net = load_native(sd, name='peftdump')
assert net is not None and net.mismatch == 0
assert_same_deltas(native_deltas(net), identity_deltas({k.replace('.default.', '.'): v for k, v in sd.items()}))
assert len(net.modules) == len(LINEAR_SHAPES)
return True
def test_diffusers_names_with_kohya_suffixes_bind():
"""Diffusers names with lora_down / lora_up suffixes and no alpha bind at scale 1."""
sd = synth_diffusers(suffix=KOHYA)
net = load_native(sd, name='kohyadiff')
assert net is not None and net.mismatch == 0
assert_same_deltas(native_deltas(net), identity_deltas(sd))
return True
def test_metadata_alpha_scales_an_alphaless_file():
"""A file-level alpha in the safetensors metadata scales every module by alpha / rank."""
sd = synth_diffusers(infix='.default')
net = load_native(sd, name='metaalpha', metadata={'alpha': '2'})
assert_same_deltas(native_deltas(net), identity_deltas({k.replace('.default.', '.'): v for k, v in sd.items()}, scale=2 / RANK))
return True
def test_metadata_alpha_yields_to_alpha_tensors():
"""A file carrying any alpha tensor keeps its own scaling and ignores the metadata alpha."""
sd = synth_diffusers()
sd['proj_out.alpha'] = torch.tensor(RANK / 2)
net = load_native(sd, name='mixedalpha', metadata={'alpha': '2'})
got = native_deltas(net)
expected = identity_deltas({k: v for k, v in sd.items() if not k.endswith('.alpha')})
expected['lora_transformer_proj_out'] = expected['lora_transformer_proj_out'] * 0.5
assert_same_deltas(got, expected)
return True
def test_non_numeric_metadata_alpha_is_ignored():
"""A metadata alpha that is not a number leaves the file at alpha == rank."""
sd = synth_diffusers()
net = load_native(sd, name='badalpha', metadata={'alpha': 'n/a'})
assert_same_deltas(native_deltas(net), identity_deltas(sd))
return True
# ============================================================
# Tests - real files
# ============================================================
REAL_FILES = sorted(glob.glob(os.path.join(shared.opts.lora_dir, 'MiniMax H3', '*.safetensors')))
def test_real_files_match_reference_loaders():
"""Every local MiniMax H3 LoRA maps onto the same targets as its reference loader, with the same factors."""
if not REAL_FILES:
log.warning(' no local MiniMax H3 LoRA files, skipped')
return True
for path in REAL_FILES:
name = os.path.basename(path)
with safetensors.safe_open(path, framework='pt') as f:
metadata = f.metadata() or {}
network_alpha = M.file_alpha(_MockNetworkOnDisk(path, name, metadata))
sd = safetensors.torch.load_file(path)
native = native_mapping(sd, network_alpha)
oracle, ignored = oracle_mapping(sd, network_alpha)
probes = assert_same_factors(native, oracle, name)
note = f' alpha=file:{network_alpha}' if network_alpha is not None else (' alpha=keys' if probes else '')
note += f' ignored={len(ignored)}' if ignored else ''
log.info(f' {name}: targets={len(native)} {"probed" if probes else "identical"}{note}')
del sd, native, oracle
return True
# ============================================================
# Runner
# ============================================================
def run_tests():
t0 = time.time()
log.warning('=== MiniMax H3 native adapter tests ===')
log.warning('=== Resolution ===')
for fn in [
test_every_reference_module_resolves_to_a_real_linear,
test_fc1_output_halves_swapped,
test_fc1_row_extras_follow_the_swap,
test_qkv_split_order,
test_diffusers_peft_keys_bind_verbatim,
]:
run_test(CAT_RESOLVE, fn)
log.warning('=== Loader vs reference loaders ===')
for fn in [
test_comfy_layout_matches_converter,
test_bare_reference_layout_matches_converter,
test_musubi_kohya_alpha_matches_converter,
test_peft_wrapped_dit_layout_matches_converter,
test_block_diagonal_qkv_with_tripled_alpha_matches_converter,
test_peft_dump_layout_binds_verbatim,
test_diffusers_names_with_kohya_suffixes_bind,
test_metadata_alpha_scales_an_alphaless_file,
test_metadata_alpha_yields_to_alpha_tensors,
test_non_numeric_metadata_alpha_is_ignored,
]:
run_test(CAT_LOADER, fn)
log.warning('=== Real files ===')
run_test(CAT_REAL, test_real_files_match_reference_loaders)
elapsed = time.time() - t0
log.warning('=== Results ===')
total_pass = 0
total_fail = 0
for cat, info in results.items():
status = 'PASS' if info['failed'] == 0 else 'FAIL'
log.info(f' {cat}: {info["passed"]} passed, {info["failed"]} failed [{status}]')
total_pass += info['passed']
total_fail += info['failed']
log.warning(f'Total: {total_pass} passed, {total_fail} failed in {elapsed:.2f}s')
return total_fail == 0
if __name__ == '__main__':
ok = run_tests()
sys.exit(0 if ok else 1)