#!/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)