435 lines
13 KiB
Python
435 lines
13 KiB
Python
import os
|
|
from logging import warnings
|
|
import torch
|
|
from typing import Union
|
|
from types import SimpleNamespace
|
|
from ...animatediff.models.unet import UNet3DConditionModel
|
|
from transformers import CLIPTextModel
|
|
from ...animatediff.utils.convert_diffusers_to_original_ms_text_to_video import convert_unet_state_dict, convert_text_enc_state_dict_v20
|
|
|
|
from .lora import (
|
|
extract_lora_ups_down,
|
|
inject_trainable_lora_extended,
|
|
save_lora_weight,
|
|
save_lora_safetensors,
|
|
train_patch_pipe,
|
|
monkeypatch_or_replace_lora,
|
|
monkeypatch_or_replace_lora_extended
|
|
)
|
|
|
|
from ...animatediff.stable_lora.lora import (
|
|
activate_lora_train,
|
|
add_lora_to,
|
|
save_lora,
|
|
load_lora,
|
|
set_mode_group
|
|
)
|
|
|
|
FILE_BASENAMES = ['unet', 'text_encoder']
|
|
LORA_FILE_TYPES = ['.pt', '.safetensors']
|
|
CLONE_OF_SIMO_KEYS = ['model', 'loras', 'target_replace_module', 'r']
|
|
STABLE_LORA_KEYS = [
|
|
'model',
|
|
'target_module',
|
|
'search_class',
|
|
'r',
|
|
'dropout',
|
|
'lora_bias',
|
|
'scale'
|
|
]
|
|
|
|
lora_versions = dict(
|
|
stable_lora = "stable_lora",
|
|
cloneofsimo = "cloneofsimo"
|
|
)
|
|
|
|
lora_func_types = dict(
|
|
loader = "loader",
|
|
injector = "injector"
|
|
)
|
|
|
|
lora_args = dict(
|
|
model = None,
|
|
loras = None,
|
|
target_replace_module = [],
|
|
target_module = [],
|
|
r = 4,
|
|
search_class = [torch.nn.Linear],
|
|
dropout = 0,
|
|
lora_bias = 'none',
|
|
scale = 0
|
|
)
|
|
|
|
LoraVersions = SimpleNamespace(**lora_versions)
|
|
LoraFuncTypes = SimpleNamespace(**lora_func_types)
|
|
|
|
LORA_VERSIONS = [LoraVersions.stable_lora, LoraVersions.cloneofsimo]
|
|
LORA_FUNC_TYPES = [LoraFuncTypes.loader, LoraFuncTypes.injector]
|
|
|
|
def filter_dict(_dict, keys=[]):
|
|
if len(keys) == 0:
|
|
assert "Keys cannot empty for filtering return dict."
|
|
|
|
for k in keys:
|
|
if k not in lora_args.keys():
|
|
assert f"{k} does not exist in available LoRA arguments"
|
|
|
|
return {k: v for k, v in _dict.items() if k in keys}
|
|
|
|
class LoraHandler(object):
|
|
def __init__(
|
|
self,
|
|
version: LORA_VERSIONS = LoraVersions.cloneofsimo,
|
|
use_unet_lora: bool = False,
|
|
use_text_lora: bool = False,
|
|
save_for_webui: bool = False,
|
|
only_for_webui: bool = False,
|
|
lora_bias: str = 'none',
|
|
unet_replace_modules: list = ['UNet3DConditionModel'],
|
|
text_encoder_replace_modules: list = ['CLIPEncoderLayer']
|
|
):
|
|
self.version = version
|
|
self.lora_loader = self.get_lora_func(func_type=LoraFuncTypes.loader)
|
|
self.lora_injector = self.get_lora_func(func_type=LoraFuncTypes.injector)
|
|
self.lora_bias = lora_bias
|
|
self.use_unet_lora = use_unet_lora
|
|
self.use_text_lora = use_text_lora
|
|
self.save_for_webui = save_for_webui
|
|
self.only_for_webui = only_for_webui
|
|
self.unet_replace_modules = unet_replace_modules
|
|
self.text_encoder_replace_modules = text_encoder_replace_modules
|
|
self.use_lora = any([use_text_lora, use_unet_lora])
|
|
|
|
if self.use_lora:
|
|
print(f"Using LoRA Version: {self.version}")
|
|
|
|
def is_cloneofsimo_lora(self):
|
|
return self.version == LoraVersions.cloneofsimo
|
|
|
|
def is_stable_lora(self):
|
|
return self.version == LoraVersions.stable_lora
|
|
|
|
def get_lora_func(self, func_type: LORA_FUNC_TYPES = LoraFuncTypes.loader):
|
|
|
|
if self.is_cloneofsimo_lora():
|
|
|
|
if func_type == LoraFuncTypes.loader:
|
|
return monkeypatch_or_replace_lora_extended
|
|
|
|
if func_type == LoraFuncTypes.injector:
|
|
return inject_trainable_lora_extended
|
|
|
|
if self.is_stable_lora():
|
|
|
|
if func_type == LoraFuncTypes.loader:
|
|
return load_lora
|
|
|
|
if func_type == LoraFuncTypes.injector:
|
|
return add_lora_to
|
|
|
|
assert "LoRA Version does not exist."
|
|
|
|
def check_lora_ext(self, lora_file: str):
|
|
return lora_file.endswith(tuple(LORA_FILE_TYPES))
|
|
|
|
def get_lora_file_path(
|
|
self,
|
|
lora_path: str,
|
|
model: Union[UNet3DConditionModel, CLIPTextModel]
|
|
):
|
|
if os.path.exists(lora_path):
|
|
lora_filenames = [fns for fns in os.listdir(lora_path)]
|
|
is_lora = self.check_lora_ext(lora_path)
|
|
|
|
is_unet = isinstance(model, UNet3DConditionModel)
|
|
is_text = isinstance(model, CLIPTextModel)
|
|
idx = 0 if is_unet else 1
|
|
|
|
base_name = FILE_BASENAMES[idx]
|
|
|
|
for lora_filename in lora_filenames:
|
|
is_lora = self.check_lora_ext(lora_filename)
|
|
if not is_lora:
|
|
continue
|
|
|
|
if base_name in lora_filename:
|
|
return os.path.join(lora_path, lora_filename)
|
|
|
|
return None
|
|
|
|
def handle_lora_load(self, file_name:str, lora_loader_args: dict = None):
|
|
self.lora_loader(**lora_loader_args)
|
|
print(f"Successfully loaded LoRA from: {file_name}")
|
|
|
|
def load_lora(self, model, lora_path: str = '', lora_loader_args: dict = None, *args, **kwargs):
|
|
try:
|
|
lora_file = self.get_lora_file_path(lora_path, model)
|
|
|
|
if lora_file is not None:
|
|
lora_loader_args.update({"lora_path": lora_file})
|
|
self.handle_lora_load(lora_file, lora_loader_args)
|
|
|
|
else:
|
|
print(f"Could not load LoRAs for {model.__class__.__name__}. Injecting new ones instead...")
|
|
|
|
except Exception as e:
|
|
print(f"An error occured while loading a LoRA file: {e}")
|
|
|
|
def get_lora_func_args(
|
|
self,
|
|
lora_path,
|
|
use_lora,
|
|
model,
|
|
replace_modules,
|
|
r,
|
|
dropout,
|
|
lora_bias,
|
|
scale
|
|
):
|
|
return_dict = lora_args.copy()
|
|
|
|
if self.is_cloneofsimo_lora():
|
|
return_dict = filter_dict(return_dict, keys=CLONE_OF_SIMO_KEYS)
|
|
return_dict.update({
|
|
"model": model,
|
|
"loras": self.get_lora_file_path(lora_path, model),
|
|
"target_replace_module": replace_modules,
|
|
"r": r
|
|
})
|
|
|
|
if self.is_stable_lora():
|
|
KEYS = ['model', 'lora_path', 'scale']
|
|
return_dict = filter_dict(return_dict, KEYS)
|
|
|
|
return_dict.update({'model': model, 'lora_path': lora_path, 'scale': scale})
|
|
|
|
return return_dict
|
|
|
|
def do_lora_injection(
|
|
self,
|
|
model,
|
|
replace_modules,
|
|
bias='none',
|
|
dropout=0,
|
|
r=4,
|
|
scale=0,
|
|
lora_loader_args=None,
|
|
):
|
|
REPLACE_MODULES = replace_modules
|
|
|
|
params = None
|
|
negation = None
|
|
is_injection_hybrid = False
|
|
|
|
if self.is_cloneofsimo_lora():
|
|
is_injection_hybrid = True
|
|
injector_args = lora_loader_args
|
|
|
|
params, negation = self.lora_injector(**injector_args)
|
|
for _up, _down in extract_lora_ups_down(
|
|
model,
|
|
target_replace_module=REPLACE_MODULES):
|
|
|
|
if all(x is not None for x in [_up, _down]):
|
|
print(f"Lora successfully injected into {model.__class__.__name__}.")
|
|
|
|
break
|
|
|
|
return params, negation, is_injection_hybrid
|
|
|
|
if self.is_stable_lora():
|
|
injector_args = lora_args.copy()
|
|
injector_args = filter_dict(injector_args, keys=STABLE_LORA_KEYS)
|
|
|
|
SEARCH_CLASS = [torch.nn.Linear, torch.nn.Conv2d, torch.nn.Conv3d, torch.nn.Embedding]
|
|
|
|
injector_args.update({
|
|
"model": model,
|
|
"target_module": REPLACE_MODULES,
|
|
"search_class": SEARCH_CLASS,
|
|
"r": r,
|
|
"dropout": dropout,
|
|
"lora_bias": self.lora_bias,
|
|
"scale": scale
|
|
})
|
|
activator = self.lora_injector(**injector_args)
|
|
activator()
|
|
|
|
return params, negation, is_injection_hybrid
|
|
|
|
def add_lora_to_model(self, use_lora, model, replace_modules, dropout=0.0, lora_path='', r=16, scale=0):
|
|
|
|
params = None
|
|
negation = None
|
|
|
|
lora_loader_args = self.get_lora_func_args(
|
|
lora_path,
|
|
use_lora,
|
|
model,
|
|
replace_modules,
|
|
r,
|
|
dropout,
|
|
self.lora_bias,
|
|
scale
|
|
)
|
|
|
|
if use_lora:
|
|
params, negation, is_injection_hybrid = self.do_lora_injection(
|
|
model,
|
|
replace_modules,
|
|
bias=self.lora_bias,
|
|
lora_loader_args=lora_loader_args,
|
|
dropout=dropout,
|
|
r=r,
|
|
scale=scale
|
|
)
|
|
|
|
if not is_injection_hybrid:
|
|
self.load_lora(model, lora_path=lora_path, lora_loader_args=lora_loader_args)
|
|
|
|
params = model if params is None else params
|
|
return params, negation
|
|
|
|
|
|
def deactivate_lora_train(self, models, deactivate=True):
|
|
"""
|
|
Usage: Use before and after sampling previews.
|
|
Currently only available for Stable LoRA.
|
|
"""
|
|
if self.is_stable_lora():
|
|
set_mode_group(models, not deactivate)
|
|
|
|
def save_cloneofsimo_lora(
|
|
self,
|
|
model,
|
|
save_path,
|
|
step,
|
|
use_safetensors=True,
|
|
lora_rank="",
|
|
lora_name="",
|
|
use_motion_lora_format=False
|
|
):
|
|
|
|
# Same arguments as top level method
|
|
def save_lora(
|
|
model,
|
|
name,
|
|
condition,
|
|
replace_modules,
|
|
step,
|
|
save_path,
|
|
use_safetensors=True,
|
|
lora_rank="",
|
|
use_motion_lora_format=False
|
|
):
|
|
if condition and replace_modules is not None:
|
|
|
|
save_path = f"{save_path}/{step}_{name}"
|
|
|
|
if not use_safetensors:
|
|
save_lora_weight(
|
|
model,
|
|
save_path + ".pt",
|
|
replace_modules,
|
|
self.lora_r,
|
|
use_motion_lora_format=use_motion_lora_format
|
|
)
|
|
else:
|
|
save_lora_safetensors(
|
|
model,
|
|
save_path + ".safetensors",
|
|
target_replace_module=replace_modules,
|
|
lora_rank=lora_rank,
|
|
use_motion_lora_format=use_motion_lora_format
|
|
)
|
|
|
|
save_lora(
|
|
model.unet,
|
|
f"{lora_name}_{FILE_BASENAMES[0]}",
|
|
self.use_unet_lora,
|
|
self.unet_replace_modules,
|
|
step,
|
|
save_path,
|
|
use_safetensors,
|
|
lora_rank,
|
|
use_motion_lora_format
|
|
)
|
|
save_lora(
|
|
model.text_encoder,
|
|
f"{lora_name}_{FILE_BASENAMES[1]}",
|
|
self.use_text_lora,
|
|
self.text_encoder_replace_modules,
|
|
step,
|
|
save_path,
|
|
use_safetensors,
|
|
lora_rank,
|
|
use_motion_lora_format
|
|
)
|
|
|
|
train_patch_pipe(model, self.use_unet_lora, self.use_text_lora)
|
|
|
|
def save_stable_lora(
|
|
self,
|
|
model,
|
|
step,
|
|
name,
|
|
save_path = '',
|
|
save_for_webui=False,
|
|
only_for_webui=False
|
|
):
|
|
import uuid
|
|
|
|
save_filename = f"{step}_{name}"
|
|
lora_metadata = metadata = {
|
|
"stable_lora_text_to_video": "v1",
|
|
"lora_name": name + "_" + uuid.uuid4().hex.lower()[:5]
|
|
}
|
|
save_lora(
|
|
unet=model.unet,
|
|
text_encoder=model.text_encoder,
|
|
save_text_weights=self.use_text_lora,
|
|
output_dir=save_path,
|
|
lora_filename=save_filename,
|
|
lora_bias=self.lora_bias,
|
|
save_for_webui=self.save_for_webui,
|
|
only_webui=self.only_for_webui,
|
|
metadata=lora_metadata,
|
|
unet_dict_converter=convert_unet_state_dict,
|
|
text_dict_converter=convert_text_enc_state_dict_v20
|
|
)
|
|
|
|
def save_lora_weights(
|
|
self,
|
|
model: None,
|
|
save_path: str ='',
|
|
step: str = '',
|
|
use_safetensors: bool = True,
|
|
lora_rank="Not Logged",
|
|
lora_name="",
|
|
use_motion_lora_format=False
|
|
):
|
|
save_path = f"{save_path}"
|
|
os.makedirs(save_path, exist_ok=True)
|
|
|
|
if self.is_cloneofsimo_lora():
|
|
if any([self.save_for_webui, self.only_for_webui]):
|
|
warnings.warn(
|
|
"""
|
|
You have 'save_for_webui' enabled, but are using cloneofsimo's LoRA implemention.
|
|
Only 'stable_lora' is supported for saving to a compatible webui file.
|
|
"""
|
|
)
|
|
self.save_cloneofsimo_lora(
|
|
model,
|
|
save_path,
|
|
step,
|
|
use_safetensors=use_safetensors,
|
|
lora_rank=lora_rank,
|
|
lora_name=lora_name,
|
|
use_motion_lora_format=use_motion_lora_format
|
|
)
|
|
|
|
if self.is_stable_lora():
|
|
name = 'lora_text_to_video'
|
|
self.save_stable_lora(model, step, name, save_path)
|
|
|