1052 lines
44 KiB
Python
Executable File
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
|