Files
aigc-apps-VideoX-Fun/videox_fun/utils/lora_utils.py
T

1052 lines
44 KiB
Python
Executable File

# LoRA network module
# reference:
# https://github.com/microsoft/LoRA/blob/main/loralib/layers.py
# https://github.com/cloneofsimo/lora/blob/master/lora_diffusion/lora.py
# https://github.com/bmaltais/kohya_ss
import hashlib
import json
import math
import os
from collections import defaultdict
from dataclasses import dataclass
from io import BytesIO
from typing import Any, Dict, List, Mapping, Optional, Tuple, Type, Union
import safetensors.torch
import torch
import torch.utils.checkpoint
from diffusers.models.lora import LoRACompatibleConv, LoRACompatibleLinear
from safetensors.torch import load_file
from transformers import T5EncoderModel
from videox_fun.utils.group_offload import (_get_top_level_group_offload_hook,
_is_group_offload_enabled,
register_auto_device_hook,
safe_enable_group_offload,
safe_remove_group_offloading)
class LoRAModule(torch.nn.Module):
"""
replaces forward method of the original Linear, instead of replacing the original Linear module.
"""
def __init__(
self,
lora_name,
org_module: torch.nn.Module,
multiplier=1.0,
lora_dim=4,
alpha=1,
dropout=None,
rank_dropout=None,
module_dropout=None,
):
"""if alpha == 0 or None, alpha is rank (no scaling)."""
super().__init__()
self.lora_name = lora_name
if org_module.__class__.__name__ == "Conv2d":
in_dim = org_module.in_channels
out_dim = org_module.out_channels
elif org_module.__class__.__name__ == "Conv3d":
in_dim = org_module.in_channels
out_dim = org_module.out_channels
else:
in_dim = org_module.in_features
out_dim = org_module.out_features
self.lora_dim = lora_dim
if org_module.__class__.__name__ == "Conv2d":
kernel_size = org_module.kernel_size
stride = org_module.stride
padding = org_module.padding
self.lora_down = torch.nn.Conv2d(in_dim, self.lora_dim, kernel_size, stride, padding, bias=False)
self.lora_up = torch.nn.Conv2d(self.lora_dim, out_dim, (1, 1), (1, 1), bias=False)
elif org_module.__class__.__name__ == "Conv3d":
kernel_size = org_module.kernel_size
stride = org_module.stride
padding = org_module.padding
self.lora_down = torch.nn.Conv3d(in_dim, self.lora_dim, kernel_size, stride, padding, bias=False)
self.lora_up = torch.nn.Conv3d(self.lora_dim, out_dim, (1, 1, 1), (1, 1, 1), bias=False)
else:
self.lora_down = torch.nn.Linear(in_dim, self.lora_dim, bias=False)
self.lora_up = torch.nn.Linear(self.lora_dim, out_dim, bias=False)
if type(alpha) == torch.Tensor:
alpha = alpha.detach().float().numpy() # without casting, bf16 causes error
alpha = self.lora_dim if alpha is None or alpha == 0 else alpha
self.scale = alpha / self.lora_dim
self.register_buffer("alpha", torch.tensor(alpha))
# same as microsoft's
torch.nn.init.kaiming_uniform_(self.lora_down.weight, a=math.sqrt(5))
torch.nn.init.zeros_(self.lora_up.weight)
self.multiplier = multiplier
self.org_module = org_module # remove in applying
self.dropout = dropout
self.rank_dropout = rank_dropout
self.module_dropout = module_dropout
def apply_to(self):
self.org_forward = self.org_module.forward
self.org_module.forward = self.forward
del self.org_module
def forward(self, x, *args, **kwargs):
weight_dtype = x.dtype
org_forwarded = self.org_forward(x)
# module dropout
if self.module_dropout is not None and self.training:
if torch.rand(1) < self.module_dropout:
return org_forwarded
lx = self.lora_down(x.to(self.lora_down.weight.dtype))
# normal dropout
if self.dropout is not None and self.training:
lx = torch.nn.functional.dropout(lx, p=self.dropout)
# rank dropout
if self.rank_dropout is not None and self.training:
mask = torch.rand((lx.size(0), self.lora_dim), device=lx.device) > self.rank_dropout
if len(lx.size()) == 3:
mask = mask.unsqueeze(1) # for Text Encoder
elif len(lx.size()) == 4:
mask = mask.unsqueeze(-1).unsqueeze(-1) # for Conv2d
elif len(lx.size()) == 5:
mask = mask.unsqueeze(-1).unsqueeze(-1).unsqueeze(-1) # for Conv3d
lx = lx * mask
# scaling for rank dropout: treat as if the rank is changed
scale = self.scale * (1.0 / (1.0 - self.rank_dropout)) # redundant for readability
else:
scale = self.scale
lx = self.lora_up(lx)
# Fused `out = org + alpha * lx`: the unfused `org + lx * multiplier * scale` materializes one extra
# full-size transient (`lx * multiplier * scale`) on top of the result. When dtypes already match we go one
# step further and add in place (mirroring peft's `result += ...`), keeping only two full-size tensors
# alive per LoRA'ed layer — the wide feed-forward projections are the activation peak of LoRA training.
# `org_forwarded` is the base layer's freshly produced output with no other consumer, so the in-place add
# is autograd-safe. The `.to(weight_dtype)` fallback keeps mixed-dtype setups (fp32 LoRA weights over bf16
# activations) correct.
alpha = self.multiplier * scale
if org_forwarded.dtype != weight_dtype or lx.dtype != weight_dtype:
return org_forwarded.to(weight_dtype) + lx.to(weight_dtype) * alpha
return org_forwarded.add_(lx, alpha=alpha)
def addnet_hash_legacy(b):
"""Old model hash used by sd-webui-additional-networks for .safetensors format files"""
m = hashlib.sha256()
b.seek(0x100000)
m.update(b.read(0x10000))
return m.hexdigest()[0:8]
def addnet_hash_safetensors(b):
"""New model hash used by sd-webui-additional-networks for .safetensors format files"""
hash_sha256 = hashlib.sha256()
blksize = 1024 * 1024
b.seek(0)
header = b.read(8)
n = int.from_bytes(header, "little")
offset = n + 8
b.seek(offset)
for chunk in iter(lambda: b.read(blksize), b""):
hash_sha256.update(chunk)
return hash_sha256.hexdigest()
def precalculate_safetensors_hashes(tensors, metadata):
"""Precalculate the model hashes needed by sd-webui-additional-networks to
save time on indexing the model later."""
# Because writing user metadata to the file can change the result of
# sd_models.model_hash(), only retain the training metadata for purposes of
# calculating the hash, as they are meant to be immutable
metadata = {k: v for k, v in metadata.items() if k.startswith("ss_")}
bytes = safetensors.torch.save(tensors, metadata)
b = BytesIO(bytes)
model_hash = addnet_hash_safetensors(b)
legacy_hash = addnet_hash_legacy(b)
return model_hash, legacy_hash
class LoRANetwork(torch.nn.Module):
TRANSFORMER_TARGET_REPLACE_MODULE = [
"CogVideoXTransformer3DModel", "WanTransformer3DModel", \
"Wan2_2Transformer3DModel", "FluxTransformer2DModel", "QwenImageTransformer2DModel", \
"Wan2_2Transformer3DModel_Animate", "Wan2_2Transformer3DModel_S2V", "FantasyTalkingTransformer3DModel", \
"HunyuanVideoTransformer3DModel", "Flux2Transformer2DModel", "ZImageTransformer2DModel", \
"LongCatVideoTransformer3DModel", "LongCatVideoAvatarTransformer3DModel", "TurboWanTransformer3DModel", \
"LTX2VideoTransformer3DModel", "InfiniteTalkTransformer3DModel", "WanAudioTransformer3DModel", \
"MOVADualTowerConditionalBridge", "FlashHeadTransformer3DModel", "LensTransformer2DModel", \
"MiniMaxH3Transformer3DModel"
]
TEXT_ENCODER_TARGET_REPLACE_MODULE = ["T5LayerSelfAttention", "T5LayerFF", "BertEncoder", "T5SelfAttention", "T5CrossAttention"]
LORA_PREFIX_TRANSFORMER = "lora_unet"
LORA_PREFIX_TEXT_ENCODER = "lora_te"
def __init__(
self,
text_encoder: Union[List[T5EncoderModel], T5EncoderModel],
unet,
multiplier: float = 1.0,
lora_dim: int = 4,
alpha: float = 1,
dropout: Optional[float] = None,
module_class: Type[object] = LoRAModule,
skip_name: str = None,
target_name: str = None,
varbose: Optional[bool] = False,
) -> None:
super().__init__()
self.multiplier = multiplier
self.lora_dim = lora_dim
self.alpha = alpha
self.dropout = dropout
print(f"create LoRA network. base dim (rank): {lora_dim}, alpha: {alpha}")
print(f"neuron dropout: p={self.dropout}")
# create module instances
def create_modules(
is_unet: bool,
root_module: torch.nn.Module,
target_replace_modules: List[torch.nn.Module],
) -> List[LoRAModule]:
prefix = (
self.LORA_PREFIX_TRANSFORMER
if is_unet
else self.LORA_PREFIX_TEXT_ENCODER
)
loras = []
skipped = []
for name, module in root_module.named_modules():
if module.__class__.__name__ in target_replace_modules:
for child_name, child_module in module.named_modules():
is_linear = child_module.__class__.__name__ == "Linear" or child_module.__class__.__name__ == "LoRACompatibleLinear"
is_conv2d = child_module.__class__.__name__ == "Conv2d" or child_module.__class__.__name__ == "LoRACompatibleConv"
is_conv2d_1x1 = is_conv2d and child_module.kernel_size == (1, 1)
is_conv3d = child_module.__class__.__name__ == "Conv3d"
is_conv3d_1x1x1 = is_conv3d and child_module.kernel_size == (1, 1, 1)
skip_names = skip_name.split(',') if skip_name is not None else []
target_names = target_name.split(',') if target_name is not None else []
skip_names = [name.strip() for name in skip_names if name.strip()]
target_names = [name.strip() for name in target_names if name.strip()]
if skip_names and any(skip_n in child_name for skip_n in skip_names):
continue
if target_names and not any(target_n in child_name for target_n in target_names):
continue
if is_linear or is_conv2d or is_conv3d:
lora_name = prefix + "." + name + "." + child_name
lora_name = lora_name.replace(".", "_")
dim = None
alpha = None
if is_linear or is_conv2d_1x1 or is_conv3d:
dim = self.lora_dim
alpha = self.alpha
if dim is None or dim == 0:
if is_linear or is_conv2d_1x1 or is_conv3d:
skipped.append(lora_name)
continue
lora = module_class(
lora_name,
child_module,
self.multiplier,
dim,
alpha,
dropout=dropout,
)
loras.append(lora)
return loras, skipped
text_encoders = text_encoder if type(text_encoder) == list else [text_encoder]
self.text_encoder_loras = []
skipped_te = []
for i, text_encoder in enumerate(text_encoders):
if text_encoder is not None:
text_encoder_loras, skipped = create_modules(False, text_encoder, LoRANetwork.TEXT_ENCODER_TARGET_REPLACE_MODULE)
self.text_encoder_loras.extend(text_encoder_loras)
skipped_te += skipped
print(f"create LoRA for Text Encoder: {len(self.text_encoder_loras)} modules.")
self.unet_loras, skipped_un = create_modules(True, unet, LoRANetwork.TRANSFORMER_TARGET_REPLACE_MODULE)
print(f"create LoRA for U-Net: {len(self.unet_loras)} modules.")
# assertion
names = set()
for lora in self.text_encoder_loras + self.unet_loras:
assert lora.lora_name not in names, f"duplicated lora name: {lora.lora_name}"
names.add(lora.lora_name)
def apply_to(self, text_encoder, unet, apply_text_encoder=True, apply_unet=True):
if apply_text_encoder:
print("enable LoRA for text encoder")
else:
self.text_encoder_loras = []
if apply_unet:
print("enable LoRA for U-Net")
else:
self.unet_loras = []
for lora in self.text_encoder_loras + self.unet_loras:
lora.apply_to()
self.add_module(lora.lora_name, lora)
def set_multiplier(self, multiplier):
self.multiplier = multiplier
for lora in self.text_encoder_loras + self.unet_loras:
lora.multiplier = self.multiplier
def load_weights(self, file):
if os.path.splitext(file)[1] == ".safetensors":
from safetensors.torch import load_file
weights_sd = load_file(file)
else:
weights_sd = torch.load(file, map_location="cpu")
info = self.load_state_dict(weights_sd, False)
return info
def prepare_optimizer_params(self, text_encoder_lr, unet_lr, default_lr):
self.requires_grad_(True)
all_params = []
def enumerate_params(loras):
params = []
for lora in loras:
params.extend(lora.parameters())
return params
if self.text_encoder_loras:
param_data = {"params": enumerate_params(self.text_encoder_loras)}
if text_encoder_lr is not None:
param_data["lr"] = text_encoder_lr
all_params.append(param_data)
if self.unet_loras:
param_data = {"params": enumerate_params(self.unet_loras)}
if unet_lr is not None:
param_data["lr"] = unet_lr
all_params.append(param_data)
return all_params
def enable_gradient_checkpointing(self):
pass
def get_trainable_params(self):
return self.parameters()
def save_weights(self, file, dtype, metadata):
if metadata is not None and len(metadata) == 0:
metadata = None
state_dict = self.state_dict()
if dtype is not None:
for key in list(state_dict.keys()):
v = state_dict[key]
v = v.detach().clone().to("cpu").to(dtype)
state_dict[key] = v
if os.path.splitext(file)[1] == ".safetensors":
from safetensors.torch import save_file
# Precalculate model hashes to save time on indexing
if metadata is None:
metadata = {}
model_hash, legacy_hash = precalculate_safetensors_hashes(state_dict, metadata)
metadata["sshs_model_hash"] = model_hash
metadata["sshs_legacy_hash"] = legacy_hash
save_file(state_dict, file, metadata)
else:
torch.save(state_dict, file)
def create_network(
multiplier: float,
network_dim: Optional[int],
network_alpha: Optional[float],
text_encoder: Union[T5EncoderModel, List[T5EncoderModel]],
transformer,
neuron_dropout: Optional[float] = None,
skip_name: str = None,
target_name: str = None,
**kwargs,
):
if network_dim is None:
network_dim = 4 # default
if network_alpha is None:
network_alpha = 1.0
network = LoRANetwork(
text_encoder,
transformer,
multiplier=multiplier,
lora_dim=network_dim,
alpha=network_alpha,
dropout=neuron_dropout,
skip_name=skip_name,
target_name=target_name,
varbose=True,
)
return network
def convert_peft_lora_to_kohya_lora(state_dict):
new_state_dict = {}
for key, value in state_dict.items():
if "diffusion_model." in key:
key = key.replace("diffusion_model.", "")
if "lora_unet__" not in key:
key = "lora_unet__" + key
key = key.replace(".lora_A.default.", ".lora_down.")
key = key.replace(".lora_B.default.", ".lora_up.")
key = key.replace(".lora_A.", ".lora_down.")
key = key.replace(".lora_B.", ".lora_up.")
key = key.replace(".", "_")
if key.endswith("_lora_up_weight"):
key = key[:-15] + ".lora_up.weight"
if key.endswith("_lora_down_weight"):
key = key[:-17] + ".lora_down.weight"
new_state_dict[key] = value
return new_state_dict
# ---------------------------------------------------------------------------
# Official TaoMate-H3 adapter support
# ---------------------------------------------------------------------------
_TAOMATE_H3_LORA_TARGET_SUFFIXES = (
"attn.qkv_proj",
"attn.out_proj",
"mlp.fc1",
"mlp.fc2",
)
_TAOMATE_H3_REFINER_BLOCKS = 2
_TAOMATE_H3_MAIN_BLOCKS = 50
class TaomateH3LoRACheckpointError(RuntimeError):
"""The adapter directory does not match the released TaoMate-H3 contract."""
def canonical_taomate_h3_lora_targets() -> Tuple[str, ...]:
"""Return the adapter's exact 52-block x 4-projection inventory."""
blocks = [f"token_refiner.blocks.{index}" for index in range(_TAOMATE_H3_REFINER_BLOCKS)]
blocks.extend(f"blocks.{index}" for index in range(_TAOMATE_H3_MAIN_BLOCKS))
return tuple(
f"{block}.{suffix}" for block in blocks for suffix in _TAOMATE_H3_LORA_TARGET_SUFFIXES
)
@dataclass(frozen=True)
class LoadedTaomateH3LoRA:
"""The validated official adapter, in its own `{module}.lora_a/lora_b` layout."""
rank: int
alpha: float
state: Mapping[str, torch.Tensor]
loaded_tensor_count: int
@property
def scale(self) -> float:
return self.alpha / self.rank
def load_taomate_h3_adapter(adapter_dir: str) -> LoadedTaomateH3LoRA:
r"""
Load and validate the official TaoMate-H3 LoRA adapter from a directory.
The directory must hold `config.json` (or `adapter_config.json`) with the
integer `rank` and the positive `alpha`, and `adapter_model.safetensors`
whose key set is exactly the canonical 52-block x 4-projection inventory
with `.lora_a` / `.lora_b` leaves. Every tensor must be float32 with
`lora_a` shaped `(rank, in_features)` and `lora_b` shaped
`(out_features, rank)`.
Args:
adapter_dir (`str`): The adapter directory.
Returns:
[`LoadedTaomateH3LoRA`]
"""
from safetensors import safe_open
root = os.path.expanduser(str(adapter_dir))
if not os.path.isdir(root):
raise TaomateH3LoRACheckpointError(f"adapter directory does not exist: {root}")
config_path = os.path.join(root, "config.json")
if not os.path.isfile(config_path):
config_path = os.path.join(root, "adapter_config.json")
try:
with open(config_path, "r", encoding="utf-8") as handle:
config = json.load(handle)
except (OSError, UnicodeError, json.JSONDecodeError) as error:
raise TaomateH3LoRACheckpointError(f"cannot read adapter config: {error}") from error
if not isinstance(config, Mapping):
raise TaomateH3LoRACheckpointError("adapter config must be a JSON object")
rank = config.get("rank")
alpha = config.get("alpha")
if (
isinstance(rank, bool)
or not isinstance(rank, int)
or rank <= 0
or isinstance(alpha, bool)
or not isinstance(alpha, (int, float))
or not math.isfinite(float(alpha))
or float(alpha) <= 0
):
raise TaomateH3LoRACheckpointError(f"adapter rank/alpha is invalid: {config}")
weights_path = os.path.join(root, "adapter_model.safetensors")
if not os.path.isfile(weights_path):
raise TaomateH3LoRACheckpointError(f"adapter weights are absent: {weights_path}")
targets = canonical_taomate_h3_lora_targets()
leaves = ("lora_a", "lora_b")
expected = {f"{target}.{leaf}" for target in targets for leaf in leaves}
selected: Dict[str, torch.Tensor] = {}
with safe_open(weights_path, framework="pt", device="cpu") as handle:
present = set(handle.keys())
if present != expected:
raise TaomateH3LoRACheckpointError(
"adapter tensor inventory differs: "
f"missing={len(expected - present)}, unexpected={len(present - expected)}"
)
for name in sorted(expected):
tensor = handle.get_tensor(name)
leaf = name.rsplit(".", 1)[-1]
if (
tensor.ndim != 2
or tensor.dtype != torch.float32
or (leaf == "lora_a" and int(tensor.shape[0]) != rank)
or (leaf == "lora_b" and int(tensor.shape[1]) != rank)
):
raise TaomateH3LoRACheckpointError(f"adapter tensor shape/dtype differs: {name}")
selected[name] = tensor.contiguous()
return LoadedTaomateH3LoRA(
rank=int(rank),
alpha=float(alpha),
state=selected,
loaded_tensor_count=len(selected),
)
def _taomate_h3_diffusers_block_path(official_block: str) -> str:
"""Map an official block stem onto this repository's module tree."""
if official_block.startswith("token_refiner.blocks."):
index = official_block[len("token_refiner.blocks.") :]
return f"token_refiner.refiner_blocks.{index}"
if official_block.startswith("blocks."):
index = official_block[len("blocks.") :]
return f"transformer_blocks.{index}"
raise TaomateH3LoRACheckpointError(f"unknown adapter block stem: {official_block}")
def official_to_kohya_lora_state_dict(loaded: LoadedTaomateH3LoRA) -> Dict[str, torch.Tensor]:
r"""
Convert the official adapter onto the kohya layout this module consumes.
Keys follow `LoRANetwork`'s naming: `lora_unet` + the diffusers module
path with dots replaced by underscores, e.g.
`lora_unet_transformer_blocks_0_attn_to_q.lora_down.weight`. The effective
update of every entry is `(alpha / rank) * up @ down(x)`, which is the
official `(alpha / rank) * lora_b @ lora_a`.
Every fused projection is split from its own tensor shape — the qkv
`lora_b` into three contiguous thirds, the gated fc1 `lora_b` into two
halves — so no per-block dimension table is involved. The released
adapter ships one geometry for all 52 blocks, the 50 main and the 2
token-refiner blocks alike: hidden size 5376, attention inner dim 7168
and ffn 14336, hence `lora_b` row counts of 3 * 7168 (qkv), 5376
(out_proj / fc2) and 2 * 14336 (fc1).
Args:
loaded ([`LoadedTaomateH3LoRA`]): The validated official adapter.
Returns:
`dict[str, torch.Tensor]`: the kohya state dict, float32 on CPU.
"""
alpha = loaded.alpha
state_dict: Dict[str, torch.Tensor] = {}
def emit(diffusers_path: str, lora_a: torch.Tensor, lora_b: torch.Tensor) -> None:
stem = "lora_unet_" + diffusers_path.replace(".", "_")
state_dict[f"{stem}.lora_down.weight"] = lora_a.contiguous()
state_dict[f"{stem}.lora_up.weight"] = lora_b.contiguous()
state_dict[f"{stem}.alpha"] = torch.tensor(alpha, dtype=torch.float32)
for target in canonical_taomate_h3_lora_targets():
# Targets end with one of the known dotted projections; split on that,
# not on the last dot ("attn.qkv_proj" itself contains a dot).
suffix = next(
candidate
for candidate in _TAOMATE_H3_LORA_TARGET_SUFFIXES
if target.endswith("." + candidate)
)
block_stem = target[: -(len(suffix) + 1)]
lora_a = loaded.state[f"{target}.lora_a"]
lora_b = loaded.state[f"{target}.lora_b"]
path = _taomate_h3_diffusers_block_path(block_stem)
if suffix == "attn.qkv_proj":
# Contiguous thirds, exactly like the checkpoint's fused qkv weight
# (`minimax_h3_conversion.split_fused_qkv`). The three projections
# share `lora_a` and the thirds are views of one storage, so clone
# every entry to keep the kohya tensors standalone (safetensors
# refuses overlapping storages).
fused_rows = int(lora_b.shape[0])
if fused_rows % 3 != 0:
raise TaomateH3LoRACheckpointError(
f"fused qkv lora_b rows do not split into three equal parts: {target}")
inner_dim = fused_rows // 3
query_b, key_b, value_b = lora_b.split(inner_dim, dim=0)
emit(f"{path}.attn.to_q", lora_a.clone(), query_b.clone())
emit(f"{path}.attn.to_k", lora_a.clone(), key_b.clone())
emit(f"{path}.attn.to_v", lora_a.clone(), value_b.clone())
elif suffix == "attn.out_proj":
# `out_proj` maps the attention inner dim back to the hidden size and
# diffusers' `to_out.0` is the matching projection, so both leaves
# carry over as they are.
emit(f"{path}.attn.to_out.0", lora_a, lora_b)
elif suffix == "mlp.fc1":
if int(lora_b.shape[0]) % 2 != 0:
raise TaomateH3LoRACheckpointError(
f"fused fc1 lora_b rows do not split into two equal halves: {target}")
# The reference fuses `[gate; value]`; diffusers' `SwiGLU` reads
# `[value; gate]` — swap the halves, as `minimax_h3_conversion` does.
gate_b, value_b = lora_b.chunk(2, dim=0)
emit(f"{path}.ff.net.0.proj", lora_a, torch.cat([value_b, gate_b], dim=0))
elif suffix == "mlp.fc2":
# `fc2` maps the ffn width back to the hidden size and diffusers'
# `ff.net.2` is the same projection, so both leaves carry over.
emit(f"{path}.ff.net.2", lora_a, lora_b)
else: # pragma: no cover - the suffix list above is closed
raise TaomateH3LoRACheckpointError(f"unknown adapter projection: {suffix}")
return state_dict
def convert_taomate_h3_adapter(
adapter_dir: str,
output_path: str,
) -> Dict[str, Any]:
r"""
Convert the official adapter directory into a kohya safetensors file.
The produced file plugs straight into `merge_lora` / `unmerge_lora`, so a
converted adapter behaves like any other LoRA checkpoint of this
repository.
Args:
adapter_dir (`str`): The official adapter directory.
output_path (`str`): Destination `.safetensors` path.
Returns:
`dict`: a small receipt with the rank/alpha and tensor count.
"""
from safetensors.torch import save_file
loaded = load_taomate_h3_adapter(adapter_dir)
state_dict = official_to_kohya_lora_state_dict(loaded)
output_path = os.path.expanduser(str(output_path))
parent = os.path.dirname(output_path)
if parent:
os.makedirs(parent, exist_ok=True)
save_file(state_dict, output_path)
return {
"rank": loaded.rank,
"alpha": loaded.alpha,
"scale": loaded.scale,
"tensor_count": len(state_dict),
"output_path": output_path,
}
def merge_lora(pipeline, lora_path, multiplier, device='cpu', dtype=torch.float32, state_dict=None, transformer_only=False, sub_transformer_name="transformer"):
if lora_path is None:
return pipeline
print(f"[LoRA Merge] Starting to merge LoRA from: {lora_path if lora_path else 'state_dict'}")
print(f"[LoRA Merge] Multiplier: {multiplier}, Device: {device}, Dtype: {dtype}")
LORA_PREFIX_TRANSFORMER = "lora_unet"
LORA_PREFIX_TEXT_ENCODER = "lora_te"
if state_dict is None:
if os.path.isdir(lora_path):
# An official TaoMate-H3 adapter directory: validated and converted to the kohya layout
# in memory, then merged like any other checkpoint.
print(f"[LoRA Merge] Detected an official TaoMate-H3 adapter directory, converting...")
loaded = load_taomate_h3_adapter(lora_path)
state_dict = official_to_kohya_lora_state_dict(loaded)
elif lora_path.endswith("safetensors"):
print(f"[LoRA Merge] Loading safetensors file...")
state_dict = load_file(lora_path)
else:
print(f"[LoRA Merge] Loading pytorch file...")
state_dict = torch.load(lora_path, map_location="cpu")
else:
print(f"[LoRA Merge] Using provided state_dict")
state_dict = state_dict
updates = defaultdict(dict)
for key, value in state_dict.items():
if "diffusion_model." in key:
key = key.replace("diffusion_model.", "")
if "lora_unet__" not in key:
key = "lora_unet__" + key
key = key.replace(".", "_")
if key.endswith("_lora_up_weight"):
key = key[:-15] + ".lora_up.weight"
if key.endswith("_lora_down_weight"):
key = key[:-17] + ".lora_down.weight"
if key.endswith("_lora_A_default_weight"):
key = key[:-22] + ".lora_A.weight"
if key.endswith("_lora_B_default_weight"):
key = key[:-22] + ".lora_B.weight"
if key.endswith("_lora_A_weight"):
key = key[:-14] + ".lora_A.weight"
if key.endswith("_lora_B_weight"):
key = key[:-14] + ".lora_B.weight"
if key.endswith("_alpha"):
key = key[:-6] + ".alpha"
key = key.replace(".lora_A.default.", ".lora_down.")
key = key.replace(".lora_B.default.", ".lora_up.")
key = key.replace(".lora_A.", ".lora_down.")
key = key.replace(".lora_B.", ".lora_up.")
layer, elem = key.split('.', 1)
updates[layer][elem] = value
print(f"[LoRA Merge] Organized into {len(updates)} layers")
sequential_cpu_offload_flag = False
if pipeline.transformer.device == torch.device(type="meta"):
print(f"[LoRA Merge] Removing hooks for meta device...")
pipeline.remove_all_hooks()
sequential_cpu_offload_flag = True
offload_device = pipeline._offload_device
merged_count = 0
skipped_count = 0
error_count = 0
for layer, elems in updates.items():
if "lora_te" in layer:
if transformer_only:
skipped_count += 1
continue
else:
layer_infos = layer.split(LORA_PREFIX_TEXT_ENCODER + "_")[-1].split("_")
curr_layer = pipeline.text_encoder
else:
layer_infos = layer.split(LORA_PREFIX_TRANSFORMER + "_")[-1].split("_")
curr_layer = getattr(pipeline, sub_transformer_name)
try:
curr_layer = curr_layer.__getattr__("_".join(layer_infos[1:]))
except Exception:
temp_name = layer_infos.pop(0)
try:
while len(layer_infos) > -1:
try:
curr_layer = curr_layer.__getattr__(temp_name + "_" + "_".join(layer_infos))
break
except Exception:
try:
curr_layer = curr_layer.__getattr__(temp_name)
if len(layer_infos) > 0:
temp_name = layer_infos.pop(0)
elif len(layer_infos) == 0:
break
except Exception:
if len(layer_infos) == 0:
print(f'[LoRA Merge] Warning: Error loading layer in front search: {layer}. Try it in back search.')
if len(temp_name) > 0:
temp_name += "_" + layer_infos.pop(0)
else:
temp_name = layer_infos.pop(0)
except Exception:
if "lora_te" in layer:
if transformer_only:
skipped_count += 1
continue
else:
layer_infos = layer.split(LORA_PREFIX_TEXT_ENCODER + "_")[-1].split("_")
curr_layer = pipeline.text_encoder
else:
layer_infos = layer.split(LORA_PREFIX_TRANSFORMER + "_")[-1].split("_")
curr_layer = getattr(pipeline, sub_transformer_name)
len_layer_infos = len(layer_infos)
start_index = 0 if len_layer_infos >= 1 and len(layer_infos[0]) > 0 else 1
end_indx = len_layer_infos
error_flag = False if len_layer_infos >= 1 else True
while start_index < len_layer_infos:
try:
if start_index >= end_indx:
print(f'[LoRA Merge] Error: Failed to load layer in back search: {layer}')
error_flag = True
break
curr_layer = curr_layer.__getattr__("_".join(layer_infos[start_index:end_indx]))
start_index = end_indx
end_indx = len_layer_infos
except Exception:
end_indx -= 1
if error_flag:
error_count += 1
continue
origin_dtype = curr_layer.weight.data.dtype
origin_device = curr_layer.weight.data.device
curr_layer = curr_layer.to(device, dtype)
weight_up = elems['lora_up.weight'].to(device, dtype)
weight_down = elems['lora_down.weight'].to(device, dtype)
if 'alpha' in elems.keys():
alpha = elems['alpha'].item() / weight_up.shape[1]
else:
alpha = 1.0
if len(weight_up.shape) == 4:
curr_layer.weight.data += multiplier * alpha * torch.mm(
weight_up.squeeze(3).squeeze(2), weight_down.squeeze(3).squeeze(2)
).unsqueeze(2).unsqueeze(3)
elif len(weight_up.shape) == 5:
# lora_up: (out_dim, lora_dim, 1, 1, 1) → (out_dim, lora_dim)
# lora_down: (lora_dim, in_dim, kD, kH, kW) → (lora_dim, in_dim*kD*kH*kW)
up_2d = weight_up.reshape(weight_up.shape[0], weight_up.shape[1])
down_2d = weight_down.reshape(weight_down.shape[0], -1)
curr_layer.weight.data += multiplier * alpha * torch.mm(up_2d, down_2d).reshape(curr_layer.weight.data.shape)
else:
curr_layer.weight.data += multiplier * alpha * torch.mm(weight_up, weight_down)
curr_layer = curr_layer.to(origin_device, origin_dtype)
merged_count += 1
print(f"[LoRA Merge] Completed: {merged_count} layers merged, {skipped_count} layers skipped, {error_count} errors")
if sequential_cpu_offload_flag:
print(f"[LoRA Merge] Re-enabling sequential CPU offload...")
pipeline.enable_sequential_cpu_offload(device=offload_device)
else:
# When group offload is active, remove and re-apply it on the whole pipeline
# so that all ModuleGroup.cpu_param_dict references are rebuilt from the
# freshly-merged weights. A simple _maybe_remove_and_reapply is not enough
# because .to() during merge creates new tensor objects that break cached refs.
try:
local_transformer = getattr(pipeline, sub_transformer_name)
if _is_group_offload_enabled(local_transformer):
print(f"[LoRA Merge] Removing group offload hooks from pipeline...")
safe_remove_group_offloading(pipeline)
print(f"[LoRA Merge] Re-applying group offload hooks to pipeline...")
register_auto_device_hook(getattr(pipeline, sub_transformer_name))
safe_enable_group_offload(
pipeline,
onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True
)
print(pipeline._execution_device)
except Exception as e:
print(f"[LoRA Merge] Warning: Failed to refresh group offload: {e}")
print(f"[LoRA Merge] ✓ LoRA merge finished successfully")
return pipeline
# TODO: Refactor with merge_lora.
def unmerge_lora(pipeline, lora_path, multiplier=1, device="cpu", dtype=torch.float32, sub_transformer_name="transformer"):
if lora_path is None:
return pipeline
print(f"[LoRA Unmerge] Starting to unmerge LoRA from: {lora_path}")
print(f"[LoRA Unmerge] Multiplier: {multiplier}, Device: {device}, Dtype: {dtype}")
"""Unmerge state_dict in LoRANetwork from the pipeline in diffusers."""
LORA_PREFIX_UNET = "lora_unet"
LORA_PREFIX_TEXT_ENCODER = "lora_te"
if os.path.isdir(lora_path):
# Same detection as `merge_lora`: an official TaoMate-H3 adapter directory is converted to
# the kohya layout in memory first.
print(f"[LoRA Unmerge] Detected an official TaoMate-H3 adapter directory, converting...")
loaded = load_taomate_h3_adapter(lora_path)
state_dict = official_to_kohya_lora_state_dict(loaded)
elif lora_path.endswith("safetensors"):
print(f"[LoRA Unmerge] Loading safetensors file...")
state_dict = load_file(lora_path)
else:
print(f"[LoRA Unmerge] Loading pytorch file...")
state_dict = torch.load(lora_path, map_location="cpu")
updates = defaultdict(dict)
for key, value in state_dict.items():
if "diffusion_model." in key:
key = key.replace("diffusion_model.", "")
if "lora_unet__" not in key:
key = "lora_unet__" + key
key = key.replace(".", "_")
if key.endswith("_lora_up_weight"):
key = key[:-15] + ".lora_up.weight"
if key.endswith("_lora_down_weight"):
key = key[:-17] + ".lora_down.weight"
if key.endswith("_lora_A_default_weight"):
key = key[:-22] + ".lora_A.weight"
if key.endswith("_lora_B_default_weight"):
key = key[:-22] + ".lora_B.weight"
if key.endswith("_lora_A_weight"):
key = key[:-14] + ".lora_A.weight"
if key.endswith("_lora_B_weight"):
key = key[:-14] + ".lora_B.weight"
if key.endswith("_alpha"):
key = key[:-6] + ".alpha"
key = key.replace(".lora_A.default.", ".lora_down.")
key = key.replace(".lora_B.default.", ".lora_up.")
key = key.replace(".lora_A.", ".lora_down.")
key = key.replace(".lora_B.", ".lora_up.")
layer, elem = key.split('.', 1)
updates[layer][elem] = value
print(f"[LoRA Unmerge] Organized into {len(updates)} layers")
sequential_cpu_offload_flag = False
if pipeline.transformer.device == torch.device(type="meta"):
print(f"[LoRA Unmerge] Removing hooks for meta device...")
pipeline.remove_all_hooks()
sequential_cpu_offload_flag = True
unmerged_count = 0
error_count = 0
for layer, elems in updates.items():
if "lora_te" in layer:
layer_infos = layer.split(LORA_PREFIX_TEXT_ENCODER + "_")[-1].split("_")
curr_layer = pipeline.text_encoder
else:
layer_infos = layer.split(LORA_PREFIX_UNET + "_")[-1].split("_")
curr_layer = getattr(pipeline, sub_transformer_name)
try:
curr_layer = curr_layer.__getattr__("_".join(layer_infos[1:]))
except Exception:
temp_name = layer_infos.pop(0)
try:
while len(layer_infos) > -1:
try:
curr_layer = curr_layer.__getattr__(temp_name + "_" + "_".join(layer_infos))
break
except Exception:
try:
curr_layer = curr_layer.__getattr__(temp_name)
if len(layer_infos) > 0:
temp_name = layer_infos.pop(0)
elif len(layer_infos) == 0:
break
except Exception:
if len(layer_infos) == 0:
print(f'[LoRA Unmerge] Warning: Error loading layer in front search: {layer}. Try it in back search.')
if len(temp_name) > 0:
temp_name += "_" + layer_infos.pop(0)
else:
temp_name = layer_infos.pop(0)
except Exception:
if "lora_te" in layer:
layer_infos = layer.split(LORA_PREFIX_TEXT_ENCODER + "_")[-1].split("_")
curr_layer = pipeline.text_encoder
else:
layer_infos = layer.split(LORA_PREFIX_UNET + "_")[-1].split("_")
curr_layer = getattr(pipeline, sub_transformer_name)
len_layer_infos = len(layer_infos)
start_index = 0 if len_layer_infos >= 1 and len(layer_infos[0]) > 0 else 1
end_indx = len_layer_infos
error_flag = False if len_layer_infos >= 1 else True
while start_index < len_layer_infos:
try:
if start_index >= end_indx:
print(f'[LoRA Unmerge] Error: Failed to load layer in back search: {layer}')
error_flag = True
break
curr_layer = curr_layer.__getattr__("_".join(layer_infos[start_index:end_indx]))
start_index = end_indx
end_indx = len_layer_infos
except Exception:
end_indx -= 1
if error_flag:
error_count += 1
continue
origin_dtype = curr_layer.weight.data.dtype
origin_device = curr_layer.weight.data.device
curr_layer = curr_layer.to(device, dtype)
weight_up = elems['lora_up.weight'].to(device, dtype)
weight_down = elems['lora_down.weight'].to(device, dtype)
if 'alpha' in elems.keys():
alpha = elems['alpha'].item() / weight_up.shape[1]
else:
alpha = 1.0
if len(weight_up.shape) == 4:
curr_layer.weight.data -= multiplier * alpha * torch.mm(
weight_up.squeeze(3).squeeze(2), weight_down.squeeze(3).squeeze(2)
).unsqueeze(2).unsqueeze(3)
elif len(weight_up.shape) == 5:
# lora_up: (out_dim, lora_dim, 1, 1, 1) → (out_dim, lora_dim)
# lora_down: (lora_dim, in_dim, kD, kH, kW) → (lora_dim, in_dim*kD*kH*kW)
up_2d = weight_up.reshape(weight_up.shape[0], weight_up.shape[1])
down_2d = weight_down.reshape(weight_down.shape[0], -1)
curr_layer.weight.data -= multiplier * alpha * torch.mm(up_2d, down_2d).reshape(curr_layer.weight.data.shape)
else:
curr_layer.weight.data -= multiplier * alpha * torch.mm(weight_up, weight_down)
curr_layer = curr_layer.to(origin_device, origin_dtype)
unmerged_count += 1
print(f"[LoRA Unmerge] Completed: {unmerged_count} layers unmerged, {error_count} errors")
if sequential_cpu_offload_flag:
print(f"[LoRA Unmerge] Re-enabling sequential CPU offload...")
pipeline.enable_sequential_cpu_offload(device=device)
else:
# Same as merge_lora: remove and re-apply group offload on the whole pipeline
# to rebuild cpu_param_dict with the post-unmerge weights.
try:
local_transformer = getattr(pipeline, sub_transformer_name)
if _is_group_offload_enabled(local_transformer):
print(f"[LoRA Unmerge] Removing group offload hooks from pipeline...")
safe_remove_group_offloading(pipeline)
print(f"[LoRA Unmerge] Re-applying group offload hooks to pipeline...")
register_auto_device_hook(getattr(pipeline, sub_transformer_name))
safe_enable_group_offload(
pipeline,
onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True
)
except Exception as e:
print(f"[LoRA Unmerge] Warning: Failed to refresh group offload: {e}")
print(f"[LoRA Unmerge] ✓ LoRA unmerge finished successfully")
return pipeline