Files
automatic/test/test-anima-native-adapters.py
CalamitousFelicitousness 5875cdd29e feat(lora): remap anima 1.0 lora block indices onto depth-expanded checkpoints
Anima 2.9B interleaves twelve new blocks among the 28 of Anima 1.0, so a block
index trained against 1.0 names a different block on the expanded model. Every
such key still resolves, since blocks 0 to 27 exist either way, so the
mismatch was silent. The Anima loader now shifts base-depth indices onto the
blocks that carry those weights, keyed by (base depth, expanded depth) and
applied only when the transformer is expanded and the LoRA stays inside the
base depth. Transformer keys move; llm_adapter and text encoder keys keep
their own numbering.
2026-09-05 02:24:04 +01:00

1551 lines
62 KiB
Python

#!/usr/bin/env python
"""
Offline unit tests for Anima native adapter loaders.
Anima is the only native arch with a multi-component network namespace: keys
route into ``lora_transformer_*`` (Cosmos 2.0 DiT), ``lora_llm_adapter_*`` (a
custom Qwen3-projection MLP), or ``lora_te_*`` (Qwen3 text encoder). Routing
is parameterized in ``modules.lora.native_adapter`` via the ``network_prefix``
callable that ``pipelines.anima.anima_lora`` supplies.
Covers the families exposed through native_adapter's generics (LoRA,
LoKR, LoHA, OFT, IA3, GLoRA, Norm, Full), focused on:
- LoRA across all five recognized prefixes (BFL transformer / BFL llm_adapter /
BFL text_encoder / kohya unet / kohya te)
- LoHA on both kohya and BFL transformer paths (the on-disk format of
``scenery-anima-base.safetensors``)
- Cosmos 2.0 path rename direct unit tests (every entry in
``COSMOS_2_FLAT_RENAME``)
- network_prefix_for component routing
- DoRA threading via the universal NetworkModule.finalize_updown hook
- calc_updown shape sanity for LoRA / LoHA across components
Save formats are cross-referenced against real Anima LoRAs in the wild:
- kohya transformer (``lora_unet_blocks_0_self_attn_q_proj.lora_down.weight``):
e.g. ``BlueArcStyle``, every entry under the community ``Anima-Preview*`` set
- kohya text-encoder (``lora_te_layers_0_self_attn_q_proj.lora_down.weight``):
e.g. ``BlueArcStyle`` (mixed unet+te kohya saves)
- kohya LoHA (``lora_unet_blocks_0_cross_attn_q_proj.hada_w1_a``):
e.g. ``scenery-anima-base``
- BFL / AI-toolkit transformer (``diffusion_model.blocks.0.self_attn.q_proj.lora_A.weight``):
community AI-toolkit / sd-scripts saves before kohya-LoCon adoption
The diffusers Cosmos 2.0 transformer layout
(``transformer_blocks[i].attn1.{to_q,to_k,to_v,to_out[0],norm_q,norm_k}``,
``transformer_blocks[i].attn2.{to_q,to_k,to_v,to_out[0]}``,
``transformer_blocks[i].ff.net.{0.proj,2}``, six adaLN linears across
``norm1/norm2/norm3.linear_1/_2``, plus top-level ``time_embed.t_embedder``,
``time_embed.norm``, ``patch_embed.proj``, ``norm_out.linear_1/_2``, and
``proj_out``) is taken straight from
``diffusers.models.transformers.transformer_cosmos.Cosmos2TransformerModel``.
Qwen3 layout (``layers[i].self_attn.{q,k,v,o}_proj``,
``layers[i].mlp.{gate,up,down}_proj``) is taken from
``transformers.models.qwen3.Qwen3Model``.
No running server required.
Usage:
python test/test-anima-native-adapters.py
"""
import os
import sys
import tempfile
import time
import torch
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 ( # pylint: disable=wrong-import-position
network, network_lora, network_hada,
)
from modules.lora import lora_common as l_common # pylint: disable=wrong-import-position
from pipelines.anima import anima_lora as A # 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()
# ============================================================
# Mock Anima pipeline (transformer + llm_adapter + text_encoder)
# ============================================================
# Test scale preserves Anima's cross-attention asymmetry (image hidden vs
# text-encoder hidden differ): HIDDEN is the DiT width, CROSS_DIM is the
# text-encoder projection width used as the kv-source for cross-attention.
# Real Anima-Preview uses HIDDEN=2048, CROSS_DIM=1024; this test uses
# proportional small values so safetensors fixtures stay cheap.
HIDDEN = 96 # DiT hidden width
CROSS_DIM = 48 # cross-attention kv-source width (text encoder output)
HEAD_DIM = 32
MLP_HIDDEN = 256 # transformer feed-forward hidden
ADALN_OUT = HIDDEN # each of norm{1,2,3}.linear_{1,2} outputs HIDDEN
TE_HIDDEN = 64
TE_MLP_HIDDEN = 128
# AnimaLLMAdapter: source_dim=target_dim=model_dim=1024 by default. Tests use
# 64 for size: ADAPTER_DIM is source/target/model_dim, ADAPTER_MLP is
# model_dim * mlp_ratio (4.0 in upstream).
ADAPTER_DIM = 64
ADAPTER_HEAD_DIM = 16
ADAPTER_MLP = ADAPTER_DIM * 4
ADAPTER_VOCAB = 50 # AnimaLLMAdapter.embed num_embeddings (upstream 32128)
N_BLOCKS = 2
N_TE_LAYERS = 2
N_ADAPTER_BLOCKS = 2 # AnimaLLMAdapter num_layers (upstream default 6)
# pylint: disable=attribute-defined-outside-init
class _Holder(torch.nn.Module):
"""Empty container module - children attached dynamically."""
def build_cosmos_block():
"""Mirror diffusers' Cosmos2TransformerBlock layout."""
block = _Holder()
# Self-attention (attn1)
block.attn1 = _Holder()
block.attn1.to_q = torch.nn.Linear(HIDDEN, HIDDEN, bias=False)
block.attn1.to_k = torch.nn.Linear(HIDDEN, HIDDEN, bias=False)
block.attn1.to_v = torch.nn.Linear(HIDDEN, HIDDEN, bias=False)
block.attn1.to_out = torch.nn.ModuleList([
torch.nn.Linear(HIDDEN, HIDDEN, bias=False),
torch.nn.Dropout(0.0),
])
block.attn1.norm_q = torch.nn.RMSNorm(HEAD_DIM)
block.attn1.norm_k = torch.nn.RMSNorm(HEAD_DIM)
# Cross-attention (attn2) - kv source is the text encoder output
block.attn2 = _Holder()
block.attn2.to_q = torch.nn.Linear(HIDDEN, HIDDEN, bias=False)
block.attn2.to_k = torch.nn.Linear(CROSS_DIM, HIDDEN, bias=False)
block.attn2.to_v = torch.nn.Linear(CROSS_DIM, HIDDEN, bias=False)
block.attn2.to_out = torch.nn.ModuleList([
torch.nn.Linear(HIDDEN, HIDDEN, bias=False),
torch.nn.Dropout(0.0),
])
# Feed-forward (diffusers FeedForward style with GELU(proj))
block.ff = _Holder()
block.ff.net = torch.nn.ModuleList()
proj_act = _Holder()
proj_act.proj = torch.nn.Linear(HIDDEN, MLP_HIDDEN, bias=True)
block.ff.net.append(proj_act)
block.ff.net.append(torch.nn.Dropout(0.0))
block.ff.net.append(torch.nn.Linear(MLP_HIDDEN, HIDDEN, bias=True))
# adaLN modulation - six separate Linear modules across norm1/norm2/norm3
for norm_name in ('norm1', 'norm2', 'norm3'):
norm_holder = _Holder()
norm_holder.linear_1 = torch.nn.Linear(HIDDEN, ADALN_OUT, bias=True)
norm_holder.linear_2 = torch.nn.Linear(HIDDEN, ADALN_OUT, bias=True)
setattr(block, norm_name, norm_holder)
return block
def build_mock_transformer(n_blocks=N_BLOCKS):
"""Mirror diffusers' Cosmos2TransformerModel top-level layout."""
transformer = _Holder()
transformer.transformer_blocks = torch.nn.ModuleList([build_cosmos_block() for _ in range(n_blocks)])
# Time embedding
transformer.time_embed = _Holder()
transformer.time_embed.t_embedder = torch.nn.Linear(HIDDEN, HIDDEN, bias=True)
transformer.time_embed.norm = torch.nn.RMSNorm(HIDDEN)
# Patch embedding
transformer.patch_embed = _Holder()
transformer.patch_embed.proj = torch.nn.Linear(HIDDEN, HIDDEN, bias=True)
# Output projection
transformer.norm_out = _Holder()
transformer.norm_out.linear_1 = torch.nn.Linear(HIDDEN, ADALN_OUT, bias=True)
transformer.norm_out.linear_2 = torch.nn.Linear(HIDDEN, ADALN_OUT, bias=True)
transformer.proj_out = torch.nn.Linear(HIDDEN, HIDDEN, bias=True)
return transformer
def _build_adapter_attention(query_dim, context_dim):
"""Mirror the ``Attention`` submodule inside AnimaLLMAdapter's TransformerBlock.
Six learnable submodules per attention: q_proj / q_norm / k_proj / k_norm
/ v_proj / o_proj. q_norm and k_norm are per-head RMSNorms; the projections
use bias=False per upstream.
"""
attn = _Holder()
attn.q_proj = torch.nn.Linear(query_dim, ADAPTER_DIM, bias=False)
attn.q_norm = torch.nn.RMSNorm(ADAPTER_HEAD_DIM, eps=1e-6)
attn.k_proj = torch.nn.Linear(context_dim, ADAPTER_DIM, bias=False)
attn.k_norm = torch.nn.RMSNorm(ADAPTER_HEAD_DIM, eps=1e-6)
attn.v_proj = torch.nn.Linear(context_dim, ADAPTER_DIM, bias=False)
attn.o_proj = torch.nn.Linear(ADAPTER_DIM, query_dim, bias=False)
return attn
def _build_adapter_block():
"""Mirror AnimaLLMAdapter.TransformerBlock with use_self_attn=True.
Layout (from modeling_llm_adapter.py in the Anima-Preview-3 repo):
norm_self_attn (RMSNorm), self_attn (Attention),
norm_cross_attn (RMSNorm), cross_attn (Attention, kv from source),
norm_mlp (RMSNorm), mlp (Sequential[Linear, GELU, Linear]).
Cross-attention's k/v take ``source_dim`` (= ADAPTER_DIM in this mock);
in the real adapter ``source_dim`` is the Qwen3 hidden (1024).
"""
block = _Holder()
block.norm_self_attn = torch.nn.RMSNorm(ADAPTER_DIM, eps=1e-6)
block.self_attn = _build_adapter_attention(query_dim=ADAPTER_DIM, context_dim=ADAPTER_DIM)
block.norm_cross_attn = torch.nn.RMSNorm(ADAPTER_DIM, eps=1e-6)
block.cross_attn = _build_adapter_attention(query_dim=ADAPTER_DIM, context_dim=ADAPTER_DIM)
block.norm_mlp = torch.nn.RMSNorm(ADAPTER_DIM, eps=1e-6)
# nn.Sequential -> named children are '0', '1', '2'. Real LoRAs target the
# two Linears at indices 0 and 2.
block.mlp = torch.nn.Sequential(
torch.nn.Linear(ADAPTER_DIM, ADAPTER_MLP),
torch.nn.GELU(),
torch.nn.Linear(ADAPTER_MLP, ADAPTER_DIM),
)
return block
def build_mock_llm_adapter():
"""Mirror AnimaLLMAdapter (modeling_llm_adapter.py).
Top-level: embed (nn.Embedding), in_proj (Linear or Identity), blocks[N]
(TransformerBlock), out_proj (Linear), norm (RMSNorm). The embed table is a
real LoRA target (e.g. Anima_colorfix_v1): its LoRA decomposes the
[vocab, dim] weight as up@down, identical in shape to a Linear delta, so the
merge path applies it unchanged. named_modules() yields embed as a single
leaf, so it adds exactly one network_name entry (not one per vocab row).
"""
adapter = _Holder()
adapter.embed = torch.nn.Embedding(ADAPTER_VOCAB, ADAPTER_DIM)
adapter.in_proj = torch.nn.Linear(ADAPTER_DIM, ADAPTER_DIM, bias=True)
adapter.blocks = torch.nn.ModuleList([_build_adapter_block() for _ in range(N_ADAPTER_BLOCKS)])
adapter.out_proj = torch.nn.Linear(ADAPTER_DIM, ADAPTER_DIM, bias=True)
adapter.norm = torch.nn.RMSNorm(ADAPTER_DIM, eps=1e-6)
return adapter
def build_mock_text_encoder():
"""Mirror the Qwen3 text-encoder layout."""
te = _Holder()
te.layers = torch.nn.ModuleList()
for _ in range(N_TE_LAYERS):
layer = _Holder()
layer.self_attn = _Holder()
layer.self_attn.q_proj = torch.nn.Linear(TE_HIDDEN, TE_HIDDEN, bias=False)
layer.self_attn.k_proj = torch.nn.Linear(TE_HIDDEN, TE_HIDDEN, bias=False)
layer.self_attn.v_proj = torch.nn.Linear(TE_HIDDEN, TE_HIDDEN, bias=False)
layer.self_attn.o_proj = torch.nn.Linear(TE_HIDDEN, TE_HIDDEN, bias=False)
layer.mlp = _Holder()
layer.mlp.gate_proj = torch.nn.Linear(TE_HIDDEN, TE_MLP_HIDDEN, bias=False)
layer.mlp.up_proj = torch.nn.Linear(TE_HIDDEN, TE_MLP_HIDDEN, bias=False)
layer.mlp.down_proj = torch.nn.Linear(TE_MLP_HIDDEN, TE_HIDDEN, bias=False)
layer.input_layernorm = torch.nn.RMSNorm(TE_HIDDEN)
layer.post_attention_layernorm = torch.nn.RMSNorm(TE_HIDDEN)
te.layers.append(layer)
te.norm = torch.nn.RMSNorm(TE_HIDDEN)
return te
class _MockAnimaPipeline:
"""Class name carries 'Anima' so name-based model-type dispatch routes correctly.
Important: NO ``text_encoder_2`` attribute - its presence would flip the
TE namespace from ``lora_te_`` to ``lora_te1_`` in
``assign_network_names_to_compvis_modules``.
"""
def __init__(self, transformer, llm_adapter, text_encoder):
self.transformer = transformer
self.llm_adapter = llm_adapter
self.text_encoder = text_encoder
class _MockAnimaSdModel:
"""Outer wrapper exposing ``pipe`` + ``network_layer_mapping`` for
``lora_convert.assign_network_names_to_compvis_modules`` to populate."""
def __init__(self, pipe):
self.pipe = pipe
self.network_layer_mapping = {}
self.embedding_db = None
self.__class__.__name__ = 'AnimaTextToImagePipeline'
def install_mock_pipe(n_blocks=N_BLOCKS):
"""Set shared.sd_model to a mock exposing an Anima-shaped 3-component pipeline.
Each test re-installs so stamped ``network_layer_name`` attributes from
prior tests do not leak across runs. ``n_blocks`` sizes the DiT, which the
depth-expansion block remap reads.
"""
transformer = build_mock_transformer(n_blocks)
llm_adapter = build_mock_llm_adapter()
text_encoder = build_mock_text_encoder()
pipe = _MockAnimaPipeline(transformer, llm_adapter, text_encoder)
sd_model = _MockAnimaSdModel(pipe)
from modules.modeldata import model_data
model_data.sd_model = sd_model
return sd_model
# ============================================================
# State-dict synthesizers (one per family/format/component)
# ============================================================
RANK = 8
# --- LoRA, transformer, kohya prefix ---
def sd_lora_kohya_self_attn_q():
"""Kohya transformer LoRA on self-attention.
On-disk: lora_unet_blocks_0_self_attn_q_proj
Cosmos rename -> transformer_blocks_0_attn1_to_q
Network key -> lora_transformer_transformer_blocks_0_attn1_to_q
"""
return {
'lora_unet_blocks_0_self_attn_q_proj.lora_down.weight': torch.randn(RANK, HIDDEN),
'lora_unet_blocks_0_self_attn_q_proj.lora_up.weight': torch.randn(HIDDEN, RANK),
'lora_unet_blocks_0_self_attn_q_proj.alpha': torch.tensor(float(RANK)),
}
def sd_lora_kohya_cross_attn_k():
"""Kohya transformer LoRA on cross-attention k_proj (k from text encoder).
On-disk: lora_unet_blocks_0_cross_attn_k_proj
Cosmos rename -> transformer_blocks_0_attn2_to_k
"""
return {
'lora_unet_blocks_0_cross_attn_k_proj.lora_down.weight': torch.randn(RANK, CROSS_DIM),
'lora_unet_blocks_0_cross_attn_k_proj.lora_up.weight': torch.randn(HIDDEN, RANK),
'lora_unet_blocks_0_cross_attn_k_proj.alpha': torch.tensor(float(RANK)),
}
def sd_lora_kohya_self_attn_output():
"""Kohya transformer LoRA on self_attn.output_proj.
On-disk: lora_unet_blocks_0_self_attn_output_proj
Cosmos rename -> transformer_blocks_0_attn1_to_out_0
"""
return {
'lora_unet_blocks_0_self_attn_output_proj.lora_down.weight': torch.randn(RANK, HIDDEN),
'lora_unet_blocks_0_self_attn_output_proj.lora_up.weight': torch.randn(HIDDEN, RANK),
}
def sd_lora_kohya_mlp_layer1():
"""Kohya transformer LoRA on mlp_layer1.
On-disk: lora_unet_blocks_0_mlp_layer1
Cosmos rename -> transformer_blocks_0_ff_net_0_proj
"""
return {
'lora_unet_blocks_0_mlp_layer1.lora_down.weight': torch.randn(RANK, HIDDEN),
'lora_unet_blocks_0_mlp_layer1.lora_up.weight': torch.randn(MLP_HIDDEN, RANK),
}
def sd_lora_kohya_mlp_layer2():
"""Kohya transformer LoRA on mlp_layer2 (the down-projection).
On-disk: lora_unet_blocks_0_mlp_layer2
Cosmos rename -> transformer_blocks_0_ff_net_2
"""
return {
'lora_unet_blocks_0_mlp_layer2.lora_down.weight': torch.randn(RANK, MLP_HIDDEN),
'lora_unet_blocks_0_mlp_layer2.lora_up.weight': torch.randn(HIDDEN, RANK),
}
def sd_lora_kohya_adaln_self_attn_1():
"""Kohya transformer LoRA on adaLN self-attn modulation linear 1.
On-disk: lora_unet_blocks_0_adaln_modulation_self_attn_1
Cosmos rename -> transformer_blocks_0_norm1_linear_1
"""
return {
'lora_unet_blocks_0_adaln_modulation_self_attn_1.lora_down.weight': torch.randn(RANK, HIDDEN),
'lora_unet_blocks_0_adaln_modulation_self_attn_1.lora_up.weight': torch.randn(ADALN_OUT, RANK),
}
# --- LoRA, text encoder, kohya prefix ---
def sd_lora_kohya_te_q():
"""Kohya TE LoRA - this branch was silently dropped by the legacy resolver
(no ``lora_te_`` kohya prefix recognized) and is newly handled.
On-disk: lora_te_layers_0_self_attn_q_proj
Network key -> lora_te_layers_0_self_attn_q_proj
"""
return {
'lora_te_layers_0_self_attn_q_proj.lora_down.weight': torch.randn(RANK, TE_HIDDEN),
'lora_te_layers_0_self_attn_q_proj.lora_up.weight': torch.randn(TE_HIDDEN, RANK),
'lora_te_layers_0_self_attn_q_proj.alpha': torch.tensor(float(RANK)),
}
def sd_lora_kohya_te_mlp_gate():
"""Kohya TE LoRA on mlp.gate_proj.
On-disk: lora_te_layers_1_mlp_gate_proj
Network key -> lora_te_layers_1_mlp_gate_proj
"""
return {
'lora_te_layers_1_mlp_gate_proj.lora_down.weight': torch.randn(RANK, TE_HIDDEN),
'lora_te_layers_1_mlp_gate_proj.lora_up.weight': torch.randn(TE_MLP_HIDDEN, RANK),
}
# --- LoRA, transformer, BFL prefix ---
def sd_lora_bfl_self_attn_q():
"""BFL/AI-toolkit transformer LoRA on self-attention.
On-disk: diffusion_model.blocks.0.self_attn.q_proj
Cosmos rename of flattened base -> transformer_blocks_0_attn1_to_q
"""
return {
'diffusion_model.blocks.0.self_attn.q_proj.lora_A.weight': torch.randn(RANK, HIDDEN),
'diffusion_model.blocks.0.self_attn.q_proj.lora_B.weight': torch.randn(HIDDEN, RANK),
}
def sd_lora_bfl_x_embedder():
"""BFL transformer LoRA on the patch embedder.
On-disk: diffusion_model.x_embedder_proj_1
Cosmos rename -> patch_embed_proj
"""
return {
'diffusion_model.x_embedder_proj_1.lora_A.weight': torch.randn(RANK, HIDDEN),
'diffusion_model.x_embedder_proj_1.lora_B.weight': torch.randn(HIDDEN, RANK),
}
def sd_lora_bfl_final_layer_linear():
"""BFL transformer LoRA on the final projection.
On-disk: diffusion_model.final_layer_linear
Cosmos rename -> proj_out
"""
return {
'diffusion_model.final_layer_linear.lora_A.weight': torch.randn(RANK, HIDDEN),
'diffusion_model.final_layer_linear.lora_B.weight': torch.randn(HIDDEN, RANK),
}
# --- LoRA, llm_adapter, BFL prefix ---
#
# Paths mirror AnimaLLMAdapter's actual module tree (see modeling_llm_adapter.py
# in the Anima-Preview-3 repo). Each block contains two Attention submodules
# (self_attn, cross_attn) each with q/k/v/o projections plus q/k RMSNorms,
# plus a 3-element Sequential MLP.
def sd_lora_bfl_llm_adapter_in_proj():
"""BFL llm_adapter LoRA on the top-level Qwen3->model_dim projection.
On-disk: diffusion_model.llm_adapter.in_proj
Network key -> lora_llm_adapter_in_proj
"""
return {
'diffusion_model.llm_adapter.in_proj.lora_A.weight': torch.randn(RANK, ADAPTER_DIM),
'diffusion_model.llm_adapter.in_proj.lora_B.weight': torch.randn(ADAPTER_DIM, RANK),
'diffusion_model.llm_adapter.in_proj.alpha': torch.tensor(float(RANK)),
}
def sd_lora_bfl_llm_adapter_self_attn_q():
"""BFL llm_adapter LoRA on a block's self-attention q_proj.
On-disk: diffusion_model.llm_adapter.blocks.0.self_attn.q_proj
Network key -> lora_llm_adapter_blocks_0_self_attn_q_proj
"""
return {
'diffusion_model.llm_adapter.blocks.0.self_attn.q_proj.lora_A.weight': torch.randn(RANK, ADAPTER_DIM),
'diffusion_model.llm_adapter.blocks.0.self_attn.q_proj.lora_B.weight': torch.randn(ADAPTER_DIM, RANK),
}
def sd_lora_bfl_llm_adapter_cross_attn_v():
"""BFL llm_adapter LoRA on a block's cross-attention v_proj.
Cross-attention's v consumes the Qwen3 hidden states (the source the
adapter is built around), so this is a common LoRA target.
On-disk: diffusion_model.llm_adapter.blocks.1.cross_attn.v_proj
Network key -> lora_llm_adapter_blocks_1_cross_attn_v_proj
"""
return {
'diffusion_model.llm_adapter.blocks.1.cross_attn.v_proj.lora_A.weight': torch.randn(RANK, ADAPTER_DIM),
'diffusion_model.llm_adapter.blocks.1.cross_attn.v_proj.lora_B.weight': torch.randn(ADAPTER_DIM, RANK),
}
def sd_lora_bfl_llm_adapter_mlp_in():
"""BFL llm_adapter LoRA on the MLP's input Linear (mlp.0 inside Sequential).
On-disk: diffusion_model.llm_adapter.blocks.0.mlp.0
Network key -> lora_llm_adapter_blocks_0_mlp_0
"""
return {
'diffusion_model.llm_adapter.blocks.0.mlp.0.lora_A.weight': torch.randn(RANK, ADAPTER_DIM),
'diffusion_model.llm_adapter.blocks.0.mlp.0.lora_B.weight': torch.randn(ADAPTER_MLP, RANK),
}
def sd_lora_bfl_llm_adapter_out_proj():
"""BFL llm_adapter LoRA on the top-level output projection.
On-disk: diffusion_model.llm_adapter.out_proj
Network key -> lora_llm_adapter_out_proj
"""
return {
'diffusion_model.llm_adapter.out_proj.lora_A.weight': torch.randn(RANK, ADAPTER_DIM),
'diffusion_model.llm_adapter.out_proj.lora_B.weight': torch.randn(ADAPTER_DIM, RANK),
}
def sd_lora_bfl_llm_adapter_embed():
"""BFL llm_adapter LoRA on the top-level token-embedding table.
The embed weight is an nn.Embedding [vocab, dim]; the LoRA decomposes that
table as up@down (up [vocab, rank], down [rank, dim]). Mirrors the on-disk
layout of Anima_colorfix_v1, which stamps lora_down/lora_up directly under
the BFL diffusion_model.llm_adapter. prefix.
On-disk: diffusion_model.llm_adapter.embed
Network key -> lora_llm_adapter_embed
"""
return {
'diffusion_model.llm_adapter.embed.lora_down.weight': torch.randn(RANK, ADAPTER_DIM),
'diffusion_model.llm_adapter.embed.lora_up.weight': torch.randn(ADAPTER_VOCAB, RANK),
}
def sd_lora_bfl_llm_adapter_mlp_in_bias():
"""BFL llm_adapter LoRA on mlp.0 carrying a companion diff_b bias delta.
Mirrors Anima_colorfix_v1, where the mlp Linears pair a weight LoRA with a bias delta.
On-disk: diffusion_model.llm_adapter.blocks.0.mlp.0 (lora_down/up + diff_b)
Network key -> lora_llm_adapter_blocks_0_mlp_0
"""
return {
'diffusion_model.llm_adapter.blocks.0.mlp.0.lora_down.weight': torch.randn(RANK, ADAPTER_DIM),
'diffusion_model.llm_adapter.blocks.0.mlp.0.lora_up.weight': torch.randn(ADAPTER_MLP, RANK),
'diffusion_model.llm_adapter.blocks.0.mlp.0.diff_b': torch.randn(ADAPTER_MLP),
}
# --- LoRA, text encoder, BFL prefix ---
def sd_lora_bfl_te_q():
"""BFL TE LoRA (the ``text_encoders.qwen3_06b.transformer.model.`` prefix).
On-disk: text_encoders.qwen3_06b.transformer.model.layers.0.self_attn.q_proj
Network key -> lora_te_layers_0_self_attn_q_proj
"""
return {
'text_encoders.qwen3_06b.transformer.model.layers.0.self_attn.q_proj.lora_A.weight': torch.randn(RANK, TE_HIDDEN),
'text_encoders.qwen3_06b.transformer.model.layers.0.self_attn.q_proj.lora_B.weight': torch.randn(TE_HIDDEN, RANK),
}
# --- LoRA combined: transformer + TE in one file ---
def sd_lora_combined_kohya():
"""Kohya save with both transformer and TE LoRAs (mirrors BlueArcStyle)."""
sd = sd_lora_kohya_self_attn_q()
sd.update(sd_lora_kohya_te_q())
return sd
# --- DoRA ---
def sd_lora_kohya_with_dora():
"""Kohya LoRA carrying a dora_scale companion."""
return {
'lora_unet_blocks_0_self_attn_v_proj.lora_down.weight': torch.randn(RANK, HIDDEN),
'lora_unet_blocks_0_self_attn_v_proj.lora_up.weight': torch.randn(HIDDEN, RANK),
'lora_unet_blocks_0_self_attn_v_proj.dora_scale': torch.randn(HIDDEN),
}
# --- LoHA ---
def sd_loha_kohya_cross_attn_q():
"""Kohya LoHA on cross-attention q_proj (mirrors scenery-anima-base layout)."""
return {
'lora_unet_blocks_0_cross_attn_q_proj.hada_w1_a': torch.randn(HIDDEN, RANK),
'lora_unet_blocks_0_cross_attn_q_proj.hada_w1_b': torch.randn(RANK, HIDDEN),
'lora_unet_blocks_0_cross_attn_q_proj.hada_w2_a': torch.randn(HIDDEN, RANK),
'lora_unet_blocks_0_cross_attn_q_proj.hada_w2_b': torch.randn(RANK, HIDDEN),
'lora_unet_blocks_0_cross_attn_q_proj.alpha': torch.tensor(float(RANK)),
}
def sd_loha_kohya_self_attn_v():
"""Kohya LoHA on self-attention v_proj."""
return {
'lora_unet_blocks_1_self_attn_v_proj.hada_w1_a': torch.randn(HIDDEN, RANK),
'lora_unet_blocks_1_self_attn_v_proj.hada_w1_b': torch.randn(RANK, HIDDEN),
'lora_unet_blocks_1_self_attn_v_proj.hada_w2_a': torch.randn(HIDDEN, RANK),
'lora_unet_blocks_1_self_attn_v_proj.hada_w2_b': torch.randn(RANK, HIDDEN),
'lora_unet_blocks_1_self_attn_v_proj.alpha': torch.tensor(float(RANK)),
}
def sd_loha_bfl_transformer():
"""BFL transformer LoHA - same Cosmos rename path as kohya, different file
prefix. Less common in the wild but the resolver handles it."""
return {
'diffusion_model.blocks.0.self_attn.q_proj.hada_w1_a': torch.randn(HIDDEN, RANK),
'diffusion_model.blocks.0.self_attn.q_proj.hada_w1_b': torch.randn(RANK, HIDDEN),
'diffusion_model.blocks.0.self_attn.q_proj.hada_w2_a': torch.randn(HIDDEN, RANK),
'diffusion_model.blocks.0.self_attn.q_proj.hada_w2_b': torch.randn(RANK, HIDDEN),
}
# ============================================================
# Helpers: write state dict to disk, mock NetworkOnDisk
# ============================================================
class TempLora:
"""Context manager: writes a state dict to a temp safetensors file and
yields a ``_MockNetworkOnDisk`` pointing at it. Cleans up on exit."""
def __init__(self, state_dict, name='test'):
self.state_dict = state_dict
self.name = name
self.path = None
def __enter__(self):
sd = {k: v.contiguous() if isinstance(v, torch.Tensor) else v 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)
def __exit__(self, exc_type, exc_val, exc_tb):
if self.path and os.path.exists(self.path):
os.unlink(self.path)
class _MockNetworkOnDisk:
"""Stand-in for ``network.NetworkOnDisk`` exposing only the attributes the
native loaders read."""
def __init__(self, filename, name):
self.filename = filename
self.name = name
self.shorthash = ''
self.sd_version = 'unknown'
def assert_shape(t: torch.Tensor, expected_shape, label=''):
actual = tuple(t.shape)
assert actual == tuple(expected_shape), f'{label}: shape {actual}, expected {tuple(expected_shape)}'
def make_network_for_module(net_module: network.NetworkModule, te_mul: float = 1.0, unet_mul: float = 1.0):
net_module.network.te_multiplier = te_mul
net_module.network.unet_multiplier = unet_mul
return net_module
# ============================================================
# Tests - parsing primitives + Cosmos rename
# ============================================================
CAT_PARSE = category('parse')
def test_parse_key_all_prefixes():
"""parse_key recognizes each of the five Anima prefixes in priority order.
Order matters: ``diffusion_model.llm_adapter.`` must precede
``diffusion_model.`` and ``text_encoders.qwen3_06b.transformer.model.``
must precede ``text_encoders.*`` (none of which exist in Anima but the
priority rule is the same generic mechanism).
"""
cases = [
('lora_unet_blocks_0_self_attn_q_proj.lora_down.weight',
A.LORA_SUFFIXES,
('lora_unet_', 'blocks_0_self_attn_q_proj', 'lora_down.weight')),
('lora_te_layers_0_self_attn_q_proj.lora_down.weight',
A.LORA_SUFFIXES,
('lora_te_', 'layers_0_self_attn_q_proj', 'lora_down.weight')),
('diffusion_model.blocks.0.self_attn.q_proj.lora_A.weight',
A.LORA_SUFFIXES,
('diffusion_model.', 'blocks.0.self_attn.q_proj', 'lora_down.weight')),
('diffusion_model.llm_adapter.input_proj.lora_A.weight',
A.LORA_SUFFIXES,
('diffusion_model.llm_adapter.', 'input_proj', 'lora_down.weight')),
('text_encoders.qwen3_06b.transformer.model.layers.0.self_attn.q_proj.lora_A.weight',
A.LORA_SUFFIXES,
('text_encoders.qwen3_06b.transformer.model.', 'layers.0.self_attn.q_proj', 'lora_down.weight')),
('random.unrelated.key', A.LORA_SUFFIXES, None),
]
for key, suffixes, expected in cases:
got = A.parse_key(key, suffixes)
assert got == expected, f'parse_key({key!r}) = {got}, expected {expected}'
return True
def test_marker_disambiguation():
"""Family markers must reject other families' files."""
pure_lora = {
'lora_unet_blocks_0_self_attn_q_proj.lora_down.weight': torch.zeros(1, 1),
'lora_unet_blocks_0_self_attn_q_proj.lora_up.weight': torch.zeros(1, 1),
}
assert A.has_marker(pure_lora, A.LORA_MARKERS)
assert not A.has_marker(pure_lora, A.LOHA_MARKERS)
assert not A.has_marker(pure_lora, A.LOKR_MARKERS)
assert not A.has_marker(pure_lora, A.OFT_MARKERS)
pure_loha = {
'lora_unet_blocks_0_self_attn_q_proj.hada_w1_a': torch.zeros(1, 1),
'lora_unet_blocks_0_self_attn_q_proj.hada_w1_b': torch.zeros(1, 1),
'lora_unet_blocks_0_self_attn_q_proj.hada_w2_a': torch.zeros(1, 1),
'lora_unet_blocks_0_self_attn_q_proj.hada_w2_b': torch.zeros(1, 1),
}
assert A.has_marker(pure_loha, A.LOHA_MARKERS)
assert not A.has_marker(pure_loha, A.LORA_MARKERS)
return True
def test_cosmos_rename_full_coverage():
"""Every entry in COSMOS_2_FLAT_RENAME applies as a substring rename.
Confirms order-sensitivity guards the longer-substring-first rule (e.g.
``t_embedder_1`` precedes any bare ``t_embedder`` lookalike).
"""
cases = [
# Module-level
('t_embedder_1', 'time_embed_t_embedder'),
('t_embedding_norm', 'time_embed_norm'),
('x_embedder_proj_1', 'patch_embed_proj'),
('final_layer_linear', 'proj_out'),
('final_layer_adaln_modulation_1', 'norm_out_linear_1'),
('final_layer_adaln_modulation_2', 'norm_out_linear_2'),
# Block-level prefix
('blocks_0_self_attn_q_proj', 'transformer_blocks_0_attn1_to_q'),
('blocks_5_self_attn_k_proj', 'transformer_blocks_5_attn1_to_k'),
('blocks_3_self_attn_v_proj', 'transformer_blocks_3_attn1_to_v'),
('blocks_2_self_attn_output_proj', 'transformer_blocks_2_attn1_to_out_0'),
('blocks_1_cross_attn_q_proj', 'transformer_blocks_1_attn2_to_q'),
('blocks_4_cross_attn_k_proj', 'transformer_blocks_4_attn2_to_k'),
('blocks_0_cross_attn_v_proj', 'transformer_blocks_0_attn2_to_v'),
('blocks_0_cross_attn_output_proj', 'transformer_blocks_0_attn2_to_out_0'),
# qk_norm RMSNorms
('blocks_0_self_attn_q_norm', 'transformer_blocks_0_attn1_norm_q'),
('blocks_0_self_attn_k_norm', 'transformer_blocks_0_attn1_norm_k'),
# MLP
('blocks_0_mlp_layer1', 'transformer_blocks_0_ff_net_0_proj'),
('blocks_0_mlp_layer2', 'transformer_blocks_0_ff_net_2'),
# adaLN modulation - six separate Linear modules
('blocks_0_adaln_modulation_self_attn_1', 'transformer_blocks_0_norm1_linear_1'),
('blocks_0_adaln_modulation_self_attn_2', 'transformer_blocks_0_norm1_linear_2'),
('blocks_0_adaln_modulation_cross_attn_1', 'transformer_blocks_0_norm2_linear_1'),
('blocks_0_adaln_modulation_cross_attn_2', 'transformer_blocks_0_norm2_linear_2'),
('blocks_0_adaln_modulation_mlp_1', 'transformer_blocks_0_norm3_linear_1'),
('blocks_0_adaln_modulation_mlp_2', 'transformer_blocks_0_norm3_linear_2'),
# Depth-expanded checkpoints (Anima-2.9B carries 40 blocks)
('blocks_39_self_attn_q_proj', 'transformer_blocks_39_attn1_to_q'),
('blocks_39_cross_attn_output_proj', 'transformer_blocks_39_attn2_to_out_0'),
('blocks_39_mlp_layer2', 'transformer_blocks_39_ff_net_2'),
]
for src, expected in cases:
got = A.cosmos_rename_flat(src)
assert got == expected, f'cosmos_rename_flat({src!r}) = {got!r}, expected {expected!r}'
return True
def test_resolve_targets_per_prefix():
"""resolve_targets emits one (path, None) per call - Anima has no fused QKV."""
# Each call returns exactly one target with no ChunkSpec.
assert A.resolve_targets('lora_unet_', 'blocks_0_self_attn_q_proj') == [
('transformer_blocks_0_attn1_to_q', None),
]
assert A.resolve_targets('diffusion_model.', 'blocks.0.self_attn.q_proj') == [
('transformer_blocks_0_attn1_to_q', None),
]
assert A.resolve_targets('diffusion_model.', 'blocks.39.cross_attn.k_proj') == [
('transformer_blocks_39_attn2_to_k', None),
]
assert A.resolve_targets('diffusion_model.llm_adapter.', 'input_proj') == [
('input_proj', None),
]
assert A.resolve_targets('text_encoders.qwen3_06b.transformer.model.', 'layers.0.self_attn.q_proj') == [
('layers_0_self_attn_q_proj', None),
]
assert A.resolve_targets('lora_te_', 'layers_0_self_attn_q_proj') == [
('layers_0_self_attn_q_proj', None),
]
# Unknown prefix yields empty target list.
assert A.resolve_targets('unknown_prefix', 'foo') == []
return True
def test_network_prefix_for_routing():
"""network_prefix_for picks the namespace per matched prefix."""
assert A.network_prefix_for('diffusion_model.llm_adapter.') == 'lora_llm_adapter_'
assert A.network_prefix_for('text_encoders.qwen3_06b.transformer.model.') == 'lora_te_'
assert A.network_prefix_for('lora_te_') == 'lora_te_'
assert A.network_prefix_for('lora_unet_') == 'lora_transformer_'
assert A.network_prefix_for('diffusion_model.') == 'lora_transformer_'
# Anything else falls through to transformer (default).
assert A.network_prefix_for(None) == 'lora_transformer_'
return True
# ============================================================
# Tests - loaders end-to-end
# ============================================================
CAT_LOADER = category('loader')
def _load_via(try_fn, state_dict, name='test'):
install_mock_pipe()
with TempLora(state_dict, name=name) as nod:
return try_fn(name, nod, lora_scale=1.0)
def test_lora_kohya_self_attn_q():
"""Kohya transformer LoRA on self_attn.q_proj binds via Cosmos rename."""
net = _load_via(A.try_load_lora, sd_lora_kohya_self_attn_q())
assert net is not None and len(net.modules) == 1, f'got {net.modules if net else None}'
assert 'lora_transformer_transformer_blocks_0_attn1_to_q' in net.modules
mod = next(iter(net.modules.values()))
assert isinstance(mod, network_lora.NetworkModuleLora)
return True
def test_lora_kohya_cross_attn_k():
"""Kohya transformer LoRA on cross_attn.k_proj (asymmetric in/out shape)."""
net = _load_via(A.try_load_lora, sd_lora_kohya_cross_attn_k())
assert net is not None and len(net.modules) == 1
assert 'lora_transformer_transformer_blocks_0_attn2_to_k' in net.modules
return True
def test_lora_kohya_self_attn_output():
"""Kohya transformer LoRA on self_attn.output_proj renames to attn1.to_out.0."""
net = _load_via(A.try_load_lora, sd_lora_kohya_self_attn_output())
assert net is not None and len(net.modules) == 1
assert 'lora_transformer_transformer_blocks_0_attn1_to_out_0' in net.modules
return True
def test_lora_kohya_mlp_layer1():
"""Kohya transformer LoRA on mlp_layer1 renames to ff.net.0.proj."""
net = _load_via(A.try_load_lora, sd_lora_kohya_mlp_layer1())
assert net is not None and len(net.modules) == 1
assert 'lora_transformer_transformer_blocks_0_ff_net_0_proj' in net.modules
return True
def test_lora_kohya_mlp_layer2():
"""Kohya transformer LoRA on mlp_layer2 renames to ff.net.2."""
net = _load_via(A.try_load_lora, sd_lora_kohya_mlp_layer2())
assert net is not None and len(net.modules) == 1
assert 'lora_transformer_transformer_blocks_0_ff_net_2' in net.modules
return True
def test_lora_kohya_adaln_self_attn_1():
"""Kohya transformer LoRA on adaLN self-attn modulation linear 1.
On-disk ``adaln_modulation_self_attn_1`` is renamed to ``norm1_linear_1``
because the diffusers Cosmos2 transformer block exposes its self-attention
adaLN as ``norm1.linear_1``.
"""
net = _load_via(A.try_load_lora, sd_lora_kohya_adaln_self_attn_1())
assert net is not None and len(net.modules) == 1
assert 'lora_transformer_transformer_blocks_0_norm1_linear_1' in net.modules
return True
def test_lora_kohya_te_q():
"""Kohya text-encoder LoRA (the branch the legacy resolver silently dropped).
Real fixture: ``BlueArcStyle`` contained 588 such keys; the pre-refactor
loader recognized only ``lora_unet_blocks_*`` and ignored everything else.
"""
net = _load_via(A.try_load_lora, sd_lora_kohya_te_q())
assert net is not None and len(net.modules) == 1, f'got {net.modules if net else None}'
assert 'lora_te_layers_0_self_attn_q_proj' in net.modules
return True
def test_lora_kohya_te_mlp_gate():
"""Kohya TE LoRA on mlp.gate_proj."""
net = _load_via(A.try_load_lora, sd_lora_kohya_te_mlp_gate())
assert net is not None and len(net.modules) == 1
assert 'lora_te_layers_1_mlp_gate_proj' in net.modules
return True
def test_lora_bfl_self_attn_q():
"""BFL transformer LoRA converges to the same network key as the kohya form."""
net = _load_via(A.try_load_lora, sd_lora_bfl_self_attn_q())
assert net is not None and len(net.modules) == 1
assert 'lora_transformer_transformer_blocks_0_attn1_to_q' in net.modules
return True
def test_lora_bfl_x_embedder():
"""BFL transformer LoRA on x_embedder_proj_1 renames to patch_embed.proj."""
net = _load_via(A.try_load_lora, sd_lora_bfl_x_embedder())
assert net is not None and len(net.modules) == 1
assert 'lora_transformer_patch_embed_proj' in net.modules
return True
def test_lora_bfl_final_layer_linear():
"""BFL transformer LoRA on final_layer_linear renames to proj_out."""
net = _load_via(A.try_load_lora, sd_lora_bfl_final_layer_linear())
assert net is not None and len(net.modules) == 1
assert 'lora_transformer_proj_out' in net.modules
return True
def test_lora_bfl_llm_adapter_in_proj():
"""BFL llm_adapter LoRA on the top-level Qwen3->model_dim projection.
Confirms top-level (non-block) paths route into ``lora_llm_adapter_*``.
"""
net = _load_via(A.try_load_lora, sd_lora_bfl_llm_adapter_in_proj())
assert net is not None and len(net.modules) == 1, f'got {net.modules if net else None}'
assert 'lora_llm_adapter_in_proj' in net.modules
return True
def test_lora_bfl_llm_adapter_self_attn_q():
"""BFL llm_adapter LoRA on a block's self_attn.q_proj.
Confirms block-nested attention projections route correctly through the
underscore-flattened path (no Cosmos rename on the adapter branch).
"""
net = _load_via(A.try_load_lora, sd_lora_bfl_llm_adapter_self_attn_q())
assert net is not None and len(net.modules) == 1
assert 'lora_llm_adapter_blocks_0_self_attn_q_proj' in net.modules
return True
def test_lora_bfl_llm_adapter_cross_attn_v():
"""BFL llm_adapter LoRA on cross_attn.v_proj.
Cross-attention is what the adapter does: v_proj consumes the Qwen3
hidden states from the source. Common LoRA target.
"""
net = _load_via(A.try_load_lora, sd_lora_bfl_llm_adapter_cross_attn_v())
assert net is not None and len(net.modules) == 1
assert 'lora_llm_adapter_blocks_1_cross_attn_v_proj' in net.modules
return True
def test_lora_bfl_llm_adapter_mlp_in():
"""BFL llm_adapter LoRA on mlp.0 (the Linear inside Sequential).
Confirms numeric Sequential indices flatten correctly: ``mlp.0`` -> ``mlp_0``.
"""
net = _load_via(A.try_load_lora, sd_lora_bfl_llm_adapter_mlp_in())
assert net is not None and len(net.modules) == 1
assert 'lora_llm_adapter_blocks_0_mlp_0' in net.modules
return True
def test_lora_bfl_llm_adapter_out_proj():
"""BFL llm_adapter LoRA on the top-level output projection."""
net = _load_via(A.try_load_lora, sd_lora_bfl_llm_adapter_out_proj())
assert net is not None and len(net.modules) == 1
assert 'lora_llm_adapter_out_proj' in net.modules
return True
def test_lora_bfl_llm_adapter_embed():
"""BFL llm_adapter LoRA on the embedding table binds to lora_llm_adapter_embed.
Regression: nn.Embedding is a valid LoRA target (Anima_colorfix_v1).
"""
net = _load_via(A.try_load_lora, sd_lora_bfl_llm_adapter_embed())
assert net is not None and len(net.modules) == 1, f'got {net.modules if net else None}'
assert 'lora_llm_adapter_embed' in net.modules
mod = next(iter(net.modules.values()))
assert isinstance(mod, network_lora.NetworkModuleLora)
return True
def test_lora_llm_adapter_mlp_bias_delta():
"""mlp.0 LoRA + diff_b binds one module carrying both weight and bias deltas.
Through the umbrella try_load: diff_b folds into the LoRA module as ex_bias;
try_load_full must not spawn a second module on the same network key.
"""
net = _load_via(A.try_load, sd_lora_bfl_llm_adapter_mlp_in_bias())
assert net is not None and len(net.modules) == 1, f'got {net.modules if net else None}'
assert 'lora_llm_adapter_blocks_0_mlp_0' in net.modules
mod = next(iter(net.modules.values()))
assert isinstance(mod, network_lora.NetworkModuleLora)
assert mod.ex_bias is not None, 'diff_b not folded into the LoRA module as ex_bias'
return True
def test_lora_bfl_te_q():
"""BFL TE LoRA via the ``text_encoders.qwen3_06b.transformer.model.`` prefix.
Converges to the same network key as the kohya ``lora_te_`` form.
"""
net = _load_via(A.try_load_lora, sd_lora_bfl_te_q())
assert net is not None and len(net.modules) == 1
assert 'lora_te_layers_0_self_attn_q_proj' in net.modules
return True
def test_lora_combined_kohya_transformer_plus_te():
"""A single file with both transformer and TE LoRAs binds modules into both
namespaces in one load (mirrors the BlueArcStyle layout)."""
net = _load_via(A.try_load_lora, sd_lora_combined_kohya())
assert net is not None and len(net.modules) == 2, f'got {net.modules if net else None}'
assert 'lora_transformer_transformer_blocks_0_attn1_to_q' in net.modules
assert 'lora_te_layers_0_self_attn_q_proj' in net.modules
return True
def test_lora_dora_threading():
"""dora_scale flows into NetworkModuleLora.dora_scale via NetworkWeights."""
net = _load_via(A.try_load_lora, sd_lora_kohya_with_dora())
assert net is not None and len(net.modules) == 1
mod = next(iter(net.modules.values()))
assert mod.dora_scale is not None, 'dora_scale not threaded into NetworkModule'
return True
def test_loha_kohya_cross_attn_q():
"""Kohya LoHA on cross_attn.q_proj (mirrors scenery-anima-base layout)."""
net = _load_via(A.try_load_loha, sd_loha_kohya_cross_attn_q())
assert net is not None and len(net.modules) == 1, f'got {net.modules if net else None}'
assert 'lora_transformer_transformer_blocks_0_attn2_to_q' in net.modules
mod = next(iter(net.modules.values()))
assert isinstance(mod, network_hada.NetworkModuleHada)
return True
def test_loha_kohya_self_attn_v():
"""Kohya LoHA on self_attn.v_proj."""
net = _load_via(A.try_load_loha, sd_loha_kohya_self_attn_v())
assert net is not None and len(net.modules) == 1
assert 'lora_transformer_transformer_blocks_1_attn1_to_v' in net.modules
return True
def test_loha_bfl_transformer():
"""BFL-format LoHA produces the same network key as the kohya form."""
net = _load_via(A.try_load_loha, sd_loha_bfl_transformer())
assert net is not None and len(net.modules) == 1
assert 'lora_transformer_transformer_blocks_0_attn1_to_q' in net.modules
return True
def test_loha_marker_rejection():
"""A pure LoRA file produces no LoHA modules (and vice versa).
Loaders gate on family-specific markers; cross-family routing yields None
rather than a partial network.
"""
# LoRA file routed through LoHA loader -> no LoHA markers -> None
net = _load_via(A.try_load_loha, sd_lora_kohya_self_attn_q())
assert net is None, f'LoHA loader matched a LoRA file: {net.modules if net else None}'
# LoHA file routed through LoRA loader -> no LoRA markers -> None
net = _load_via(A.try_load_lora, sd_loha_kohya_cross_attn_q())
assert net is None, f'LoRA loader matched a LoHA file: {net.modules if net else None}'
return True
def test_try_load_chain_routes_lora_and_loha():
"""try_load (the umbrella) dispatches to whichever family matches."""
# Pure LoRA file
net = _load_via(A.try_load, sd_lora_kohya_self_attn_q())
assert net is not None and len(net.modules) == 1
assert isinstance(next(iter(net.modules.values())), network_lora.NetworkModuleLora)
# Pure LoHA file
net = _load_via(A.try_load, sd_loha_kohya_cross_attn_q())
assert net is not None and len(net.modules) == 1
assert isinstance(next(iter(net.modules.values())), network_hada.NetworkModuleHada)
return True
# ============================================================
# Tests - calc_updown shape sanity
# ============================================================
CAT_MATH = category('math')
def test_lora_calc_updown_transformer():
"""NetworkModuleLora.calc_updown shape sanity against a transformer target."""
net = _load_via(A.try_load_lora, sd_lora_kohya_self_attn_q())
mod = make_network_for_module(next(iter(net.modules.values())))
target = torch.randn(HIDDEN, HIDDEN)
updown, _ = mod.calc_updown(target)
assert_shape(updown, target.shape, label='LoRA transformer calc_updown')
return True
def test_lora_calc_updown_cross_attn_asymmetric():
"""Cross-attention k_proj has Linear(CROSS_DIM, HIDDEN); the LoRA update
must match that asymmetric in/out shape."""
net = _load_via(A.try_load_lora, sd_lora_kohya_cross_attn_k())
mod = make_network_for_module(next(iter(net.modules.values())))
target = torch.randn(HIDDEN, CROSS_DIM)
updown, _ = mod.calc_updown(target)
assert_shape(updown, target.shape, label='LoRA cross-attn asymmetric calc_updown')
return True
def test_lora_calc_updown_llm_adapter():
"""LoRA calc_updown shape sanity against an llm_adapter target."""
net = _load_via(A.try_load_lora, sd_lora_bfl_llm_adapter_self_attn_q())
mod = make_network_for_module(next(iter(net.modules.values())))
target = torch.randn(ADAPTER_DIM, ADAPTER_DIM)
updown, _ = mod.calc_updown(target)
assert_shape(updown, target.shape, label='LoRA llm_adapter calc_updown')
return True
def test_lora_calc_updown_llm_adapter_embed():
"""LoRA calc_updown against the embedding table yields a [vocab, dim] delta.
up@down (up [vocab, rank] @ down [rank, dim]) reconstructs the nn.Embedding
weight shape so network_add_weights can sum it into the table.
"""
net = _load_via(A.try_load_lora, sd_lora_bfl_llm_adapter_embed())
mod = make_network_for_module(next(iter(net.modules.values())))
target = torch.randn(ADAPTER_VOCAB, ADAPTER_DIM)
updown, _ = mod.calc_updown(target)
assert_shape(updown, target.shape, label='LoRA llm_adapter embed calc_updown')
return True
def test_lora_calc_updown_with_bias_delta():
"""calc_updown on a LoRA+diff_b module returns a weight updown and a bias ex_bias.
The bias delta rides the second tuple element (ex_bias), which network_calc_weights
accumulates into batch_ex_bias for application to the module's bias.
"""
net = _load_via(A.try_load, sd_lora_bfl_llm_adapter_mlp_in_bias())
mod = make_network_for_module(next(iter(net.modules.values())))
target = torch.randn(ADAPTER_MLP, ADAPTER_DIM) # mlp.0 weight shape [out, in]
updown, ex_bias = mod.calc_updown(target)
assert_shape(updown, target.shape, label='LoRA+bias updown')
assert ex_bias is not None, 'ex_bias missing for diff_b module'
assert_shape(ex_bias, (ADAPTER_MLP,), label='LoRA+bias ex_bias')
return True
def test_lora_calc_updown_text_encoder():
"""LoRA calc_updown shape sanity against a TE target."""
net = _load_via(A.try_load_lora, sd_lora_kohya_te_q())
mod = make_network_for_module(next(iter(net.modules.values())))
target = torch.randn(TE_HIDDEN, TE_HIDDEN)
updown, _ = mod.calc_updown(target)
assert_shape(updown, target.shape, label='LoRA text-encoder calc_updown')
return True
def test_loha_calc_updown_transformer():
"""NetworkModuleHada.calc_updown shape sanity against a transformer target."""
net = _load_via(A.try_load_loha, sd_loha_kohya_cross_attn_q())
mod = make_network_for_module(next(iter(net.modules.values())))
target = torch.randn(HIDDEN, HIDDEN)
updown, _ = mod.calc_updown(target)
assert_shape(updown, target.shape, label='LoHA transformer calc_updown')
return True
# === DoRA convention tests (apply_weight_decompose dual-path) ===
#
# DoRA stores a per-axis magnitude vector that the apply path renormalizes
# against. Two conventions exist in the wild:
#
# - per-input (DoRA paper / kohya): dora_scale shape ``(1, in)`` or ``(in,)``,
# ``W' = W * (m / ||W||_col)`` rescales each column to magnitude m[i].
# - per-output (LyCORIS / PEFT): dora_scale shape ``(out, 1)`` or ``(out,)``,
# ``W' = W * (m / ||W||_row)`` rescales each row to magnitude m[o].
#
# The pre-fix apply_weight_decompose computed per-input norm regardless and
# silently broadcast against per-output dora_scale via ``(out, 1) / (1, in)
# -> (out, in)``, scrambling the update. These tests guard the dual-path
# detection from regressing.
def _per_input_dora_sd():
"""Synthetic LoRA + per-input dora_scale (kohya convention) on a transformer target."""
return {
'lora_unet_blocks_0_self_attn_q_proj.lora_down.weight': torch.randn(RANK, HIDDEN),
'lora_unet_blocks_0_self_attn_q_proj.lora_up.weight': torch.randn(HIDDEN, RANK),
'lora_unet_blocks_0_self_attn_q_proj.dora_scale': torch.randn(1, HIDDEN).abs() + 0.1,
}
def _per_output_dora_sd():
"""Synthetic LoRA + per-output dora_scale (LyCORIS / PEFT convention).
Targets cross_attn.k_proj which has asymmetric shape ``(HIDDEN, CROSS_DIM)``,
so the ``(HIDDEN, 1)`` dora_scale is structurally distinguishable from
a ``(1, CROSS_DIM)`` per-input scale.
"""
return {
'lora_unet_blocks_0_cross_attn_k_proj.lora_down.weight': torch.randn(RANK, CROSS_DIM),
'lora_unet_blocks_0_cross_attn_k_proj.lora_up.weight': torch.randn(HIDDEN, RANK),
'lora_unet_blocks_0_cross_attn_k_proj.dora_scale': torch.randn(HIDDEN, 1).abs() + 0.1,
}
def test_dora_per_input_convention():
"""Per-input dora_scale (1, in): merged weight columns match dora_scale.
Standard DoRA paper convention. ``W' = W * (m / ||W||_col)`` produces
new column norms equal to dora_scale.
"""
net = _load_via(A.try_load_lora, _per_input_dora_sd())
mod = make_network_for_module(next(iter(net.modules.values())))
target = torch.randn(HIDDEN, HIDDEN) * 0.05
updown, _ = mod.calc_updown(target)
new_weight = target + updown
dora_scale = mod.dora_scale.squeeze() # (in,)
col_norms = new_weight.norm(dim=0) # (in,)
# Allow some numerical slack; the calc_scale factor (alpha/dim) may multiply
# updown after the DoRA path. Without alpha, calc_scale returns 1.0 and the
# match is exact to numerical precision.
rel_err = (col_norms - dora_scale).abs() / dora_scale.abs().clamp_min(1e-6)
assert rel_err.max() < 1e-3, f'per-input col norms diverge from dora_scale: max rel err {rel_err.max():.2e}'
return True
def test_dora_per_output_convention():
"""Per-output dora_scale (out, 1): merged weight rows match dora_scale.
LyCORIS / PEFT convention. ``W' = W * (m / ||W||_row)`` produces new row
norms equal to dora_scale. Regression target for the LoKR+DoRA file
(カードキャプターさくらAnimaPreview3Base) that surfaced this bug.
"""
net = _load_via(A.try_load_lora, _per_output_dora_sd())
mod = make_network_for_module(next(iter(net.modules.values())))
target = torch.randn(HIDDEN, CROSS_DIM) * 0.05
updown, _ = mod.calc_updown(target)
new_weight = target + updown
dora_scale = mod.dora_scale.squeeze() # (out,)
row_norms = new_weight.norm(dim=1) # (out,)
rel_err = (row_norms - dora_scale).abs() / dora_scale.abs().clamp_min(1e-6)
assert rel_err.max() < 1e-3, f'per-output row norms diverge from dora_scale: max rel err {rel_err.max():.2e}'
return True
def test_dora_square_weight_1d_defaults_to_per_input():
"""When the target is square (out == in) AND dora_scale is 1D, the length
matches both axes and shape can't disambiguate. Default to per-input to
preserve the dominant kohya save behavior.
2D dora_scale on a square target is structurally unambiguous and is
handled correctly by ``test_dora_per_output_convention`` against an
asymmetric target. This test guards the 1D-on-square ambiguity branch.
"""
sd = {
'lora_unet_blocks_0_self_attn_q_proj.lora_down.weight': torch.randn(RANK, HIDDEN),
'lora_unet_blocks_0_self_attn_q_proj.lora_up.weight': torch.randn(HIDDEN, RANK),
'lora_unet_blocks_0_self_attn_q_proj.dora_scale': torch.randn(HIDDEN).abs() + 0.1, # 1D, length HIDDEN
}
net = _load_via(A.try_load_lora, sd)
mod = make_network_for_module(next(iter(net.modules.values())))
target = torch.randn(HIDDEN, HIDDEN) * 0.05
updown, _ = mod.calc_updown(target)
new_weight = target + updown
dora_scale = mod.dora_scale.squeeze()
# Per-input branch: column norms match dora_scale.
col_norms = new_weight.norm(dim=0)
rel_err_col = (col_norms - dora_scale).abs() / dora_scale.abs().clamp_min(1e-6)
assert rel_err_col.max() < 1e-3, f'1D square dora_scale should default to per-input: max rel err {rel_err_col.max():.2e}'
return True
# ============================================================
# Tests - depth-expanded checkpoint block remap
# ============================================================
CAT_REMAP = category('remap')
BASE_DEPTH = 28
EXPANDED_DEPTH = 40
def sd_lora_kohya_block(index):
"""Kohya transformer LoRA on self_attn.q_proj at an arbitrary block index."""
stem = f'lora_unet_blocks_{index}_self_attn_q_proj'
return {
f'{stem}.lora_down.weight': torch.randn(RANK, HIDDEN),
f'{stem}.lora_up.weight': torch.randn(HIDDEN, RANK),
f'{stem}.alpha': torch.tensor(float(RANK)),
}
def _load_at_depth(try_fn, state_dict, n_blocks, name='test'):
install_mock_pipe(n_blocks)
with TempLora(state_dict, name=name) as nod:
return try_fn(name, nod, lora_scale=1.0)
def test_expansion_map_matches_manifest():
"""The 28->40 table skips the author's insertion positions in order."""
table = A.expansion_map(BASE_DEPTH, EXPANDED_DEPTH)
inserted = set(A.BLOCK_EXPANSIONS[(BASE_DEPTH, EXPANDED_DEPTH)])
assert len(table) == BASE_DEPTH
assert set(table.values()).isdisjoint(inserted), 'no base block may land on an inserted block'
assert (table[0], table[1], table[2], table[27]) == (0, 1, 3, 39)
assert list(table.values()) == sorted(table.values()), 'order must be preserved'
return True
def test_expansion_map_covers_every_base_block_uniquely():
"""All base indices land on distinct blocks inside the expanded depth."""
table = A.expansion_map(BASE_DEPTH, EXPANDED_DEPTH)
assert len(set(table.values())) == BASE_DEPTH
assert max(table.values()) < EXPANDED_DEPTH
return True
def test_remap_binds_last_base_block_to_top_of_expanded():
"""A 1.0 LoRA on block 27 binds to block 39 of a 40-block transformer."""
net = _load_at_depth(A.try_load_lora, sd_lora_kohya_block(27), EXPANDED_DEPTH)
assert net is not None and len(net.modules) == 1, f'got {net.modules if net else None}'
assert 'lora_transformer_transformer_blocks_39_attn1_to_q' in net.modules
return True
def test_remap_leaves_blocks_before_first_insertion():
"""Blocks 0 and 1 precede the first insertion, so they do not move."""
for i in (0, 1):
net = _load_at_depth(A.try_load_lora, sd_lora_kohya_block(i), EXPANDED_DEPTH)
assert f'lora_transformer_transformer_blocks_{i}_attn1_to_q' in net.modules
return True
def test_no_remap_on_base_depth_model():
"""A 28-block transformer matches no expansion entry, so indices pass through."""
net = _load_at_depth(A.try_load_lora, sd_lora_kohya_block(1), BASE_DEPTH)
assert 'lora_transformer_transformer_blocks_1_attn1_to_q' in net.modules
return True
def test_no_remap_when_lora_is_native_to_expanded_depth():
"""A LoRA reaching past the base depth was trained on the expanded model."""
sd = sd_lora_kohya_block(5)
sd.update(sd_lora_kohya_block(39))
net = _load_at_depth(A.try_load_lora, sd, EXPANDED_DEPTH)
assert 'lora_transformer_transformer_blocks_5_attn1_to_q' in net.modules
assert 'lora_transformer_transformer_blocks_39_attn1_to_q' in net.modules
return True
def test_remap_leaves_adapter_and_te_indices_alone():
"""llm_adapter blocks and text encoder layers carry unrelated numbering."""
install_mock_pipe(EXPANDED_DEPTH)
out = A.remap_blocks({
('lora_unet_', 'blocks_27_self_attn_q_proj'): {},
('diffusion_model.llm_adapter.', 'blocks.1.self_attn.q_proj'): {},
('lora_te_', 'layers_1_self_attn_q_proj'): {},
})
assert ('lora_unet_', 'blocks_39_self_attn_q_proj') in out
assert ('diffusion_model.llm_adapter.', 'blocks.1.self_attn.q_proj') in out
assert ('lora_te_', 'layers_1_self_attn_q_proj') in out
return True
def test_remap_handles_dotted_bfl_keys():
"""BFL keys arrive dotted rather than underscore-flattened."""
install_mock_pipe(EXPANDED_DEPTH)
out = A.remap_blocks({('diffusion_model.', 'blocks.27.self_attn.q_proj'): {}})
assert ('diffusion_model.', 'blocks.39.self_attn.q_proj') in out
return True
# ============================================================
# Test runner
# ============================================================
def run_tests():
t0 = time.time()
log.warning('=== Parsing primitives + Cosmos rename ===')
for fn in [
test_parse_key_all_prefixes,
test_marker_disambiguation,
test_cosmos_rename_full_coverage,
test_resolve_targets_per_prefix,
test_network_prefix_for_routing,
]:
run_test(CAT_PARSE, fn)
log.warning('=== Loaders ===')
for fn in [
# Kohya transformer (Cosmos rename)
test_lora_kohya_self_attn_q,
test_lora_kohya_cross_attn_k,
test_lora_kohya_self_attn_output,
test_lora_kohya_mlp_layer1,
test_lora_kohya_mlp_layer2,
test_lora_kohya_adaln_self_attn_1,
# Kohya text-encoder (newly supported)
test_lora_kohya_te_q,
test_lora_kohya_te_mlp_gate,
# BFL transformer
test_lora_bfl_self_attn_q,
test_lora_bfl_x_embedder,
test_lora_bfl_final_layer_linear,
# BFL llm_adapter (mirrors AnimaLLMAdapter's real module tree)
test_lora_bfl_llm_adapter_in_proj,
test_lora_bfl_llm_adapter_self_attn_q,
test_lora_bfl_llm_adapter_cross_attn_v,
test_lora_bfl_llm_adapter_mlp_in,
test_lora_bfl_llm_adapter_out_proj,
test_lora_bfl_llm_adapter_embed,
test_lora_llm_adapter_mlp_bias_delta,
# BFL text-encoder
test_lora_bfl_te_q,
# Combined transformer + TE
test_lora_combined_kohya_transformer_plus_te,
# DoRA threading
test_lora_dora_threading,
# LoHA (new family)
test_loha_kohya_cross_attn_q,
test_loha_kohya_self_attn_v,
test_loha_bfl_transformer,
# Cross-family rejection
test_loha_marker_rejection,
# Umbrella dispatcher
test_try_load_chain_routes_lora_and_loha,
]:
run_test(CAT_LOADER, fn)
log.warning('=== calc_updown shape sanity ===')
for fn in [
test_lora_calc_updown_transformer,
test_lora_calc_updown_cross_attn_asymmetric,
test_lora_calc_updown_llm_adapter,
test_lora_calc_updown_llm_adapter_embed,
test_lora_calc_updown_with_bias_delta,
test_lora_calc_updown_text_encoder,
test_loha_calc_updown_transformer,
# DoRA convention (apply_weight_decompose dual-path)
test_dora_per_input_convention,
test_dora_per_output_convention,
test_dora_square_weight_1d_defaults_to_per_input,
]:
run_test(CAT_MATH, fn)
log.warning('=== depth-expanded block remap ===')
for fn in [
test_expansion_map_matches_manifest,
test_expansion_map_covers_every_base_block_uniquely,
test_remap_binds_last_base_block_to_top_of_expanded,
test_remap_leaves_blocks_before_first_insertion,
test_no_remap_on_base_depth_model,
test_no_remap_when_lora_is_native_to_expanded_depth,
test_remap_leaves_adapter_and_te_indices_alone,
test_remap_handles_dotted_bfl_keys,
]:
run_test(CAT_REMAP, fn)
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)