12 Commits
Author SHA1 Message Date
kijai 09fef403d1 Use prodigy-plus-schedule-free pypi package instead, fix typo in config 2025-04-02 10:35:43 +03:00
Jukka Seppänen 639b3e80ba Merge pull request #130 from Mikhael-Danilov/patch-1
Fix FluxTrainAndValidateLoop.validate
2025-02-03 23:27:28 +02:00
Mikhael-Danilov 74611324dc Fix FluxTrainAndValidateLoop.validate 2025-02-03 23:11:31 +03:00
kijai f7025638fa remove prints 2025-02-02 16:32:20 +02:00
kijai 6a91611a2b Fix network args, update example folder path to support template loader 2025-02-02 16:29:04 +02:00
kijai f6af45a169 update prodigyplusschedulefree, support Flex with arg bypass_flux_guidance
https://github.com/kohya-ss/sd-scripts/pull/1893
2025-01-31 21:52:02 +02:00
kijai 5f254225c7 Update pyproject.toml 2025-01-19 17:57:36 +02:00
kijai 580bd8bb06 Add prodigyplusschedulefree license and link
apologies for not initially including this
2025-01-19 17:57:17 +02:00
kijai 4343f2060a Add lycoris license and mention
sorry for the oversight
2025-01-19 12:19:50 +02:00
kijai 998968f5ff Support LyCORIS
https://github.com/KohakuBlueleaf/Lycoris
2025-01-16 19:36:18 +02:00
kijai 30cea9e372 Support SDXL
Still experimental, some nodes overlap and should be renamed for clarity.
2025-01-11 17:46:41 +02:00
kijai 136697a655 Update prodigy-plus-schedule-free 2025-01-10 12:12:58 +02:00
55 changed files with 15372 additions and 1432 deletions
+8
View File
@@ -1,5 +1,11 @@
# ComfyUI Flux Trainer
Wrapper for slightly modified kohya's training scripts: https://github.com/kohya-ss/sd-scripts
Including code from: https://github.com/KohakuBlueleaf/Lycoris
And https://github.com/LoganBooker/prodigy-plus-schedule-free
## DISCLAIMER:
I have **very** little previous experience in training anything, Flux is basically first model I've been inspired to learn. Previously I've only trained AnimateDiff Motion Loras, and built similar training nodes for it.
@@ -42,5 +48,7 @@ For full model training the fp16 version of the main model needs to be used.
Currently supports LoRA training, and untested full finetune with code from kohya's scripts: https://github.com/kohya-ss/sd-scripts
Experimental support for LyCORIS training has been added as well, using code from: https://github.com/KohakuBlueleaf/Lycoris
![Screenshot 2024-08-21 020207](https://github.com/user-attachments/assets/1686b180-90c8-41d0-8c96-63e76ebc2475)
+4
View File
@@ -1,8 +1,12 @@
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
from .nodes_sd3 import NODE_CLASS_MAPPINGS as NODE_CLASS_MAPPINGS_SD3
from .nodes_sd3 import NODE_DISPLAY_NAME_MAPPINGS as NODE_DISPLAY_NAME_MAPPINGS_SD3
from .nodes_sdxl import NODE_CLASS_MAPPINGS as NODE_CLASS_MAPPINGS_SDXL
from .nodes_sdxl import NODE_DISPLAY_NAME_MAPPINGS as NODE_DISPLAY_NAME_MAPPINGS_SDXL
NODE_CLASS_MAPPINGS.update(NODE_CLASS_MAPPINGS_SD3)
NODE_CLASS_MAPPINGS.update(NODE_CLASS_MAPPINGS_SDXL)
NODE_DISPLAY_NAME_MAPPINGS.update(NODE_DISPLAY_NAME_MAPPINGS_SD3)
NODE_DISPLAY_NAME_MAPPINGS.update(NODE_DISPLAY_NAME_MAPPINGS_SDXL)
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
File diff suppressed because it is too large Load Diff
Binary file not shown.

Before

Width:  |  Height:  |  Size: 2.5 MiB

+6
View File
@@ -669,6 +669,9 @@ class FluxTrainer:
if not args.apply_t5_attn_mask:
t5_attn_mask = None
if args.bypass_flux_guidance:
flux_utils.bypass_flux_guidance(flux)
with accelerator.autocast():
# YiYi notes: divide it by 1000 for now because we scale it by 1000 in the transformer model (we should not keep it but I want to keep the inputs same for the model for testing)
model_pred = flux(
@@ -685,6 +688,9 @@ class FluxTrainer:
# unpack latents
model_pred = flux_utils.unpack_latents(model_pred, packed_latent_height, packed_latent_width)
if args.bypass_flux_guidance:
flux_utils.restore_flux_guidance(flux)
# apply model prediction type
model_pred, weighting = flux_train_utils.apply_model_prediction_type(args, model_pred, noisy_model_input, sigmas)
+6
View File
@@ -356,6 +356,9 @@ class FluxNetworkTrainer(NetworkTrainer):
"""
return model_pred
if args.bypass_flux_guidance:
flux_utils.bypass_flux_guidance(unet)
model_pred = call_dit(
img=packed_noisy_model_input,
@@ -371,6 +374,9 @@ class FluxNetworkTrainer(NetworkTrainer):
# unpack latents
model_pred = flux_utils.unpack_latents(model_pred, packed_latent_height, packed_latent_width)
if args.bypass_flux_guidance: #for flex
flux_utils.restore_flux_guidance(unet)
# apply model prediction type
model_pred, weighting = flux_train_utils.apply_model_prediction_type(args, model_pred, noisy_model_input, sigmas)
+5
View File
@@ -578,3 +578,8 @@ def add_flux_train_arguments(parser: argparse.ArgumentParser):
default=3.0,
help="Discrete flow shift for the Euler Discrete Scheduler, default is 3.0. / Euler Discrete Schedulerの離散フローシフト、デフォルトは3.0。",
)
parser.add_argument(
"--bypass_flux_guidance"
, action="store_true"
, help="bypass flux guidance module for Flex.1-Alpha Training"
)
+7
View File
@@ -21,7 +21,14 @@ MODEL_VERSION_FLUX_V1 = "flux1"
MODEL_NAME_DEV = "dev"
MODEL_NAME_SCHNELL = "schnell"
# bypass guidance
def bypass_flux_guidance(transformer):
transformer.params.guidance_embed = False
# restore the forward function
def restore_flux_guidance(transformer):
transformer.params.guidance_embed = True
def analyze_checkpoint_state(ckpt_path: str) -> Tuple[bool, bool, Tuple[int, int], List[str]]:
"""
チェックポイントの状態を分析し、DiffusersかBFLか、devかschnellか、ブロック数を計算して返す。
File diff suppressed because it is too large Load Diff
+583
View File
@@ -0,0 +1,583 @@
import torch
import safetensors
from accelerate import init_empty_weights
from accelerate.utils.modeling import set_module_tensor_to_device
from safetensors.torch import load_file, save_file
from transformers import CLIPTextModel, CLIPTextConfig, CLIPTextModelWithProjection, CLIPTokenizer
from typing import List
from diffusers import AutoencoderKL, EulerDiscreteScheduler, UNet2DConditionModel
from . import model_util
from . import sdxl_original_unet
from .utils import setup_logging
setup_logging()
import logging
logger = logging.getLogger(__name__)
VAE_SCALE_FACTOR = 0.13025
MODEL_VERSION_SDXL_BASE_V1_0 = "sdxl_base_v1-0"
# Diffusersの設定を読み込むための参照モデル
DIFFUSERS_REF_MODEL_ID_SDXL = "stabilityai/stable-diffusion-xl-base-1.0"
DIFFUSERS_SDXL_UNET_CONFIG = {
"act_fn": "silu",
"addition_embed_type": "text_time",
"addition_embed_type_num_heads": 64,
"addition_time_embed_dim": 256,
"attention_head_dim": [5, 10, 20],
"block_out_channels": [320, 640, 1280],
"center_input_sample": False,
"class_embed_type": None,
"class_embeddings_concat": False,
"conv_in_kernel": 3,
"conv_out_kernel": 3,
"cross_attention_dim": 2048,
"cross_attention_norm": None,
"down_block_types": ["DownBlock2D", "CrossAttnDownBlock2D", "CrossAttnDownBlock2D"],
"downsample_padding": 1,
"dual_cross_attention": False,
"encoder_hid_dim": None,
"encoder_hid_dim_type": None,
"flip_sin_to_cos": True,
"freq_shift": 0,
"in_channels": 4,
"layers_per_block": 2,
"mid_block_only_cross_attention": None,
"mid_block_scale_factor": 1,
"mid_block_type": "UNetMidBlock2DCrossAttn",
"norm_eps": 1e-05,
"norm_num_groups": 32,
"num_attention_heads": None,
"num_class_embeds": None,
"only_cross_attention": False,
"out_channels": 4,
"projection_class_embeddings_input_dim": 2816,
"resnet_out_scale_factor": 1.0,
"resnet_skip_time_act": False,
"resnet_time_scale_shift": "default",
"sample_size": 128,
"time_cond_proj_dim": None,
"time_embedding_act_fn": None,
"time_embedding_dim": None,
"time_embedding_type": "positional",
"timestep_post_act": None,
"transformer_layers_per_block": [1, 2, 10],
"up_block_types": ["CrossAttnUpBlock2D", "CrossAttnUpBlock2D", "UpBlock2D"],
"upcast_attention": False,
"use_linear_projection": True,
}
def convert_sdxl_text_encoder_2_checkpoint(checkpoint, max_length):
SDXL_KEY_PREFIX = "conditioner.embedders.1.model."
# SD2のと、基本的には同じ。logit_scaleを後で使うので、それを追加で返す
# logit_scaleはcheckpointの保存時に使用する
def convert_key(key):
# common conversion
key = key.replace(SDXL_KEY_PREFIX + "transformer.", "text_model.encoder.")
key = key.replace(SDXL_KEY_PREFIX, "text_model.")
if "resblocks" in key:
# resblocks conversion
key = key.replace(".resblocks.", ".layers.")
if ".ln_" in key:
key = key.replace(".ln_", ".layer_norm")
elif ".mlp." in key:
key = key.replace(".c_fc.", ".fc1.")
key = key.replace(".c_proj.", ".fc2.")
elif ".attn.out_proj" in key:
key = key.replace(".attn.out_proj.", ".self_attn.out_proj.")
elif ".attn.in_proj" in key:
key = None # 特殊なので後で処理する
else:
raise ValueError(f"unexpected key in SD: {key}")
elif ".positional_embedding" in key:
key = key.replace(".positional_embedding", ".embeddings.position_embedding.weight")
elif ".text_projection" in key:
key = key.replace("text_model.text_projection", "text_projection.weight")
elif ".logit_scale" in key:
key = None # 後で処理する
elif ".token_embedding" in key:
key = key.replace(".token_embedding.weight", ".embeddings.token_embedding.weight")
elif ".ln_final" in key:
key = key.replace(".ln_final", ".final_layer_norm")
# ckpt from comfy has this key: text_model.encoder.text_model.embeddings.position_ids
elif ".embeddings.position_ids" in key:
key = None # remove this key: position_ids is not used in newer transformers
return key
keys = list(checkpoint.keys())
new_sd = {}
for key in keys:
new_key = convert_key(key)
if new_key is None:
continue
new_sd[new_key] = checkpoint[key]
# attnの変換
for key in keys:
if ".resblocks" in key and ".attn.in_proj_" in key:
# 三つに分割
values = torch.chunk(checkpoint[key], 3)
key_suffix = ".weight" if "weight" in key else ".bias"
key_pfx = key.replace(SDXL_KEY_PREFIX + "transformer.resblocks.", "text_model.encoder.layers.")
key_pfx = key_pfx.replace("_weight", "")
key_pfx = key_pfx.replace("_bias", "")
key_pfx = key_pfx.replace(".attn.in_proj", ".self_attn.")
new_sd[key_pfx + "q_proj" + key_suffix] = values[0]
new_sd[key_pfx + "k_proj" + key_suffix] = values[1]
new_sd[key_pfx + "v_proj" + key_suffix] = values[2]
# logit_scale はDiffusersには含まれないが、保存時に戻したいので別途返す
logit_scale = checkpoint.get(SDXL_KEY_PREFIX + "logit_scale", None)
# temporary workaround for text_projection.weight.weight for Playground-v2
if "text_projection.weight.weight" in new_sd:
logger.info("convert_sdxl_text_encoder_2_checkpoint: convert text_projection.weight.weight to text_projection.weight")
new_sd["text_projection.weight"] = new_sd["text_projection.weight.weight"]
del new_sd["text_projection.weight.weight"]
return new_sd, logit_scale
# load state_dict without allocating new tensors
def _load_state_dict_on_device(model, state_dict, device, dtype=None):
# dtype will use fp32 as default
missing_keys = list(model.state_dict().keys() - state_dict.keys())
unexpected_keys = list(state_dict.keys() - model.state_dict().keys())
# similar to model.load_state_dict()
if not missing_keys and not unexpected_keys:
for k in list(state_dict.keys()):
set_module_tensor_to_device(model, k, device, value=state_dict.pop(k), dtype=dtype)
return "<All keys matched successfully>"
# error_msgs
error_msgs: List[str] = []
if missing_keys:
error_msgs.insert(0, "Missing key(s) in state_dict: {}. ".format(", ".join('"{}"'.format(k) for k in missing_keys)))
if unexpected_keys:
error_msgs.insert(0, "Unexpected key(s) in state_dict: {}. ".format(", ".join('"{}"'.format(k) for k in unexpected_keys)))
raise RuntimeError("Error(s) in loading state_dict for {}:\n\t{}".format(model.__class__.__name__, "\n\t".join(error_msgs)))
def load_models_from_sdxl_checkpoint(model_version, ckpt_path, map_location, dtype=None, disable_mmap=False):
# model_version is reserved for future use
# dtype is used for full_fp16/bf16 integration. Text Encoder will remain fp32, because it runs on CPU when caching
# Load the state dict
if model_util.is_safetensors(ckpt_path):
checkpoint = None
if disable_mmap:
state_dict = safetensors.torch.load(open(ckpt_path, "rb").read())
else:
try:
state_dict = load_file(ckpt_path, device=map_location)
except:
state_dict = load_file(ckpt_path) # prevent device invalid Error
epoch = None
global_step = None
else:
checkpoint = torch.load(ckpt_path, map_location=map_location)
if "state_dict" in checkpoint:
state_dict = checkpoint["state_dict"]
epoch = checkpoint.get("epoch", 0)
global_step = checkpoint.get("global_step", 0)
else:
state_dict = checkpoint
epoch = 0
global_step = 0
checkpoint = None
# U-Net
logger.info("building U-Net")
with init_empty_weights():
unet = sdxl_original_unet.SdxlUNet2DConditionModel()
logger.info("loading U-Net from checkpoint")
unet_sd = {}
for k in list(state_dict.keys()):
if k.startswith("model.diffusion_model."):
unet_sd[k.replace("model.diffusion_model.", "")] = state_dict.pop(k)
info = _load_state_dict_on_device(unet, unet_sd, device=map_location, dtype=dtype)
logger.info(f"U-Net: {info}")
# Text Encoders
logger.info("building text encoders")
# Text Encoder 1 is same to Stability AI's SDXL
text_model1_cfg = CLIPTextConfig(
vocab_size=49408,
hidden_size=768,
intermediate_size=3072,
num_hidden_layers=12,
num_attention_heads=12,
max_position_embeddings=77,
hidden_act="quick_gelu",
layer_norm_eps=1e-05,
dropout=0.0,
attention_dropout=0.0,
initializer_range=0.02,
initializer_factor=1.0,
pad_token_id=1,
bos_token_id=0,
eos_token_id=2,
model_type="clip_text_model",
projection_dim=768,
# torch_dtype="float32",
# transformers_version="4.25.0.dev0",
)
with init_empty_weights():
text_model1 = CLIPTextModel._from_config(text_model1_cfg)
# Text Encoder 2 is different from Stability AI's SDXL. SDXL uses open clip, but we use the model from HuggingFace.
# Note: Tokenizer from HuggingFace is different from SDXL. We must use open clip's tokenizer.
text_model2_cfg = CLIPTextConfig(
vocab_size=49408,
hidden_size=1280,
intermediate_size=5120,
num_hidden_layers=32,
num_attention_heads=20,
max_position_embeddings=77,
hidden_act="gelu",
layer_norm_eps=1e-05,
dropout=0.0,
attention_dropout=0.0,
initializer_range=0.02,
initializer_factor=1.0,
pad_token_id=1,
bos_token_id=0,
eos_token_id=2,
model_type="clip_text_model",
projection_dim=1280,
# torch_dtype="float32",
# transformers_version="4.25.0.dev0",
)
with init_empty_weights():
text_model2 = CLIPTextModelWithProjection(text_model2_cfg)
logger.info("loading text encoders from checkpoint")
te1_sd = {}
te2_sd = {}
for k in list(state_dict.keys()):
if k.startswith("conditioner.embedders.0.transformer."):
te1_sd[k.replace("conditioner.embedders.0.transformer.", "")] = state_dict.pop(k)
elif k.startswith("conditioner.embedders.1.model."):
te2_sd[k] = state_dict.pop(k)
# 最新の transformers では position_ids を含むとエラーになるので削除 / remove position_ids for latest transformers
if "text_model.embeddings.position_ids" in te1_sd:
te1_sd.pop("text_model.embeddings.position_ids")
info1 = _load_state_dict_on_device(text_model1, te1_sd, device=map_location) # remain fp32
logger.info(f"text encoder 1: {info1}")
converted_sd, logit_scale = convert_sdxl_text_encoder_2_checkpoint(te2_sd, max_length=77)
info2 = _load_state_dict_on_device(text_model2, converted_sd, device=map_location) # remain fp32
logger.info(f"text encoder 2: {info2}")
# prepare vae
logger.info("building VAE")
vae_config = model_util.create_vae_diffusers_config()
with init_empty_weights():
vae = AutoencoderKL(**vae_config)
logger.info("loading VAE from checkpoint")
converted_vae_checkpoint = model_util.convert_ldm_vae_checkpoint(state_dict, vae_config)
info = _load_state_dict_on_device(vae, converted_vae_checkpoint, device=map_location, dtype=dtype)
logger.info(f"VAE: {info}")
ckpt_info = (epoch, global_step) if epoch is not None else None
return text_model1, text_model2, vae, unet, logit_scale, ckpt_info
def make_unet_conversion_map():
unet_conversion_map_layer = []
for i in range(3): # num_blocks is 3 in sdxl
# loop over downblocks/upblocks
for j in range(2):
# loop over resnets/attentions for downblocks
hf_down_res_prefix = f"down_blocks.{i}.resnets.{j}."
sd_down_res_prefix = f"input_blocks.{3*i + j + 1}.0."
unet_conversion_map_layer.append((sd_down_res_prefix, hf_down_res_prefix))
if i < 3:
# no attention layers in down_blocks.3
hf_down_atn_prefix = f"down_blocks.{i}.attentions.{j}."
sd_down_atn_prefix = f"input_blocks.{3*i + j + 1}.1."
unet_conversion_map_layer.append((sd_down_atn_prefix, hf_down_atn_prefix))
for j in range(3):
# loop over resnets/attentions for upblocks
hf_up_res_prefix = f"up_blocks.{i}.resnets.{j}."
sd_up_res_prefix = f"output_blocks.{3*i + j}.0."
unet_conversion_map_layer.append((sd_up_res_prefix, hf_up_res_prefix))
# if i > 0: commentout for sdxl
# no attention layers in up_blocks.0
hf_up_atn_prefix = f"up_blocks.{i}.attentions.{j}."
sd_up_atn_prefix = f"output_blocks.{3*i + j}.1."
unet_conversion_map_layer.append((sd_up_atn_prefix, hf_up_atn_prefix))
if i < 3:
# no downsample in down_blocks.3
hf_downsample_prefix = f"down_blocks.{i}.downsamplers.0.conv."
sd_downsample_prefix = f"input_blocks.{3*(i+1)}.0.op."
unet_conversion_map_layer.append((sd_downsample_prefix, hf_downsample_prefix))
# no upsample in up_blocks.3
hf_upsample_prefix = f"up_blocks.{i}.upsamplers.0."
sd_upsample_prefix = f"output_blocks.{3*i + 2}.{2}." # change for sdxl
unet_conversion_map_layer.append((sd_upsample_prefix, hf_upsample_prefix))
hf_mid_atn_prefix = "mid_block.attentions.0."
sd_mid_atn_prefix = "middle_block.1."
unet_conversion_map_layer.append((sd_mid_atn_prefix, hf_mid_atn_prefix))
for j in range(2):
hf_mid_res_prefix = f"mid_block.resnets.{j}."
sd_mid_res_prefix = f"middle_block.{2*j}."
unet_conversion_map_layer.append((sd_mid_res_prefix, hf_mid_res_prefix))
unet_conversion_map_resnet = [
# (stable-diffusion, HF Diffusers)
("in_layers.0.", "norm1."),
("in_layers.2.", "conv1."),
("out_layers.0.", "norm2."),
("out_layers.3.", "conv2."),
("emb_layers.1.", "time_emb_proj."),
("skip_connection.", "conv_shortcut."),
]
unet_conversion_map = []
for sd, hf in unet_conversion_map_layer:
if "resnets" in hf:
for sd_res, hf_res in unet_conversion_map_resnet:
unet_conversion_map.append((sd + sd_res, hf + hf_res))
else:
unet_conversion_map.append((sd, hf))
for j in range(2):
hf_time_embed_prefix = f"time_embedding.linear_{j+1}."
sd_time_embed_prefix = f"time_embed.{j*2}."
unet_conversion_map.append((sd_time_embed_prefix, hf_time_embed_prefix))
for j in range(2):
hf_label_embed_prefix = f"add_embedding.linear_{j+1}."
sd_label_embed_prefix = f"label_emb.0.{j*2}."
unet_conversion_map.append((sd_label_embed_prefix, hf_label_embed_prefix))
unet_conversion_map.append(("input_blocks.0.0.", "conv_in."))
unet_conversion_map.append(("out.0.", "conv_norm_out."))
unet_conversion_map.append(("out.2.", "conv_out."))
return unet_conversion_map
def convert_diffusers_unet_state_dict_to_sdxl(du_sd):
unet_conversion_map = make_unet_conversion_map()
conversion_map = {hf: sd for sd, hf in unet_conversion_map}
return convert_unet_state_dict(du_sd, conversion_map)
def convert_unet_state_dict(src_sd, conversion_map):
converted_sd = {}
for src_key, value in src_sd.items():
# さすがに全部回すのは時間がかかるので右から要素を削りつつprefixを探す
src_key_fragments = src_key.split(".")[:-1] # remove weight/bias
while len(src_key_fragments) > 0:
src_key_prefix = ".".join(src_key_fragments) + "."
if src_key_prefix in conversion_map:
converted_prefix = conversion_map[src_key_prefix]
converted_key = converted_prefix + src_key[len(src_key_prefix) :]
converted_sd[converted_key] = value
break
src_key_fragments.pop(-1)
assert len(src_key_fragments) > 0, f"key {src_key} not found in conversion map"
return converted_sd
def convert_sdxl_unet_state_dict_to_diffusers(sd):
unet_conversion_map = make_unet_conversion_map()
conversion_dict = {sd: hf for sd, hf in unet_conversion_map}
return convert_unet_state_dict(sd, conversion_dict)
def convert_text_encoder_2_state_dict_to_sdxl(checkpoint, logit_scale):
def convert_key(key):
# position_idsの除去
if ".position_ids" in key:
return None
# common
key = key.replace("text_model.encoder.", "transformer.")
key = key.replace("text_model.", "")
if "layers" in key:
# resblocks conversion
key = key.replace(".layers.", ".resblocks.")
if ".layer_norm" in key:
key = key.replace(".layer_norm", ".ln_")
elif ".mlp." in key:
key = key.replace(".fc1.", ".c_fc.")
key = key.replace(".fc2.", ".c_proj.")
elif ".self_attn.out_proj" in key:
key = key.replace(".self_attn.out_proj.", ".attn.out_proj.")
elif ".self_attn." in key:
key = None # 特殊なので後で処理する
else:
raise ValueError(f"unexpected key in DiffUsers model: {key}")
elif ".position_embedding" in key:
key = key.replace("embeddings.position_embedding.weight", "positional_embedding")
elif ".token_embedding" in key:
key = key.replace("embeddings.token_embedding.weight", "token_embedding.weight")
elif "text_projection" in key: # no dot in key
key = key.replace("text_projection.weight", "text_projection")
elif "final_layer_norm" in key:
key = key.replace("final_layer_norm", "ln_final")
return key
keys = list(checkpoint.keys())
new_sd = {}
for key in keys:
new_key = convert_key(key)
if new_key is None:
continue
new_sd[new_key] = checkpoint[key]
# attnの変換
for key in keys:
if "layers" in key and "q_proj" in key:
# 三つを結合
key_q = key
key_k = key.replace("q_proj", "k_proj")
key_v = key.replace("q_proj", "v_proj")
value_q = checkpoint[key_q]
value_k = checkpoint[key_k]
value_v = checkpoint[key_v]
value = torch.cat([value_q, value_k, value_v])
new_key = key.replace("text_model.encoder.layers.", "transformer.resblocks.")
new_key = new_key.replace(".self_attn.q_proj.", ".attn.in_proj_")
new_sd[new_key] = value
if logit_scale is not None:
new_sd["logit_scale"] = logit_scale
return new_sd
def save_stable_diffusion_checkpoint(
output_file,
text_encoder1,
text_encoder2,
unet,
epochs,
steps,
ckpt_info,
vae,
logit_scale,
metadata,
save_dtype=None,
):
state_dict = {}
def update_sd(prefix, sd):
for k, v in sd.items():
key = prefix + k
if save_dtype is not None:
v = v.detach().clone().to("cpu").to(save_dtype)
state_dict[key] = v
# Convert the UNet model
update_sd("model.diffusion_model.", unet.state_dict())
# Convert the text encoders
update_sd("conditioner.embedders.0.transformer.", text_encoder1.state_dict())
text_enc2_dict = convert_text_encoder_2_state_dict_to_sdxl(text_encoder2.state_dict(), logit_scale)
update_sd("conditioner.embedders.1.model.", text_enc2_dict)
# Convert the VAE
vae_dict = model_util.convert_vae_state_dict(vae.state_dict())
update_sd("first_stage_model.", vae_dict)
# Put together new checkpoint
key_count = len(state_dict.keys())
new_ckpt = {"state_dict": state_dict}
# epoch and global_step are sometimes not int
if ckpt_info is not None:
epochs += ckpt_info[0]
steps += ckpt_info[1]
new_ckpt["epoch"] = epochs
new_ckpt["global_step"] = steps
if model_util.is_safetensors(output_file):
save_file(state_dict, output_file, metadata)
else:
torch.save(new_ckpt, output_file)
return key_count
def save_diffusers_checkpoint(
output_dir, text_encoder1, text_encoder2, unet, pretrained_model_name_or_path, vae=None, use_safetensors=False, save_dtype=None
):
from diffusers import StableDiffusionXLPipeline
# convert U-Net
unet_sd = unet.state_dict()
du_unet_sd = convert_sdxl_unet_state_dict_to_diffusers(unet_sd)
diffusers_unet = UNet2DConditionModel(**DIFFUSERS_SDXL_UNET_CONFIG)
if save_dtype is not None:
diffusers_unet.to(save_dtype)
diffusers_unet.load_state_dict(du_unet_sd)
# create pipeline to save
if pretrained_model_name_or_path is None:
pretrained_model_name_or_path = DIFFUSERS_REF_MODEL_ID_SDXL
scheduler = EulerDiscreteScheduler.from_pretrained(pretrained_model_name_or_path, subfolder="scheduler")
tokenizer1 = CLIPTokenizer.from_pretrained(pretrained_model_name_or_path, subfolder="tokenizer")
tokenizer2 = CLIPTokenizer.from_pretrained(pretrained_model_name_or_path, subfolder="tokenizer_2")
if vae is None:
vae = AutoencoderKL.from_pretrained(pretrained_model_name_or_path, subfolder="vae")
# prevent local path from being saved
def remove_name_or_path(model):
if hasattr(model, "config"):
model.config._name_or_path = None
model.config._name_or_path = None
remove_name_or_path(diffusers_unet)
remove_name_or_path(text_encoder1)
remove_name_or_path(text_encoder2)
remove_name_or_path(scheduler)
remove_name_or_path(tokenizer1)
remove_name_or_path(tokenizer2)
remove_name_or_path(vae)
pipeline = StableDiffusionXLPipeline(
unet=diffusers_unet,
text_encoder=text_encoder1,
text_encoder_2=text_encoder2,
vae=vae,
scheduler=scheduler,
tokenizer=tokenizer1,
tokenizer_2=tokenizer2,
)
if save_dtype is not None:
pipeline.to(None, save_dtype)
pipeline.save_pretrained(output_dir, safe_serialization=use_safetensors)
+272
View File
@@ -0,0 +1,272 @@
# some parts are modified from Diffusers library (Apache License 2.0)
import math
from types import SimpleNamespace
from typing import Any, Optional
import torch
import torch.utils.checkpoint
from torch import nn
from torch.nn import functional as F
from einops import rearrange
from .utils import setup_logging
setup_logging()
import logging
logger = logging.getLogger(__name__)
from . import sdxl_original_unet
from .sdxl_model_util import convert_sdxl_unet_state_dict_to_diffusers, convert_diffusers_unet_state_dict_to_sdxl
class ControlNetConditioningEmbedding(nn.Module):
def __init__(self):
super().__init__()
dims = [16, 32, 96, 256]
self.conv_in = nn.Conv2d(3, dims[0], kernel_size=3, padding=1)
self.blocks = nn.ModuleList([])
for i in range(len(dims) - 1):
channel_in = dims[i]
channel_out = dims[i + 1]
self.blocks.append(nn.Conv2d(channel_in, channel_in, kernel_size=3, padding=1))
self.blocks.append(nn.Conv2d(channel_in, channel_out, kernel_size=3, padding=1, stride=2))
self.conv_out = nn.Conv2d(dims[-1], 320, kernel_size=3, padding=1)
nn.init.zeros_(self.conv_out.weight) # zero module weight
nn.init.zeros_(self.conv_out.bias) # zero module bias
def forward(self, x):
x = self.conv_in(x)
x = F.silu(x)
for block in self.blocks:
x = block(x)
x = F.silu(x)
x = self.conv_out(x)
return x
class SdxlControlNet(sdxl_original_unet.SdxlUNet2DConditionModel):
def __init__(self, multiplier: Optional[float] = None, **kwargs):
super().__init__(**kwargs)
self.multiplier = multiplier
# remove unet layers
self.output_blocks = nn.ModuleList([])
del self.out
self.controlnet_cond_embedding = ControlNetConditioningEmbedding()
dims = [320, 320, 320, 320, 640, 640, 640, 1280, 1280]
self.controlnet_down_blocks = nn.ModuleList([])
for dim in dims:
self.controlnet_down_blocks.append(nn.Conv2d(dim, dim, kernel_size=1))
nn.init.zeros_(self.controlnet_down_blocks[-1].weight) # zero module weight
nn.init.zeros_(self.controlnet_down_blocks[-1].bias) # zero module bias
self.controlnet_mid_block = nn.Conv2d(1280, 1280, kernel_size=1)
nn.init.zeros_(self.controlnet_mid_block.weight) # zero module weight
nn.init.zeros_(self.controlnet_mid_block.bias) # zero module bias
def init_from_unet(self, unet: sdxl_original_unet.SdxlUNet2DConditionModel):
unet_sd = unet.state_dict()
unet_sd = {k: v for k, v in unet_sd.items() if not k.startswith("out")}
sd = super().state_dict()
sd.update(unet_sd)
info = super().load_state_dict(sd, strict=True, assign=True)
return info
def load_state_dict(self, state_dict: dict, strict: bool = True, assign: bool = True) -> Any:
# convert state_dict to SAI format
unet_sd = {}
for k in list(state_dict.keys()):
if not k.startswith("controlnet_"):
unet_sd[k] = state_dict.pop(k)
unet_sd = convert_diffusers_unet_state_dict_to_sdxl(unet_sd)
state_dict.update(unet_sd)
super().load_state_dict(state_dict, strict=strict, assign=assign)
def state_dict(self, destination=None, prefix="", keep_vars=False):
# convert state_dict to Diffusers format
state_dict = super().state_dict(destination, prefix, keep_vars)
control_net_sd = {}
for k in list(state_dict.keys()):
if k.startswith("controlnet_"):
control_net_sd[k] = state_dict.pop(k)
state_dict = convert_sdxl_unet_state_dict_to_diffusers(state_dict)
state_dict.update(control_net_sd)
return state_dict
def forward(
self,
x: torch.Tensor,
timesteps: Optional[torch.Tensor] = None,
context: Optional[torch.Tensor] = None,
y: Optional[torch.Tensor] = None,
cond_image: Optional[torch.Tensor] = None,
**kwargs,
) -> torch.Tensor:
# broadcast timesteps to batch dimension
timesteps = timesteps.expand(x.shape[0])
t_emb = sdxl_original_unet.get_timestep_embedding(timesteps, self.model_channels, downscale_freq_shift=0)
t_emb = t_emb.to(x.dtype)
emb = self.time_embed(t_emb)
assert x.shape[0] == y.shape[0], f"batch size mismatch: {x.shape[0]} != {y.shape[0]}"
assert x.dtype == y.dtype, f"dtype mismatch: {x.dtype} != {y.dtype}"
emb = emb + self.label_emb(y)
def call_module(module, h, emb, context):
x = h
for layer in module:
if isinstance(layer, sdxl_original_unet.ResnetBlock2D):
x = layer(x, emb)
elif isinstance(layer, sdxl_original_unet.Transformer2DModel):
x = layer(x, context)
else:
x = layer(x)
return x
h = x
multiplier = self.multiplier if self.multiplier is not None else 1.0
hs = []
for i, module in enumerate(self.input_blocks):
h = call_module(module, h, emb, context)
if i == 0:
h = self.controlnet_cond_embedding(cond_image) + h
hs.append(self.controlnet_down_blocks[i](h) * multiplier)
h = call_module(self.middle_block, h, emb, context)
h = self.controlnet_mid_block(h) * multiplier
return hs, h
class SdxlControlledUNet(sdxl_original_unet.SdxlUNet2DConditionModel):
"""
This class is for training purpose only.
"""
def __init__(self, **kwargs):
super().__init__(**kwargs)
def forward(self, x, timesteps=None, context=None, y=None, input_resi_add=None, mid_add=None, **kwargs):
# broadcast timesteps to batch dimension
timesteps = timesteps.expand(x.shape[0])
hs = []
t_emb = sdxl_original_unet.get_timestep_embedding(timesteps, self.model_channels, downscale_freq_shift=0)
t_emb = t_emb.to(x.dtype)
emb = self.time_embed(t_emb)
assert x.shape[0] == y.shape[0], f"batch size mismatch: {x.shape[0]} != {y.shape[0]}"
assert x.dtype == y.dtype, f"dtype mismatch: {x.dtype} != {y.dtype}"
emb = emb + self.label_emb(y)
def call_module(module, h, emb, context):
x = h
for layer in module:
if isinstance(layer, sdxl_original_unet.ResnetBlock2D):
x = layer(x, emb)
elif isinstance(layer, sdxl_original_unet.Transformer2DModel):
x = layer(x, context)
else:
x = layer(x)
return x
h = x
for module in self.input_blocks:
h = call_module(module, h, emb, context)
hs.append(h)
h = call_module(self.middle_block, h, emb, context)
h = h + mid_add
for module in self.output_blocks:
resi = hs.pop() + input_resi_add.pop()
h = torch.cat([h, resi], dim=1)
h = call_module(module, h, emb, context)
h = h.type(x.dtype)
h = call_module(self.out, h, emb, context)
return h
if __name__ == "__main__":
import time
logger.info("create unet")
unet = SdxlControlledUNet()
unet.to("cuda", torch.bfloat16)
unet.set_use_sdpa(True)
unet.set_gradient_checkpointing(True)
unet.train()
logger.info("create control_net")
control_net = SdxlControlNet()
control_net.to("cuda")
control_net.set_use_sdpa(True)
control_net.set_gradient_checkpointing(True)
control_net.train()
logger.info("Initialize control_net from unet")
control_net.init_from_unet(unet)
unet.requires_grad_(False)
control_net.requires_grad_(True)
# 使用メモリ量確認用の疑似学習ループ
logger.info("preparing optimizer")
# optimizer = torch.optim.SGD(unet.parameters(), lr=1e-3, nesterov=True, momentum=0.9) # not working
import bitsandbytes
optimizer = bitsandbytes.adam.Adam8bit(control_net.parameters(), lr=1e-3) # not working
# optimizer = bitsandbytes.optim.RMSprop8bit(unet.parameters(), lr=1e-3) # working at 23.5 GB with torch2
# optimizer=bitsandbytes.optim.Adagrad8bit(unet.parameters(), lr=1e-3) # working at 23.5 GB with torch2
# import transformers
# optimizer = transformers.optimization.Adafactor(unet.parameters(), relative_step=True) # working at 22.2GB with torch2
scaler = torch.cuda.amp.GradScaler(enabled=True)
logger.info("start training")
steps = 10
batch_size = 1
for step in range(steps):
logger.info(f"step {step}")
if step == 1:
time_start = time.perf_counter()
x = torch.randn(batch_size, 4, 128, 128).cuda() # 1024x1024
t = torch.randint(low=0, high=1000, size=(batch_size,), device="cuda")
txt = torch.randn(batch_size, 77, 2048).cuda()
vector = torch.randn(batch_size, sdxl_original_unet.ADM_IN_CHANNELS).cuda()
cond_img = torch.rand(batch_size, 3, 1024, 1024).cuda()
with torch.cuda.amp.autocast(enabled=True, dtype=torch.bfloat16):
input_resi_add, mid_add = control_net(x, t, txt, vector, cond_img)
output = unet(x, t, txt, vector, input_resi_add, mid_add)
target = torch.randn_like(output)
loss = torch.nn.functional.mse_loss(output, target)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
optimizer.zero_grad(set_to_none=True)
time_end = time.perf_counter()
logger.info(f"elapsed time: {time_end - time_start} [sec] for last {steps - 1} steps")
logger.info("finish training")
sd = control_net.state_dict()
from safetensors.torch import save_file
save_file(sd, r"E:\Work\SD\Tmp\sdxl\ctrl\control_net.safetensors")
File diff suppressed because it is too large Load Diff
+381
View File
@@ -0,0 +1,381 @@
import argparse
import math
import os
from typing import Optional
import torch
from .device_utils import init_ipex, clean_memory_on_device
init_ipex()
from accelerate import init_empty_weights
from tqdm import tqdm
from transformers import CLIPTokenizer
from . import model_util, sdxl_model_util, train_util, sdxl_original_unet
from .sdxl_lpw_stable_diffusion import SdxlStableDiffusionLongPromptWeightingPipeline
from .utils import setup_logging
setup_logging()
import logging
logger = logging.getLogger(__name__)
TOKENIZER1_PATH = "openai/clip-vit-large-patch14"
TOKENIZER2_PATH = "laion/CLIP-ViT-bigG-14-laion2B-39B-b160k"
# DEFAULT_NOISE_OFFSET = 0.0357
def load_target_model(args, accelerator, model_version: str, weight_dtype):
model_dtype = match_mixed_precision(args, weight_dtype) # prepare fp16/bf16
for pi in range(accelerator.state.num_processes):
if pi == accelerator.state.local_process_index:
logger.info(f"loading model for process {accelerator.state.local_process_index}/{accelerator.state.num_processes}")
(
load_stable_diffusion_format,
text_encoder1,
text_encoder2,
vae,
unet,
logit_scale,
ckpt_info,
) = _load_target_model(
args.pretrained_model_name_or_path,
args.vae,
model_version,
weight_dtype,
accelerator.device if args.lowram else "cpu",
model_dtype,
args.disable_mmap_load_safetensors,
)
# work on low-ram device
if args.lowram:
text_encoder1.to(accelerator.device)
text_encoder2.to(accelerator.device)
unet.to(accelerator.device)
vae.to(accelerator.device)
clean_memory_on_device(accelerator.device)
accelerator.wait_for_everyone()
return load_stable_diffusion_format, text_encoder1, text_encoder2, vae, unet, logit_scale, ckpt_info
def _load_target_model(
name_or_path: str, vae_path: Optional[str], model_version: str, weight_dtype, device="cpu", model_dtype=None, disable_mmap=False
):
# model_dtype only work with full fp16/bf16
name_or_path = os.readlink(name_or_path) if os.path.islink(name_or_path) else name_or_path
load_stable_diffusion_format = os.path.isfile(name_or_path) # determine SD or Diffusers
if load_stable_diffusion_format:
logger.info(f"load StableDiffusion checkpoint: {name_or_path}")
(
text_encoder1,
text_encoder2,
vae,
unet,
logit_scale,
ckpt_info,
) = sdxl_model_util.load_models_from_sdxl_checkpoint(model_version, name_or_path, device, model_dtype, disable_mmap)
else:
# Diffusers model is loaded to CPU
from diffusers import StableDiffusionXLPipeline
variant = "fp16" if weight_dtype == torch.float16 else None
logger.info(f"load Diffusers pretrained models: {name_or_path}, variant={variant}")
try:
try:
pipe = StableDiffusionXLPipeline.from_pretrained(
name_or_path, torch_dtype=model_dtype, variant=variant, tokenizer=None
)
except EnvironmentError as ex:
if variant is not None:
logger.info("try to load fp32 model")
pipe = StableDiffusionXLPipeline.from_pretrained(name_or_path, variant=None, tokenizer=None)
else:
raise ex
except EnvironmentError as ex:
logger.error(
f"model is not found as a file or in Hugging Face, perhaps file name is wrong? / 指定したモデル名のファイル、またはHugging Faceのモデルが見つかりません。ファイル名が誤っているかもしれません: {name_or_path}"
)
raise ex
text_encoder1 = pipe.text_encoder
text_encoder2 = pipe.text_encoder_2
# convert to fp32 for cache text_encoders outputs
if text_encoder1.dtype != torch.float32:
text_encoder1 = text_encoder1.to(dtype=torch.float32)
if text_encoder2.dtype != torch.float32:
text_encoder2 = text_encoder2.to(dtype=torch.float32)
vae = pipe.vae
unet = pipe.unet
del pipe
# Diffusers U-Net to original U-Net
state_dict = sdxl_model_util.convert_diffusers_unet_state_dict_to_sdxl(unet.state_dict())
with init_empty_weights():
unet = sdxl_original_unet.SdxlUNet2DConditionModel() # overwrite unet
sdxl_model_util._load_state_dict_on_device(unet, state_dict, device=device, dtype=model_dtype)
logger.info("U-Net converted to original U-Net")
logit_scale = None
ckpt_info = None
# VAEを読み込む
if vae_path is not None:
vae = model_util.load_vae(vae_path, weight_dtype)
logger.info("additional VAE loaded")
return load_stable_diffusion_format, text_encoder1, text_encoder2, vae, unet, logit_scale, ckpt_info
def load_tokenizers(args: argparse.Namespace):
logger.info("prepare tokenizers")
original_paths = [TOKENIZER1_PATH, TOKENIZER2_PATH]
tokeniers = []
for i, original_path in enumerate(original_paths):
tokenizer: CLIPTokenizer = None
if args.tokenizer_cache_dir:
local_tokenizer_path = os.path.join(args.tokenizer_cache_dir, original_path.replace("/", "_"))
if os.path.exists(local_tokenizer_path):
logger.info(f"load tokenizer from cache: {local_tokenizer_path}")
tokenizer = CLIPTokenizer.from_pretrained(local_tokenizer_path)
if tokenizer is None:
tokenizer = CLIPTokenizer.from_pretrained(original_path)
if args.tokenizer_cache_dir and not os.path.exists(local_tokenizer_path):
logger.info(f"save Tokenizer to cache: {local_tokenizer_path}")
tokenizer.save_pretrained(local_tokenizer_path)
if i == 1:
tokenizer.pad_token_id = 0 # fix pad token id to make same as open clip tokenizer
tokeniers.append(tokenizer)
if hasattr(args, "max_token_length") and args.max_token_length is not None:
logger.info(f"update token length: {args.max_token_length}")
return tokeniers
def match_mixed_precision(args, weight_dtype):
if args.full_fp16:
assert (
weight_dtype == torch.float16
), "full_fp16 requires mixed precision='fp16' / full_fp16を使う場合はmixed_precision='fp16'を指定してください。"
return weight_dtype
elif args.full_bf16:
assert (
weight_dtype == torch.bfloat16
), "full_bf16 requires mixed precision='bf16' / full_bf16を使う場合はmixed_precision='bf16'を指定してください。"
return weight_dtype
else:
return None
def timestep_embedding(timesteps, dim, max_period=10000):
"""
Create sinusoidal timestep embeddings.
:param timesteps: a 1-D Tensor of N indices, one per batch element.
These may be fractional.
:param dim: the dimension of the output.
:param max_period: controls the minimum frequency of the embeddings.
:return: an [N x dim] Tensor of positional embeddings.
"""
half = dim // 2
freqs = torch.exp(-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half).to(
device=timesteps.device
)
args = timesteps[:, None].float() * freqs[None]
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
if dim % 2:
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
return embedding
def get_timestep_embedding(x, outdim):
assert len(x.shape) == 2
b, dims = x.shape[0], x.shape[1]
x = torch.flatten(x)
emb = timestep_embedding(x, outdim)
emb = torch.reshape(emb, (b, dims * outdim))
return emb
def get_size_embeddings(orig_size, crop_size, target_size, device):
emb1 = get_timestep_embedding(orig_size, 256)
emb2 = get_timestep_embedding(crop_size, 256)
emb3 = get_timestep_embedding(target_size, 256)
vector = torch.cat([emb1, emb2, emb3], dim=1).to(device)
return vector
def save_sd_model_on_train_end(
args: argparse.Namespace,
src_path: str,
save_stable_diffusion_format: bool,
use_safetensors: bool,
save_dtype: torch.dtype,
epoch: int,
global_step: int,
text_encoder1,
text_encoder2,
unet,
vae,
logit_scale,
ckpt_info,
):
def sd_saver(ckpt_file, epoch_no, global_step):
sai_metadata = train_util.get_sai_model_spec(None, args, True, False, False, is_stable_diffusion_ckpt=True)
sdxl_model_util.save_stable_diffusion_checkpoint(
ckpt_file,
text_encoder1,
text_encoder2,
unet,
epoch_no,
global_step,
ckpt_info,
vae,
logit_scale,
sai_metadata,
save_dtype,
)
def diffusers_saver(out_dir):
sdxl_model_util.save_diffusers_checkpoint(
out_dir,
text_encoder1,
text_encoder2,
unet,
src_path,
vae,
use_safetensors=use_safetensors,
save_dtype=save_dtype,
)
train_util.save_sd_model_on_train_end_common(
args, save_stable_diffusion_format, use_safetensors, epoch, global_step, sd_saver, diffusers_saver
)
# epochとstepの保存、メタデータにepoch/stepが含まれ引数が同じになるため、統合している
# on_epoch_end: Trueならepoch終了時、Falseならstep経過時
def save_sd_model_on_epoch_end_or_stepwise(
args: argparse.Namespace,
on_epoch_end: bool,
accelerator,
src_path,
save_stable_diffusion_format: bool,
use_safetensors: bool,
save_dtype: torch.dtype,
epoch: int,
num_train_epochs: int,
global_step: int,
text_encoder1,
text_encoder2,
unet,
vae,
logit_scale,
ckpt_info,
):
def sd_saver(ckpt_file, epoch_no, global_step):
sai_metadata = train_util.get_sai_model_spec(None, args, True, False, False, is_stable_diffusion_ckpt=True)
sdxl_model_util.save_stable_diffusion_checkpoint(
ckpt_file,
text_encoder1,
text_encoder2,
unet,
epoch_no,
global_step,
ckpt_info,
vae,
logit_scale,
sai_metadata,
save_dtype,
)
def diffusers_saver(out_dir):
sdxl_model_util.save_diffusers_checkpoint(
out_dir,
text_encoder1,
text_encoder2,
unet,
src_path,
vae,
use_safetensors=use_safetensors,
save_dtype=save_dtype,
)
train_util.save_sd_model_on_epoch_end_or_stepwise_common(
args,
on_epoch_end,
accelerator,
save_stable_diffusion_format,
use_safetensors,
epoch,
num_train_epochs,
global_step,
sd_saver,
diffusers_saver,
)
def add_sdxl_training_arguments(parser: argparse.ArgumentParser, support_text_encoder_caching: bool = True):
parser.add_argument(
"--cache_text_encoder_outputs", action="store_true", help="cache text encoder outputs / text encoderの出力をキャッシュする"
)
parser.add_argument(
"--cache_text_encoder_outputs_to_disk",
action="store_true",
help="cache text encoder outputs to disk / text encoderの出力をディスクにキャッシュする",
)
parser.add_argument(
"--disable_mmap_load_safetensors",
action="store_true",
help="disable mmap load for safetensors. Speed up model loading in WSL environment / safetensorsのmmapロードを無効にする。WSL環境等でモデル読み込みを高速化できる",
)
def verify_sdxl_training_args(args: argparse.Namespace, supportTextEncoderCaching: bool = True):
assert not args.v2, "v2 cannot be enabled in SDXL training / SDXL学習ではv2を有効にすることはできません"
if args.v_parameterization:
logger.warning("v_parameterization will be unexpected / SDXL学習ではv_parameterizationは想定外の動作になります")
if args.clip_skip is not None:
logger.warning("clip_skip will be unexpected / SDXL学習ではclip_skipは動作しません")
# if args.multires_noise_iterations:
# logger.info(
# f"Warning: SDXL has been trained with noise_offset={DEFAULT_NOISE_OFFSET}, but noise_offset is disabled due to multires_noise_iterations / SDXLはnoise_offset={DEFAULT_NOISE_OFFSET}で学習されていますが、multires_noise_iterationsが有効になっているためnoise_offsetは無効になります"
# )
# else:
# if args.noise_offset is None:
# args.noise_offset = DEFAULT_NOISE_OFFSET
# elif args.noise_offset != DEFAULT_NOISE_OFFSET:
# logger.info(
# f"Warning: SDXL has been trained with noise_offset={DEFAULT_NOISE_OFFSET} / SDXLはnoise_offset={DEFAULT_NOISE_OFFSET}で学習されています"
# )
# logger.info(f"noise_offset is set to {args.noise_offset} / noise_offsetが{args.noise_offset}に設定されました")
# assert (
# not hasattr(args, "weighted_captions") or not args.weighted_captions
# ), "weighted_captions cannot be enabled in SDXL training currently / SDXL学習では今のところweighted_captionsを有効にすることはできません"
if supportTextEncoderCaching:
if args.cache_text_encoder_outputs_to_disk and not args.cache_text_encoder_outputs:
args.cache_text_encoder_outputs = True
logger.warning(
"cache_text_encoder_outputs is enabled because cache_text_encoder_outputs_to_disk is enabled / "
+ "cache_text_encoder_outputs_to_diskが有効になっているためcache_text_encoder_outputsが有効になりました"
)
def sample_images(*args, **kwargs):
return train_util.sample_images_common(SdxlStableDiffusionLongPromptWeightingPipeline, *args, **kwargs)
+306
View File
@@ -0,0 +1,306 @@
import os
from typing import Any, List, Optional, Tuple, Union
import numpy as np
import torch
from transformers import CLIPTokenizer, CLIPTextModel, CLIPTextModelWithProjection
from .strategy_base import TokenizeStrategy, TextEncodingStrategy, TextEncoderOutputsCachingStrategy
from .utils import setup_logging
setup_logging()
import logging
logger = logging.getLogger(__name__)
TOKENIZER1_PATH = "openai/clip-vit-large-patch14"
TOKENIZER2_PATH = "laion/CLIP-ViT-bigG-14-laion2B-39B-b160k"
class SdxlTokenizeStrategy(TokenizeStrategy):
def __init__(self, max_length: Optional[int], tokenizer_cache_dir: Optional[str] = None) -> None:
self.tokenizer1 = self._load_tokenizer(CLIPTokenizer, TOKENIZER1_PATH, tokenizer_cache_dir=tokenizer_cache_dir)
self.tokenizer2 = self._load_tokenizer(CLIPTokenizer, TOKENIZER2_PATH, tokenizer_cache_dir=tokenizer_cache_dir)
self.tokenizer2.pad_token_id = 0 # use 0 as pad token for tokenizer2
if max_length is None:
self.max_length = self.tokenizer1.model_max_length
else:
self.max_length = max_length + 2
def tokenize(self, text: Union[str, List[str]]) -> List[torch.Tensor]:
text = [text] if isinstance(text, str) else text
return (
torch.stack([self._get_input_ids(self.tokenizer1, t, self.max_length) for t in text], dim=0),
torch.stack([self._get_input_ids(self.tokenizer2, t, self.max_length) for t in text], dim=0),
)
def tokenize_with_weights(self, text: str | List[str]) -> Tuple[List[torch.Tensor]]:
text = [text] if isinstance(text, str) else text
tokens1_list, tokens2_list = [], []
weights1_list, weights2_list = [], []
for t in text:
tokens1, weights1 = self._get_input_ids(self.tokenizer1, t, self.max_length, weighted=True)
tokens2, weights2 = self._get_input_ids(self.tokenizer2, t, self.max_length, weighted=True)
tokens1_list.append(tokens1)
tokens2_list.append(tokens2)
weights1_list.append(weights1)
weights2_list.append(weights2)
return [torch.stack(tokens1_list, dim=0), torch.stack(tokens2_list, dim=0)], [
torch.stack(weights1_list, dim=0),
torch.stack(weights2_list, dim=0),
]
class SdxlTextEncodingStrategy(TextEncodingStrategy):
def __init__(self) -> None:
pass
def _pool_workaround(
self, text_encoder: CLIPTextModelWithProjection, last_hidden_state: torch.Tensor, input_ids: torch.Tensor, eos_token_id: int
):
r"""
workaround for CLIP's pooling bug: it returns the hidden states for the max token id as the pooled output
instead of the hidden states for the EOS token
If we use Textual Inversion, we need to use the hidden states for the EOS token as the pooled output
Original code from CLIP's pooling function:
\# text_embeds.shape = [batch_size, sequence_length, transformer.width]
\# take features from the eot embedding (eot_token is the highest number in each sequence)
\# casting to torch.int for onnx compatibility: argmax doesn't support int64 inputs with opset 14
pooled_output = last_hidden_state[
torch.arange(last_hidden_state.shape[0], device=last_hidden_state.device),
input_ids.to(dtype=torch.int, device=last_hidden_state.device).argmax(dim=-1),
]
"""
# input_ids: b*n,77
# find index for EOS token
# Following code is not working if one of the input_ids has multiple EOS tokens (very odd case)
# eos_token_index = torch.where(input_ids == eos_token_id)[1]
# eos_token_index = eos_token_index.to(device=last_hidden_state.device)
# Create a mask where the EOS tokens are
eos_token_mask = (input_ids == eos_token_id).int()
# Use argmax to find the last index of the EOS token for each element in the batch
eos_token_index = torch.argmax(eos_token_mask, dim=1) # this will be 0 if there is no EOS token, it's fine
eos_token_index = eos_token_index.to(device=last_hidden_state.device)
# get hidden states for EOS token
pooled_output = last_hidden_state[
torch.arange(last_hidden_state.shape[0], device=last_hidden_state.device), eos_token_index
]
# apply projection: projection may be of different dtype than last_hidden_state
pooled_output = text_encoder.text_projection(pooled_output.to(text_encoder.text_projection.weight.dtype))
pooled_output = pooled_output.to(last_hidden_state.dtype)
return pooled_output
def _get_hidden_states_sdxl(
self,
input_ids1: torch.Tensor,
input_ids2: torch.Tensor,
tokenizer1: CLIPTokenizer,
tokenizer2: CLIPTokenizer,
text_encoder1: Union[CLIPTextModel, torch.nn.Module],
text_encoder2: Union[CLIPTextModelWithProjection, torch.nn.Module],
unwrapped_text_encoder2: Optional[CLIPTextModelWithProjection] = None,
):
# input_ids: b,n,77 -> b*n, 77
b_size = input_ids1.size()[0]
if input_ids1.size()[1] == 1:
max_token_length = None
else:
max_token_length = input_ids1.size()[1] * input_ids1.size()[2]
input_ids1 = input_ids1.reshape((-1, tokenizer1.model_max_length)) # batch_size*n, 77
input_ids2 = input_ids2.reshape((-1, tokenizer2.model_max_length)) # batch_size*n, 77
input_ids1 = input_ids1.to(text_encoder1.device)
input_ids2 = input_ids2.to(text_encoder2.device)
# text_encoder1
enc_out = text_encoder1(input_ids1, output_hidden_states=True, return_dict=True)
hidden_states1 = enc_out["hidden_states"][11]
# text_encoder2
enc_out = text_encoder2(input_ids2, output_hidden_states=True, return_dict=True)
hidden_states2 = enc_out["hidden_states"][-2] # penuultimate layer
# pool2 = enc_out["text_embeds"]
unwrapped_text_encoder2 = unwrapped_text_encoder2 or text_encoder2
pool2 = self._pool_workaround(unwrapped_text_encoder2, enc_out["last_hidden_state"], input_ids2, tokenizer2.eos_token_id)
# b*n, 77, 768 or 1280 -> b, n*77, 768 or 1280
n_size = 1 if max_token_length is None else max_token_length // 75
hidden_states1 = hidden_states1.reshape((b_size, -1, hidden_states1.shape[-1]))
hidden_states2 = hidden_states2.reshape((b_size, -1, hidden_states2.shape[-1]))
if max_token_length is not None:
# bs*3, 77, 768 or 1024
# encoder1: <BOS>...<EOS> の三連を <BOS>...<EOS> へ戻す
states_list = [hidden_states1[:, 0].unsqueeze(1)] # <BOS>
for i in range(1, max_token_length, tokenizer1.model_max_length):
states_list.append(hidden_states1[:, i : i + tokenizer1.model_max_length - 2]) # <BOS> の後から <EOS> の前まで
states_list.append(hidden_states1[:, -1].unsqueeze(1)) # <EOS>
hidden_states1 = torch.cat(states_list, dim=1)
# v2: <BOS>...<EOS> <PAD> ... の三連を <BOS>...<EOS> <PAD> ... へ戻す 正直この実装でいいのかわからん
states_list = [hidden_states2[:, 0].unsqueeze(1)] # <BOS>
for i in range(1, max_token_length, tokenizer2.model_max_length):
chunk = hidden_states2[:, i : i + tokenizer2.model_max_length - 2] # <BOS> の後から 最後の前まで
# this causes an error:
# RuntimeError: one of the variables needed for gradient computation has been modified by an inplace operation
# if i > 1:
# for j in range(len(chunk)): # batch_size
# if input_ids2[n_index + j * n_size, 1] == tokenizer2.eos_token_id: # 空、つまり <BOS> <EOS> <PAD> ...のパターン
# chunk[j, 0] = chunk[j, 1] # 次の <PAD> の値をコピーする
states_list.append(chunk) # <BOS> の後から <EOS> の前まで
states_list.append(hidden_states2[:, -1].unsqueeze(1)) # <EOS> か <PAD> のどちらか
hidden_states2 = torch.cat(states_list, dim=1)
# pool はnの最初のものを使う
pool2 = pool2[::n_size]
return hidden_states1, hidden_states2, pool2
def encode_tokens(
self, tokenize_strategy: TokenizeStrategy, models: List[Any], tokens: List[torch.Tensor]
) -> List[torch.Tensor]:
"""
Args:
tokenize_strategy: TokenizeStrategy
models: List of models, [text_encoder1, text_encoder2, unwrapped text_encoder2 (optional)].
If text_encoder2 is wrapped by accelerate, unwrapped_text_encoder2 is required
tokens: List of tokens, for text_encoder1 and text_encoder2
"""
if len(models) == 2:
text_encoder1, text_encoder2 = models
unwrapped_text_encoder2 = None
else:
text_encoder1, text_encoder2, unwrapped_text_encoder2 = models
tokens1, tokens2 = tokens
sdxl_tokenize_strategy = tokenize_strategy # type: SdxlTokenizeStrategy
tokenizer1, tokenizer2 = sdxl_tokenize_strategy.tokenizer1, sdxl_tokenize_strategy.tokenizer2
hidden_states1, hidden_states2, pool2 = self._get_hidden_states_sdxl(
tokens1, tokens2, tokenizer1, tokenizer2, text_encoder1, text_encoder2, unwrapped_text_encoder2
)
return [hidden_states1, hidden_states2, pool2]
def encode_tokens_with_weights(
self,
tokenize_strategy: TokenizeStrategy,
models: List[Any],
tokens_list: List[torch.Tensor],
weights_list: List[torch.Tensor],
) -> List[torch.Tensor]:
hidden_states1, hidden_states2, pool2 = self.encode_tokens(tokenize_strategy, models, tokens_list)
weights_list = [weights.to(hidden_states1.device) for weights in weights_list]
# apply weights
if weights_list[0].shape[1] == 1: # no max_token_length
# weights: ((b, 1, 77), (b, 1, 77)), hidden_states: (b, 77, 768), (b, 77, 768)
hidden_states1 = hidden_states1 * weights_list[0].squeeze(1).unsqueeze(2)
hidden_states2 = hidden_states2 * weights_list[1].squeeze(1).unsqueeze(2)
else:
# weights: ((b, n, 77), (b, n, 77)), hidden_states: (b, n*75+2, 768), (b, n*75+2, 768)
for weight, hidden_states in zip(weights_list, [hidden_states1, hidden_states2]):
for i in range(weight.shape[1]):
hidden_states[:, i * 75 + 1 : i * 75 + 76] = hidden_states[:, i * 75 + 1 : i * 75 + 76] * weight[
:, i, 1:-1
].unsqueeze(-1)
return [hidden_states1, hidden_states2, pool2]
class SdxlTextEncoderOutputsCachingStrategy(TextEncoderOutputsCachingStrategy):
SDXL_TEXT_ENCODER_OUTPUTS_NPZ_SUFFIX = "_te_outputs.npz"
def __init__(
self,
cache_to_disk: bool,
batch_size: int,
skip_disk_cache_validity_check: bool,
is_partial: bool = False,
is_weighted: bool = False,
) -> None:
super().__init__(cache_to_disk, batch_size, skip_disk_cache_validity_check, is_partial, is_weighted)
def get_outputs_npz_path(self, image_abs_path: str) -> str:
return os.path.splitext(image_abs_path)[0] + SdxlTextEncoderOutputsCachingStrategy.SDXL_TEXT_ENCODER_OUTPUTS_NPZ_SUFFIX
def is_disk_cached_outputs_expected(self, npz_path: str):
if not self.cache_to_disk:
return False
if not os.path.exists(npz_path):
return False
if self.skip_disk_cache_validity_check:
return True
try:
npz = np.load(npz_path)
if "hidden_state1" not in npz or "hidden_state2" not in npz or "pool2" not in npz:
return False
except Exception as e:
logger.error(f"Error loading file: {npz_path}")
raise e
return True
def load_outputs_npz(self, npz_path: str) -> List[np.ndarray]:
data = np.load(npz_path)
hidden_state1 = data["hidden_state1"]
hidden_state2 = data["hidden_state2"]
pool2 = data["pool2"]
return [hidden_state1, hidden_state2, pool2]
def cache_batch_outputs(
self, tokenize_strategy: TokenizeStrategy, models: List[Any], text_encoding_strategy: TextEncodingStrategy, infos: List
):
sdxl_text_encoding_strategy = text_encoding_strategy # type: SdxlTextEncodingStrategy
captions = [info.caption for info in infos]
if self.is_weighted:
tokens_list, weights_list = tokenize_strategy.tokenize_with_weights(captions)
with torch.no_grad():
hidden_state1, hidden_state2, pool2 = sdxl_text_encoding_strategy.encode_tokens_with_weights(
tokenize_strategy, models, tokens_list, weights_list
)
else:
tokens1, tokens2 = tokenize_strategy.tokenize(captions)
with torch.no_grad():
hidden_state1, hidden_state2, pool2 = sdxl_text_encoding_strategy.encode_tokens(
tokenize_strategy, models, [tokens1, tokens2]
)
if hidden_state1.dtype == torch.bfloat16:
hidden_state1 = hidden_state1.float()
if hidden_state2.dtype == torch.bfloat16:
hidden_state2 = hidden_state2.float()
if pool2.dtype == torch.bfloat16:
pool2 = pool2.float()
hidden_state1 = hidden_state1.cpu().numpy()
hidden_state2 = hidden_state2.cpu().numpy()
pool2 = pool2.cpu().numpy()
for i, info in enumerate(infos):
hidden_state1_i = hidden_state1[i]
hidden_state2_i = hidden_state2[i]
pool2_i = pool2[i]
if self.cache_to_disk:
np.savez(
info.text_encoder_outputs_npz,
hidden_state1=hidden_state1_i,
hidden_state2=hidden_state2_i,
pool2=pool2_i,
)
else:
info.text_encoder_outputs = [hidden_state1_i, hidden_state2_i, pool2_i]
+54 -75
View File
@@ -4858,7 +4858,7 @@ def get_optimizer(args, trainable_params):
elif optimizer_type.endswith("schedulefree".lower()):
if optimizer_type.lower() == "ProdigyPlusScheduleFree".lower():
from ..prodigyplusschedulefree.prodigy_plus_schedulefree import ProdigyPlusScheduleFree
from prodigyplus.prodigy_plus_schedulefree import ProdigyPlusScheduleFree
optimizer_class = ProdigyPlusScheduleFree
logger.info(f"use ProdigyScheduleFree optimizer | {optimizer_kwargs}")
else:
@@ -6016,33 +6016,18 @@ def sample_images_common(
tokenizer,
text_encoder,
unet,
validation_settings=None,
prompt_replacement=None,
controlnet=None,
):
"""
StableDiffusionLongPromptWeightingPipelineの改造版を使うようにしたので、clip skipおよびプロンプトの重みづけに対応した
TODO Use strategies here
"""
if steps == 0:
if not args.sample_at_first:
return
else:
if args.sample_every_n_steps is None and args.sample_every_n_epochs is None:
return
if args.sample_every_n_epochs is not None:
# sample_every_n_steps は無視する
if epoch is None or epoch % args.sample_every_n_epochs != 0:
return
else:
if steps % args.sample_every_n_steps != 0 or epoch is not None: # steps is not divisible or end of epoch
return
logger.info("")
logger.info(f"generating sample images at step / サンプル画像生成 ステップ: {steps}")
if not os.path.isfile(args.sample_prompts):
logger.error(f"No prompt file / プロンプトファイルがありません: {args.sample_prompts}")
return
logger.info(f"generating sample images at step: {steps}")
distributed_state = PartialState() # for multi gpu distributed inference. this is a singleton, so it's safe to use it here
@@ -6056,18 +6041,25 @@ def sample_images_common(
else:
text_encoder = accelerator.unwrap_model(text_encoder)
# read prompts
if args.sample_prompts.endswith(".txt"):
with open(args.sample_prompts, "r", encoding="utf-8") as f:
lines = f.readlines()
prompts = [line.strip() for line in lines if len(line.strip()) > 0 and line[0] != "#"]
elif args.sample_prompts.endswith(".toml"):
with open(args.sample_prompts, "r", encoding="utf-8") as f:
data = toml.load(f)
prompts = [dict(**data["prompt"], **subset) for subset in data["prompt"]["subset"]]
elif args.sample_prompts.endswith(".json"):
with open(args.sample_prompts, "r", encoding="utf-8") as f:
prompts = json.load(f)
prompts = []
for line in args.sample_prompts:
line = line.strip()
if len(line) > 0 and line[0] != "#":
prompts.append(line)
# preprocess prompts
for i in range(len(prompts)):
prompt_dict = prompts[i]
if isinstance(prompt_dict, str):
from .train_util import line_to_prompt_dict
prompt_dict = line_to_prompt_dict(prompt_dict)
prompts[i] = prompt_dict
assert isinstance(prompt_dict, dict)
# Adds an enumerator to the dict based on prompt position. Used later to name image files. Also cleanup of extra data in original prompt dict.
prompt_dict["enum"] = i
prompt_dict.pop("subset", None)
default_scheduler = get_my_scheduler(sample_sampler=args.sample_sampler, v_parameterization=args.v_parameterization)
@@ -6086,18 +6078,6 @@ def sample_images_common(
save_dir = args.output_dir + "/sample"
os.makedirs(save_dir, exist_ok=True)
# preprocess prompts
for i in range(len(prompts)):
prompt_dict = prompts[i]
if isinstance(prompt_dict, str):
prompt_dict = line_to_prompt_dict(prompt_dict)
prompts[i] = prompt_dict
assert isinstance(prompt_dict, dict)
# Adds an enumerator to the dict based on prompt position. Used later to name image files. Also cleanup of extra data in original prompt dict.
prompt_dict["enum"] = i
prompt_dict.pop("subset", None)
# save random state to restore later
rng_state = torch.get_rng_state()
cuda_rng_state = None
@@ -6106,26 +6086,13 @@ def sample_images_common(
except Exception:
pass
if distributed_state.num_processes <= 1:
# If only one device is available, just use the original prompt list. We don't need to care about the distribution of prompts.
with torch.no_grad():
for prompt_dict in prompts:
sample_image_inference(
accelerator, args, pipeline, save_dir, prompt_dict, epoch, steps, prompt_replacement, controlnet=controlnet
)
else:
# Creating list with N elements, where each element is a list of prompt_dicts, and N is the number of processes available (number of devices available)
# prompt_dicts are assigned to lists based on order of processes, to attempt to time the image creation time to match enum order. Probably only works when steps and sampler are identical.
per_process_prompts = [] # list of lists
for i in range(distributed_state.num_processes):
per_process_prompts.append(prompts[i :: distributed_state.num_processes])
with torch.no_grad():
with distributed_state.split_between_processes(per_process_prompts) as prompt_dict_lists:
for prompt_dict in prompt_dict_lists[0]:
sample_image_inference(
accelerator, args, pipeline, save_dir, prompt_dict, epoch, steps, prompt_replacement, controlnet=controlnet
)
with torch.no_grad():
image_tensor_list = []
for prompt_dict in prompts:
image_tensor = sample_image_inference(
accelerator, args, pipeline, save_dir, prompt_dict, epoch, steps, prompt_replacement, controlnet=controlnet, validation_settings=validation_settings
)
image_tensor_list.append(image_tensor)
# clear pipeline and cache to reduce vram usage
del pipeline
@@ -6136,6 +6103,7 @@ def sample_images_common(
vae.to(org_vae_device)
clean_memory_on_device(accelerator.device)
return torch.cat(image_tensor_list, dim=0)
def sample_image_inference(
@@ -6146,19 +6114,29 @@ def sample_image_inference(
prompt_dict,
epoch,
steps,
prompt_replacement,
prompt_replacement=None,
controlnet=None,
validation_settings=None,
):
assert isinstance(prompt_dict, dict)
negative_prompt = prompt_dict.get("negative_prompt")
sample_steps = prompt_dict.get("sample_steps", 30)
width = prompt_dict.get("width", 512)
height = prompt_dict.get("height", 512)
scale = prompt_dict.get("scale", 7.5)
seed = prompt_dict.get("seed")
controlnet_image = prompt_dict.get("controlnet_image")
if validation_settings is not None:
sample_steps = validation_settings["steps"]
width = validation_settings["width"]
height = validation_settings["height"]
scale = validation_settings["guidance_scale"]
sampler_name = validation_settings["sampler"]
seed = validation_settings["seed"]
controlnet_image=None
else:
sample_steps = prompt_dict.get("sample_steps", 30)
width = prompt_dict.get("width", 512)
height = prompt_dict.get("height", 512)
scale = prompt_dict.get("scale", 7.5)
seed = prompt_dict.get("seed")
controlnet_image = prompt_dict.get("controlnet_image")
sampler_name: str = prompt_dict.get("sample_sampler", args.sample_sampler)
prompt: str = prompt_dict.get("prompt", "")
sampler_name: str = prompt_dict.get("sample_sampler", args.sample_sampler)
negative_prompt = prompt_dict.get("negative_prompt")
if prompt_replacement is not None:
prompt = prompt.replace(prompt_replacement[0], prompt_replacement[1])
@@ -6212,7 +6190,7 @@ def sample_image_inference(
with torch.cuda.device(torch.cuda.current_device()):
torch.cuda.empty_cache()
image = pipeline.latents_to_image(latents)[0]
image, tensors = pipeline.latents_to_image(latents)
# adding accelerator.wait_for_everyone() here should sync up and ensure that sample images are saved in the same order as the original prompt list
# but adding 'enum' to the filename should be enough
@@ -6222,7 +6200,8 @@ def sample_image_inference(
seed_suffix = "" if seed is None else f"_{seed}"
i: int = prompt_dict["enum"]
img_filename = f"{'' if args.output_name is None else args.output_name + '_'}{num_suffix}_{i:02d}_{ts_str}{seed_suffix}.png"
image.save(os.path.join(save_dir, img_filename))
image[0].save(os.path.join(save_dir, img_filename))
# send images to wandb if enabled
if "wandb" in [tracker.name for tracker in accelerator.trackers]:
@@ -6232,7 +6211,7 @@ def sample_image_inference(
# not to commit images to avoid inconsistency between training and logging steps
wandb_tracker.log({f"sample_{i}": wandb.Image(image, caption=prompt)}, commit=False) # positive prompt as a caption
return tensors.permute(0, 2, 3, 1)
# endregion
+201
View File
@@ -0,0 +1,201 @@
Apache License
Version 2.0, January 2004
http://www.apache.org/licenses/
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
1. Definitions.
"License" shall mean the terms and conditions for use, reproduction,
and distribution as defined by Sections 1 through 9 of this document.
"Licensor" shall mean the copyright owner or entity authorized by
the copyright owner that is granting the License.
"Legal Entity" shall mean the union of the acting entity and all
other entities that control, are controlled by, or are under common
control with that entity. For the purposes of this definition,
"control" means (i) the power, direct or indirect, to cause the
direction or management of such entity, whether by contract or
otherwise, or (ii) ownership of fifty percent (50%) or more of the
outstanding shares, or (iii) beneficial ownership of such entity.
"You" (or "Your") shall mean an individual or Legal Entity
exercising permissions granted by this License.
"Source" form shall mean the preferred form for making modifications,
including but not limited to software source code, documentation
source, and configuration files.
"Object" form shall mean any form resulting from mechanical
transformation or translation of a Source form, including but
not limited to compiled object code, generated documentation,
and conversions to other media types.
"Work" shall mean the work of authorship, whether in Source or
Object form, made available under the License, as indicated by a
copyright notice that is included in or attached to the work
(an example is provided in the Appendix below).
"Derivative Works" shall mean any work, whether in Source or Object
form, that is based on (or derived from) the Work and for which the
editorial revisions, annotations, elaborations, or other modifications
represent, as a whole, an original work of authorship. For the purposes
of this License, Derivative Works shall not include works that remain
separable from, or merely link (or bind by name) to the interfaces of,
the Work and Derivative Works thereof.
"Contribution" shall mean any work of authorship, including
the original version of the Work and any modifications or additions
to that Work or Derivative Works thereof, that is intentionally
submitted to Licensor for inclusion in the Work by the copyright owner
or by an individual or Legal Entity authorized to submit on behalf of
the copyright owner. For the purposes of this definition, "submitted"
means any form of electronic, verbal, or written communication sent
to the Licensor or its representatives, including but not limited to
communication on electronic mailing lists, source code control systems,
and issue tracking systems that are managed by, or on behalf of, the
Licensor for the purpose of discussing and improving the Work, but
excluding communication that is conspicuously marked or otherwise
designated in writing by the copyright owner as "Not a Contribution."
"Contributor" shall mean Licensor and any individual or Legal Entity
on behalf of whom a Contribution has been received by Licensor and
subsequently incorporated within the Work.
2. Grant of Copyright License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
copyright license to reproduce, prepare Derivative Works of,
publicly display, publicly perform, sublicense, and distribute the
Work and such Derivative Works in Source or Object form.
3. Grant of Patent License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
(except as stated in this section) patent license to make, have made,
use, offer to sell, sell, import, and otherwise transfer the Work,
where such license applies only to those patent claims licensable
by such Contributor that are necessarily infringed by their
Contribution(s) alone or by combination of their Contribution(s)
with the Work to which such Contribution(s) was submitted. If You
institute patent litigation against any entity (including a
cross-claim or counterclaim in a lawsuit) alleging that the Work
or a Contribution incorporated within the Work constitutes direct
or contributory patent infringement, then any patent licenses
granted to You under this License for that Work shall terminate
as of the date such litigation is filed.
4. Redistribution. You may reproduce and distribute copies of the
Work or Derivative Works thereof in any medium, with or without
modifications, and in Source or Object form, provided that You
meet the following conditions:
(a) You must give any other recipients of the Work or
Derivative Works a copy of this License; and
(b) You must cause any modified files to carry prominent notices
stating that You changed the files; and
(c) You must retain, in the Source form of any Derivative Works
that You distribute, all copyright, patent, trademark, and
attribution notices from the Source form of the Work,
excluding those notices that do not pertain to any part of
the Derivative Works; and
(d) If the Work includes a "NOTICE" text file as part of its
distribution, then any Derivative Works that You distribute must
include a readable copy of the attribution notices contained
within such NOTICE file, excluding those notices that do not
pertain to any part of the Derivative Works, in at least one
of the following places: within a NOTICE text file distributed
as part of the Derivative Works; within the Source form or
documentation, if provided along with the Derivative Works; or,
within a display generated by the Derivative Works, if and
wherever such third-party notices normally appear. The contents
of the NOTICE file are for informational purposes only and
do not modify the License. You may add Your own attribution
notices within Derivative Works that You distribute, alongside
or as an addendum to the NOTICE text from the Work, provided
that such additional attribution notices cannot be construed
as modifying the License.
You may add Your own copyright statement to Your modifications and
may provide additional or different license terms and conditions
for use, reproduction, or distribution of Your modifications, or
for any such Derivative Works as a whole, provided Your use,
reproduction, and distribution of the Work otherwise complies with
the conditions stated in this License.
5. Submission of Contributions. Unless You explicitly state otherwise,
any Contribution intentionally submitted for inclusion in the Work
by You to the Licensor shall be under the terms and conditions of
this License, without any additional terms or conditions.
Notwithstanding the above, nothing herein shall supersede or modify
the terms of any separate license agreement you may have executed
with Licensor regarding such Contributions.
6. Trademarks. This License does not grant permission to use the trade
names, trademarks, service marks, or product names of the Licensor,
except as required for reasonable and customary use in describing the
origin of the Work and reproducing the content of the NOTICE file.
7. Disclaimer of Warranty. Unless required by applicable law or
agreed to in writing, Licensor provides the Work (and each
Contributor provides its Contributions) on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
implied, including, without limitation, any warranties or conditions
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
PARTICULAR PURPOSE. You are solely responsible for determining the
appropriateness of using or redistributing the Work and assume any
risks associated with Your exercise of permissions under this License.
8. Limitation of Liability. In no event and under no legal theory,
whether in tort (including negligence), contract, or otherwise,
unless required by applicable law (such as deliberate and grossly
negligent acts) or agreed to in writing, shall any Contributor be
liable to You for damages, including any direct, indirect, special,
incidental, or consequential damages of any character arising as a
result of this License or out of the use or inability to use the
Work (including but not limited to damages for loss of goodwill,
work stoppage, computer failure or malfunction, or any and all
other commercial damages or losses), even if such Contributor
has been advised of the possibility of such damages.
9. Accepting Warranty or Additional Liability. While redistributing
the Work or Derivative Works thereof, You may choose to offer,
and charge a fee for, acceptance of support, warranty, indemnity,
or other liability obligations and/or rights consistent with this
License. However, in accepting such obligations, You may act only
on Your own behalf and on Your sole responsibility, not on behalf
of any other Contributor, and only if You agree to indemnify,
defend, and hold each Contributor harmless for any liability
incurred by, or claims asserted against, such Contributor by reason
of your accepting any such warranty or additional liability.
END OF TERMS AND CONDITIONS
APPENDIX: How to apply the Apache License to your work.
To apply the Apache License to your work, attach the following
boilerplate notice, with the fields enclosed by brackets "[]"
replaced with your own identifying information. (Don't include
the brackets!) The text should be enclosed in the appropriate
comment syntax for the file format. We also recommend that a
file or class name and description of purpose be included on the
same "printed page" as the copyright notice for easier
identification within third-party archives.
Copyright 2023 KohakuBlueLeaf
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
+28
View File
@@ -0,0 +1,28 @@
#source https://github.com/KohakuBlueleaf/Lycoris
# try:
# from . import kohya
# except Exception:
# pass
# from . import (
# modules,
# utils,
# )
# from .modules.locon import LoConModule
# from .modules.loha import LohaModule
# from .modules.lokr import LokrModule
# from .modules.dylora import DyLoraModule
# from .modules.glora import GLoRAModule
# from .modules.norms import NormModule
# from .modules.full import FullModule
# from .modules.diag_oft import DiagOFTModule
# from .modules import make_module
# from .wrapper import (
# LycorisNetwork,
# create_lycoris,
# create_lycoris_from_weights,
# )
# from .logging import logger
+151
View File
@@ -0,0 +1,151 @@
PRESET = {
"full": {
"enable_conv": True,
"unet_target_module": [
"Transformer2DModel",
"ResnetBlock2D",
"Downsample2D",
"Upsample2D",
"HunYuanDiTBlock", #HunYuanDiT
"DoubleStreamBlock", #Flux
"SingleStreamBlock", #Flux
"SingleDiTBlock", #SD3.5
"MMDoubleStreamBlock", #HunYuanVideo
"MMSingleStreamBlock", #HunYuanVideo
],
"unet_target_name": [
"conv_in",
"conv_out",
"time_embedding.linear_1",
"time_embedding.linear_2",
],
"text_encoder_target_module": [
"CLIPAttention",
"CLIPSdpaAttention",
"CLIPMLP",
"MT5Block",
"BertLayer",
],
"text_encoder_target_name": [],
},
"full-lin": {
"enable_conv": False,
"unet_target_module": [
"Transformer2DModel",
"ResnetBlock2D",
"HunYuanDiTBlock",
"DoubleStreamBlock",
"SingleStreamBlock",
"SingleDiTBlock",
"MMDoubleStreamBlock", #HunYuanVideo
"MMSingleStreamBlock", #HunYuanVideo
],
"unet_target_name": [
"time_embedding.linear_1",
"time_embedding.linear_2",
],
"text_encoder_target_module": [
"CLIPAttention",
"CLIPSdpaAttention",
"CLIPMLP",
"MT5Block",
"BertLayer",
],
"text_encoder_target_name": [],
},
"attn-mlp": {
"enable_conv": False,
"unet_target_module": [
"Transformer2DModel",
"HunYuanDiTBlock",
"DoubleStreamBlock",
"SingleStreamBlock",
"SingleDiTBlock",
"MMDoubleStreamBlock", #HunYuanVideo
"MMSingleStreamBlock", #HunYuanVideo
],
"unet_target_name": [],
"text_encoder_target_module": [
"CLIPAttention",
"CLIPSdpaAttention",
"CLIPMLP",
"MT5Block",
"BertLayer",
],
"text_encoder_target_name": [],
},
"attn-only": {
"enable_conv": False,
"unet_target_module": [
"CrossAttention",
"SelfAttention",
],
"unet_target_name": [],
"text_encoder_target_module": [
"CLIPAttention",
"CLIPSdpaAttention",
"BertAttention",
"MT5LayerSelfAttention",
],
"text_encoder_target_name": [],
},
"unet-only": {
"enable_conv": True,
"unet_target_module": [
"Transformer2DModel",
"ResnetBlock2D",
"Downsample2D",
"Upsample2D",
"HunYuanDiTBlock",
"DoubleStreamBlock",
"SingleStreamBlock",
"SingleDiTBlock",
"MMDoubleStreamBlock", #HunYuanVideo
"MMSingleStreamBlock", #HunYuanVideo
],
"unet_target_name": [
"conv_in",
"conv_out",
"time_embedding.linear_1",
"time_embedding.linear_2",
],
"text_encoder_target_module": [],
"text_encoder_target_name": [],
},
"unet-transformer-only": {
"enable_conv": False,
"unet_target_module": [
"Transformer2DModel",
"HunYuanDiTBlock",
"DoubleStreamBlock",
"SingleStreamBlock",
"SingleDiTBlock",
"MMDoubleStreamBlock", #HunYuanVideo
"MMSingleStreamBlock", #HunYuanVideo
],
"unet_target_name": [],
"text_encoder_target_module": [],
"text_encoder_target_name": [],
},
"unet-convblock-only": {
"enable_conv": True,
"unet_target_module": ["ResnetBlock2D", "Downsample2D", "Upsample2D"],
"unet_target_name": [
"conv_in",
"conv_out",
],
"text_encoder_target_module": [],
"text_encoder_target_name": [],
},
"ia3": {
"enable_conv": False,
"unet_target_module": [],
"unet_target_name": ["to_k", "to_v", "ff.net.2"],
"text_encoder_target_module": [],
"text_encoder_target_name": ["k_proj", "v_proj", "mlp.fc2"],
"name_algo_map": {
"mlp.fc2": {"train_on_input": True},
"ff.net.2": {"train_on_input": True},
},
},
}
+9
View File
@@ -0,0 +1,9 @@
from .general import (
rebuild_tucker,
factorization,
power2factorization,
FUNC_LIST,
tucker_weight,
tucker_weight_from_conv,
apply_dora_scale,
)
+122
View File
@@ -0,0 +1,122 @@
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from .general import power2factorization, FUNC_LIST
from .diag_oft import get_r
def weight_gen(org_weight, max_block_size, boft_m=-1, rescale=False):
"""### boft_weight_gen
Args:
org_weight (torch.Tensor): the weight tensor
max_block_size (int): max block size
rescale (bool, optional): whether to rescale the weight. Defaults to False.
Returns:
torch.Tensor: oft_blocks[, rescale_weight]
"""
out_dim, *rest = org_weight.shape
block_size, block_num = power2factorization(out_dim, max_block_size)
max_boft_m = sum(int(i) for i in f"{block_num-1:b}") + 1
if boft_m == -1:
boft_m = max_boft_m
boft_m = min(boft_m, max_boft_m)
oft_blocks = torch.zeros(boft_m, block_num, block_size, block_size)
if rescale is not None:
return oft_blocks, torch.ones(out_dim, *[1] * len(rest))
else:
return oft_blocks, None
def diff_weight(org_weight, *weights, constraint=None):
"""### boft_diff_weight
Args:
org_weight (torch.Tensor): the weight tensor of original model
weights (tuple[torch.Tensor]): (oft_blocks[, rescale_weight])
constraint (float, optional): constraint for oft
Returns:
torch.Tensor: ΔW
"""
oft_blocks, rescale = weights
m, num, b, _ = oft_blocks.shape
r_b = b // 2
I = torch.eye(b, device=oft_blocks.device)
r = get_r(oft_blocks, I, constraint)
inp = org = org_weight.to(dtype=r.dtype)
for i in range(m):
bi = r[i] # b_num, b_size, b_size
g = 2
k = 2**i * r_b
inp = (
inp.unflatten(-1, (-1, g, k))
.transpose(-2, -1)
.flatten(-3)
.unflatten(-1, (-1, b))
)
inp = torch.einsum("b i j, b j ... -> b i ...", bi, inp)
inp = inp.flatten(-2).unflatten(-1, (-1, k, g)).transpose(-2, -1).flatten(-3)
if rescale is not None:
inp = inp * rescale
return inp - org
def bypass_forward_diff(org_out, *weights, constraint=None, need_transpose=False):
"""### boft_bypass_forward_diff
Args:
x (torch.Tensor): the input tensor for original model
org_out (torch.Tensor): the output tensor from original model
weights (tuple[torch.Tensor]): (oft_blocks[, rescale_weight])
constraint (float, optional): constraint for oft
need_transpose (bool, optional):
whether to transpose the input and output,
set to `True` if the original model have "dim" not in the last axis.
For example: Convolution layers
Returns:
torch.Tensor: output tensor
"""
oft_blocks, rescale = weights
m, num, b, _ = oft_blocks.shape
r_b = b // 2
I = torch.eye(b, device=oft_blocks.device)
r = get_r(oft_blocks, I, constraint)
inp = org = org_out.to(dtype=r.dtype)
if need_transpose:
inp = org = inp.transpose(1, -1)
for i in range(m):
bi = r[i] # b_num, b_size, b_size
g = 2
k = 2**i * r_b
# ... (c g k) ->... (c k g)
# ... (d b) -> ... d b
inp = (
inp.unflatten(-1, (-1, g, k))
.transpose(-2, -1)
.flatten(-3)
.unflatten(-1, (-1, b))
)
inp = torch.einsum("b i j, ... b j -> ... b i", bi, inp)
# ... d b -> ... (d b)
# ... (c k g) -> ... (c g k)
inp = inp.flatten(-2).unflatten(-1, (-1, k, g)).transpose(-2, -1).flatten(-3)
if rescale is not None:
inp = inp * rescale.transpose(0, -1)
inp = inp - org
if need_transpose:
inp = inp.transpose(1, -1)
return inp
+112
View File
@@ -0,0 +1,112 @@
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from .general import factorization, FUNC_LIST
def get_r(oft_blocks, I=None, constraint=0):
if I is None:
I = torch.eye(oft_blocks.shape[-1], device=oft_blocks.device)
if I.ndim < oft_blocks.ndim:
for _ in range(oft_blocks.ndim - I.ndim):
I = I.unsqueeze(0)
# for Q = -Q^T
q = oft_blocks - oft_blocks.transpose(-1, -2)
normed_q = q
if constraint is not None and constraint > 0:
q_norm = torch.norm(q) + 1e-8
if q_norm > constraint:
normed_q = q * constraint / q_norm
# use float() to prevent unsupported type
r = (I + normed_q) @ (I - normed_q).float().inverse()
return r
def weight_gen(org_weight, max_block_size=-1, rescale=False):
"""### weight_gen
Args:
org_weight (torch.Tensor): the weight tensor
max_block_size (int): max block size
rescale (bool, optional): whether to rescale the weight. Defaults to False.
Returns:
torch.Tensor: oft_blocks[, rescale_weight]
"""
out_dim, *rest = org_weight.shape
block_size, block_num = factorization(out_dim, max_block_size)
oft_blocks = torch.zeros(block_num, block_size, block_size)
if rescale:
return oft_blocks, torch.ones(out_dim, *[1] * len(rest))
else:
return oft_blocks, None
def diff_weight(org_weight, *weights, constraint=None):
"""### diff_weight
Args:
org_weight (torch.Tensor): the weight tensor of original model
weights (tuple[torch.Tensor]): (oft_blocks[, rescale_weight])
constraint (float, optional): constraint for oft
Returns:
torch.Tensor: ΔW
"""
oft_blocks, rescale = weights
I = torch.eye(oft_blocks.shape[1], device=oft_blocks.device)
r = get_r(oft_blocks, I, constraint)
block_num, block_size, _ = oft_blocks.shape
_, *shape = org_weight.shape
org_weight = org_weight.to(dtype=r.dtype)
org_weight = org_weight.view(block_num, block_size, *shape)
# Init R=0, so add I on it to ensure the output of step0 is original model output
weight = torch.einsum(
"k n m, k n ... -> k m ...",
r - I,
org_weight,
).view(-1, *shape)
if rescale is not None:
weight = rescale * weight
weight = weight + (rescale - 1) * org_weight
return weight
def bypass_forward_diff(x, org_out, *weights, constraint=None, need_transpose=False):
"""### bypass_forward_diff
Args:
x (torch.Tensor): the input tensor for original model
org_out (torch.Tensor): the output tensor from original model
weights (tuple[torch.Tensor]): (oft_blocks[, rescale_weight])
constraint (float, optional): constraint for oft
need_transpose (bool, optional):
whether to transpose the input and output,
set to `True` if the original model have "dim" not in the last axis.
For example: Convolution layers
Returns:
torch.Tensor: output tensor
"""
oft_blocks, rescale = weights
block_num, block_size, _ = oft_blocks.shape
I = torch.eye(block_size, device=oft_blocks.device)
r = get_r(oft_blocks, I, constraint)
if need_transpose:
org_out = org_out.transpose(1, -1)
org_out = org_out.to(dtype=r.dtype)
*shape, _ = org_out.shape
oft_out = torch.einsum(
"k n m, ... k n -> ... k m", r - I, org_out.view(*shape, block_num, block_size)
)
out = oft_out.view(*shape, -1)
if rescale is not None:
out = rescale.transpose(-1, 0) * out
out = out + (rescale - 1).transpose(-1, 0) * org_out
if need_transpose:
out = out.transpose(1, -1)
return out
+108
View File
@@ -0,0 +1,108 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
FUNC_LIST = [None, None, F.linear, F.conv1d, F.conv2d, F.conv3d]
def rebuild_tucker(t, wa, wb):
rebuild2 = torch.einsum("i j ..., i p, j r -> p r ...", t, wa, wb)
return rebuild2
def factorization(dimension: int, factor: int = -1) -> tuple[int, int]:
"""
return a tuple of two value of input dimension decomposed by the number closest to factor
second value is higher or equal than first value.
In LoRA with Kroneckor Product, first value is a value for weight scale.
second value is a value for weight.
Because of non-commutative property, A⊗B ≠ B⊗A. Meaning of two matrices is slightly different.
examples)
factor
-1 2 4 8 16 ...
127 -> 1, 127 127 -> 1, 127 127 -> 1, 127 127 -> 1, 127 127 -> 1, 127
128 -> 8, 16 128 -> 2, 64 128 -> 4, 32 128 -> 8, 16 128 -> 8, 16
250 -> 10, 25 250 -> 2, 125 250 -> 2, 125 250 -> 5, 50 250 -> 10, 25
360 -> 8, 45 360 -> 2, 180 360 -> 4, 90 360 -> 8, 45 360 -> 12, 30
512 -> 16, 32 512 -> 2, 256 512 -> 4, 128 512 -> 8, 64 512 -> 16, 32
1024 -> 32, 32 1024 -> 2, 512 1024 -> 4, 256 1024 -> 8, 128 1024 -> 16, 64
"""
if factor > 0 and (dimension % factor) == 0:
m = factor
n = dimension // factor
if m > n:
n, m = m, n
return m, n
if factor < 0:
factor = dimension
m, n = 1, dimension
length = m + n
while m < n:
new_m = m + 1
while dimension % new_m != 0:
new_m += 1
new_n = dimension // new_m
if new_m + new_n > length or new_m > factor:
break
else:
m, n = new_m, new_n
if m > n:
n, m = m, n
return m, n
def power2factorization(dimension: int, factor: int = -1) -> tuple[int, int]:
"""
m = 2k
n = 2**p
m*n = dim
"""
if factor == -1:
factor = dimension
# Find the first solution and check if it is even doable
m = n = 0
while m <= factor:
m += 2
while dimension % m != 0 and m < dimension:
m += 2
if m > factor:
break
if sum(int(i) for i in f"{dimension//m:b}") == 1:
n = dimension // m
if n == 0:
return None, n
return dimension // n, n
def tucker_weight_from_conv(up, down, mid):
up = up.reshape(up.size(0), up.size(1))
down = down.reshape(down.size(0), down.size(1))
return torch.einsum("m n ..., i m, n j -> i j ...", mid, up, down)
def tucker_weight(wa, wb, t):
temp = torch.einsum("i j ..., j r -> i r ...", t, wb)
return torch.einsum("i j ..., i r -> r j ...", temp, wa)
def apply_dora_scale(org_weight, rebuild, dora_scale, scale):
dora_norm_dims = org_weight.dim() - 1
weight = org_weight + rebuild
weight = weight.to(dora_scale.dtype)
weight_norm = (
weight.transpose(0, 1)
.reshape(weight.shape[1], -1)
.norm(dim=1, keepdim=True)
.reshape(weight.shape[1], *[1] * dora_norm_dims)
.transpose(0, 1)
)
merged_scale1 = weight / weight_norm * dora_scale
diff_weight = merged_scale1 - org_weight
return org_weight + diff_weight * scale
+85
View File
@@ -0,0 +1,85 @@
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from .general import rebuild_tucker, FUNC_LIST
def weight_gen(org_weight, rank, tucker=True):
"""### weight_gen
Args:
org_weight (torch.Tensor): the weight tensor
rank (int): low rank
Returns:
torch.Tensor: down, up[, mid]
"""
out_dim, in_dim, *k = org_weight.shape
if k and tucker:
down = torch.empty(rank, in_dim, *(1 for _ in k))
up = torch.empty(out_dim, rank, *(1 for _ in k))
mid = torch.empty(rank, rank, *k)
nn.init.kaiming_uniform_(down, a=math.sqrt(5))
nn.init.constant_(up, 0)
nn.init.kaiming_uniform_(mid, a=math.sqrt(5))
return down, up, mid
else:
down = torch.empty(rank, in_dim)
up = torch.empty(out_dim, rank)
nn.init.kaiming_uniform_(down, a=math.sqrt(5))
nn.init.constant_(up, 0)
return down, up, None
def diff_weight(*weights: tuple[torch.Tensor], gamma=1.0):
"""### diff_weight
Get ΔW = BA, where BA is low rank decomposition
Args:
weights (tuple[torch.Tensor]): (down, up[, mid])
gamma (float, optional): scale factor, normally alpha/rank here
Returns:
torch.Tensor: ΔW
"""
d, u, m = weights
R, I, *k = d.shape
O, R, *_ = u.shape
u = u * gamma
if m is None:
result = u.reshape(-1, u.size(1)) @ d.reshape(d.size(0), -1)
else:
R, R, *k = m.shape
u = u.reshape(u.size(0), -1).transpose(0, 1)
d = d.reshape(d.size(0), -1)
result = rebuild_tucker(m, u, d)
return result.reshape(O, I, *k)
def bypass_forward_diff(x, org_out, *weights, gamma=1.0, extra_args={}):
"""### bypass_forward_diff
Args:
x (torch.Tensor): input tensor
weights (tuple[torch.Tensor]): (down, up[, mid])
gamma (float, optional): scale factor, normally alpha/rank here
extra_args (dict, optional): extra args for forward func, \
e.g. padding, stride for Conv1/2/3d
Returns:
torch.Tensor: output tensor
"""
d, u, m = weights
if m is not None:
down = FUNC_LIST[d.dim()](x, d)
mid = FUNC_LIST[d.dim()](down, m, **extra_args)
up = FUNC_LIST[d.dim()](mid, u)
else:
down = FUNC_LIST[d.dim()](x, d, **extra_args)
up = FUNC_LIST[d.dim()](down, u)
return up * gamma
+165
View File
@@ -0,0 +1,165 @@
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from .general import FUNC_LIST
class HadaWeight(torch.autograd.Function):
@staticmethod
def forward(ctx, w1d, w1u, w2d, w2u, scale=torch.tensor(1)):
ctx.save_for_backward(w1d, w1u, w2d, w2u, scale)
diff_weight = ((w1u @ w1d) * (w2u @ w2d)) * scale
return diff_weight
@staticmethod
def backward(ctx, grad_out):
(w1d, w1u, w2d, w2u, scale) = ctx.saved_tensors
grad_out = grad_out * scale
temp = grad_out * (w2u @ w2d)
grad_w1u = temp @ w1d.T
grad_w1d = w1u.T @ temp
temp = grad_out * (w1u @ w1d)
grad_w2u = temp @ w2d.T
grad_w2d = w2u.T @ temp
del temp
return grad_w1d, grad_w1u, grad_w2d, grad_w2u, None
class HadaWeightTucker(torch.autograd.Function):
@staticmethod
def forward(ctx, t1, w1d, w1u, t2, w2d, w2u, scale=torch.tensor(1)):
ctx.save_for_backward(t1, w1d, w1u, t2, w2d, w2u, scale)
rebuild1 = torch.einsum("i j ..., j r, i p -> p r ...", t1, w1d, w1u)
rebuild2 = torch.einsum("i j ..., j r, i p -> p r ...", t2, w2d, w2u)
return rebuild1 * rebuild2 * scale
@staticmethod
def backward(ctx, grad_out):
(t1, w1d, w1u, t2, w2d, w2u, scale) = ctx.saved_tensors
grad_out = grad_out * scale
temp = torch.einsum("i j ..., j r -> i r ...", t2, w2d)
rebuild = torch.einsum("i j ..., i r -> r j ...", temp, w2u)
grad_w = rebuild * grad_out
del rebuild
grad_w1u = torch.einsum("r j ..., i j ... -> r i", temp, grad_w)
grad_temp = torch.einsum("i j ..., i r -> r j ...", grad_w, w1u.T)
del grad_w, temp
grad_w1d = torch.einsum("i r ..., i j ... -> r j", t1, grad_temp)
grad_t1 = torch.einsum("i j ..., j r -> i r ...", grad_temp, w1d.T)
del grad_temp
temp = torch.einsum("i j ..., j r -> i r ...", t1, w1d)
rebuild = torch.einsum("i j ..., i r -> r j ...", temp, w1u)
grad_w = rebuild * grad_out
del rebuild
grad_w2u = torch.einsum("r j ..., i j ... -> r i", temp, grad_w)
grad_temp = torch.einsum("i j ..., i r -> r j ...", grad_w, w2u.T)
del grad_w, temp
grad_w2d = torch.einsum("i r ..., i j ... -> r j", t2, grad_temp)
grad_t2 = torch.einsum("i j ..., j r -> i r ...", grad_temp, w2d.T)
del grad_temp
return grad_t1, grad_w1d, grad_w1u, grad_t2, grad_w2d, grad_w2u, None
def make_weight(w1d, w1u, w2d, w2u, scale):
return HadaWeight.apply(w1d, w1u, w2d, w2u, scale)
def make_weight_tucker(t1, w1d, w1u, t2, w2d, w2u, scale):
return HadaWeightTucker.apply(t1, w1d, w1u, t2, w2d, w2u, scale)
def weight_gen(org_weight, rank, tucker=True):
"""### weight_gen
Args:
org_weight (torch.Tensor): the weight tensor
rank (int): low rank
Returns:
torch.Tensor: w1d, w2d, w1u, w2u[, t1, t2]
"""
out_dim, in_dim, *k = org_weight.shape
if k and tucker:
w1d = torch.empty(rank, in_dim)
w1u = torch.empty(rank, out_dim)
t1 = torch.empty(rank, rank, *k)
w2d = torch.empty(rank, in_dim)
w2u = torch.empty(rank, out_dim)
t2 = torch.empty(rank, rank, *k)
nn.init.normal_(t1, std=0.1)
nn.init.normal_(t2, std=0.1)
else:
w1d = torch.empty(rank, in_dim)
w1u = torch.empty(out_dim, rank)
w2d = torch.empty(rank, in_dim)
w2u = torch.empty(out_dim, rank)
t1 = t2 = None
nn.init.normal_(w1d, std=1)
nn.init.constant_(w1u, 0)
nn.init.normal_(w2d, std=1)
nn.init.normal_(w2u, std=0.1)
return w1d, w1u, w2d, w2u, t1, t2
def diff_weight(*weights, gamma=1.0):
"""### diff_weight
Get ΔW = BA, where BA is low rank decomposition
Args:
wegihts (tuple[torch.Tensor]): (w1d, w2d, w1u, w2u[, t1, t2])
gamma (float, optional): scale factor, normally alpha/rank here
Returns:
torch.Tensor: ΔW
"""
w1d, w1u, w2d, w2u, t1, t2 = weights
if t1 is not None and t2 is not None:
R, I = w1d.shape
R, O = w1u.shape
R, R, *k = t1.shape
result = make_weight_tucker(t1, w1d, w1u, t2, w2d, w2u, gamma)
else:
R, I, *k = w1d.shape
O, R, *_ = w1u.shape
w1d = w1d.reshape(w1d.size(0), -1)
w1u = w1u.reshape(-1, w1u.size(1))
w2d = w2d.reshape(w2d.size(0), -1)
w2u = w2u.reshape(-1, w2u.size(1))
result = make_weight(w1d, w1u, w2d, w2u, gamma)
result = result.reshape(O, I, *k)
return result
def bypass_forward_diff(x, org_out, *weights, gamma=1.0, extra_args={}):
"""### bypass_forward_diff
Args:
x (torch.Tensor): input tensor
weights (tuple[torch.Tensor]): (w1d, w2d, w1u, w2u[, t1, t2])
gamma (float, optional): scale factor, normally alpha/rank here
extra_args (dict, optional): extra args for forward func, \
e.g. padding, stride for Conv1/2/3d
Returns:
torch.Tensor: output tensor
"""
w1d, w1u, w2d, w2u, t1, t2 = weights
diff_w = diff_weight(w1d, w1u, w2d, w2u, t1, t2, gamma)
return FUNC_LIST[w1d.dim() if t1 is None else t1.dim()](x, diff_w, **extra_args)
+247
View File
@@ -0,0 +1,247 @@
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from .general import rebuild_tucker, FUNC_LIST
from .general import factorization
def make_kron(w1, w2, scale):
for _ in range(w2.dim() - w1.dim()):
w1 = w1.unsqueeze(-1)
w2 = w2.contiguous()
rebuild = torch.kron(w1, w2)
if scale != 1:
rebuild = rebuild * scale
return rebuild
def weight_gen(
org_weight,
rank,
tucker=True,
factor=-1,
decompose_both=False,
full_matrix=False,
unbalanced_factorization=False,
):
"""### weight_gen
Args:
org_weight (torch.Tensor): the weight tensor
rank (int): low rank
Returns:
torch.Tensor | None: w1, w1a, w1b, w2, w2a, w2b, t2
"""
out_dim, in_dim, *k = org_weight.shape
w1 = w1a = w1b = None
w2 = w2a = w2b = None
t2 = None
use_w1 = use_w2 = False
if k:
k_size = k
shape = (out_dim, in_dim, *k_size)
in_m, in_n = factorization(in_dim, factor)
out_l, out_k = factorization(out_dim, factor)
if unbalanced_factorization:
out_l, out_k = out_k, out_l
shape = ((out_l, out_k), (in_m, in_n), *k_size) # ((a, b), (c, d), *k_size)
tucker = tucker and any(i != 1 for i in k_size)
if (
decompose_both
and rank < max(shape[0][0], shape[1][0]) / 2
and not full_matrix
):
w1a = torch.empty(shape[0][0], rank)
w1b = torch.empty(rank, shape[1][0])
else:
use_w1 = True
w1 = torch.empty(shape[0][0], shape[1][0]) # a*c, 1-mode
if rank >= max(shape[0][1], shape[1][1]) / 2 or full_matrix:
use_w2 = True
w2 = torch.empty(shape[0][1], shape[1][1], *k_size)
elif tucker:
t2 = torch.empty(rank, rank, *shape[2:])
w2a = torch.empty(rank, shape[0][1]) # b, 1-mode
w2b = torch.empty(rank, shape[1][1]) # d, 2-mode
else: # Conv2d not tucker
# bigger part. weight and LoRA. [b, dim] x [dim, d*k1*k2]
w2a = torch.empty(shape[0][1], rank)
w2b = torch.empty(rank, shape[1][1], *shape[2:])
# w1 ⊗ (w2a x w2b) = (a, b)⊗((c, dim)x(dim, d*k1*k2)) = (a, b)⊗(c, d*k1*k2) = (ac, bd*k1*k2)
else: # Linear
shape = (out_dim, in_dim)
in_m, in_n = factorization(in_dim, factor)
out_l, out_k = factorization(out_dim, factor)
if unbalanced_factorization:
out_l, out_k = out_k, out_l
shape = (
(out_l, out_k),
(in_m, in_n),
) # ((a, b), (c, d)), out_dim = a*c, in_dim = b*d
# smaller part. weight scale
if decompose_both and rank < max(shape[0][0], shape[1][0]) / 2:
w1a = torch.empty(shape[0][0], rank)
w1b = torch.empty(rank, shape[1][0])
else:
use_w1 = True
w1 = torch.empty(shape[0][0], shape[1][0]) # a*c, 1-mode
if rank < max(shape[0][1], shape[1][1]) / 2:
# bigger part. weight and LoRA. [b, dim] x [dim, d]
w2a = torch.empty(shape[0][1], rank)
w2b = torch.empty(rank, shape[1][1])
# w1 ⊗ (w2a x w2b) = (a, b)⊗((c, dim)x(dim, d)) = (a, b)⊗(c, d) = (ac, bd)
else:
use_w2 = True
w2 = torch.empty(shape[0][1], shape[1][1])
if use_w2:
torch.nn.init.constant_(w2, 1)
else:
if tucker:
torch.nn.init.kaiming_uniform_(t2, a=math.sqrt(5))
torch.nn.init.kaiming_uniform_(w2a, a=math.sqrt(5))
torch.nn.init.constant_(w2b, 1)
if use_w1:
torch.nn.init.kaiming_uniform_(w1, a=math.sqrt(5))
else:
torch.nn.init.kaiming_uniform_(w1a, a=math.sqrt(5))
torch.nn.init.kaiming_uniform_(w1b, a=math.sqrt(5))
return w1, w1a, w1b, w2, w2a, w2b, t2
def diff_weight(*weights, gamma=1.0):
"""### diff_weight
Args:
weights (tuple[torch.Tensor]): (w1, w1a, w1b, w2, w2a, w2b, t)
gamma (float, optional): scale factor, normally alpha/rank here
Returns:
torch.Tensor: ΔW
"""
w1, w1a, w1b, w2, w2a, w2b, t = weights
if w1a is not None:
rank = w1a.shape[1]
elif w2a is not None:
rank = w2a.shape[1]
else:
rank = gamma
scale = gamma / rank
if w1 is None:
w1 = w1a @ w1b
if w2 is None:
if t is None:
r, o, *k = w2b.shape
w2 = w2a @ w2b.view(r, -1)
w2 = w2.view(-1, o, *k)
else:
w2 = rebuild_tucker(t, w2a, w2b)
return make_kron(w1, w2, scale)
def bypass_forward_diff(h, org_out, *weights, gamma=1.0, extra_args={}):
"""### bypass_forward_diff
Args:
weights (tuple[torch.Tensor]): (w1, w1a, w1b, w2, w2a, w2b, t)
gamma (float, optional): scale factor, normally alpha/rank here
extra_args (dict, optional): extra args for forward func, \
e.g. padding, stride for Conv1/2/3d
Returns:
torch.Tensor: output tensor
"""
w1, w1a, w1b, w2, w2a, w2b, t = weights
use_w1 = w1 is not None
use_w2 = w2 is not None
tucker = t is not None
dim = t.dim() if tucker else w2.dim() if w2 is not None else w2b.dim()
rank = w1b.size(0) if not use_w1 else w2b.size(0) if not use_w2 else gamma
scale = gamma / rank
is_conv = dim > 2
op = FUNC_LIST[dim]
if is_conv:
kw_dict = extra_args
else:
kw_dict = {}
if use_w2:
ba = w2
else:
a = w2b
b = w2a
if t is not None:
a = a.view(*a.shape, *[1] * (dim - 2))
b = b.view(*b.shape, *[1] * (dim - 2))
elif is_conv:
b = b.view(*b.shape, *[1] * (dim - 2))
if use_w1:
c = w1
else:
c = w1a @ w1b
uq = c.size(1)
if is_conv:
# (b, uq), vq, ...
B, _, *rest = h.shape
h_in_group = h.reshape(B * uq, -1, *rest)
else:
# b, ..., uq, vq
h_in_group = h.reshape(*h.shape[:-1], uq, -1)
if use_w2:
hb = op(h_in_group, ba, **kw_dict)
else:
if is_conv:
if tucker:
ha = op(h_in_group, a)
ht = op(ha, t, **kw_dict)
hb = op(ht, b)
else:
ha = op(h_in_group, a, **kw_dict)
hb = op(ha, b)
else:
ha = op(h_in_group, a, **kw_dict)
hb = op(ha, b)
if is_conv:
# (b, uq), vp, ..., f
# -> b, uq, vp, ..., f
# -> b, f, vp, ..., uq
hb = hb.view(B, -1, *hb.shape[1:])
h_cross_group = hb.transpose(1, -1)
else:
# b, ..., uq, vq
# -> b, ..., vq, uq
h_cross_group = hb.transpose(-1, -2)
hc = F.linear(h_cross_group, c)
if is_conv:
# b, f, vp, ..., up
# -> b, up, vp, ... ,f
# -> b, c, ..., f
hc = hc.transpose(1, -1)
h = hc.reshape(B, -1, *hc.shape[3:])
else:
# b, ..., vp, up
# -> b, ..., up, vp
# -> b, ..., c
hc = hc.transpose(-1, -2)
h = hc.reshape(*hc.shape[:-2], -1)
return h * scale
+676
View File
@@ -0,0 +1,676 @@
import os
import fnmatch
import re
import logging
from typing import Any, List
import torch
from .utils import precalculate_safetensors_hashes
from .wrapper import LycorisNetwork, network_module_dict, deprecated_arg_dict
from .modules.locon import LoConModule
from .modules.loha import LohaModule
from .modules.ia3 import IA3Module
from .modules.lokr import LokrModule
from .modules.dylora import DyLoraModule
from .modules.glora import GLoRAModule
from .modules.norms import NormModule
from .modules.full import FullModule
from .modules.diag_oft import DiagOFTModule
from .modules.boft import ButterflyOFTModule
from .modules import make_module, get_module
from .config import PRESET
from .utils.preset import read_preset
from .utils import str_bool
from .logging import logger
def create_network(
multiplier, network_dim, network_alpha, vae, text_encoder, unet, **kwargs
):
for key, value in list(kwargs.items()):
if key in deprecated_arg_dict:
logger.warning(
f"{key} is deprecated. Please use {deprecated_arg_dict[key]} instead.",
stacklevel=2,
)
kwargs[deprecated_arg_dict[key]] = value
if network_dim is None:
network_dim = 4 # default
conv_dim = int(kwargs.get("conv_dim", network_dim) or network_dim)
conv_alpha = float(kwargs.get("conv_alpha", network_alpha) or network_alpha)
dropout = float(kwargs.get("dropout", 0.0) or 0.0)
rank_dropout = float(kwargs.get("rank_dropout", 0.0) or 0.0)
module_dropout = float(kwargs.get("module_dropout", 0.0) or 0.0)
algo = (kwargs.get("algo", "lora") or "lora").lower()
use_tucker = str_bool(
not kwargs.get("disable_conv_cp", True)
or kwargs.get("use_conv_cp", False)
or kwargs.get("use_cp", False)
or kwargs.get("use_tucker", False)
)
use_scalar = str_bool(kwargs.get("use_scalar", False))
block_size = int(kwargs.get("block_size", None) or 4)
train_norm = str_bool(kwargs.get("train_norm", False))
constraint = float(kwargs.get("constraint", None) or 0)
rescaled = str_bool(kwargs.get("rescaled", False))
weight_decompose = str_bool(kwargs.get("dora_wd", False))
wd_on_output = str_bool(kwargs.get("wd_on_output", False))
full_matrix = str_bool(kwargs.get("full_matrix", False))
bypass_mode = str_bool(kwargs.get("bypass_mode", None))
rs_lora = str_bool(kwargs.get("rs_lora", False))
unbalanced_factorization = str_bool(kwargs.get("unbalanced_factorization", False))
train_t5xxl = str_bool(kwargs.get("train_t5xxl", False))
if unbalanced_factorization:
logger.info("Unbalanced factorization for LoKr is enabled")
if bypass_mode:
logger.info("Bypass mode is enabled")
if weight_decompose:
logger.info("Weight decomposition is enabled")
if full_matrix:
logger.info("Full matrix mode for LoKr is enabled")
preset_str = kwargs.get("preset", "full")
if preset_str not in PRESET:
preset = read_preset(preset_str)
else:
preset = PRESET[preset_str]
assert preset is not None
LycorisNetworkKohya.apply_preset(preset)
logger.info(f"Using rank adaptation algo: {algo}")
if algo == "ia3" and preset_str != "ia3":
logger.warning("It is recommended to use preset ia3 for IA^3 algorithm")
network = LycorisNetworkKohya(
text_encoder,
unet,
multiplier=multiplier,
lora_dim=network_dim,
conv_lora_dim=conv_dim,
alpha=network_alpha,
conv_alpha=conv_alpha,
dropout=dropout,
rank_dropout=rank_dropout,
module_dropout=module_dropout,
use_tucker=use_tucker,
use_scalar=use_scalar,
network_module=algo,
train_norm=train_norm,
decompose_both=kwargs.get("decompose_both", False),
factor=kwargs.get("factor", -1),
block_size=block_size,
constraint=constraint,
rescaled=rescaled,
weight_decompose=weight_decompose,
wd_on_out=wd_on_output,
full_matrix=full_matrix,
bypass_mode=bypass_mode,
rs_lora=rs_lora,
unbalanced_factorization=unbalanced_factorization,
train_t5xxl=train_t5xxl,
)
return network
def create_network_from_weights(
multiplier,
file,
vae,
text_encoder,
unet,
weights_sd=None,
for_inference=False,
**kwargs,
):
if weights_sd is None:
if os.path.splitext(file)[1] == ".safetensors":
from safetensors.torch import load_file, safe_open
weights_sd = load_file(file)
else:
weights_sd = torch.load(file, map_location="cpu")
# get dim/alpha mapping
unet_loras = {}
te_loras = {}
for key, value in weights_sd.items():
if "." not in key:
continue
lora_name = key.split(".")[0]
if lora_name.startswith(LycorisNetworkKohya.LORA_PREFIX_UNET):
unet_loras[lora_name] = None
elif lora_name.startswith(LycorisNetworkKohya.LORA_PREFIX_TEXT_ENCODER):
te_loras[lora_name] = None
for name, modules in unet.named_modules():
lora_name = f"{LycorisNetworkKohya.LORA_PREFIX_UNET}_{name}".replace(".", "_")
if lora_name in unet_loras:
unet_loras[lora_name] = modules
if isinstance(text_encoder, list):
text_encoders = text_encoder
use_index = True
else:
text_encoders = [text_encoder]
use_index = False
for idx, te in enumerate(text_encoders):
if use_index:
prefix = f"{LycorisNetworkKohya.LORA_PREFIX_TEXT_ENCODER}{idx+1}"
else:
prefix = LycorisNetworkKohya.LORA_PREFIX_TEXT_ENCODER
for name, modules in te.named_modules():
lora_name = f"{prefix}_{name}".replace(".", "_")
if lora_name in te_loras:
te_loras[lora_name] = modules
original_level = logger.level
logger.setLevel(logging.ERROR)
network = LycorisNetworkKohya(text_encoder, unet)
network.unet_loras = []
network.text_encoder_loras = []
logger.setLevel(original_level)
logger.info("Loading UNet Modules from state dict...")
for lora_name, orig_modules in unet_loras.items():
if orig_modules is None:
continue
lyco_type, params = get_module(weights_sd, lora_name)
module = make_module(lyco_type, params, lora_name, orig_modules)
if module is not None:
network.unet_loras.append(module)
logger.info(f"{len(network.unet_loras)} Modules Loaded")
logger.info("Loading TE Modules from state dict...")
for lora_name, orig_modules in te_loras.items():
if orig_modules is None:
continue
lyco_type, params = get_module(weights_sd, lora_name)
module = make_module(lyco_type, params, lora_name, orig_modules)
if module is not None:
network.text_encoder_loras.append(module)
logger.info(f"{len(network.text_encoder_loras)} Modules Loaded")
for lora in network.unet_loras + network.text_encoder_loras:
lora.multiplier = multiplier
return network, weights_sd
class LycorisNetworkKohya(LycorisNetwork):
"""
LoRA + LoCon
"""
# Ignore proj_in or proj_out, their channels is only a few.
ENABLE_CONV = True
UNET_TARGET_REPLACE_MODULE = [
"Transformer2DModel",
"ResnetBlock2D",
"Downsample2D",
"Upsample2D",
"HunYuanDiTBlock",
"DoubleStreamBlock",
"SingleStreamBlock",
"SingleDiTBlock",
"MMDoubleStreamBlock", #HunYuanVideo
"MMSingleStreamBlock", #HunYuanVideo
]
UNET_TARGET_REPLACE_NAME = [
"conv_in",
"conv_out",
"time_embedding.linear_1",
"time_embedding.linear_2",
]
TEXT_ENCODER_TARGET_REPLACE_MODULE = [
"CLIPAttention",
"CLIPSdpaAttention",
"CLIPMLP",
"MT5Block",
"BertLayer",
]
TEXT_ENCODER_TARGET_REPLACE_NAME = []
LORA_PREFIX_UNET = "lora_unet"
LORA_PREFIX_TEXT_ENCODER = "lora_te"
MODULE_ALGO_MAP = {}
NAME_ALGO_MAP = {}
USE_FNMATCH = False
@classmethod
def apply_preset(cls, preset):
if "enable_conv" in preset:
cls.ENABLE_CONV = preset["enable_conv"]
if "unet_target_module" in preset:
cls.UNET_TARGET_REPLACE_MODULE = preset["unet_target_module"]
if "unet_target_name" in preset:
cls.UNET_TARGET_REPLACE_NAME = preset["unet_target_name"]
if "text_encoder_target_module" in preset:
cls.TEXT_ENCODER_TARGET_REPLACE_MODULE = preset[
"text_encoder_target_module"
]
if "text_encoder_target_name" in preset:
cls.TEXT_ENCODER_TARGET_REPLACE_NAME = preset["text_encoder_target_name"]
if "module_algo_map" in preset:
cls.MODULE_ALGO_MAP = preset["module_algo_map"]
if "name_algo_map" in preset:
cls.NAME_ALGO_MAP = preset["name_algo_map"]
if "use_fnmatch" in preset:
cls.USE_FNMATCH = preset["use_fnmatch"]
return cls
def __init__(
self,
text_encoder,
unet,
multiplier=1.0,
lora_dim=4,
conv_lora_dim=4,
alpha=1,
conv_alpha=1,
use_tucker=False,
dropout=0,
rank_dropout=0,
module_dropout=0,
network_module: str = "locon",
norm_modules=NormModule,
train_norm=False,
train_t5xxl=False,
**kwargs,
) -> None:
torch.nn.Module.__init__(self)
root_kwargs = kwargs
self.multiplier = multiplier
self.lora_dim = lora_dim
self.train_t5xxl = train_t5xxl
if not self.ENABLE_CONV:
conv_lora_dim = 0
self.conv_lora_dim = int(conv_lora_dim)
if self.conv_lora_dim and self.conv_lora_dim != self.lora_dim:
logger.info("Apply different lora dim for conv layer")
logger.info(f"Conv Dim: {conv_lora_dim}, Linear Dim: {lora_dim}")
elif self.conv_lora_dim == 0:
logger.info("Disable conv layer")
self.alpha = alpha
self.conv_alpha = float(conv_alpha)
if self.conv_lora_dim and self.alpha != self.conv_alpha:
logger.info("Apply different alpha value for conv layer")
logger.info(f"Conv alpha: {conv_alpha}, Linear alpha: {alpha}")
if 1 >= dropout >= 0:
logger.info(f"Use Dropout value: {dropout}")
self.dropout = dropout
self.rank_dropout = rank_dropout
self.module_dropout = module_dropout
self.use_tucker = use_tucker
def create_single_module(
lora_name: str,
module: torch.nn.Module,
algo_name,
dim=None,
alpha=None,
use_tucker=self.use_tucker,
**kwargs,
):
for k, v in root_kwargs.items():
if k in kwargs:
continue
kwargs[k] = v
if train_norm and "Norm" in module.__class__.__name__:
return norm_modules(
lora_name,
module,
self.multiplier,
self.rank_dropout,
self.module_dropout,
**kwargs,
)
lora = None
if isinstance(module, torch.nn.Linear) and lora_dim > 0:
dim = dim or lora_dim
alpha = alpha or self.alpha
elif isinstance(
module, (torch.nn.Conv1d, torch.nn.Conv2d, torch.nn.Conv3d)
):
k_size, *_ = module.kernel_size
if k_size == 1 and lora_dim > 0:
dim = dim or lora_dim
alpha = alpha or self.alpha
elif conv_lora_dim > 0 or dim:
dim = dim or conv_lora_dim
alpha = alpha or self.conv_alpha
else:
return None
else:
return None
lora = network_module_dict[algo_name](
lora_name,
module,
self.multiplier,
dim,
alpha,
self.dropout,
self.rank_dropout,
self.module_dropout,
use_tucker,
**kwargs,
)
return lora
def create_modules_(
prefix: str,
root_module: torch.nn.Module,
algo,
configs={},
):
loras = {}
lora_names = []
for name, module in root_module.named_modules():
module_name = module.__class__.__name__
if module_name in self.MODULE_ALGO_MAP and module is not root_module:
next_config = self.MODULE_ALGO_MAP[module_name]
next_algo = next_config.get("algo", algo)
new_loras, new_lora_names = create_modules_(
f"{prefix}_{name}", module, next_algo, next_config
)
for lora_name, lora in zip(new_lora_names, new_loras):
if lora_name not in loras:
loras[lora_name] = lora
lora_names.append(lora_name)
continue
if name:
lora_name = prefix + "." + name
else:
lora_name = prefix
lora_name = lora_name.replace(".", "_")
if lora_name in loras:
continue
lora = create_single_module(lora_name, module, algo, **configs)
if lora is not None:
loras[lora_name] = lora
lora_names.append(lora_name)
return [loras[lora_name] for lora_name in lora_names], lora_names
# create module instances
def create_modules(
prefix,
root_module: torch.nn.Module,
target_replace_modules,
target_replace_names=[],
) -> List:
logger.info("Create LyCORIS Module")
loras = []
next_config = {}
for name, module in root_module.named_modules():
module_name = module.__class__.__name__
if module_name in target_replace_modules and not any(
self.match_fn(t, name) for t in target_replace_names
):
if module_name in self.MODULE_ALGO_MAP:
next_config = self.MODULE_ALGO_MAP[module_name]
algo = next_config.get("algo", network_module)
else:
algo = network_module
loras.extend(
create_modules_(f"{prefix}_{name}", module, algo, next_config)[
0
]
)
next_config = {}
elif name in target_replace_names or any(
self.match_fn(t, name) for t in target_replace_names
):
conf_from_name = self.find_conf_for_name(name)
if conf_from_name is not None:
next_config = conf_from_name
algo = next_config.get("algo", network_module)
elif module_name in self.MODULE_ALGO_MAP:
next_config = self.MODULE_ALGO_MAP[module_name]
algo = next_config.get("algo", network_module)
else:
algo = network_module
lora_name = prefix + "." + name
lora_name = lora_name.replace(".", "_")
lora = create_single_module(lora_name, module, algo, **next_config)
next_config = {}
if lora is not None:
loras.append(lora)
return loras
if network_module == GLoRAModule:
logger.info("GLoRA enabled, only train transformer")
# only train transformer (for GLoRA)
LycorisNetworkKohya.UNET_TARGET_REPLACE_MODULE = [
"Transformer2DModel",
"Attention",
]
LycorisNetworkKohya.UNET_TARGET_REPLACE_NAME = []
self.text_encoder_loras = []
if text_encoder:
if isinstance(text_encoder, list):
text_encoders = text_encoder
use_index = True
else:
text_encoders = [text_encoder]
use_index = False
for i, te in enumerate(text_encoders):
self.text_encoder_loras.extend(
create_modules(
LycorisNetworkKohya.LORA_PREFIX_TEXT_ENCODER
+ (f"{i+1}" if use_index else ""),
te,
LycorisNetworkKohya.TEXT_ENCODER_TARGET_REPLACE_MODULE,
LycorisNetworkKohya.TEXT_ENCODER_TARGET_REPLACE_NAME,
)
)
logger.info(
f"create LyCORIS for Text Encoder: {len(self.text_encoder_loras)} modules."
)
self.unet_loras = create_modules(
LycorisNetworkKohya.LORA_PREFIX_UNET,
unet,
LycorisNetworkKohya.UNET_TARGET_REPLACE_MODULE,
LycorisNetworkKohya.UNET_TARGET_REPLACE_NAME,
)
logger.info(f"create LyCORIS for U-Net: {len(self.unet_loras)} modules.")
algo_table = {}
for lora in self.text_encoder_loras + self.unet_loras:
algo_table[lora.__class__.__name__] = (
algo_table.get(lora.__class__.__name__, 0) + 1
)
logger.info(f"module type table: {algo_table}")
self.weights_sd = None
self.loras = self.text_encoder_loras + self.unet_loras
# assertion
names = set()
for lora in self.loras:
assert (
lora.lora_name not in names
), f"duplicated lora name: {lora.lora_name}"
names.add(lora.lora_name)
def match_fn(self, pattern: str, name: str) -> bool:
if self.USE_FNMATCH:
return fnmatch.fnmatch(name, pattern)
return re.match(pattern, name)
def find_conf_for_name(
self,
name: str,
) -> dict[str, Any]:
if name in self.NAME_ALGO_MAP.keys():
return self.NAME_ALGO_MAP[name]
for key, value in self.NAME_ALGO_MAP.items():
if self.match_fn(key, name):
return value
return None
def load_weights(self, file):
if os.path.splitext(file)[1] == ".safetensors":
from safetensors.torch import load_file, safe_open
self.weights_sd = load_file(file)
else:
self.weights_sd = torch.load(file, map_location="cpu")
missing, unexpected = self.load_state_dict(self.weights_sd, strict=False)
state = {}
if missing:
state["missing keys"] = missing
if unexpected:
state["unexpected keys"] = unexpected
return state
def apply_to(self, text_encoder, unet, apply_text_encoder=None, apply_unet=None):
assert (
apply_text_encoder is not None and apply_unet is not None
), f"internal error: flag not set"
if apply_text_encoder:
logger.info("enable LyCORIS for text encoder")
else:
self.text_encoder_loras = []
if apply_unet:
logger.info("enable LyCORIS for U-Net")
else:
self.unet_loras = []
self.loras = self.text_encoder_loras + self.unet_loras
for lora in self.loras:
lora.apply_to()
self.add_module(lora.lora_name, lora)
if self.weights_sd:
# if some weights are not in state dict, it is ok because initial LoRA does nothing (lora_up is initialized by zeros)
info = self.load_state_dict(self.weights_sd, False)
logger.info(f"weights are loaded: {info}")
# TODO refactor to common function with apply_to
def merge_to(self, text_encoder, unet, weights_sd, dtype, device):
apply_text_encoder = apply_unet = False
for key in weights_sd.keys():
if key.startswith(LycorisNetworkKohya.LORA_PREFIX_TEXT_ENCODER):
apply_text_encoder = True
elif key.startswith(LycorisNetworkKohya.LORA_PREFIX_UNET):
apply_unet = True
if apply_text_encoder:
logger.info("enable LoRA for text encoder")
else:
self.text_encoder_loras = []
if apply_unet:
logger.info("enable LoRA for U-Net")
else:
self.unet_loras = []
self.loras = self.text_encoder_loras + self.unet_loras
super().merge_to(1)
def apply_max_norm_regularization(self, max_norm_value, device):
key_scaled = 0
norms = []
for module in self.unet_loras + self.text_encoder_loras:
scaled, norm = module.apply_max_norm(max_norm_value, device)
if scaled is None:
continue
norms.append(norm)
key_scaled += scaled
if key_scaled == 0:
return 0, 0, 0
return key_scaled, sum(norms) / len(norms), max(norms)
def prepare_optimizer_params(self, text_encoder_lr=None, unet_lr: float = 1e-4, learning_rate=None):
def enumerate_params(loras):
params = []
for lora in loras:
params.extend(lora.parameters())
return params
self.requires_grad_(True)
all_params = []
lr_descriptions = []
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)
lr_descriptions.append("text_encoder")
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)
lr_descriptions.append("unet")
return all_params, lr_descriptions
def enable_gradient_checkpointing(self):
# not supported
pass
def prepare_grad_etc(self, text_encoder, unet):
self.requires_grad_(True)
def on_epoch_start(self, text_encoder, unet):
self.train()
#def on_step_start(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 = precalculate_safetensors_hashes(state_dict)
metadata["sshs_model_hash"] = model_hash
save_file(state_dict, file, metadata)
else:
torch.save(state_dict, file)
+52
View File
@@ -0,0 +1,52 @@
import sys
import copy
import logging
from functools import cache
class ColoredFormatter(logging.Formatter):
COLORS = {
"DEBUG": "\033[0;36m", # CYAN
"INFO": "\033[0;32m", # GREEN
"WARNING": "\033[0;33m", # YELLOW
"ERROR": "\033[0;31m", # RED
"CRITICAL": "\033[0;37;41m", # WHITE ON RED
"RESET": "\033[0m", # RESET COLOR
}
def format(self, record):
colored_record = copy.copy(record)
levelname = colored_record.levelname
seq = self.COLORS.get(levelname, self.COLORS["RESET"])
colored_record.levelname = f"{seq}{levelname}{self.COLORS['RESET']}"
return super().format(colored_record)
logger = logging.getLogger("LyCORIS")
logger.propagate = False
logger.setLevel(logging.INFO)
if not logger.handlers:
handler = logging.StreamHandler(sys.stdout)
handler.setFormatter(
ColoredFormatter(
"%(asctime)s|[%(name)s]-%(levelname)s: %(message)s", "%Y-%m-%d %H:%M:%S"
)
)
logger.addHandler(handler)
@cache
def info_once(msg):
logger.info(msg)
@cache
def warning_once(msg):
logger.warning(msg)
@cache
def error_once(msg):
logger.error(msg)
+46
View File
@@ -0,0 +1,46 @@
import torch
import torch.nn as nn
from .base import LycorisBaseModule
from .locon import LoConModule
from .loha import LohaModule
from .lokr import LokrModule
from .full import FullModule
from .norms import NormModule
from .diag_oft import DiagOFTModule
from .boft import ButterflyOFTModule
from .glora import GLoRAModule
from .dylora import DyLoraModule
from .ia3 import IA3Module
from ..functional.general import factorization
MODULE_LIST = [
LoConModule,
LohaModule,
IA3Module,
LokrModule,
FullModule,
NormModule,
DiagOFTModule,
ButterflyOFTModule,
GLoRAModule,
DyLoraModule,
]
def get_module(lyco_state_dict, lora_name):
for module in MODULE_LIST:
if module.algo_check(lyco_state_dict, lora_name):
return module, tuple(module.extract_state_dict(lyco_state_dict, lora_name))
return None, None
@torch.no_grad()
def make_module(lyco_type: LycorisBaseModule, params, lora_name, orig_module):
try:
module = lyco_type.make_module_from_state_dict(lora_name, orig_module, *params)
except NotImplementedError:
module = None
return module
+315
View File
@@ -0,0 +1,315 @@
from collections import OrderedDict
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.nn.utils.parametrize as parametrize
from ..utils.quant import QuantLinears, log_bypass, log_suspect
class ModuleCustomSD(nn.Module):
def __init__(self):
super().__init__()
self._register_load_state_dict_pre_hook(self.load_weight_prehook)
self.register_load_state_dict_post_hook(self.load_weight_hook)
def load_weight_prehook(
self,
state_dict,
prefix,
local_metadata,
strict,
missing_keys,
unexpected_keys,
error_msgs,
):
pass
def load_weight_hook(self, module, incompatible_keys):
pass
def custom_state_dict(self):
return None
def state_dict(self, *args, destination=None, prefix="", keep_vars=False):
# TODO: Remove `args` and the parsing logic when BC allows.
if len(args) > 0:
if destination is None:
destination = args[0]
if len(args) > 1 and prefix == "":
prefix = args[1]
if len(args) > 2 and keep_vars is False:
keep_vars = args[2]
# DeprecationWarning is ignored by default
if destination is None:
destination = OrderedDict()
destination._metadata = OrderedDict()
local_metadata = dict(version=self._version)
if hasattr(destination, "_metadata"):
destination._metadata[prefix[:-1]] = local_metadata
if (custom_sd := self.custom_state_dict()) is not None:
for k, v in custom_sd.items():
destination[f"{prefix}{k}"] = v
return destination
else:
return super().state_dict(
*args, destination=destination, prefix=prefix, keep_vars=keep_vars
)
class LycorisBaseModule(ModuleCustomSD):
name: str
dtype_tensor: torch.Tensor
support_module = {}
weight_list = []
weight_list_det = []
def __init__(
self,
lora_name,
org_module: nn.Module,
multiplier=1.0,
dropout=0.0,
rank_dropout=0.0,
module_dropout=0.0,
rank_dropout_scale=False,
bypass_mode=None,
**kwargs,
):
"""if alpha == 0 or None, alpha is rank (no scaling)."""
super().__init__()
self.lora_name = lora_name
self.not_supported = False
self.module = type(org_module)
if isinstance(org_module, nn.Linear):
self.module_type = "linear"
self.shape = (org_module.out_features, org_module.in_features)
self.op = F.linear
self.dim = org_module.out_features
self.kw_dict = {}
elif isinstance(org_module, nn.Conv1d):
self.module_type = "conv1d"
self.shape = (
org_module.out_channels,
org_module.in_channels,
*org_module.kernel_size,
)
self.op = F.conv1d
self.dim = org_module.out_channels
self.kw_dict = {
"stride": org_module.stride,
"padding": org_module.padding,
"dilation": org_module.dilation,
"groups": org_module.groups,
}
elif isinstance(org_module, nn.Conv2d):
self.module_type = "conv2d"
self.shape = (
org_module.out_channels,
org_module.in_channels,
*org_module.kernel_size,
)
self.op = F.conv2d
self.dim = org_module.out_channels
self.kw_dict = {
"stride": org_module.stride,
"padding": org_module.padding,
"dilation": org_module.dilation,
"groups": org_module.groups,
}
elif isinstance(org_module, nn.Conv3d):
self.module_type = "conv3d"
self.shape = (
org_module.out_channels,
org_module.in_channels,
*org_module.kernel_size,
)
self.op = F.conv3d
self.dim = org_module.out_channels
self.kw_dict = {
"stride": org_module.stride,
"padding": org_module.padding,
"dilation": org_module.dilation,
"groups": org_module.groups,
}
elif isinstance(org_module, nn.LayerNorm):
self.module_type = "layernorm"
self.shape = tuple(org_module.normalized_shape)
self.op = F.layer_norm
self.dim = org_module.normalized_shape[0]
self.kw_dict = {
"normalized_shape": org_module.normalized_shape,
"eps": org_module.eps,
}
elif isinstance(org_module, nn.GroupNorm):
self.module_type = "groupnorm"
self.shape = (org_module.num_channels,)
self.op = F.group_norm
self.group_num = org_module.num_groups
self.dim = org_module.num_channels
self.kw_dict = {"num_groups": org_module.num_groups, "eps": org_module.eps}
else:
self.not_supported = True
self.module_type = "unknown"
self.register_buffer("dtype_tensor", torch.tensor(0.0), persistent=False)
self.is_quant = False
if isinstance(org_module, QuantLinears):
if not bypass_mode:
log_bypass()
self.is_quant = True
bypass_mode = True
if (
isinstance(org_module, nn.Linear)
and org_module.__class__.__name__ != "Linear"
):
if bypass_mode is None:
log_suspect()
bypass_mode = True
if bypass_mode == True:
self.is_quant = True
self.bypass_mode = bypass_mode
self.dropout = dropout
self.rank_dropout = rank_dropout
self.rank_dropout_scale = rank_dropout_scale
self.module_dropout = module_dropout
## Dropout things
# Since LoKr/LoHa/OFT/BOFT are hard to follow the rank_dropout definition from kohya
# We redefine the dropout procedure here.
# g(x) = WX + drop(Brank_drop(AX)) for LoCon(lora), bypass
# g(x) = WX + drop(ΔWX) for any algo except LoCon(lora), bypass
# g(x) = (W + Brank_drop(A))X for LoCon(lora), rebuid
# g(x) = (W + rank_drop(ΔW))X for any algo except LoCon(lora), rebuild
self.drop = nn.Identity() if dropout == 0 else nn.Dropout(dropout)
self.rank_drop = (
nn.Identity() if rank_dropout == 0 else nn.Dropout(rank_dropout)
)
self.multiplier = multiplier
self.org_forward = org_module.forward
self.org_module = [org_module]
@classmethod
def parametrize(cls, org_module, attr, *args, **kwargs):
from .full import FullModule
if cls is FullModule:
raise RuntimeError("FullModule cannot be used for parametrize.")
target_param = getattr(org_module, attr)
kwargs["bypass_mode"] = False
if target_param.dim() == 2:
proxy_module = nn.Linear(
target_param.shape[0], target_param.shape[1], bias=False
)
proxy_module.weight = target_param
elif target_param.dim() > 2:
module_type = [
None,
None,
None,
nn.Conv1d,
nn.Conv2d,
nn.Conv3d,
None,
None,
][target_param.dim()]
proxy_module = module_type(
target_param.shape[0],
target_param.shape[1],
*target_param.shape[2:],
bias=False,
)
proxy_module.weight = target_param
module_obj = cls("", proxy_module, *args, **kwargs)
module_obj.forward = module_obj.parametrize_forward
module_obj.to(target_param)
parametrize.register_parametrization(org_module, attr, module_obj)
return module_obj
@classmethod
def algo_check(cls, state_dict, lora_name):
return any(f"{lora_name}.{k}" in state_dict for k in cls.weight_list_det)
@classmethod
def extract_state_dict(cls, state_dict, lora_name):
return [state_dict.get(f"{lora_name}.{k}", None) for k in cls.weight_list]
@classmethod
def make_module_from_state_dict(cls, lora_name, orig_module, *weights):
raise NotImplementedError
@property
def dtype(self):
return self.dtype_tensor.dtype
@property
def device(self):
return self.dtype_tensor.device
@property
def org_weight(self):
return self.org_module[0].weight
@org_weight.setter
def org_weight(self, value):
self.org_module[0].weight.data.copy_(value)
def apply_to(self, **kwargs):
if self.not_supported:
return
self.org_forward = self.org_module[0].forward
self.org_module[0].forward = self.forward
def restore(self):
if self.not_supported:
return
self.org_module[0].forward = self.org_forward
def merge_to(self, multiplier=1.0):
if self.not_supported:
return
self_device = next(self.parameters()).device
self_dtype = next(self.parameters()).dtype
self.to(self.org_weight)
weight, bias = self.get_merged_weight(
multiplier, self.org_weight.shape, self.org_weight.device
)
self.org_weight = weight.to(self.org_weight)
if bias is not None:
bias = bias.to(self.org_weight)
if self.org_module[0].bias is not None:
self.org_module[0].bias.data.copy_(bias)
else:
self.org_module[0].bias = nn.Parameter(bias)
self.to(self_device, self_dtype)
def get_diff_weight(self, multiplier=1.0, shape=None, device=None):
raise NotImplementedError
def get_merged_weight(self, multiplier=1.0, shape=None, device=None):
raise NotImplementedError
@torch.no_grad()
def apply_max_norm(self, max_norm, device=None):
return None, None
def bypass_forward_diff(self, x, scale=1):
raise NotImplementedError
def bypass_forward(self, x, scale=1):
raise NotImplementedError
def parametrize_forward(self, x: torch.Tensor, *args, **kwargs):
return self.get_merged_weight(
multiplier=self.multiplier, shape=x.shape, device=x.device
)[0].to(x.dtype)
def forward(self, *args, **kwargs):
raise NotImplementedError
+255
View File
@@ -0,0 +1,255 @@
from functools import cache
from math import log2
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from .base import LycorisBaseModule
from ..functional import power2factorization
from ..logging import logger
@cache
def log_butterfly_factorize(dim, factor, result):
logger.info(
f"Use BOFT({int(log2(result[1]))}, {result[0]//2})"
f" (equivalent to factor={result[0]}) "
f"for {dim=} and {factor=}"
)
def butterfly_factor(dimension: int, factor: int = -1) -> tuple[int, int]:
m, n = power2factorization(dimension, factor)
if n == 0:
raise ValueError(
f"It is impossible to decompose {dimension} with factor {factor} under BOFT constraints."
)
log_butterfly_factorize(dimension, factor, (m, n))
return m, n
class ButterflyOFTModule(LycorisBaseModule):
name = "boft"
support_module = {
"linear",
"conv1d",
"conv2d",
"conv3d",
}
weight_list = [
"oft_blocks",
"rescale",
"alpha",
]
weight_list_det = ["oft_blocks"]
def __init__(
self,
lora_name,
org_module: nn.Module,
multiplier=1.0,
lora_dim=4,
alpha=1,
dropout=0.0,
rank_dropout=0.0,
module_dropout=0.0,
use_tucker=False,
use_scalar=False,
rank_dropout_scale=False,
constraint=0,
rescaled=False,
bypass_mode=None,
**kwargs,
):
super().__init__(
lora_name,
org_module,
multiplier,
dropout,
rank_dropout,
module_dropout,
rank_dropout_scale,
bypass_mode,
)
if self.module_type not in self.support_module:
raise ValueError(f"{self.module_type} is not supported in BOFT algo.")
out_dim = self.dim
b, m_exp = butterfly_factor(out_dim, lora_dim)
self.block_size = b
self.block_num = m_exp
# BOFT(m, b)
self.boft_b = b
self.boft_m = sum(int(i) for i in f"{m_exp-1:b}") + 1
# block_num > block_size
self.rescaled = rescaled
self.constraint = constraint * out_dim
self.register_buffer("alpha", torch.tensor(constraint))
self.oft_blocks = nn.Parameter(
torch.zeros(self.boft_m, self.block_num, self.block_size, self.block_size)
)
if rescaled:
self.rescale = nn.Parameter(
torch.ones(out_dim, *(1 for _ in range(org_module.weight.dim() - 1)))
)
@classmethod
def algo_check(cls, state_dict, lora_name):
if f"{lora_name}.oft_blocks" in state_dict:
oft_blocks = state_dict[f"{lora_name}.oft_blocks"]
if oft_blocks.ndim == 4:
return True
return False
@classmethod
def make_module_from_state_dict(
cls, lora_name, orig_module, oft_blocks, rescale, alpha
):
m, n, s, _ = oft_blocks.shape
module = cls(
lora_name,
orig_module,
1,
lora_dim=s,
constraint=float(alpha),
rescaled=rescale is not None,
)
module.oft_blocks.copy_(oft_blocks)
if rescale is not None:
module.rescale.copy_(rescale)
return module
@property
def I(self):
return torch.eye(self.block_size, device=self.device)
def get_r(self):
I = self.I
# for Q = -Q^T
q = self.oft_blocks - self.oft_blocks.transpose(-1, -2)
normed_q = q
# Diag OFT style constrain
if self.constraint > 0:
q_norm = torch.norm(q) + 1e-8
if q_norm > self.constraint:
normed_q = q * self.constraint / q_norm
# use float() to prevent unsupported type
r = (I + normed_q) @ (I - normed_q).float().inverse()
return r
def make_weight(self, scale=1, device=None, diff=False):
m = self.boft_m
b = self.boft_b
r_b = b // 2
r = self.get_r()
inp = org = self.org_weight.to(device, dtype=r.dtype)
for i in range(m):
bi = r[i] # b_num, b_size, b_size
g = 2
k = 2**i * r_b
if scale != 1:
bi = bi * scale + (1 - scale) * self.I
inp = (
inp.unflatten(-1, (-1, g, k))
.transpose(-2, -1)
.flatten(-3)
.unflatten(-1, (-1, b))
)
inp = torch.einsum("b i j, b j ... -> b i ...", bi, inp)
inp = (
inp.flatten(-2).unflatten(-1, (-1, k, g)).transpose(-2, -1).flatten(-3)
)
if self.rescaled:
inp = inp * self.rescale
if diff:
inp = inp - org
return inp.to(self.oft_blocks.dtype)
def get_diff_weight(self, multiplier=1, shape=None, device=None):
diff = self.make_weight(scale=multiplier, device=device, diff=True)
if shape is not None:
diff = diff.view(shape)
return diff, None
def get_merged_weight(self, multiplier=1, shape=None, device=None):
diff = self.make_weight(scale=multiplier, device=device)
if shape is not None:
diff = diff.view(shape)
return diff, None
@torch.no_grad()
def apply_max_norm(self, max_norm, device=None):
orig_norm = self.oft_blocks.to(device).norm()
norm = torch.clamp(orig_norm, max_norm / 2)
desired = torch.clamp(norm, max=max_norm)
ratio = desired / norm
scaled = norm != desired
if scaled:
self.oft_blocks *= ratio
return scaled, orig_norm * ratio
def _bypass_forward(self, x, scale=1, diff=False):
m = self.boft_m
b = self.boft_b
r_b = b // 2
r = self.get_r()
inp = org = self.org_forward(x)
if self.op in {F.conv2d, F.conv1d, F.conv3d}:
inp = inp.transpose(1, -1)
for i in range(m):
bi = r[i] # b_num, b_size, b_size
g = 2
k = 2**i * r_b
if scale != 1:
bi = bi * scale + (1 - scale) * self.I
inp = (
inp.unflatten(-1, (-1, g, k))
.transpose(-2, -1)
.flatten(-3)
.unflatten(-1, (-1, b))
)
inp = torch.einsum("b i j, ... b j -> ... b i", bi, inp)
inp = (
inp.flatten(-2).unflatten(-1, (-1, k, g)).transpose(-2, -1).flatten(-3)
)
if self.rescaled:
inp = inp * self.rescale.transpose(0, -1)
if self.op in {F.conv2d, F.conv1d, F.conv3d}:
inp = inp.transpose(1, -1)
if diff:
inp = inp - org
return inp
def bypass_forward_diff(self, x, scale=1):
return self._bypass_forward(x, scale, diff=True)
def bypass_forward(self, x, scale=1):
return self._bypass_forward(x, scale, diff=False)
def forward(self, x, *args, **kwargs):
if self.module_dropout and self.training:
if torch.rand(1) < self.module_dropout:
return self.org_forward(x)
scale = self.multiplier
if self.bypass_mode:
return self.bypass_forward(x, scale)
else:
w = self.make_weight(scale, x.device)
kw_dict = self.kw_dict | {"weight": w, "bias": self.org_module[0].bias}
return self.op(x, **kw_dict)
+217
View File
@@ -0,0 +1,217 @@
from functools import cache
import torch
import torch.nn as nn
import torch.nn.functional as F
from .base import LycorisBaseModule
from ..functional import factorization
from ..logging import logger
@cache
def log_oft_factorize(dim, factor, num, bdim):
logger.info(
f"Use OFT(block num: {num}, block dim: {bdim})"
f" (equivalent to lora_dim={num}) "
f"for {dim=} and lora_dim={factor=}"
)
class DiagOFTModule(LycorisBaseModule):
name = "diag-oft"
support_module = {
"linear",
"conv1d",
"conv2d",
"conv3d",
}
weight_list = [
"oft_blocks",
"rescale",
"alpha",
]
weight_list_det = ["oft_blocks"]
def __init__(
self,
lora_name,
org_module: nn.Module,
multiplier=1.0,
lora_dim=4,
alpha=1,
dropout=0.0,
rank_dropout=0.0,
module_dropout=0.0,
use_tucker=False,
use_scalar=False,
rank_dropout_scale=False,
constraint=0,
rescaled=False,
bypass_mode=None,
**kwargs,
):
super().__init__(
lora_name,
org_module,
multiplier,
dropout,
rank_dropout,
module_dropout,
rank_dropout_scale,
bypass_mode,
)
if self.module_type not in self.support_module:
raise ValueError(f"{self.module_type} is not supported in Diag-OFT algo.")
out_dim = self.dim
self.block_size, self.block_num = factorization(out_dim, lora_dim)
# block_num > block_size
self.rescaled = rescaled
self.constraint = constraint * out_dim
self.register_buffer("alpha", torch.tensor(constraint))
self.oft_blocks = nn.Parameter(
torch.zeros(self.block_num, self.block_size, self.block_size)
)
if rescaled:
self.rescale = nn.Parameter(
torch.ones(out_dim, *(1 for _ in range(org_module.weight.dim() - 1)))
)
log_oft_factorize(
dim=out_dim,
factor=lora_dim,
num=self.block_num,
bdim=self.block_size,
)
@classmethod
def algo_check(cls, state_dict, lora_name):
if f"{lora_name}.oft_blocks" in state_dict:
oft_blocks = state_dict[f"{lora_name}.oft_blocks"]
if oft_blocks.ndim == 3:
return True
return False
@classmethod
def make_module_from_state_dict(
cls, lora_name, orig_module, oft_blocks, rescale, alpha
):
n, s, _ = oft_blocks.shape
module = cls(
lora_name,
orig_module,
1,
lora_dim=s,
constraint=float(alpha),
rescaled=rescale is not None,
)
module.oft_blocks.copy_(oft_blocks)
if rescale is not None:
module.rescale.copy_(rescale)
return module
@property
def I(self):
return torch.eye(self.block_size, device=self.device)
def get_r(self):
I = self.I
# for Q = -Q^T
q = self.oft_blocks - self.oft_blocks.transpose(1, 2)
normed_q = q
if self.constraint > 0:
q_norm = torch.norm(q) + 1e-8
if q_norm > self.constraint:
normed_q = q * self.constraint / q_norm
# use float() to prevent unsupported type
r = (I + normed_q) @ (I - normed_q).float().inverse()
return r
def make_weight(self, scale=1, device=None, diff=False):
r = self.get_r()
_, *shape = self.org_weight.shape
org_weight = self.org_weight.to(device, dtype=r.dtype)
org_weight = org_weight.view(self.block_num, self.block_size, *shape)
# Init R=0, so add I on it to ensure the output of step0 is original model output
weight = torch.einsum(
"k n m, k n ... -> k m ...",
self.rank_drop(r * scale) - scale * self.I + (0 if diff else self.I),
org_weight,
).view(-1, *shape)
if self.rescaled:
weight = self.rescale * weight
if diff:
weight = weight + (self.rescale - 1) * org_weight
return weight.to(self.oft_blocks.dtype)
def get_diff_weight(self, multiplier=1, shape=None, device=None):
diff = self.make_weight(scale=multiplier, device=device, diff=True)
if shape is not None:
diff = diff.view(shape)
return diff, None
def get_merged_weight(self, multiplier=1, shape=None, device=None):
diff = self.make_weight(scale=multiplier, device=device)
if shape is not None:
diff = diff.view(shape)
return diff, None
@torch.no_grad()
def apply_max_norm(self, max_norm, device=None):
orig_norm = self.oft_blocks.to(device).norm()
norm = torch.clamp(orig_norm, max_norm / 2)
desired = torch.clamp(norm, max=max_norm)
ratio = desired / norm
scaled = norm != desired
if scaled:
self.oft_blocks *= ratio
return scaled, orig_norm * ratio
def _bypass_forward(self, x, scale=1, diff=False):
r = self.get_r()
org_out = self.org_forward(x)
if self.op in {F.conv2d, F.conv1d, F.conv3d}:
org_out = org_out.transpose(1, -1)
*shape, _ = org_out.shape
org_out = org_out.view(*shape, self.block_num, self.block_size)
mask = neg_mask = 1
if self.dropout != 0 and self.training:
mask = torch.ones_like(org_out)
mask = self.drop(mask)
neg_mask = torch.max(mask) - mask
oft_out = torch.einsum(
"k n m, ... k n -> ... k m",
r * scale * mask + (1 - scale) * self.I * neg_mask,
org_out,
)
if diff:
out = out - org_out
out = oft_out.view(*shape, -1)
if self.rescaled:
out = self.rescale.transpose(-1, 0) * out
out = out + (self.rescale.transpose(-1, 0) - 1) * org_out
if self.op in {F.conv2d, F.conv1d, F.conv3d}:
out = out.transpose(1, -1)
return out
def bypass_forward_diff(self, x, scale=1):
return self._bypass_forward(x, scale, diff=True)
def bypass_forward(self, x, scale=1):
return self._bypass_forward(x, scale, diff=False)
def forward(self, x: torch.Tensor, *args, **kwargs):
if self.module_dropout and self.training:
if torch.rand(1) < self.module_dropout:
return self.org_forward(x)
scale = self.multiplier
if self.bypass_mode:
return self.bypass_forward(x, scale)
else:
w = self.make_weight(scale, x.device)
kw_dict = self.kw_dict | {"weight": w, "bias": self.org_module[0].bias}
return self.op(x, **kw_dict)
+156
View File
@@ -0,0 +1,156 @@
import math
import random
import torch
import torch.nn as nn
from .base import LycorisBaseModule
from ..utils import product
class DyLoraModule(LycorisBaseModule):
support_module = {
"linear",
"conv1d",
"conv2d",
"conv3d",
}
def __init__(
self,
lora_name,
org_module: nn.Module,
multiplier=1.0,
lora_dim=4,
alpha=1,
dropout=0.0,
rank_dropout=0.0,
module_dropout=0.0,
use_tucker=False,
block_size=4,
use_scalar=False,
rank_dropout_scale=False,
weight_decompose=False,
bypass_mode=None,
rs_lora=False,
train_on_input=False,
**kwargs,
):
"""if alpha == 0 or None, alpha is rank (no scaling)."""
super().__init__(
lora_name,
org_module,
multiplier,
dropout,
rank_dropout,
module_dropout,
rank_dropout_scale,
bypass_mode,
)
if self.module_type not in self.support_module:
raise ValueError(f"{self.module_type} is not supported in IA^3 algo.")
assert lora_dim % block_size == 0, "lora_dim must be a multiple of block_size"
self.block_count = lora_dim // block_size
self.block_size = block_size
shape = (
self.shape[0],
product(self.shape[1:]),
)
self.lora_dim = lora_dim
self.up_list = nn.ParameterList(
[torch.empty(shape[0], self.block_size) for i in range(self.block_count)]
)
self.down_list = nn.ParameterList(
[torch.empty(self.block_size, shape[1]) for i in range(self.block_count)]
)
if type(alpha) == torch.Tensor:
alpha = alpha.detach().float().numpy() # without casting, bf16 causes error
alpha = lora_dim if alpha is None or alpha == 0 else alpha
self.scale = alpha / self.lora_dim
self.register_buffer("alpha", torch.tensor(alpha)) # 定数として扱える
# Need more experiences on init method
for v in self.down_list:
torch.nn.init.kaiming_uniform_(v, a=math.sqrt(5))
for v in self.up_list:
torch.nn.init.zeros_(v)
def load_state_dict(self, state_dict, strict: bool = True, assign: bool = False):
return
def custom_state_dict(self):
destination = {}
destination["alpha"] = self.alpha
destination["lora_up.weight"] = nn.Parameter(
torch.concat(list(self.up_list), dim=1)
)
destination["lora_down.weight"] = nn.Parameter(
torch.concat(list(self.down_list)).reshape(
self.lora_dim, -1, *self.shape[2:]
)
)
return destination
def get_weight(self, rank):
b = math.ceil(rank / self.block_size)
down = torch.concat(
list(i.data for i in self.down_list[:b]) + list(self.down_list[b : (b + 1)])
)
up = torch.concat(
list(i.data for i in self.up_list[:b]) + list(self.up_list[b : (b + 1)]),
dim=1,
)
return down, up, self.alpha / (b + 1)
def get_random_rank_weight(self):
b = random.randint(0, self.block_count - 1)
return self.get_weight(b * self.block_size)
def get_diff_weight(self, multiplier=1, shape=None, device=None, rank=None):
if rank is None:
down, up, scale = self.get_random_rank_weight()
else:
down, up, scale = self.get_weight(rank)
w = up @ (down * (scale * multiplier))
if device is not None:
w = w.to(device)
if shape is not None:
w = w.view(shape)
else:
w = w.view(self.shape)
return w, None
def get_merged_weight(self, multiplier=1, shape=None, device=None, rank=None):
diff, _ = self.get_diff_weight(multiplier, shape, device, rank)
return diff + self.org_weight, None
def bypass_forward_diff(self, x, scale=1, rank=None):
if rank is None:
down, up, gamma = self.get_random_rank_weight()
else:
down, up, scale = self.get_weight(rank)
down = down.view(self.lora_dim, -1, *self.shape[2:])
up = up.view(-1, self.lora_dim, *(1 for _ in self.shape[2:]))
scale = scale * gamma
return self.op(self.op(x, down, **self.kw_dict), up)
def bypass_forward(self, x, scale=1, rank=None):
return self.org_forward(x) + self.bypass_forward_diff(x, scale, rank)
def forward(self, x, *args, **kwargs):
if self.module_dropout and self.training:
if torch.rand(1) < self.module_dropout:
return self.org_forward(x)
if self.bypass_mode:
return self.bypass_forward(x, self.multiplier)
else:
weight = self.get_merged_weight(multiplier=self.multiplier)[0]
bias = (
None
if self.org_module[0].bias is None
else self.org_module[0].bias.data
)
return self.op(x, weight, bias, **self.kw_dict)
+214
View File
@@ -0,0 +1,214 @@
from functools import cache
import torch
import torch.nn as nn
from .base import LycorisBaseModule
from ..logging import logger
@cache
def log_bypass_override():
return logger.warning(
"Automatic Bypass-Mode detected in algo=full, "
"override with bypass_mode=False since algo=full not support bypass mode. "
"If you are using quantized model which require bypass mode, please don't use algo=full. "
)
class FullModule(LycorisBaseModule):
name = "full"
support_module = {
"linear",
"conv1d",
"conv2d",
"conv3d",
}
weight_list = ["diff", "diff_b"]
weight_list_det = ["diff"]
def __init__(
self,
lora_name,
org_module: nn.Module,
multiplier=1.0,
lora_dim=4,
alpha=1,
dropout=0.0,
rank_dropout=0.0,
module_dropout=0.0,
use_tucker=False,
use_scalar=False,
rank_dropout_scale=False,
bypass_mode=None,
**kwargs,
):
org_bypass = bypass_mode
super().__init__(
lora_name,
org_module,
multiplier,
dropout,
rank_dropout,
module_dropout,
rank_dropout_scale,
bypass_mode,
)
if bypass_mode and org_bypass is None:
self.bypass_mode = False
log_bypass_override()
if self.module_type not in self.support_module:
raise ValueError(f"{self.module_type} is not supported in Full algo.")
if self.is_quant:
raise ValueError(
"Quant Linear is not supported and meaningless in Full algo."
)
if self.bypass_mode:
raise ValueError("bypass mode is not supported in Full algo.")
self.weight = nn.Parameter(torch.zeros_like(org_module.weight))
if org_module.bias is not None:
self.bias = nn.Parameter(torch.zeros_like(org_module.bias))
else:
self.bias = None
self.is_diff = True
self._org_weight = [self.org_module[0].weight.data.cpu().clone()]
if self.org_module[0].bias is not None:
self.org_bias = [self.org_module[0].bias.data.cpu().clone()]
else:
self.org_bias = None
@classmethod
def make_module_from_state_dict(cls, lora_name, orig_module, diff, diff_b):
module = cls(
lora_name,
orig_module,
1,
)
module.weight.copy_(diff)
if diff_b is not None:
if orig_module.bias is not None:
module.bias.copy_(diff_b)
else:
module.bias = nn.Parameter(diff_b)
module.is_diff = True
return module
@property
def org_weight(self):
return self._org_weight[0]
@org_weight.setter
def org_weight(self, value):
self.org_module[0].weight.data.copy_(value)
def apply_to(self, **kwargs):
self.org_forward = self.org_module[0].forward
self.org_module[0].forward = self.forward
self.weight.data.add_(self.org_module[0].weight.data)
self._org_weight = [self.org_module[0].weight.data.cpu().clone()]
delattr(self.org_module[0], "weight")
if self.org_module[0].bias is not None:
self.bias.data.add_(self.org_module[0].bias.data)
self.org_bias = [self.org_module[0].bias.data.cpu().clone()]
delattr(self.org_module[0], "bias")
else:
self.org_bias = None
self.is_diff = False
def restore(self):
self.org_module[0].forward = self.org_forward
self.org_module[0].weight = nn.Parameter(self._org_weight[0])
if self.org_bias is not None:
self.org_module[0].bias = nn.Parameter(self.org_bias[0])
def custom_state_dict(self):
sd = {"diff": self.weight.data.cpu() - self._org_weight[0]}
if self.bias is not None:
sd["diff_b"] = self.bias.data.cpu() - self.org_bias[0]
return sd
def load_weight_prehook(
self,
state_dict,
prefix,
local_metadata,
strict,
missing_keys,
unexpected_keys,
error_msgs,
):
diff_weight = state_dict.pop(f"{prefix}diff")
state_dict[f"{prefix}weight"] = diff_weight + self.weight.data.to(diff_weight)
if f"{prefix}diff_b" in state_dict:
diff_bias = state_dict.pop(f"{prefix}diff_b")
state_dict[f"{prefix}bias"] = diff_bias + self.bias.data.to(diff_bias)
def make_weight(self, scale=1, device=None):
drop = (
torch.rand(self.dim, device=device) > self.rank_dropout
if self.rank_dropout and self.training
else 1
)
if drop != 1 or scale != 1 or self.is_diff:
diff_w, diff_b = self.get_diff_weight(scale, device=device)
weight = self.org_weight + diff_w * drop
if self.org_bias is not None:
bias = self.org_bias + diff_b * drop
else:
bias = None
else:
weight = self.weight
bias = self.bias
return weight, bias
def get_diff_weight(self, multiplier=1, shape=None, device=None):
if self.is_diff:
diff_b = None
if self.bias is not None:
diff_b = self.bias * multiplier
return self.weight * multiplier, diff_b
org_weight = self.org_module[0].weight.to(device, dtype=self.weight.dtype)
diff = self.weight.to(device) - org_weight
diff_b = None
if shape:
diff = diff.view(shape)
if self.bias is not None:
org_bias = self.org_module[0].bias.to(device, dtype=self.bias.dtype)
diff_b = self.bias.to(device) - org_bias
if device is not None:
diff = diff.to(device)
if self.bias is not None:
diff_b = diff_b.to(device)
if multiplier != 1:
diff = diff * multiplier
if diff_b is not None:
diff_b = diff_b * multiplier
return diff * multiplier, diff_b
def get_merged_weight(self, multiplier=1, shape=None, device=None):
weight, bias = self.make_weight(multiplier, device)
if shape is not None:
weight = weight.view(shape)
if bias is not None:
bias = bias.view(shape[0])
return weight, bias
def forward(self, x: torch.Tensor, *args, **kwargs):
if (
self.module_dropout
and self.training
and torch.rand(1) < self.module_dropout
):
original = True
else:
original = False
if original:
return self.org_forward(x)
scale = self.multiplier
weight, bias = self.make_weight(scale, x.device)
kw_dict = self.kw_dict | {"weight": weight, "bias": bias}
return self.op(x, **kw_dict)
+262
View File
@@ -0,0 +1,262 @@
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from .base import LycorisBaseModule
from ..functional import tucker_weight_from_conv
class GLoRAModule(LycorisBaseModule):
name = "glora"
support_module = {
"linear",
"conv1d",
"conv2d",
"conv3d",
}
weight_list = [
"a1.weight",
"a2.weight",
"b1.weight",
"b2.weight",
"bm.weight",
"alpha",
]
weight_list_det = ["a1.weight"]
def __init__(
self,
lora_name,
org_module: nn.Module,
multiplier=1.0,
lora_dim=4,
alpha=1,
dropout=0.0,
rank_dropout=0.0,
module_dropout=0.0,
use_tucker=False,
use_scalar=False,
rank_dropout_scale=False,
weight_decompose=False,
bypass_mode=None,
rs_lora=False,
**kwargs,
):
"""
f(x) = WX + WAX + BX, where A and B are low-rank matrices
bypass_forward(x) = W(X+A(X)) + B(X)
bypass_forward_diff(x) = W(A(X)) + B(X)
get_merged_weight() = W + WA + B
get_diff_weight() = WA + B
"""
super().__init__(
lora_name,
org_module,
multiplier,
dropout,
rank_dropout,
module_dropout,
rank_dropout_scale,
bypass_mode,
)
if self.module_type not in self.support_module:
raise ValueError(f"{self.module_type} is not supported in GLoRA algo.")
self.lora_dim = lora_dim
self.tucker = False
self.rs_lora = rs_lora
if self.module_type.startswith("conv"):
self.isconv = True
# For general LoCon
in_dim = org_module.in_channels
k_size = org_module.kernel_size
stride = org_module.stride
padding = org_module.padding
out_dim = org_module.out_channels
use_tucker = use_tucker and all(i == 1 for i in k_size)
self.down_op = self.op
self.up_op = self.op
# A
self.a2 = self.module(in_dim, lora_dim, 1, bias=False)
self.a1 = self.module(lora_dim, in_dim, 1, bias=False)
# B
if use_tucker and any(i != 1 for i in k_size):
self.b2 = self.module(in_dim, lora_dim, 1, bias=False)
self.bm = self.module(
lora_dim, lora_dim, k_size, stride, padding, bias=False
)
self.tucker = True
else:
self.b2 = self.module(
in_dim, lora_dim, k_size, stride, padding, bias=False
)
self.b1 = self.module(lora_dim, out_dim, 1, bias=False)
else:
self.isconv = False
self.down_op = F.linear
self.up_op = F.linear
in_dim = org_module.in_features
out_dim = org_module.out_features
self.a2 = nn.Linear(in_dim, lora_dim, bias=False)
self.a1 = nn.Linear(lora_dim, in_dim, bias=False)
self.b2 = nn.Linear(in_dim, lora_dim, bias=False)
self.b1 = nn.Linear(lora_dim, out_dim, bias=False)
if type(alpha) == torch.Tensor:
alpha = alpha.detach().float().numpy() # without casting, bf16 causes error
alpha = lora_dim if alpha is None or alpha == 0 else alpha
r_factor = lora_dim
if self.rs_lora:
r_factor = math.sqrt(r_factor)
self.scale = alpha / r_factor
self.register_buffer("alpha", torch.tensor(alpha)) # 定数として扱える
if use_scalar:
self.scalar = nn.Parameter(torch.tensor(0.0))
else:
self.register_buffer("scalar", torch.tensor(1.0), persistent=False)
# same as microsoft's
torch.nn.init.kaiming_uniform_(self.a1.weight, a=math.sqrt(5))
torch.nn.init.kaiming_uniform_(self.b1.weight, a=math.sqrt(5))
if use_scalar:
torch.nn.init.kaiming_uniform_(self.a2.weight, a=math.sqrt(5))
torch.nn.init.kaiming_uniform_(self.b2.weight, a=math.sqrt(5))
else:
torch.nn.init.zeros_(self.a2.weight)
torch.nn.init.zeros_(self.b2.weight)
@classmethod
def make_module_from_state_dict(
cls, lora_name, orig_module, a1, a2, b1, b2, bm, alpha
):
module = cls(
lora_name,
orig_module,
1,
a2.size(0),
float(alpha),
use_tucker=bm is not None,
)
module.a1.weight.data.copy_(a1)
module.a2.weight.data.copy_(a2)
module.b1.weight.data.copy_(b1)
module.b2.weight.data.copy_(b2)
if bm is not None:
module.bm.weight.data.copy_(bm)
return module
def custom_state_dict(self):
destination = {}
destination["alpha"] = self.alpha
destination["a1.weight"] = self.a1.weight
destination["a2.weight"] = self.a2.weight * self.scalar
destination["b1.weight"] = self.b1.weight
destination["b2.weight"] = self.b2.weight * self.scalar
if self.tucker:
destination["bm.weight"] = self.bm.weight
return destination
def load_weight_hook(self, module: nn.Module, incompatible_keys):
missing_keys = incompatible_keys.missing_keys
for key in missing_keys:
if "scalar" in key:
del missing_keys[missing_keys.index(key)]
if isinstance(self.scalar, nn.Parameter):
self.scalar.data.copy_(torch.ones_like(self.scalar))
elif getattr(self, "scalar", None) is not None:
self.scalar.copy_(torch.ones_like(self.scalar))
else:
self.register_buffer(
"scalar", torch.ones_like(self.scalar), persistent=False
)
def make_weight(self, device=None):
wa1 = self.a1.weight.view(self.a1.weight.size(0), -1)
wa2 = self.a2.weight.view(self.a2.weight.size(0), -1)
orig = self.org_weight
if self.tucker:
wb = tucker_weight_from_conv(self.b1.weight, self.b2.weight, self.bm.weight)
else:
wb1 = self.b1.weight.view(self.b1.weight.size(0), -1)
wb2 = self.b2.weight.view(self.b2.weight.size(0), -1)
wb = wb1 @ wb2
wb = wb.view(*orig.shape)
if orig.dim() > 2:
w_wa1 = torch.einsum("o i ..., i j -> o j ...", orig, wa1)
w_wa2 = torch.einsum("o i ..., i j -> o j ...", w_wa1, wa2)
else:
w_wa2 = (orig @ wa1) @ wa2
return (wb + w_wa2) * self.scale * self.scalar
def get_diff_weight(self, multiplier=1.0, shape=None, device=None):
weight = self.make_weight(device) * multiplier
if shape is not None:
weight = weight.view(shape)
return weight, None
def get_merged_weight(self, multiplier=1, shape=None, device=None):
diff_w, _ = self.get_diff_weight(multiplier, shape, device)
return self.org_weight + diff_w, None
def _bypass_forward(self, x, scale=1, diff=False):
scale = self.scale * scale
ax_mid = self.a2(x) * scale
bx_mid = self.b2(x) * scale
if self.rank_dropout and self.training:
drop_a = (
torch.rand(self.lora_dim, device=ax_mid.device) < self.rank_dropout
).to(ax_mid.dtype)
drop_b = (
torch.rand(self.lora_dim, device=bx_mid.device) < self.rank_dropout
).to(bx_mid.dtype)
if self.rank_dropout_scale:
drop_a /= drop_a.mean()
drop_b /= drop_b.mean()
if (dims := len(x.shape)) == 4:
drop_a = drop_a.view(1, -1, 1, 1)
drop_b = drop_b.view(1, -1, 1, 1)
else:
drop_a = drop_a.view(*[1] * (dims - 1), -1)
drop_b = drop_b.view(*[1] * (dims - 1), -1)
ax_mid = ax_mid * drop_a
bx_mid = bx_mid * drop_b
return (
self.org_forward(
(0 if diff else x) + self.drop(self.a1(ax_mid)) * self.scale
)
+ self.drop(self.b1(bx_mid)) * self.scale
)
def bypass_forward_diff(self, x, scale=1):
return self._bypass_forward(x, scale=scale, diff=True)
def bypass_forward(self, x, scale=1):
return self._bypass_forward(x, scale=scale, diff=False)
def forward(self, x, *args, **kwargs):
if self.module_dropout and self.training:
if torch.rand(1) < self.module_dropout:
return self.org_forward(x)
if self.bypass_mode:
return self.bypass_forward(x, self.multiplier)
else:
weight = (
self.org_module[0].weight.data.to(self.dtype)
+ self.get_diff_weight(multiplier=self.multiplier)[0]
)
bias = (
None
if self.org_module[0].bias is None
else self.org_module[0].bias.data
)
return self.op(x, weight, bias, **self.kw_dict)
+142
View File
@@ -0,0 +1,142 @@
import torch
import torch.nn as nn
from .base import LycorisBaseModule
class IA3Module(LycorisBaseModule):
name = "ia3"
support_module = {
"linear",
"conv1d",
"conv2d",
"conv3d",
}
weight_list = ["weight", "on_input"]
weight_list_det = ["on_input"]
def __init__(
self,
lora_name,
org_module: nn.Module,
multiplier=1.0,
lora_dim=4,
alpha=1,
dropout=0.0,
rank_dropout=0.0,
module_dropout=0.0,
use_tucker=False,
use_scalar=False,
rank_dropout_scale=False,
weight_decompose=False,
bypass_mode=None,
rs_lora=False,
train_on_input=False,
**kwargs,
):
"""if alpha == 0 or None, alpha is rank (no scaling)."""
super().__init__(
lora_name,
org_module,
multiplier,
dropout,
rank_dropout,
module_dropout,
rank_dropout_scale,
bypass_mode,
)
if self.module_type not in self.support_module:
raise ValueError(f"{self.module_type} is not supported in IA^3 algo.")
if self.module_type.startswith("conv"):
self.isconv = True
in_dim = org_module.in_channels
out_dim = org_module.out_channels
if train_on_input:
train_dim = in_dim
else:
train_dim = out_dim
self.weight = nn.Parameter(
torch.empty(1, train_dim, *(1 for _ in self.shape[2:]))
)
else:
in_dim = org_module.in_features
out_dim = org_module.out_features
if train_on_input:
train_dim = in_dim
else:
train_dim = out_dim
self.weight = nn.Parameter(torch.empty(train_dim))
# Need more experiences on init method
torch.nn.init.constant_(self.weight, 0)
self.train_input = train_on_input
self.register_buffer("on_input", torch.tensor(int(train_on_input)))
@classmethod
def make_module_from_state_dict(cls, lora_name, orig_module, weight):
module = cls(
lora_name,
orig_module,
1,
)
module.weight.data.copy_(weight)
return module
def apply_to(self):
self.org_forward = self.org_module[0].forward
self.org_module[0].forward = self.forward
def make_weight(self, multiplier=1, shape=None, device=None, diff=False):
weight = self.weight * multiplier + int(not diff)
if self.train_input:
diff = self.org_weight * weight
else:
diff = self.org_weight.transpose(0, 1) * weight
diff = diff.transpose(0, 1)
if shape is not None:
diff = diff.view(shape)
if device is not None:
diff = diff.to(device)
return diff
def get_diff_weight(self, multiplier=1, shape=None, device=None):
diff = self.make_weight(
multiplier=multiplier, shape=shape, device=device, diff=True
)
return diff, None
def get_merged_weight(self, multiplier=1, shape=None, device=None):
diff = self.make_weight(multiplier=multiplier, shape=shape, device=device)
return diff, None
def _bypass_forward(self, x, scale=1, diff=False):
weight = self.weight * scale + int(not diff)
if self.train_input:
x = x * weight
out = self.org_forward(x)
if not self.train_input:
out = out * weight
return out
def bypass_forward_diff(self, x, scale=1):
return self._bypass_forward(x, scale, diff=True)
def bypass_forward(self, x, scale=1):
return self._bypass_forward(x, scale, diff=False)
def forward(self, x, *args, **kwargs):
if self.module_dropout and self.training:
if torch.rand(1) < self.module_dropout:
return self.org_forward(x)
if self.bypass_mode:
return self.bypass_forward(x, self.multiplier)
else:
weight = self.get_merged_weight(multiplier=self.multiplier)[0]
bias = (
None
if self.org_module[0].bias is None
else self.org_module[0].bias.data
)
return self.op(x, weight, bias, **self.kw_dict)
+332
View File
@@ -0,0 +1,332 @@
import math
from functools import cache
import torch
import torch.nn as nn
import torch.nn.functional as F
from .base import LycorisBaseModule
from ..functional.general import rebuild_tucker
from ..logging import logger
@cache
def log_wd():
return logger.warning(
"Using weight_decompose=True with LoRA (DoRA) will ignore network_dropout."
"Only rank dropout and module dropout will be applied"
)
class LoConModule(LycorisBaseModule):
name = "locon"
support_module = {
"linear",
"conv1d",
"conv2d",
"conv3d",
}
weight_list = [
"lora_up.weight",
"lora_down.weight",
"lora_mid.weight",
"alpha",
"dora_scale",
]
weight_list_det = ["lora_up.weight"]
def __init__(
self,
lora_name,
org_module: nn.Module,
multiplier=1.0,
lora_dim=4,
alpha=1,
dropout=0.0,
rank_dropout=0.0,
module_dropout=0.0,
use_tucker=False,
use_scalar=False,
rank_dropout_scale=False,
weight_decompose=False,
wd_on_out=False,
bypass_mode=None,
rs_lora=False,
**kwargs,
):
"""if alpha == 0 or None, alpha is rank (no scaling)."""
super().__init__(
lora_name,
org_module,
multiplier,
dropout,
rank_dropout,
module_dropout,
rank_dropout_scale,
bypass_mode,
)
if self.module_type not in self.support_module:
raise ValueError(f"{self.module_type} is not supported in LoRA/LoCon algo.")
self.lora_dim = lora_dim
self.tucker = False
self.rs_lora = rs_lora
if self.module_type.startswith("conv"):
self.isconv = True
# For general LoCon
in_dim = org_module.in_channels
k_size = org_module.kernel_size
stride = org_module.stride
padding = org_module.padding
out_dim = org_module.out_channels
use_tucker = use_tucker and any(i != 1 for i in k_size)
self.down_op = self.op
self.up_op = self.op
if use_tucker and any(i != 1 for i in k_size):
self.lora_down = self.module(in_dim, lora_dim, 1, bias=False)
self.lora_mid = self.module(
lora_dim, lora_dim, k_size, stride, padding, bias=False
)
self.tucker = True
else:
self.lora_down = self.module(
in_dim, lora_dim, k_size, stride, padding, bias=False
)
self.lora_up = self.module(lora_dim, out_dim, 1, bias=False)
elif isinstance(org_module, nn.Linear):
self.isconv = False
self.down_op = F.linear
self.up_op = F.linear
in_dim = org_module.in_features
out_dim = org_module.out_features
self.lora_down = nn.Linear(in_dim, lora_dim, bias=False)
self.lora_up = nn.Linear(lora_dim, out_dim, bias=False)
else:
raise NotImplementedError
self.wd = weight_decompose
self.wd_on_out = wd_on_out
if self.wd:
org_weight = org_module.weight.cpu().clone().float()
self.dora_norm_dims = org_weight.dim() - 1
if self.wd_on_out:
self.dora_scale = nn.Parameter(
torch.norm(
org_weight.reshape(org_weight.shape[0], -1),
dim=1,
keepdim=True,
).reshape(org_weight.shape[0], *[1] * self.dora_norm_dims)
).float()
else:
self.dora_scale = nn.Parameter(
torch.norm(
org_weight.transpose(1, 0).reshape(org_weight.shape[1], -1),
dim=1,
keepdim=True,
)
.reshape(org_weight.shape[1], *[1] * self.dora_norm_dims)
.transpose(1, 0)
).float()
if dropout:
self.dropout = nn.Dropout(dropout)
if self.wd:
log_wd()
else:
self.dropout = nn.Identity()
if type(alpha) == torch.Tensor:
alpha = alpha.detach().float().numpy() # without casting, bf16 causes error
alpha = lora_dim if alpha is None or alpha == 0 else alpha
r_factor = lora_dim
if self.rs_lora:
r_factor = math.sqrt(r_factor)
self.scale = alpha / r_factor
self.register_buffer("alpha", torch.tensor(alpha * (lora_dim / r_factor)))
if use_scalar:
self.scalar = nn.Parameter(torch.tensor(0.0))
else:
self.register_buffer("scalar", torch.tensor(1.0), persistent=False)
# same as microsoft's
torch.nn.init.kaiming_uniform_(self.lora_down.weight, a=math.sqrt(5))
if use_scalar:
torch.nn.init.kaiming_uniform_(self.lora_up.weight, a=math.sqrt(5))
else:
torch.nn.init.constant_(self.lora_up.weight, 0)
if self.tucker:
torch.nn.init.kaiming_uniform_(self.lora_mid.weight, a=math.sqrt(5))
@classmethod
def make_module_from_state_dict(
cls, lora_name, orig_module, up, down, mid, alpha, dora_scale
):
module = cls(
lora_name,
orig_module,
1,
down.size(0),
float(alpha),
use_tucker=mid is not None,
weight_decompose=dora_scale is not None,
)
module.lora_up.weight.data.copy_(up)
module.lora_down.weight.data.copy_(down)
if mid is not None:
module.lora_mid.weight.data.copy_(mid)
if dora_scale is not None:
module.dora_scale.copy_(dora_scale)
return module
def load_weight_hook(self, module: nn.Module, incompatible_keys):
missing_keys = incompatible_keys.missing_keys
for key in missing_keys:
if "scalar" in key:
del missing_keys[missing_keys.index(key)]
if isinstance(self.scalar, nn.Parameter):
self.scalar.data.copy_(torch.ones_like(self.scalar))
elif getattr(self, "scalar", None) is not None:
self.scalar.copy_(torch.ones_like(self.scalar))
else:
self.register_buffer(
"scalar", torch.ones_like(self.scalar), persistent=False
)
def make_weight(self, device=None):
wa = self.lora_up.weight.to(device)
wb = self.lora_down.weight.to(device)
if self.tucker:
t = self.lora_mid.weight
wa = wa.view(wa.size(0), -1).transpose(0, 1)
wb = wb.view(wb.size(0), -1)
weight = rebuild_tucker(t, wa, wb)
else:
weight = wa.view(wa.size(0), -1) @ wb.view(wb.size(0), -1)
weight = weight.view(self.shape)
if self.training and self.rank_dropout:
drop = (torch.rand(weight.size(0), device=device) > self.rank_dropout).to(
weight.dtype
)
drop = drop.view(-1, *[1] * len(weight.shape[1:]))
if self.rank_dropout_scale:
drop /= drop.mean()
weight *= drop
return weight * self.scalar.to(device)
def get_diff_weight(self, multiplier=1, shape=None, device=None):
scale = self.scale * multiplier
diff = self.make_weight(device=device) * scale
if shape is not None:
diff = diff.view(shape)
if device is not None:
diff = diff.to(device)
return diff, None
def get_merged_weight(self, multiplier=1, shape=None, device=None):
diff = self.get_diff_weight(multiplier=1, shape=shape, device=device)[0]
weight = self.org_weight
if self.wd:
merged = self.apply_weight_decompose(weight + diff, multiplier)
else:
merged = weight + diff * multiplier
return merged, None
def apply_weight_decompose(self, weight, multiplier=1):
weight = weight.to(self.dora_scale.dtype)
if self.wd_on_out:
weight_norm = (
weight.reshape(weight.shape[0], -1)
.norm(dim=1)
.reshape(weight.shape[0], *[1] * self.dora_norm_dims)
) + torch.finfo(weight.dtype).eps
else:
weight_norm = (
weight.transpose(0, 1)
.reshape(weight.shape[1], -1)
.norm(dim=1, keepdim=True)
.reshape(weight.shape[1], *[1] * self.dora_norm_dims)
.transpose(0, 1)
) + torch.finfo(weight.dtype).eps
scale = self.dora_scale.to(weight.device) / weight_norm
if multiplier != 1:
scale = multiplier * (scale - 1) + 1
return weight * scale
def custom_state_dict(self):
destination = {}
if self.wd:
destination["dora_scale"] = self.dora_scale
destination["alpha"] = self.alpha
destination["lora_up.weight"] = self.lora_up.weight * self.scalar
destination["lora_down.weight"] = self.lora_down.weight
if self.tucker:
destination["lora_mid.weight"] = self.lora_mid.weight
return destination
@torch.no_grad()
def apply_max_norm(self, max_norm, device=None):
orig_norm = self.make_weight(device).norm() * self.scale
norm = torch.clamp(orig_norm, max_norm / 2)
desired = torch.clamp(norm, max=max_norm)
ratio = desired.cpu() / norm.cpu()
scaled = norm != desired
if scaled:
self.scalar *= ratio
return scaled, orig_norm * ratio
def bypass_forward_diff(self, x, scale=1):
if self.tucker:
mid = self.lora_mid(self.lora_down(x))
else:
mid = self.lora_down(x)
if self.rank_dropout and self.training:
drop = (
torch.rand(self.lora_dim, device=mid.device) > self.rank_dropout
).to(mid.dtype)
if self.rank_dropout_scale:
drop /= drop.mean()
if (dims := len(x.shape)) == 4:
drop = drop.view(1, -1, 1, 1)
else:
drop = drop.view(*[1] * (dims - 1), -1)
mid = mid * drop
return self.dropout(self.lora_up(mid) * self.scalar * self.scale * scale)
def bypass_forward(self, x, scale=1):
return self.org_forward(x) + self.bypass_forward_diff(x, scale=scale)
def forward(self, x):
if self.module_dropout and self.training:
if torch.rand(1) < self.module_dropout:
return self.org_forward(x)
scale = self.scale
dtype = self.dtype
if not self.bypass_mode:
diff_weight = self.make_weight(x.device).to(dtype) * scale
weight = self.org_module[0].weight.data.to(dtype)
if self.wd:
weight = self.apply_weight_decompose(
weight + diff_weight, self.multiplier
)
else:
weight = weight + diff_weight * self.multiplier
bias = (
None
if self.org_module[0].bias is None
else self.org_module[0].bias.data
)
return self.op(x, weight, bias, **self.kw_dict)
else:
return self.bypass_forward(x, scale=self.multiplier)
+329
View File
@@ -0,0 +1,329 @@
import math
import torch
import torch.nn as nn
from .base import LycorisBaseModule
from ..functional.loha import diff_weight as loha_diff_weight
class LohaModule(LycorisBaseModule):
name = "loha"
support_module = {
"linear",
"conv1d",
"conv2d",
"conv3d",
}
weight_list = [
"hada_w1_a",
"hada_w1_b",
"hada_w2_a",
"hada_w2_b",
"hada_t1",
"hada_t2",
"alpha",
"dora_scale",
]
weight_list_det = ["hada_w1_a"]
def __init__(
self,
lora_name,
org_module: nn.Module,
multiplier=1.0,
lora_dim=4,
alpha=1,
dropout=0.0,
rank_dropout=0.0,
module_dropout=0.0,
use_tucker=False,
use_scalar=False,
rank_dropout_scale=False,
weight_decompose=False,
wd_on_out=False,
bypass_mode=None,
rs_lora=False,
**kwargs,
):
super().__init__(
lora_name,
org_module,
multiplier,
dropout,
rank_dropout,
module_dropout,
rank_dropout_scale,
bypass_mode,
)
if self.module_type not in self.support_module:
raise ValueError(f"{self.module_type} is not supported in LoHa algo.")
self.lora_name = lora_name
self.lora_dim = lora_dim
self.tucker = False
self.rs_lora = rs_lora
w_shape = self.shape
if self.module_type.startswith("conv"):
in_dim = org_module.in_channels
k_size = org_module.kernel_size
out_dim = org_module.out_channels
self.shape = (out_dim, in_dim, *k_size)
self.tucker = use_tucker and any(i != 1 for i in k_size)
if self.tucker:
w_shape = (out_dim, in_dim, *k_size)
else:
w_shape = (out_dim, in_dim * torch.tensor(k_size).prod().item())
if self.tucker:
self.hada_t1 = nn.Parameter(torch.empty(lora_dim, lora_dim, *w_shape[2:]))
self.hada_w1_a = nn.Parameter(
torch.empty(lora_dim, w_shape[0])
) # out_dim, 1-mode
self.hada_w1_b = nn.Parameter(
torch.empty(lora_dim, w_shape[1])
) # in_dim , 2-mode
self.hada_t2 = nn.Parameter(torch.empty(lora_dim, lora_dim, *w_shape[2:]))
self.hada_w2_a = nn.Parameter(
torch.empty(lora_dim, w_shape[0])
) # out_dim, 1-mode
self.hada_w2_b = nn.Parameter(
torch.empty(lora_dim, w_shape[1])
) # in_dim , 2-mode
else:
self.hada_w1_a = nn.Parameter(torch.empty(w_shape[0], lora_dim))
self.hada_w1_b = nn.Parameter(torch.empty(lora_dim, w_shape[1]))
self.hada_w2_a = nn.Parameter(torch.empty(w_shape[0], lora_dim))
self.hada_w2_b = nn.Parameter(torch.empty(lora_dim, w_shape[1]))
self.wd = weight_decompose
self.wd_on_out = wd_on_out
if self.wd:
org_weight = org_module.weight.cpu().clone().float()
self.dora_norm_dims = org_weight.dim() - 1
if self.wd_on_out:
self.dora_scale = nn.Parameter(
torch.norm(
org_weight.reshape(org_weight.shape[0], -1),
dim=1,
keepdim=True,
).reshape(org_weight.shape[0], *[1] * self.dora_norm_dims)
).float()
else:
self.dora_scale = nn.Parameter(
torch.norm(
org_weight.transpose(1, 0).reshape(org_weight.shape[1], -1),
dim=1,
keepdim=True,
)
.reshape(org_weight.shape[1], *[1] * self.dora_norm_dims)
.transpose(1, 0)
).float()
if self.dropout:
print("[WARN]LoHa/LoKr haven't implemented normal dropout yet.")
if type(alpha) == torch.Tensor:
alpha = alpha.detach().float().numpy() # without casting, bf16 causes error
alpha = lora_dim if alpha is None or alpha == 0 else alpha
r_factor = lora_dim
if self.rs_lora:
r_factor = math.sqrt(r_factor)
self.scale = alpha / r_factor
self.register_buffer("alpha", torch.tensor(alpha * (lora_dim / r_factor)))
if use_scalar:
self.scalar = nn.Parameter(torch.tensor(0.0))
else:
self.register_buffer("scalar", torch.tensor(1.0), persistent=False)
# Need more experiments on init method
if self.tucker:
torch.nn.init.normal_(self.hada_t1, std=0.1)
torch.nn.init.normal_(self.hada_t2, std=0.1)
torch.nn.init.normal_(self.hada_w1_b, std=1)
torch.nn.init.normal_(self.hada_w1_a, std=0.1)
torch.nn.init.normal_(self.hada_w2_b, std=1)
if use_scalar:
torch.nn.init.normal_(self.hada_w2_a, std=0.1)
else:
torch.nn.init.constant_(self.hada_w2_a, 0)
@classmethod
def make_module_from_state_dict(
cls, lora_name, orig_module, w1a, w1b, w2a, w2b, t1, t2, alpha, dora_scale
):
module = cls(
lora_name,
orig_module,
1,
w1b.size(0),
float(alpha),
use_tucker=t1 is not None,
weight_decompose=dora_scale is not None,
)
module.hada_w1_a.copy_(w1a)
module.hada_w1_b.copy_(w1b)
module.hada_w2_a.copy_(w2a)
module.hada_w2_b.copy_(w2b)
if t1 is not None:
module.hada_t1.copy_(t1)
module.hada_t2.copy_(t2)
if dora_scale is not None:
module.dora_scale.copy_(dora_scale)
return module
def load_weight_hook(self, module: nn.Module, incompatible_keys):
missing_keys = incompatible_keys.missing_keys
for key in missing_keys:
if "scalar" in key:
del missing_keys[missing_keys.index(key)]
if isinstance(self.scalar, nn.Parameter):
self.scalar.data.copy_(torch.ones_like(self.scalar))
elif getattr(self, "scalar", None) is not None:
self.scalar.copy_(torch.ones_like(self.scalar))
else:
self.register_buffer(
"scalar", torch.ones_like(self.scalar), persistent=False
)
def get_weight(self, shape):
scale = torch.tensor(
self.scale, dtype=self.hada_w1_b.dtype, device=self.hada_w1_b.device
)
if self.tucker:
weight = loha_diff_weight(
self.hada_w1_b,
self.hada_w1_a,
self.hada_w2_b,
self.hada_w2_a,
self.hada_t1,
self.hada_t2,
gamma=scale,
)
else:
weight = loha_diff_weight(
self.hada_w1_b,
self.hada_w1_a,
self.hada_w2_b,
self.hada_w2_a,
None,
None,
gamma=scale,
)
if shape is not None:
weight = weight.reshape(shape)
if self.training and self.rank_dropout:
drop = (torch.rand(weight.size(0)) > self.rank_dropout).to(weight.dtype)
drop = drop.view(-1, *[1] * len(weight.shape[1:])).to(weight.device)
if self.rank_dropout_scale:
drop /= drop.mean()
weight *= drop
return weight
def get_diff_weight(self, multiplier=1, shape=None, device=None):
scale = self.scale * multiplier
diff = self.get_weight(shape) * scale
if device is not None:
diff = diff.to(device)
return diff, None
def get_merged_weight(self, multiplier=1, shape=None, device=None):
diff = self.get_diff_weight(multiplier=1, shape=shape, device=device)[0]
weight = self.org_weight
if self.wd:
merged = self.apply_weight_decompose(weight + diff, multiplier)
else:
merged = weight + diff * multiplier
return merged, None
def apply_weight_decompose(self, weight, multiplier=1):
weight = weight.to(self.dora_scale.dtype)
if self.wd_on_out:
weight_norm = (
weight.reshape(weight.shape[0], -1)
.norm(dim=1)
.reshape(weight.shape[0], *[1] * self.dora_norm_dims)
) + torch.finfo(weight.dtype).eps
else:
weight_norm = (
weight.transpose(0, 1)
.reshape(weight.shape[1], -1)
.norm(dim=1, keepdim=True)
.reshape(weight.shape[1], *[1] * self.dora_norm_dims)
.transpose(0, 1)
) + torch.finfo(weight.dtype).eps
scale = self.dora_scale.to(weight.device) / weight_norm
if multiplier != 1:
scale = multiplier * (scale - 1) + 1
return weight * scale
def custom_state_dict(self):
destination = {}
destination["alpha"] = self.alpha
if self.wd:
destination["dora_scale"] = self.dora_scale
destination["hada_w1_a"] = self.hada_w1_a * self.scalar
destination["hada_w1_b"] = self.hada_w1_b
destination["hada_w2_a"] = self.hada_w2_a
destination["hada_w2_b"] = self.hada_w2_b
if self.tucker:
destination["hada_t1"] = self.hada_t1
destination["hada_t2"] = self.hada_t2
return destination
@torch.no_grad()
def apply_max_norm(self, max_norm, device=None):
orig_norm = (self.get_weight(self.shape) * self.scalar).norm()
norm = torch.clamp(orig_norm, max_norm / 2)
desired = torch.clamp(norm, max=max_norm)
ratio = desired.cpu() / norm.cpu()
scaled = norm != desired
if scaled:
self.scalar *= ratio
return scaled, orig_norm * ratio
def bypass_forward_diff(self, x, scale=1):
diff_weight = self.get_weight(self.shape) * self.scalar * scale
return self.drop(self.op(x, diff_weight, **self.kw_dict))
def bypass_forward(self, x, scale=1):
return self.org_forward(x) + self.bypass_forward_diff(x, scale=scale)
def forward(self, x: torch.Tensor, *args, **kwargs):
if self.module_dropout and self.training:
if torch.rand(1) < self.module_dropout:
return self.op(
x,
self.org_module[0].weight.data,
(
None
if self.org_module[0].bias is None
else self.org_module[0].bias.data
),
)
if self.bypass_mode:
return self.bypass_forward(x, scale=self.multiplier)
else:
diff_weight = self.get_weight(self.shape).to(self.dtype) * self.scalar
weight = self.org_module[0].weight.data.to(self.dtype)
if self.wd:
weight = self.apply_weight_decompose(
weight + diff_weight, self.multiplier
)
else:
weight = weight + diff_weight * self.multiplier
bias = (
None
if self.org_module[0].bias is None
else self.org_module[0].bias.data
)
return self.op(x, weight, bias, **self.kw_dict)
+609
View File
@@ -0,0 +1,609 @@
import math
from functools import cache
import torch
import torch.nn as nn
import torch.nn.functional as F
from .base import LycorisBaseModule
from ..functional import factorization, rebuild_tucker
from ..functional.lokr import make_kron
from ..logging import logger
@cache
def logging_force_full_matrix(lora_dim, dim, factor):
logger.warning(
f"lora_dim {lora_dim} is too large for"
f" dim={dim} and {factor=}"
", using full matrix mode."
)
class LokrModule(LycorisBaseModule):
name = "kron"
support_module = {
"linear",
"conv1d",
"conv2d",
"conv3d",
}
weight_list = [
"lokr_w1",
"lokr_w1_a",
"lokr_w1_b",
"lokr_w2",
"lokr_w2_a",
"lokr_w2_b",
"lokr_t1",
"lokr_t2",
"alpha",
"dora_scale",
]
weight_list_det = ["lokr_w1", "lokr_w1_a"]
def __init__(
self,
lora_name,
org_module: nn.Module,
multiplier=1.0,
lora_dim=4,
alpha=1,
dropout=0.0,
rank_dropout=0.0,
module_dropout=0.0,
use_tucker=False,
use_scalar=False,
decompose_both=False,
factor: int = -1, # factorization factor
rank_dropout_scale=False,
weight_decompose=False,
wd_on_out=False,
full_matrix=False,
bypass_mode=None,
rs_lora=False,
unbalanced_factorization=False,
**kwargs,
):
super().__init__(
lora_name,
org_module,
multiplier,
dropout,
rank_dropout,
module_dropout,
rank_dropout_scale,
bypass_mode,
)
if self.module_type not in self.support_module:
raise ValueError(f"{self.module_type} is not supported in LoKr algo.")
factor = int(factor)
self.lora_dim = lora_dim
self.tucker = False
self.use_w1 = False
self.use_w2 = False
self.full_matrix = full_matrix
self.rs_lora = rs_lora
if self.module_type.startswith("conv"):
in_dim = org_module.in_channels
k_size = org_module.kernel_size
out_dim = org_module.out_channels
self.shape = (out_dim, in_dim, *k_size)
in_m, in_n = factorization(in_dim, factor)
out_l, out_k = factorization(out_dim, factor)
if unbalanced_factorization:
out_l, out_k = out_k, out_l
shape = ((out_l, out_k), (in_m, in_n), *k_size) # ((a, b), (c, d), *k_size)
self.tucker = use_tucker and any(i != 1 for i in k_size)
if (
decompose_both
and lora_dim < max(shape[0][0], shape[1][0]) / 2
and not self.full_matrix
):
self.lokr_w1_a = nn.Parameter(torch.empty(shape[0][0], lora_dim))
self.lokr_w1_b = nn.Parameter(torch.empty(lora_dim, shape[1][0]))
else:
self.use_w1 = True
self.lokr_w1 = nn.Parameter(
torch.empty(shape[0][0], shape[1][0])
) # a*c, 1-mode
if lora_dim >= max(shape[0][1], shape[1][1]) / 2 or self.full_matrix:
if not self.full_matrix:
logging_force_full_matrix(lora_dim, max(in_dim, out_dim), factor)
self.use_w2 = True
self.lokr_w2 = nn.Parameter(
torch.empty(shape[0][1], shape[1][1], *k_size)
)
elif self.tucker:
self.lokr_t2 = nn.Parameter(torch.empty(lora_dim, lora_dim, *shape[2:]))
self.lokr_w2_a = nn.Parameter(
torch.empty(lora_dim, shape[0][1])
) # b, 1-mode
self.lokr_w2_b = nn.Parameter(
torch.empty(lora_dim, shape[1][1])
) # d, 2-mode
else: # Conv2d not tucker
# bigger part. weight and LoRA. [b, dim] x [dim, d*k1*k2]
self.lokr_w2_a = nn.Parameter(torch.empty(shape[0][1], lora_dim))
self.lokr_w2_b = nn.Parameter(
torch.empty(
lora_dim, shape[1][1] * torch.tensor(shape[2:]).prod().item()
)
)
# w1 ⊗ (w2_a x w2_b) = (a, b)⊗((c, dim)x(dim, d*k1*k2)) = (a, b)⊗(c, d*k1*k2) = (ac, bd*k1*k2)
else: # Linear
in_dim = org_module.in_features
out_dim = org_module.out_features
self.shape = (out_dim, in_dim)
in_m, in_n = factorization(in_dim, factor)
out_l, out_k = factorization(out_dim, factor)
if unbalanced_factorization:
out_l, out_k = out_k, out_l
shape = (
(out_l, out_k),
(in_m, in_n),
) # ((a, b), (c, d)), out_dim = a*c, in_dim = b*d
# smaller part. weight scale
if (
decompose_both
and lora_dim < max(shape[0][0], shape[1][0]) / 2
and not self.full_matrix
):
self.lokr_w1_a = nn.Parameter(torch.empty(shape[0][0], lora_dim))
self.lokr_w1_b = nn.Parameter(torch.empty(lora_dim, shape[1][0]))
else:
self.use_w1 = True
self.lokr_w1 = nn.Parameter(
torch.empty(shape[0][0], shape[1][0])
) # a*c, 1-mode
if lora_dim < max(shape[0][1], shape[1][1]) / 2 and not self.full_matrix:
# bigger part. weight and LoRA. [b, dim] x [dim, d]
self.lokr_w2_a = nn.Parameter(torch.empty(shape[0][1], lora_dim))
self.lokr_w2_b = nn.Parameter(torch.empty(lora_dim, shape[1][1]))
# w1 ⊗ (w2_a x w2_b) = (a, b)⊗((c, dim)x(dim, d)) = (a, b)⊗(c, d) = (ac, bd)
else:
if not self.full_matrix:
logging_force_full_matrix(lora_dim, max(in_dim, out_dim), factor)
self.use_w2 = True
self.lokr_w2 = nn.Parameter(torch.empty(shape[0][1], shape[1][1]))
self.wd = weight_decompose
self.wd_on_out = wd_on_out
if self.wd:
org_weight = org_module.weight.cpu().clone().float()
self.dora_norm_dims = org_weight.dim() - 1
if self.wd_on_out:
self.dora_scale = nn.Parameter(
torch.norm(
org_weight.reshape(org_weight.shape[0], -1),
dim=1,
keepdim=True,
).reshape(org_weight.shape[0], *[1] * self.dora_norm_dims)
).float()
else:
self.dora_scale = nn.Parameter(
torch.norm(
org_weight.transpose(1, 0).reshape(org_weight.shape[1], -1),
dim=1,
keepdim=True,
)
.reshape(org_weight.shape[1], *[1] * self.dora_norm_dims)
.transpose(1, 0)
).float()
self.dropout = dropout
if dropout:
print("[WARN]LoHa/LoKr haven't implemented normal dropout yet.")
self.rank_dropout = rank_dropout
self.rank_dropout_scale = rank_dropout_scale
self.module_dropout = module_dropout
if isinstance(alpha, torch.Tensor):
alpha = alpha.detach().float().numpy() # without casting, bf16 causes error
alpha = lora_dim if alpha is None or alpha == 0 else alpha
if self.use_w2 and self.use_w1:
# use scale = 1
alpha = lora_dim
r_factor = lora_dim
if self.rs_lora:
r_factor = math.sqrt(r_factor)
self.scale = alpha / r_factor
self.register_buffer("alpha", torch.tensor(alpha * (lora_dim / r_factor)))
if use_scalar:
self.scalar = nn.Parameter(torch.tensor(0.0))
else:
self.register_buffer("scalar", torch.tensor(1.0), persistent=False)
if self.use_w2:
if use_scalar:
torch.nn.init.kaiming_uniform_(self.lokr_w2, a=math.sqrt(5))
else:
torch.nn.init.constant_(self.lokr_w2, 0)
else:
if self.tucker:
torch.nn.init.kaiming_uniform_(self.lokr_t2, a=math.sqrt(5))
torch.nn.init.kaiming_uniform_(self.lokr_w2_a, a=math.sqrt(5))
if use_scalar:
torch.nn.init.kaiming_uniform_(self.lokr_w2_b, a=math.sqrt(5))
else:
torch.nn.init.constant_(self.lokr_w2_b, 0)
if self.use_w1:
torch.nn.init.kaiming_uniform_(self.lokr_w1, a=math.sqrt(5))
else:
torch.nn.init.kaiming_uniform_(self.lokr_w1_a, a=math.sqrt(5))
torch.nn.init.kaiming_uniform_(self.lokr_w1_b, a=math.sqrt(5))
@classmethod
def make_module_from_state_dict(
cls,
lora_name,
orig_module,
w1,
w1a,
w1b,
w2,
w2a,
w2b,
_,
t2,
alpha,
dora_scale,
):
full_matrix = False
if w1a is not None:
lora_dim = w1a.size(1)
elif w2a is not None:
lora_dim = w2a.size(1)
else:
full_matrix = True
lora_dim = 1
if w1 is None:
out_dim = w1a.size(0)
in_dim = w1b.size(1)
else:
out_dim, in_dim = w1.shape
shape_s = [out_dim, in_dim]
if w2 is None:
out_dim *= w2a.size(0)
in_dim *= w2b.size(1)
else:
out_dim *= w2.size(0)
in_dim *= w2.size(1)
if (
shape_s[0] == factorization(out_dim, -1)[0]
and shape_s[1] == factorization(in_dim, -1)[0]
):
factor = -1
else:
w1_shape = w1.shape if w1 is not None else (w1a.size(0), w1b.size(1))
w2_shape = w2.shape if w2 is not None else (w2a.size(0), w2b.size(1))
shape_group_1 = (w1_shape[0], w2_shape[0])
shape_group_2 = (w1_shape[1], w2_shape[1])
w_shape = (w1_shape[0] * w2_shape[0], w1_shape[1] * w2_shape[1])
factor1 = max(w1.shape) if w1 is not None else max(w1a.size(0), w1b.size(1))
factor2 = max(w2.shape) if w2 is not None else max(w2a.size(0), w2b.size(1))
if (
w_shape[0] % factor1 == 0
and w_shape[1] % factor1 == 0
and factor1 in shape_group_1
and factor1 in shape_group_2
):
factor = factor1
elif (
w_shape[0] % factor2 == 0
and w_shape[1] % factor2 == 0
and factor2 in shape_group_1
and factor2 in shape_group_2
):
factor = factor2
else:
factor = min(factor1, factor2)
module = cls(
lora_name,
orig_module,
1,
lora_dim,
float(alpha),
use_tucker=t2 is not None,
decompose_both=w1 is None and w2 is None,
factor=factor,
weight_decompose=dora_scale is not None,
full_matrix=full_matrix,
)
if w1 is not None:
module.lokr_w1.copy_(w1)
else:
module.lokr_w1_a.copy_(w1a)
module.lokr_w1_b.copy_(w1b)
if w2 is not None:
module.lokr_w2.copy_(w2)
else:
module.lokr_w2_a.copy_(w2a)
module.lokr_w2_b.copy_(w2b)
if t2 is not None:
module.lokr_t2.copy_(t2)
if dora_scale is not None:
module.dora_scale.copy_(dora_scale)
return module
def load_weight_hook(self, module: nn.Module, incompatible_keys):
missing_keys = incompatible_keys.missing_keys
for key in missing_keys:
if "scalar" in key:
del missing_keys[missing_keys.index(key)]
if isinstance(self.scalar, nn.Parameter):
self.scalar.data.copy_(torch.ones_like(self.scalar))
elif getattr(self, "scalar", None) is not None:
self.scalar.copy_(torch.ones_like(self.scalar))
else:
self.register_buffer(
"scalar", torch.ones_like(self.scalar), persistent=False
)
def get_weight(self, shape):
weight = make_kron(
self.lokr_w1 if self.use_w1 else self.lokr_w1_a @ self.lokr_w1_b,
(
self.lokr_w2
if self.use_w2
else (
rebuild_tucker(self.lokr_t2, self.lokr_w2_a, self.lokr_w2_b)
if self.tucker
else self.lokr_w2_a @ self.lokr_w2_b
)
),
self.scale,
)
dtype = weight.dtype
if shape is not None:
weight = weight.view(shape)
if self.training and self.rank_dropout:
drop = (torch.rand(weight.size(0)) > self.rank_dropout).to(dtype)
drop = drop.view(-1, *[1] * len(weight.shape[1:]))
if self.rank_dropout_scale:
drop /= drop.mean()
weight *= drop
return weight
def get_diff_weight(self, multiplier=1, shape=None, device=None):
scale = self.scale * multiplier
diff = self.get_weight(shape) * scale
if device is not None:
diff = diff.to(device)
return diff, None
def get_merged_weight(self, multiplier=1, shape=None, device=None):
diff = self.get_diff_weight(multiplier=1, shape=shape, device=device)[0]
weight = self.org_weight
if self.wd:
merged = self.apply_weight_decompose(weight + diff, multiplier)
else:
merged = weight + diff * multiplier
return merged, None
def apply_weight_decompose(self, weight, multiplier=1):
weight = weight.to(self.dora_scale.dtype)
if self.wd_on_out:
weight_norm = (
weight.reshape(weight.shape[0], -1)
.norm(dim=1)
.reshape(weight.shape[0], *[1] * self.dora_norm_dims)
) + torch.finfo(weight.dtype).eps
else:
weight_norm = (
weight.transpose(0, 1)
.reshape(weight.shape[1], -1)
.norm(dim=1, keepdim=True)
.reshape(weight.shape[1], *[1] * self.dora_norm_dims)
.transpose(0, 1)
) + torch.finfo(weight.dtype).eps
scale = self.dora_scale.to(weight.device) / weight_norm
if multiplier != 1:
scale = multiplier * (scale - 1) + 1
return weight * scale
def custom_state_dict(self):
destination = {}
destination["alpha"] = self.alpha
if self.wd:
destination["dora_scale"] = self.dora_scale
if self.use_w1:
destination["lokr_w1"] = self.lokr_w1 * self.scalar
else:
destination["lokr_w1_a"] = self.lokr_w1_a * self.scalar
destination["lokr_w1_b"] = self.lokr_w1_b
if self.use_w2:
destination["lokr_w2"] = self.lokr_w2
else:
destination["lokr_w2_a"] = self.lokr_w2_a
destination["lokr_w2_b"] = self.lokr_w2_b
if self.tucker:
destination["lokr_t2"] = self.lokr_t2
return destination
@torch.no_grad()
def apply_max_norm(self, max_norm, device=None):
orig_norm = self.get_weight(self.shape).norm()
norm = torch.clamp(orig_norm, max_norm / 2)
desired = torch.clamp(norm, max=max_norm)
ratio = desired.cpu() / norm.cpu()
scaled = norm != desired
if scaled:
modules = 4 - self.use_w1 - self.use_w2 + (not self.use_w2 and self.tucker)
if self.use_w1:
self.lokr_w1 *= ratio ** (1 / modules)
else:
self.lokr_w1_a *= ratio ** (1 / modules)
self.lokr_w1_b *= ratio ** (1 / modules)
if self.use_w2:
self.lokr_w2 *= ratio ** (1 / modules)
else:
if self.tucker:
self.lokr_t2 *= ratio ** (1 / modules)
self.lokr_w2_a *= ratio ** (1 / modules)
self.lokr_w2_b *= ratio ** (1 / modules)
return scaled, orig_norm * ratio
def bypass_forward_diff(self, h, scale=1):
is_conv = self.module_type.startswith("conv")
if self.use_w2:
ba = self.lokr_w2
else:
a = self.lokr_w2_b
b = self.lokr_w2_a
if self.tucker:
t = self.lokr_t2
a = a.view(*a.shape, *[1] * (len(t.shape) - 2))
b = b.view(*b.shape, *[1] * (len(t.shape) - 2))
elif is_conv:
a = a.view(*a.shape, *self.shape[2:])
b = b.view(*b.shape, *[1] * (len(self.shape) - 2))
if self.use_w1:
c = self.lokr_w1
else:
c = self.lokr_w1_a @ self.lokr_w1_b
uq = c.size(1)
if is_conv:
# (b, uq), vq, ...
b, _, *rest = h.shape
h_in_group = h.reshape(b * uq, -1, *rest)
else:
# b, ..., uq, vq
h_in_group = h.reshape(*h.shape[:-1], uq, -1)
if self.use_w2:
hb = self.op(h_in_group, ba, **self.kw_dict)
else:
if is_conv:
if self.tucker:
ha = self.op(h_in_group, a)
ht = self.op(ha, t, **self.kw_dict)
hb = self.op(ht, b)
else:
ha = self.op(h_in_group, a, **self.kw_dict)
hb = self.op(ha, b)
else:
ha = self.op(h_in_group, a, **self.kw_dict)
hb = self.op(ha, b)
if is_conv:
# (b, uq), vp, ..., f
# -> b, uq, vp, ..., f
# -> b, f, vp, ..., uq
hb = hb.view(b, -1, *hb.shape[1:])
h_cross_group = hb.transpose(1, -1)
else:
# b, ..., uq, vq
# -> b, ..., vq, uq
h_cross_group = hb.transpose(-1, -2)
hc = F.linear(h_cross_group, c)
if is_conv:
# b, f, vp, ..., up
# -> b, up, vp, ... ,f
# -> b, c, ..., f
hc = hc.transpose(1, -1)
h = hc.reshape(b, -1, *hc.shape[3:])
else:
# b, ..., vp, up
# -> b, ..., up, vp
# -> b, ..., c
hc = hc.transpose(-1, -2)
h = hc.reshape(*hc.shape[:-2], -1)
return self.drop(h * scale * self.scalar)
def bypass_forward(self, x, scale=1):
return self.org_forward(x) + self.bypass_forward_diff(x, scale=scale)
def forward(self, x: torch.Tensor, *args, **kwargs):
if self.module_dropout and self.training:
if torch.rand(1) < self.module_dropout:
return self.org_forward(x)
if self.bypass_mode:
return self.bypass_forward(x, self.multiplier)
else:
diff_weight = self.get_weight(self.shape).to(self.dtype) * self.scalar
weight = self.org_module[0].weight.data.to(self.dtype)
if self.wd:
weight = self.apply_weight_decompose(
weight + diff_weight, self.multiplier
)
elif self.multiplier == 1:
weight = weight + diff_weight
else:
weight = weight + diff_weight * self.multiplier
bias = (
None
if self.org_module[0].bias is None
else self.org_module[0].bias.data
)
return self.op(x, weight, bias, **self.kw_dict)
if __name__ == "__main__":
base = nn.Conv2d(128, 128, 3, 1, 1)
net = LokrModule(
"",
base,
multiplier=1,
lora_dim=4,
alpha=1,
weight_decompose=False,
use_tucker=False,
use_scalar=False,
decompose_both=True,
)
net.apply_to()
sd = net.state_dict()
for key in sd:
if key != "alpha":
sd[key] = torch.randn_like(sd[key])
net.load_state_dict(sd)
test_input = torch.randn(1, 128, 16, 16)
test_output = net(test_input)
print(test_output.shape)
net2 = LokrModule(
"",
base,
multiplier=1,
lora_dim=4,
alpha=1,
weight_decompose=False,
use_tucker=False,
use_scalar=False,
bypass_mode=True,
decompose_both=True,
)
net2.apply_to()
net2.load_state_dict(sd)
print(net2)
test_output2 = net(test_input)
print(F.mse_loss(test_output, test_output2))
+161
View File
@@ -0,0 +1,161 @@
import torch
import torch.nn as nn
from .base import LycorisBaseModule
from ..logging import warning_once
class NormModule(LycorisBaseModule):
name = "norm"
support_module = {
"layernorm",
"groupnorm",
}
weight_list = ["w_norm", "b_norm"]
weight_list_det = ["w_norm"]
def __init__(
self,
lora_name,
org_module: nn.Module,
multiplier=1.0,
rank_dropout=0.0,
module_dropout=0.0,
rank_dropout_scale=False,
**kwargs,
):
"""if alpha == 0 or None, alpha is rank (no scaling)."""
super().__init__(
lora_name=lora_name,
org_module=org_module,
multiplier=multiplier,
rank_dropout=rank_dropout,
module_dropout=module_dropout,
rank_dropout_scale=rank_dropout_scale,
**kwargs,
)
if self.module_type == "unknown":
if not hasattr(org_module, "weight") or not hasattr(org_module, "_norm"):
warning_once(f"{type(org_module)} is not supported in Norm algo.")
self.not_supported = True
return
else:
self.dim = org_module.weight.numel()
self.not_supported = False
elif self.module_type not in self.support_module:
warning_once(f"{self.module_type} is not supported in Norm algo.")
self.not_supported = True
return
self.w_norm = nn.Parameter(torch.zeros(self.dim))
if hasattr(org_module, "bias"):
self.b_norm = nn.Parameter(torch.zeros(self.dim))
if hasattr(org_module, "_norm"):
self.org_norm = org_module._norm
else:
self.org_norm = None
@classmethod
def make_module_from_state_dict(cls, lora_name, orig_module, w_norm, b_norm):
module = cls(
lora_name,
orig_module,
1,
)
module.w_norm.copy_(w_norm)
if b_norm is not None:
module.b_norm.copy_(b_norm)
return module
def make_weight(self, scale=1, device=None):
org_weight = self.org_module[0].weight.to(device, dtype=self.w_norm.dtype)
if hasattr(self.org_module[0], "bias"):
org_bias = self.org_module[0].bias.to(device, dtype=self.b_norm.dtype)
else:
org_bias = None
if self.rank_dropout and self.training:
drop = (torch.rand(self.dim, device=device) < self.rank_dropout).to(
self.w_norm.device
)
if self.rank_dropout_scale:
drop /= drop.mean()
else:
drop = 1
drop = (
torch.rand(self.dim, device=device) < self.rank_dropout
if self.rank_dropout and self.training
else 1
)
weight = self.w_norm.to(device) * drop * scale
if org_bias is not None:
bias = self.b_norm.to(device) * drop * scale
return org_weight + weight, org_bias + bias if org_bias is not None else None
def get_diff_weight(self, multiplier=1, shape=None, device=None):
if self.not_supported:
return 0, 0
w = self.w_norm * multiplier
if device is not None:
w = w.to(device)
if shape is not None:
w = w.view(shape)
if self.b_norm is not None:
b = self.b_norm * multiplier
if device is not None:
b = b.to(device)
if shape is not None:
b = b.view(shape)
else:
b = None
return w, b
def get_merged_weight(self, multiplier=1, shape=None, device=None):
if self.not_supported:
return None, None
diff_w, diff_b = self.get_diff_weight(multiplier, shape, device)
org_w = self.org_module[0].weight.to(device, dtype=self.w_norm.dtype)
weight = org_w + diff_w
if diff_b is not None:
org_b = self.org_module[0].bias.to(device, dtype=self.b_norm.dtype)
bias = org_b + diff_b
else:
bias = None
return weight, bias
def forward(self, x):
if self.not_supported or (
self.module_dropout
and self.training
and torch.rand(1) < self.module_dropout
):
return self.org_forward(x)
scale = self.multiplier
w, b = self.make_weight(scale, x.device)
if self.org_norm is not None:
normed = self.org_norm(x)
scaled = normed * w
if b is not None:
scaled += b
return scaled
kw_dict = self.kw_dict | {"weight": w, "bias": b}
return self.op(x, **kw_dict)
if __name__ == "__main__":
base = nn.LayerNorm(128).cuda()
norm = NormModule("test", base, 1).cuda()
print(norm)
test_input = torch.randn(1, 128).cuda()
test_output = norm(test_input)
torch.sum(test_output).backward()
print(test_output.shape)
base = nn.GroupNorm(4, 128).cuda()
norm = NormModule("test", base, 1).cuda()
print(norm)
test_input = torch.randn(1, 128, 3, 3).cuda()
test_output = norm(test_input)
torch.sum(test_output).backward()
print(test_output.shape)
+483
View File
@@ -0,0 +1,483 @@
import re
import hashlib
from io import BytesIO
from typing import Dict, Tuple, Union
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.linalg as linalg
import safetensors.torch
from tqdm import tqdm
from .general import *
def load_bytes_in_safetensors(tensors):
bytes = safetensors.torch.save(tensors)
b = BytesIO(bytes)
b.seek(0)
header = b.read(8)
n = int.from_bytes(header, "little")
offset = n + 8
b.seek(offset)
return b.read()
def precalculate_safetensors_hashes(state_dict):
# calculate each tensor one by one to reduce memory usage
hash_sha256 = hashlib.sha256()
for tensor in state_dict.values():
single_tensor_sd = {"tensor": tensor}
bytes_for_tensor = load_bytes_in_safetensors(single_tensor_sd)
hash_sha256.update(bytes_for_tensor)
return f"0x{hash_sha256.hexdigest()}"
def str_bool(val):
return str(val).lower() != "false"
def default(val, d):
return val if val is not None else d
def make_sparse(t: torch.Tensor, sparsity=0.95):
abs_t = torch.abs(t)
np_array = abs_t.detach().cpu().numpy()
quan = float(np.quantile(np_array, sparsity))
sparse_t = t.masked_fill(abs_t < quan, 0)
return sparse_t
def extract_conv(
weight: Union[torch.Tensor, nn.Parameter],
mode="fixed",
mode_param=0,
device="cpu",
is_cp=False,
) -> Tuple[nn.Parameter, nn.Parameter]:
weight = weight.to(device)
out_ch, in_ch, kernel_size, _ = weight.shape
U, S, Vh = linalg.svd(weight.reshape(out_ch, -1))
if mode == "full":
return weight, "full"
elif mode == "fixed":
lora_rank = mode_param
elif mode == "threshold":
assert mode_param >= 0
lora_rank = torch.sum(S > mode_param)
elif mode == "ratio":
assert 1 >= mode_param >= 0
min_s = torch.max(S) * mode_param
lora_rank = torch.sum(S > min_s)
elif mode == "quantile" or mode == "percentile":
assert 1 >= mode_param >= 0
s_cum = torch.cumsum(S, dim=0)
min_cum_sum = mode_param * torch.sum(S)
lora_rank = torch.sum(s_cum < min_cum_sum)
else:
raise NotImplementedError(
'Extract mode should be "fixed", "threshold", "ratio" or "quantile"'
)
lora_rank = max(1, lora_rank)
lora_rank = min(out_ch, in_ch, lora_rank)
if lora_rank >= out_ch / 2 and not is_cp:
return weight, "full"
U = U[:, :lora_rank]
S = S[:lora_rank]
U = U @ torch.diag(S).to(device)
Vh = Vh[:lora_rank, :]
diff = (weight - (U @ Vh).reshape(out_ch, in_ch, kernel_size, kernel_size)).detach()
extract_weight_A = Vh.reshape(lora_rank, in_ch, kernel_size, kernel_size).detach()
extract_weight_B = U.reshape(out_ch, lora_rank, 1, 1).detach()
del U, S, Vh, weight
return (extract_weight_A, extract_weight_B, diff), "low rank"
def extract_linear(
weight: Union[torch.Tensor, nn.Parameter],
mode="fixed",
mode_param=0,
device="cpu",
) -> Tuple[nn.Parameter, nn.Parameter]:
weight = weight.to(device)
out_ch, in_ch = weight.shape
U, S, Vh = linalg.svd(weight)
if mode == "full":
return weight, "full"
elif mode == "fixed":
lora_rank = mode_param
elif mode == "threshold":
assert mode_param >= 0
lora_rank = torch.sum(S > mode_param)
elif mode == "ratio":
assert 1 >= mode_param >= 0
min_s = torch.max(S) * mode_param
lora_rank = torch.sum(S > min_s)
elif mode == "quantile" or mode == "percentile":
assert 1 >= mode_param >= 0
s_cum = torch.cumsum(S, dim=0)
min_cum_sum = mode_param * torch.sum(S)
lora_rank = torch.sum(s_cum < min_cum_sum)
else:
raise NotImplementedError(
'Extract mode should be "fixed", "threshold", "ratio" or "quantile"'
)
lora_rank = max(1, lora_rank)
lora_rank = min(out_ch, in_ch, lora_rank)
if lora_rank >= out_ch / 2:
return weight, "full"
U = U[:, :lora_rank]
S = S[:lora_rank]
U = U @ torch.diag(S).to(device)
Vh = Vh[:lora_rank, :]
diff = (weight - U @ Vh).detach()
extract_weight_A = Vh.reshape(lora_rank, in_ch).detach()
extract_weight_B = U.reshape(out_ch, lora_rank).detach()
del U, S, Vh, weight
return (extract_weight_A, extract_weight_B, diff), "low rank"
@torch.no_grad()
def extract_diff(
base_tes,
db_tes,
base_unet,
db_unet,
mode="fixed",
linear_mode_param=0,
conv_mode_param=0,
extract_device="cpu",
use_bias=False,
sparsity=0.98,
small_conv=True,
):
UNET_TARGET_REPLACE_MODULE = [
"Linear",
"Conv2d",
"LayerNorm",
"GroupNorm",
"GroupNorm32",
]
TEXT_ENCODER_TARGET_REPLACE_MODULE = [
"Embedding",
"Linear",
"Conv2d",
"LayerNorm",
"GroupNorm",
"GroupNorm32",
]
LORA_PREFIX_UNET = "lora_unet"
LORA_PREFIX_TEXT_ENCODER = "lora_te"
def make_state_dict(
prefix,
root_module: torch.nn.Module,
target_module: torch.nn.Module,
target_replace_modules,
):
loras = {}
temp = {}
for name, module in root_module.named_modules():
if module.__class__.__name__ in target_replace_modules:
temp[name] = module
for name, module in tqdm(
list((n, m) for n, m in target_module.named_modules() if n in temp)
):
weights = temp[name]
lora_name = prefix + "." + name
lora_name = lora_name.replace(".", "_")
layer = module.__class__.__name__
if layer in {
"Linear",
"Conv2d",
"LayerNorm",
"GroupNorm",
"GroupNorm32",
"Embedding",
}:
root_weight = module.weight
if torch.allclose(root_weight, weights.weight):
continue
else:
continue
module = module.to(extract_device)
weights = weights.to(extract_device)
if mode == "full":
decompose_mode = "full"
elif layer == "Linear":
weight, decompose_mode = extract_linear(
(root_weight - weights.weight),
mode,
linear_mode_param,
device=extract_device,
)
if decompose_mode == "low rank":
extract_a, extract_b, diff = weight
elif layer == "Conv2d":
is_linear = root_weight.shape[2] == 1 and root_weight.shape[3] == 1
weight, decompose_mode = extract_conv(
(root_weight - weights.weight),
mode,
linear_mode_param if is_linear else conv_mode_param,
device=extract_device,
)
if decompose_mode == "low rank":
extract_a, extract_b, diff = weight
if small_conv and not is_linear and decompose_mode == "low rank":
dim = extract_a.size(0)
(extract_c, extract_a, _), _ = extract_conv(
extract_a.transpose(0, 1),
"fixed",
dim,
extract_device,
True,
)
extract_a = extract_a.transpose(0, 1)
extract_c = extract_c.transpose(0, 1)
loras[f"{lora_name}.lora_mid.weight"] = (
extract_c.detach().cpu().contiguous().half()
)
diff = (
(
root_weight
- torch.einsum(
"i j k l, j r, p i -> p r k l",
extract_c,
extract_a.flatten(1, -1),
extract_b.flatten(1, -1),
)
)
.detach()
.cpu()
.contiguous()
)
del extract_c
else:
module = module.to("cpu")
weights = weights.to("cpu")
continue
if decompose_mode == "low rank":
loras[f"{lora_name}.lora_down.weight"] = (
extract_a.detach().cpu().contiguous().half()
)
loras[f"{lora_name}.lora_up.weight"] = (
extract_b.detach().cpu().contiguous().half()
)
loras[f"{lora_name}.alpha"] = torch.Tensor([extract_a.shape[0]]).half()
if use_bias:
diff = diff.detach().cpu().reshape(extract_b.size(0), -1)
sparse_diff = make_sparse(diff, sparsity).to_sparse().coalesce()
indices = sparse_diff.indices().to(torch.int16)
values = sparse_diff.values().half()
loras[f"{lora_name}.bias_indices"] = indices
loras[f"{lora_name}.bias_values"] = values
loras[f"{lora_name}.bias_size"] = torch.tensor(diff.shape).to(
torch.int16
)
del extract_a, extract_b, diff
elif decompose_mode == "full":
if "Norm" in layer:
w_key = "w_norm"
b_key = "b_norm"
else:
w_key = "diff"
b_key = "diff_b"
weight_diff = module.weight - weights.weight
loras[f"{lora_name}.{w_key}"] = (
weight_diff.detach().cpu().contiguous().half()
)
if getattr(weights, "bias", None) is not None:
bias_diff = module.bias - weights.bias
loras[f"{lora_name}.{b_key}"] = (
bias_diff.detach().cpu().contiguous().half()
)
else:
raise NotImplementedError
module = module.to("cpu")
weights = weights.to("cpu")
return loras
all_loras = {}
all_loras |= make_state_dict(
LORA_PREFIX_UNET,
base_unet,
db_unet,
UNET_TARGET_REPLACE_MODULE,
)
del base_unet, db_unet
if torch.cuda.is_available():
torch.cuda.empty_cache()
for idx, (te1, te2) in enumerate(zip(base_tes, db_tes)):
if len(base_tes) > 1:
prefix = f"{LORA_PREFIX_TEXT_ENCODER}{idx+1}"
else:
prefix = LORA_PREFIX_TEXT_ENCODER
all_loras |= make_state_dict(
prefix,
te1,
te2,
TEXT_ENCODER_TARGET_REPLACE_MODULE,
)
del te1, te2
all_lora_name = set()
for k in all_loras:
lora_name, weight = k.rsplit(".", 1)
all_lora_name.add(lora_name)
print(len(all_lora_name))
return all_loras
re_digits = re.compile(r"\d+")
re_compiled = {}
suffix_conversion = {
"attentions": {},
"resnets": {
"conv1": "in_layers_2",
"conv2": "out_layers_3",
"norm1": "in_layers_0",
"norm2": "out_layers_0",
"time_emb_proj": "emb_layers_1",
"conv_shortcut": "skip_connection",
},
}
def convert_diffusers_name_to_compvis(key):
def match(match_list, regex_text):
regex = re_compiled.get(regex_text)
if regex is None:
regex = re.compile(regex_text)
re_compiled[regex_text] = regex
r = re.match(regex, key)
if not r:
return False
match_list.clear()
match_list.extend([int(x) if re.match(re_digits, x) else x for x in r.groups()])
return True
m = []
if match(m, r"lora_unet_conv_in(.*)"):
return f"lora_unet_input_blocks_0_0{m[0]}"
if match(m, r"lora_unet_conv_out(.*)"):
return f"lora_unet_out_2{m[0]}"
if match(m, r"lora_unet_time_embedding_linear_(\d+)(.*)"):
return f"lora_unet_time_embed_{m[0] * 2 - 2}{m[1]}"
if match(m, r"lora_unet_down_blocks_(\d+)_(attentions|resnets)_(\d+)_(.+)"):
suffix = suffix_conversion.get(m[1], {}).get(m[3], m[3])
return f"lora_unet_input_blocks_{1 + m[0] * 3 + m[2]}_{1 if m[1] == 'attentions' else 0}_{suffix}"
if match(m, r"lora_unet_mid_block_(attentions|resnets)_(\d+)_(.+)"):
suffix = suffix_conversion.get(m[0], {}).get(m[2], m[2])
return (
f"lora_unet_middle_block_{1 if m[0] == 'attentions' else m[1] * 2}_{suffix}"
)
if match(m, r"lora_unet_up_blocks_(\d+)_(attentions|resnets)_(\d+)_(.+)"):
suffix = suffix_conversion.get(m[1], {}).get(m[3], m[3])
return f"lora_unet_output_blocks_{m[0] * 3 + m[2]}_{1 if m[1] == 'attentions' else 0}_{suffix}"
if match(m, r"lora_unet_down_blocks_(\d+)_downsamplers_0_conv"):
return f"lora_unet_input_blocks_{3 + m[0] * 3}_0_op"
if match(m, r"lora_unet_up_blocks_(\d+)_upsamplers_0_conv"):
return f"lora_unet_output_blocks_{2 + m[0] * 3}_2_conv"
return key
@torch.no_grad()
def merge(tes, unet, lyco_state_dict, scale: float = 1.0, device="cpu"):
from ..modules import make_module, get_module
LORA_PREFIX_UNET = "lora_unet"
LORA_PREFIX_TEXT_ENCODER = "lora_te"
merged = 0
def merge_state_dict(
prefix,
root_module: torch.nn.Module,
lyco_state_dict: Dict[str, torch.Tensor],
):
nonlocal merged
for child_name, child_module in tqdm(
list(root_module.named_modules()), desc=f"Merging {prefix}"
):
lora_name = prefix + "." + child_name
lora_name = lora_name.replace(".", "_")
lyco_type, params = get_module(lyco_state_dict, lora_name)
if lyco_type is None:
continue
module = make_module(lyco_type, params, lora_name, child_module)
if module is None:
continue
module.to(device)
module.merge_to(scale)
key_dict.pop(convert_diffusers_name_to_compvis(lora_name), None)
key_dict.pop(lora_name, None)
merged += 1
key_dict = {}
for k, v in tqdm(list(lyco_state_dict.items()), desc="Converting Dtype and Device"):
module, weight_key = k.split(".", 1)
convert_key = convert_diffusers_name_to_compvis(module)
if convert_key != module and len(tes) > 1:
# kohya's format for sdxl is as same as SGM, not diffusers
del lyco_state_dict[k]
key_dict[convert_key] = key_dict.get(convert_key, []) + [k]
k = f"{convert_key}.{weight_key}"
else:
key_dict[module] = key_dict.get(module, []) + [k]
lyco_state_dict[k] = v.float().cpu()
for idx, te in enumerate(tes):
if len(tes) > 1:
prefix = LORA_PREFIX_TEXT_ENCODER + str(idx + 1)
else:
prefix = LORA_PREFIX_TEXT_ENCODER
merge_state_dict(
prefix,
te,
lyco_state_dict,
)
torch.cuda.empty_cache()
merge_state_dict(
LORA_PREFIX_UNET,
unet,
lyco_state_dict,
)
torch.cuda.empty_cache()
print(f"Unused state dict key: {key_dict}")
print(f"{merged} Modules been merged")
+5
View File
@@ -0,0 +1,5 @@
def product(xs: list[int | float]):
res = 1
for x in xs:
res *= x
return res
+35
View File
@@ -0,0 +1,35 @@
import logging
import copy
import sys
class ColoredFormatter(logging.Formatter):
COLORS = {
"DEBUG": "\033[0;36m", # CYAN
"INFO": "\033[0;32m", # GREEN
"WARNING": "\033[0;33m", # YELLOW
"ERROR": "\033[0;31m", # RED
"CRITICAL": "\033[0;37;41m", # WHITE ON RED
"RESET": "\033[0m", # RESET COLOR
}
def format(self, record):
colored_record = copy.copy(record)
levelname = colored_record.levelname
seq = self.COLORS.get(levelname, self.COLORS["RESET"])
colored_record.levelname = f"{seq}{levelname}{self.COLORS['RESET']}"
return super().format(colored_record)
# Create a new logger
logger = logging.getLogger("LyCORIS")
logger.propagate = False
# Add handler if we don't have one.
if not logger.handlers:
handler = logging.StreamHandler(sys.stdout)
handler.setFormatter(ColoredFormatter("[%(name)s]-%(levelname)s: %(message)s"))
logger.addHandler(handler)
logger.setLevel(logging.DEBUG)
logger.debug("Logger initialized.")
+9
View File
@@ -0,0 +1,9 @@
import toml
def read_preset(preset):
try:
return toml.load(preset)
except Exception as e:
print("Error: cannot read preset file. ", e)
return None
+88
View File
@@ -0,0 +1,88 @@
from functools import cache
SUPPORT_QUANT = False
try:
from bitsandbytes.nn import LinearNF4, Linear8bitLt, LinearFP4
SUPPORT_QUANT = True
except Exception:
import torch.nn as nn
class LinearNF4(nn.Linear):
pass
class Linear8bitLt(nn.Linear):
pass
class LinearFP4(nn.Linear):
pass
try:
from quanto.nn import QLinear, QConv2d, QLayerNorm
SUPPORT_QUANT = True
except Exception:
import torch.nn as nn
class QLinear(nn.Linear):
pass
class QConv2d(nn.Conv2d):
pass
class QLayerNorm(nn.LayerNorm):
pass
try:
from optimum.quanto.nn import (
QLinear as QLinearOpt,
QConv2d as QConv2dOpt,
QLayerNorm as QLayerNormOpt,
)
SUPPORT_QUANT = True
except Exception:
import torch.nn as nn
class QLinearOpt(nn.Linear):
pass
class QConv2dOpt(nn.Conv2d):
pass
class QLayerNormOpt(nn.LayerNorm):
pass
from ..logging import logger
QuantLinears = (
Linear8bitLt,
LinearFP4,
LinearNF4,
QLinear,
QConv2d,
QLayerNorm,
QLinearOpt,
QConv2dOpt,
QLayerNormOpt,
)
@cache
def log_bypass():
return logger.warning(
"Using bnb/quanto/optimum-quanto with LyCORIS will enable force-bypass mode."
)
@cache
def log_suspect():
return logger.warning(
"Non-native Linear detected but bypass_mode is not set. "
"Automatically using force-bypass mode to avoid possible issues. "
"Please set bypass_mode=False explicitly if there are no quantized layers."
)
+13
View File
@@ -0,0 +1,13 @@
memory_efficient_attention = None
try:
import xformers
except Exception:
pass
try:
from xformers.ops import memory_efficient_attention
XFORMERS_AVAIL = True
except Exception:
memory_efficient_attention = None
XFORMERS_AVAIL = False
+640
View File
@@ -0,0 +1,640 @@
# General LyCORIS wrapper based on kohya-ss/sd-scripts' style
import os
import fnmatch
import re
import logging
from typing import Any, List
import torch
import torch.nn as nn
from .modules.locon import LoConModule
from .modules.loha import LohaModule
from .modules.lokr import LokrModule
from .modules.dylora import DyLoraModule
from .modules.glora import GLoRAModule
from .modules.norms import NormModule
from .modules.full import FullModule
from .modules.diag_oft import DiagOFTModule
from .modules.boft import ButterflyOFTModule
from .modules import get_module, make_module
from .config import PRESET
from .utils.preset import read_preset
from .utils import str_bool
from .logging import logger
VALID_PRESET_KEYS = [
"enable_conv",
"target_module",
"target_name",
"module_algo_map",
"name_algo_map",
"lora_prefix",
"use_fnmatch",
"unet_target_module",
"unet_target_name",
"text_encoder_target_module",
"text_encoder_target_name",
"exclude_name",
]
network_module_dict = {
"lora": LoConModule,
"locon": LoConModule,
"loha": LohaModule,
"lokr": LokrModule,
"dylora": DyLoraModule,
"glora": GLoRAModule,
"full": FullModule,
"diag-oft": DiagOFTModule,
"boft": ButterflyOFTModule,
}
deprecated_arg_dict = {
"disable_conv_cp": "use_tucker",
"use_cp": "use_tucker",
"use_conv_cp": "use_tucker",
"constrain": "constraint",
}
def create_lycoris(module, multiplier=1.0, linear_dim=4, linear_alpha=1, **kwargs):
for key, value in list(kwargs.items()):
if key in deprecated_arg_dict:
logger.warning(
f"{key} is deprecated. Please use {deprecated_arg_dict[key]} instead.",
stacklevel=2,
)
kwargs[deprecated_arg_dict[key]] = value
if linear_dim is None:
linear_dim = 4 # default
conv_dim = int(kwargs.get("conv_dim", linear_dim) or linear_dim)
conv_alpha = float(kwargs.get("conv_alpha", linear_alpha) or linear_alpha)
dropout = float(kwargs.get("dropout", 0.0) or 0.0)
rank_dropout = float(kwargs.get("rank_dropout", 0.0) or 0.0)
module_dropout = float(kwargs.get("module_dropout", 0.0) or 0.0)
algo = (kwargs.get("algo", "lora") or "lora").lower()
use_tucker = str_bool(
not kwargs.get("disable_conv_cp", True)
or kwargs.get("use_conv_cp", False)
or kwargs.get("use_cp", False)
or kwargs.get("use_tucker", False)
)
use_scalar = str_bool(kwargs.get("use_scalar", False))
block_size = int(kwargs.get("block_size", 4) or 4)
train_norm = str_bool(kwargs.get("train_norm", False))
constraint = float(kwargs.get("constraint", 0) or 0)
rescaled = str_bool(kwargs.get("rescaled", False))
weight_decompose = str_bool(kwargs.get("dora_wd", False))
wd_on_output = str_bool(kwargs.get("wd_on_output", False))
full_matrix = str_bool(kwargs.get("full_matrix", False))
bypass_mode = str_bool(kwargs.get("bypass_mode", None))
unbalanced_factorization = str_bool(kwargs.get("unbalanced_factorization", False))
if unbalanced_factorization:
logger.info("Unbalanced factorization for LoKr is enabled")
if bypass_mode:
logger.info("Bypass mode is enabled")
if weight_decompose:
logger.info("Weight decomposition is enabled")
if full_matrix:
logger.info("Full matrix mode for LoKr is enabled")
preset = kwargs.get("preset", "full")
if preset not in PRESET:
preset = read_preset(preset)
else:
preset = PRESET[preset]
assert preset is not None
LycorisNetwork.apply_preset(preset)
logger.info(f"Using rank adaptation algo: {algo}")
network = LycorisNetwork(
module,
multiplier=multiplier,
lora_dim=linear_dim,
conv_lora_dim=conv_dim,
alpha=linear_alpha,
conv_alpha=conv_alpha,
dropout=dropout,
rank_dropout=rank_dropout,
module_dropout=module_dropout,
use_tucker=use_tucker,
use_scalar=use_scalar,
network_module=algo,
train_norm=train_norm,
decompose_both=kwargs.get("decompose_both", False),
factor=kwargs.get("factor", -1),
block_size=block_size,
constraint=constraint,
rescaled=rescaled,
weight_decompose=weight_decompose,
wd_on_out=wd_on_output,
full_matrix=full_matrix,
bypass_mode=bypass_mode,
unbalanced_factorization=unbalanced_factorization,
)
return network
def create_lycoris_from_weights(multiplier, file, module, weights_sd=None, **kwargs):
if weights_sd is None:
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")
# get dim/alpha mapping
loras = {}
for key in weights_sd:
if "." not in key:
continue
lora_name = key.split(".")[0]
loras[lora_name] = None
for name, modules in module.named_modules():
lora_name = f"{LycorisNetwork.LORA_PREFIX}_{name}".replace(".", "_")
if lora_name in loras:
loras[lora_name] = modules
original_level = logger.level
logger.setLevel(logging.ERROR)
network = LycorisNetwork(module, init_only=True)
network.multiplier = multiplier
network.loras = []
logger.setLevel(original_level)
logger.info("Loading Modules from state dict...")
for lora_name, orig_modules in loras.items():
if orig_modules is None:
continue
lyco_type, params = get_module(weights_sd, lora_name)
module = make_module(lyco_type, params, lora_name, orig_modules)
if module is not None:
network.loras.append(module)
network.algo_table[module.__class__.__name__] = (
network.algo_table.get(module.__class__.__name__, 0) + 1
)
logger.info(f"{len(network.loras)} Modules Loaded")
for lora in network.loras:
lora.multiplier = multiplier
return network, weights_sd
class LycorisNetwork(torch.nn.Module):
ENABLE_CONV = True
TARGET_REPLACE_MODULE = [
"Linear",
"Conv1d",
"Conv2d",
"Conv3d",
"GroupNorm",
"LayerNorm",
]
TARGET_REPLACE_NAME = []
LORA_PREFIX = "lycoris"
MODULE_ALGO_MAP = {}
NAME_ALGO_MAP = {}
USE_FNMATCH = False
TARGET_EXCLUDE_NAME = []
@classmethod
def apply_preset(cls, preset):
for preset_key in preset.keys():
if preset_key not in VALID_PRESET_KEYS:
raise KeyError(
f'Unknown preset key "{preset_key}". Valid keys: {VALID_PRESET_KEYS}'
)
if "enable_conv" in preset:
cls.ENABLE_CONV = preset["enable_conv"]
if "target_module" in preset:
cls.TARGET_REPLACE_MODULE = preset["target_module"]
if "target_name" in preset:
cls.TARGET_REPLACE_NAME = preset["target_name"]
if "module_algo_map" in preset:
cls.MODULE_ALGO_MAP = preset["module_algo_map"]
if "name_algo_map" in preset:
cls.NAME_ALGO_MAP = preset["name_algo_map"]
if "lora_prefix" in preset:
cls.LORA_PREFIX = preset["lora_prefix"]
if "use_fnmatch" in preset:
cls.USE_FNMATCH = preset["use_fnmatch"]
if "exclude_name" in preset:
cls.TARGET_EXCLUDE_NAME = preset["exclude_name"]
return cls
def __init__(
self,
module: nn.Module,
multiplier=1.0,
lora_dim=4,
conv_lora_dim=4,
alpha=1,
conv_alpha=1,
use_tucker=False,
dropout=0,
rank_dropout=0,
module_dropout=0,
network_module: str = "locon",
norm_modules=NormModule,
train_norm=False,
init_only=False,
**kwargs,
) -> None:
super().__init__()
root_kwargs = kwargs
self.weights_sd = None
if init_only:
self.multiplier = 1
self.lora_dim = 0
self.alpha = 1
self.conv_lora_dim = 0
self.conv_alpha = 1
self.dropout = 0
self.rank_dropout = 0
self.module_dropout = 0
self.use_tucker = False
self.loras = []
self.algo_table = {}
return
self.multiplier = multiplier
self.lora_dim = lora_dim
if not self.ENABLE_CONV:
conv_lora_dim = 0
self.conv_lora_dim = int(conv_lora_dim)
if self.conv_lora_dim and self.conv_lora_dim != self.lora_dim:
logger.info("Apply different lora dim for conv layer")
logger.info(f"Conv Dim: {conv_lora_dim}, Linear Dim: {lora_dim}")
elif self.conv_lora_dim == 0:
logger.info("Disable conv layer")
self.alpha = alpha
self.conv_alpha = float(conv_alpha)
if self.conv_lora_dim and self.alpha != self.conv_alpha:
logger.info("Apply different alpha value for conv layer")
logger.info(f"Conv alpha: {conv_alpha}, Linear alpha: {alpha}")
if 1 >= dropout >= 0:
logger.info(f"Use Dropout value: {dropout}")
self.dropout = dropout
self.rank_dropout = rank_dropout
self.module_dropout = module_dropout
self.use_tucker = use_tucker
def create_single_module(
lora_name: str,
module: torch.nn.Module,
algo_name,
dim=None,
alpha=None,
use_tucker=self.use_tucker,
**kwargs,
):
for k, v in root_kwargs.items():
if k in kwargs:
continue
kwargs[k] = v
if train_norm and "Norm" in module.__class__.__name__:
return norm_modules(
lora_name,
module,
self.multiplier,
self.rank_dropout,
self.module_dropout,
**kwargs,
)
lora = None
if isinstance(module, torch.nn.Linear) and lora_dim > 0:
dim = dim or lora_dim
alpha = alpha or self.alpha
elif isinstance(
module, (torch.nn.Conv1d, torch.nn.Conv2d, torch.nn.Conv3d)
):
k_size, *_ = module.kernel_size
if k_size == 1 and lora_dim > 0:
dim = dim or lora_dim
alpha = alpha or self.alpha
elif conv_lora_dim > 0 or dim:
dim = dim or conv_lora_dim
alpha = alpha or self.conv_alpha
else:
return None
else:
return None
lora = network_module_dict[algo_name](
lora_name,
module,
self.multiplier,
dim,
alpha,
self.dropout,
self.rank_dropout,
self.module_dropout,
use_tucker,
**kwargs,
)
return lora
def create_modules_(
prefix: str,
root_module: torch.nn.Module,
algo,
current_lora_map: dict[str, Any],
configs={},
):
assert current_lora_map is not None, "No mapping supplied"
loras = current_lora_map
lora_names = []
for name, module in root_module.named_modules():
module_name = module.__class__.__name__
if module_name in self.MODULE_ALGO_MAP and module is not root_module:
next_config = self.MODULE_ALGO_MAP[module_name]
next_algo = next_config.get("algo", algo)
new_loras, new_lora_names, new_lora_map = create_modules_(
f"{prefix}_{name}" if name else prefix,
module,
next_algo,
loras,
configs=next_config,
)
loras = {**loras, **new_lora_map}
for lora_name, lora in zip(new_lora_names, new_loras):
if lora_name not in loras and lora_name not in current_lora_map:
loras[lora_name] = lora
if lora_name not in lora_names:
lora_names.append(lora_name)
continue
if name:
lora_name = prefix + "." + name
else:
lora_name = prefix
if f"{self.LORA_PREFIX}_." in lora_name:
lora_name = lora_name.replace(
f"{self.LORA_PREFIX}_.",
f"{self.LORA_PREFIX}.",
)
lora_name = lora_name.replace(".", "_")
if lora_name in loras:
continue
lora = create_single_module(lora_name, module, algo, **configs)
if lora is not None:
loras[lora_name] = lora
lora_names.append(lora_name)
return [loras[lora_name] for lora_name in lora_names], lora_names, loras
# create module instances
def create_modules(
prefix,
root_module: torch.nn.Module,
target_replace_modules,
target_replace_names=[],
target_exclude_names=[],
) -> List:
logger.info("Create LyCORIS Module")
loras = []
lora_map = {}
next_config = {}
for name, module in root_module.named_modules():
if name in target_exclude_names or any(
self.match_fn(t, name) for t in target_exclude_names
):
continue
module_name = module.__class__.__name__
if module_name in target_replace_modules and not any(
self.match_fn(t, name) for t in target_replace_names
):
if module_name in self.MODULE_ALGO_MAP:
next_config = self.MODULE_ALGO_MAP[module_name]
algo = next_config.get("algo", network_module)
else:
algo = network_module
lora_lst, _, _lora_map = create_modules_(
f"{prefix}_{name}",
module,
algo,
lora_map,
configs=next_config,
)
lora_map = {**lora_map, **_lora_map}
loras.extend(lora_lst)
next_config = {}
elif name in target_replace_names or any(
self.match_fn(t, name) for t in target_replace_names
):
conf_from_name = self.find_conf_for_name(name)
if conf_from_name is not None:
next_config = conf_from_name
algo = next_config.get("algo", network_module)
elif module_name in self.MODULE_ALGO_MAP:
next_config = self.MODULE_ALGO_MAP[module_name]
algo = next_config.get("algo", network_module)
else:
algo = network_module
lora_name = prefix + "." + name
lora_name = lora_name.replace(".", "_")
if lora_name in lora_map:
continue
lora = create_single_module(lora_name, module, algo, **next_config)
next_config = {}
if lora is not None:
lora_map[lora.lora_name] = lora
loras.append(lora)
return loras
self.loras = create_modules(
LycorisNetwork.LORA_PREFIX,
module,
list(
set(
[
*LycorisNetwork.TARGET_REPLACE_MODULE,
*LycorisNetwork.MODULE_ALGO_MAP.keys(),
]
)
),
list(
set(
[
*LycorisNetwork.TARGET_REPLACE_NAME,
*LycorisNetwork.NAME_ALGO_MAP.keys(),
]
)
),
target_exclude_names=LycorisNetwork.TARGET_EXCLUDE_NAME,
)
logger.info(f"create LyCORIS: {len(self.loras)} modules.")
algo_table = {}
for lora in self.loras:
algo_table[lora.__class__.__name__] = (
algo_table.get(lora.__class__.__name__, 0) + 1
)
logger.info(f"module type table: {algo_table}")
# Assertion to ensure we have not accidentally wrapped some layers
# multiple times.
names = set()
for lora in self.loras:
assert (
lora.lora_name not in names
), f"duplicated lora name: {lora.lora_name}"
names.add(lora.lora_name)
def match_fn(self, pattern: str, name: str) -> bool:
if self.USE_FNMATCH:
return fnmatch.fnmatch(name, pattern)
return bool(re.match(pattern, name))
def find_conf_for_name(
self,
name: str,
) -> dict[str, Any]:
if name in self.NAME_ALGO_MAP.keys():
return self.NAME_ALGO_MAP[name]
for key, value in self.NAME_ALGO_MAP.items():
if self.match_fn(key, name):
return value
return None
def set_multiplier(self, multiplier):
self.multiplier = multiplier
for lora in self.loras:
lora.multiplier = self.multiplier
def load_weights(self, file):
if os.path.splitext(file)[1] == ".safetensors":
from safetensors.torch import load_file, safe_open
self.weights_sd = load_file(file)
else:
self.weights_sd = torch.load(file, map_location="cpu")
missing, unexpected = self.load_state_dict(self.weights_sd, strict=False)
state = {}
if missing:
state["missing keys"] = missing
if unexpected:
state["unexpected keys"] = unexpected
return state
def apply_to(self):
"""
Register to modules to the subclass so that torch sees them.
"""
for lora in self.loras:
lora.apply_to()
self.add_module(lora.lora_name, lora)
if self.weights_sd:
# if some weights are not in state dict, it is ok because initial LoRA does nothing (lora_up is initialized by zeros)
info = self.load_state_dict(self.weights_sd, False)
logger.info(f"weights are loaded: {info}")
def is_mergeable(self):
return True
def restore(self):
for lora in self.loras:
lora.restore()
def merge_to(self, weight=1.0):
for lora in self.loras:
lora.merge_to(weight)
def apply_max_norm_regularization(self, max_norm_value, device):
key_scaled = 0
norms = []
for module in self.loras:
scaled, norm = module.apply_max_norm(max_norm_value, device)
if scaled is None:
continue
norms.append(norm)
key_scaled += scaled
if key_scaled == 0:
return key_scaled, 0, 0
return key_scaled, sum(norms) / len(norms), max(norms)
def enable_gradient_checkpointing(self):
# not supported
def make_ckpt(module):
if isinstance(module, torch.nn.Module):
module.grad_ckpt = True
self.apply(make_ckpt)
pass
def prepare_optimizer_params(self, lr):
def enumerate_params(loras):
params = []
for lora in loras:
params.extend(lora.parameters())
return params
self.requires_grad_(True)
all_params = []
param_data = {"params": enumerate_params(self.loras)}
if lr is not None:
param_data["lr"] = lr
all_params.append(param_data)
return all_params
def prepare_grad_etc(self, *args):
self.requires_grad_(True)
def on_epoch_start(self, *args):
self.train()
def get_trainable_params(self, *args):
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 = {}
save_file(state_dict, file, metadata)
else:
torch.save(state_dict, file)
+1 -1
View File
@@ -11,7 +11,7 @@ from transformers import CLIPTextModel
import numpy as np
import torch
import re
from .utils import setup_logging
from ..library.utils import setup_logging
from ..library.sdxl_original_unet import SdxlUNet2DConditionModel
setup_logging()
+76 -28
View File
@@ -339,6 +339,50 @@ class OptimizerConfigProdigy:
kwargs["min_snr_gamma"] = min_snr_gamma if min_snr_gamma != 0.0 else None
return (kwargs,)
class TrainNetworkConfig:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"network_type": (["lora", "LyCORIS/LoKr", "LyCORIS/Locon", "LyCORIS/LoHa"], {"default": "lora", "tooltip": "network type"}),
"lycoris_preset": (["full", "full-lin", "attn-mlp", "attn-only"], {"default": "attn-mlp"}),
"factor": ("INT",{"default": -1, "min": -1, "max": 16, "step": 1, "tooltip": "LoKr factor"}),
"extra_network_args": ("STRING",{"multiline": True, "default": "", "tooltip": "additional network args"}),
},
}
RETURN_TYPES = ("NETWORK_CONFIG",)
RETURN_NAMES = ("network_config",)
FUNCTION = "create_config"
CATEGORY = "FluxTrainer"
def create_config(self, network_type, extra_network_args, lycoris_preset, factor):
extra_args = [arg.strip() for arg in extra_network_args.strip().split('|') if arg.strip()]
if network_type == "lora":
network_module = ".networks.lora"
elif network_type == "LyCORIS/LoKr":
network_module = ".lycoris.kohya"
algo = "lokr"
elif network_type == "LyCORIS/Locon":
network_module = ".lycoris.kohya"
algo = "locon"
elif network_type == "LyCORIS/LoHa":
network_module = ".lycoris.kohya"
algo = "loha"
network_args = [
f"algo={algo}",
f"factor={factor}",
f"preset={lycoris_preset}"
]
network_config = {
"network_module": network_module,
"network_args": network_args + extra_args
}
return (network_config,)
class OptimizerConfigProdigyPlusScheduleFree:
@classmethod
@@ -348,18 +392,20 @@ class OptimizerConfigProdigyPlusScheduleFree:
"max_grad_norm": ("FLOAT",{"default": 0.0, "min": 0.0, "tooltip": "gradient clipping"}),
"prodigy_steps": ("INT",{"default": 0, "min": 0, "tooltip": "Freeze Prodigy stepsize adjustments after a certain optimiser step."}),
"d0": ("FLOAT",{"default": 1e-6, "min": 0.0,"step": 1e-7, "tooltip": "initial learning rate"}),
"d_coeff": ("FLOAT",{"default": 1.0, "min": 0.0, "step": 1e-7, "tooltip": "Coefficient in the expression for the estimate of d (default 1.0). Values such as 0.5 and 2.0 typically work as well. Changing this parameter is the preferred way to tune the method."}),
"d_coef": ("FLOAT",{"default": 1.0, "min": 0.0, "step": 1e-7, "tooltip": "Coefficient in the expression for the estimate of d (default 1.0). Values such as 0.5 and 2.0 typically work as well. Changing this parameter is the preferred way to tune the method."}),
"split_groups": ("BOOLEAN",{"default": True, "tooltip": "Track individual adaptation values for each parameter group."}),
#"beta3": ("FLOAT",{"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.0001, "tooltip": " Coefficient for computing the Prodigy stepsize using running averages. If set to None, uses the value of square root of beta2 (default: None)."}),
#"beta4": ("FLOAT",{"default": 0, "min": 0.0, "max": 1.0, "step": 0.0001, "tooltip": "Coefficient for updating the learning rate from Prodigy's adaptive stepsize. Smooths out spikes in learning rate adjustments. If set to None, beta1 is used instead. (default 0, which disables smoothing and uses original Prodigy behaviour)."}),
"use_bias_correction": ("BOOLEAN",{"default": False, "tooltip": "Turn on Adafactor-style bias correction, which scales beta2 directly."}),
"use_bias_correction": ("BOOLEAN",{"default": False, "tooltip": "Use the RAdam variant of schedule-free"}),
"min_snr_gamma": ("FLOAT",{"default": 5.0, "min": 0.0, "step": 0.01, "tooltip": "gamma for reducing the weight of high loss timesteps. Lower numbers have stronger effect. 5 is recommended by the paper"}),
"use_stableadamw": ("BOOLEAN",{"default": True, "tooltip": "Scales parameter updates by the root-mean-square of the normalised gradient, in essence identical to Adafactor's gradient scaling. Set to False if the adaptive learning rate never improves."}),
"use_cautious" : ("BOOLEAN",{"default": False, "tooltip": "Experimental. Perform 'cautious' updates, as proposed in https://arxiv.org/pdf/2411.16085. Modifies the update to isolate and boost values that align with the current gradient."}),
"use_adopt": ("BOOLEAN",{"default": False, "tooltip": "Experimental. Performs a modified step where the second moment is updated after the parameter update, so as not to include the current gradient in the denominator. This is a partial implementation of ADOPT (https://arxiv.org/abs/2411.02853), as we don't have a first moment to use for the update."}),
"use_grams": ("BOOLEAN",{"default": False, "tooltip": "Perform 'grams' updates, as proposed in https://arxiv.org/abs/2412.17107. Modifies the update using sign operations that align with the current gradient. Note that we do not have access to a first moment, so this deviates from the paper (we apply the sign directly to the update). May have a limited effect."}),
"stochastic_rounding": ("BOOLEAN",{"default": True, "tooltip": "Use stochastic rounding for bfloat16 weights"}),
"use_orthograd": ("BOOLEAN",{"default": False, "tooltip": "Experimental. Updates weights using the component of the gradient that is orthogonal to the current weight direction, as described in (https://arxiv.org/pdf/2501.04697). Can help prevent overfitting and improve generalisation."}),
"use_focus ": ("BOOLEAN",{"default": False, "tooltip": "Experimental. Modifies the update step to better handle noise at large step sizes. (https://arxiv.org/abs/2501.12243). This method is incompatible with factorisation, Muon and Adam-atan2."}),
"extra_optimizer_args": ("STRING",{"multiline": True, "default": "", "tooltip": "additional optimizer args"}),
},
}
@@ -389,7 +435,7 @@ class InitFluxLoRATraining:
"optimizer_settings": ("ARGS",),
"output_name": ("STRING", {"default": "flux_lora", "multiline": False}),
"output_dir": ("STRING", {"default": "flux_trainer_output", "multiline": False, "tooltip": "path to dataset, root is the 'ComfyUI' folder, with windows portable 'ComfyUI_windows_portable'"}),
"network_dim": ("INT", {"default": 4, "min": 1, "max": 2048, "step": 1, "tooltip": "network dim"}),
"network_dim": ("INT", {"default": 4, "min": 1, "max": 100000, "step": 1, "tooltip": "network dim"}),
"network_alpha": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 2048.0, "step": 0.01, "tooltip": "network alpha"}),
"learning_rate": ("FLOAT", {"default": 4e-4, "min": 0.0, "max": 10.0, "step": 0.000001, "tooltip": "learning rate"}),
"max_train_steps": ("INT", {"default": 1500, "min": 1, "max": 100000, "step": 1, "tooltip": "max number of training steps"}),
@@ -422,6 +468,7 @@ class InitFluxLoRATraining:
"block_args": ("ARGS", {"default": "", "tooltip": "limit the blocks used in the LoRA"}),
"gradient_checkpointing": (["enabled", "enabled_with_cpu_offloading", "disabled"], {"default": "enabled", "tooltip": "use gradient checkpointing"}),
"loss_args": ("ARGS", {"default": "", "tooltip": "loss args"}),
"network_config": ("NETWORK_CONFIG", {"tooltip": "additional network config"}),
},
"hidden": {
"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"
@@ -435,7 +482,7 @@ class InitFluxLoRATraining:
def init_training(self, flux_models, dataset, optimizer_settings, sample_prompts, output_name, attention_mode,
gradient_dtype, save_dtype, additional_args=None, resume_args=None, train_text_encoder='disabled',
block_args=None, gradient_checkpointing="enabled", prompt=None, extra_pnginfo=None, clip_l_lr=0, T5_lr=0, loss_args=None, **kwargs):
block_args=None, gradient_checkpointing="enabled", prompt=None, extra_pnginfo=None, clip_l_lr=0, T5_lr=0, loss_args=None, network_config=None, **kwargs):
mm.soft_empty_cache()
output_dir = os.path.abspath(kwargs.get("output_dir"))
@@ -451,6 +498,8 @@ class InitFluxLoRATraining:
dataset_toml = toml.dumps(json.loads(dataset_config))
parser = train_network_setup_parser()
flux_train_utils.add_flux_train_arguments(parser)
if additional_args is not None:
print(f"additional_args: {additional_args}")
args, _ = parser.parse_known_args(args=shlex.split(additional_args))
@@ -499,7 +548,7 @@ class InitFluxLoRATraining:
"persistent_data_loader_workers": False,
"max_data_loader_n_workers": 0,
"seed": 42,
"network_module": ".networks.lora_flux",
"network_module": ".networks.lora_flux" if network_config is None else network_config["network_module"],
"dataset_config": dataset_toml,
"output_name": f"{output_name}_rank{kwargs.get('network_dim')}_{save_dtype}",
"loss_type": "l2",
@@ -508,6 +557,7 @@ class InitFluxLoRATraining:
"network_train_unet_only": True if train_text_encoder == 'disabled' else False,
"fp8_base_unet": True if "fp8" in train_text_encoder else False,
"disable_mmap_load_safetensors": False,
"network_args": None if network_config is None else network_config["network_args"],
}
attention_settings = {
"sdpa": {"mem_eff_attn": True, "xformers": False, "spda": True},
@@ -527,21 +577,6 @@ class InitFluxLoRATraining:
if T5_lr != "NaN":
config_dict["text_encoder_lr"] = [clip_l_lr, T5_lr]
#network args
additional_network_args = []
if "T5" in train_text_encoder:
additional_network_args.append("train_t5xxl=True")
if block_args:
additional_network_args.append(block_args["include"])
# Handle network_args in args Namespace
if hasattr(args, 'network_args') and isinstance(args.network_args, list):
args.network_args.extend(additional_network_args)
else:
setattr(args, 'network_args', additional_network_args)
if gradient_checkpointing == "disabled":
config_dict["gradient_checkpointing"] = False
elif gradient_checkpointing == "enabled_with_cpu_offloading":
@@ -564,6 +599,21 @@ class InitFluxLoRATraining:
for key, value in config_dict.items():
setattr(args, key, value)
#network args
additional_network_args = []
if "T5" in train_text_encoder:
additional_network_args.append("train_t5xxl=True")
if block_args:
additional_network_args.append(block_args["include"])
# Handle network_args in args Namespace
if hasattr(args, 'network_args') and isinstance(args.network_args, list):
args.network_args.extend(additional_network_args)
else:
setattr(args, 'network_args', additional_network_args)
saved_args_file_path = os.path.join(output_dir, f"{output_name}_args.json")
with open(saved_args_file_path, 'w') as f:
@@ -653,6 +703,8 @@ class InitFluxTraining:
dataset_toml = toml.dumps(json.loads(dataset_config))
parser = train_setup_parser()
flux_train_utils.add_flux_train_arguments(parser)
if additional_args is not None:
print(f"additional_args: {additional_args}")
args, _ = parser.parse_known_args(args=shlex.split(additional_args))
@@ -928,15 +980,9 @@ class FluxTrainAndValidateLoop:
return (trainer, network_trainer.global_step)
def validate(self, network_trainer, validation_settings=None):
params = (
network_trainer.accelerator,
network_trainer.args,
params = (
network_trainer.current_epoch.value,
network_trainer.global_step,
network_trainer.unet,
network_trainer.vae,
network_trainer.text_encoder,
network_trainer.sample_prompts_te_outputs,
validation_settings
)
network_trainer.optimizer_eval_fn()
@@ -1721,6 +1767,7 @@ NODE_CLASS_MAPPINGS = {
"FluxTrainAndValidateLoop": FluxTrainAndValidateLoop,
"OptimizerConfigProdigyPlusScheduleFree": OptimizerConfigProdigyPlusScheduleFree,
"FluxTrainerLossConfig": FluxTrainerLossConfig,
"TrainNetworkConfig": TrainNetworkConfig,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"InitFluxLoRATraining": "Init Flux LoRA Training",
@@ -1747,4 +1794,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"FluxTrainAndValidateLoop": "Flux Train And Validate Loop",
"OptimizerConfigProdigyPlusScheduleFree": "Optimizer Config ProdigyPlusScheduleFree",
"FluxTrainerLossConfig": "Flux Trainer Loss Config",
"TrainNetworkConfig": "Train Network Config",
}
+465
View File
@@ -0,0 +1,465 @@
import os
import torch
import folder_paths
import comfy.model_management as mm
import comfy.utils
import toml
import json
import time
import shutil
import shlex
script_directory = os.path.dirname(os.path.abspath(__file__))
from .sdxl_train_network import SdxlNetworkTrainer
from .library import sdxl_train_util
from .library.device_utils import init_ipex
init_ipex()
from .library import train_util
from .train_network import setup_parser as train_network_setup_parser
import logging
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
logger = logging.getLogger(__name__)
class SDXLModelSelect:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"checkpoint": (folder_paths.get_filename_list("checkpoints"), ),
},
"optional": {
"lora_path": ("STRING",{"multiline": True, "forceInput": True, "default": "", "tooltip": "pre-trained LoRA path to load (network_weights)"}),
}
}
RETURN_TYPES = ("TRAIN_SDXL_MODELS",)
RETURN_NAMES = ("sdxl_models",)
FUNCTION = "loadmodel"
CATEGORY = "FluxTrainer/SDXL"
def loadmodel(self, checkpoint, lora_path=""):
checkpoint_path = folder_paths.get_full_path("checkpoints", checkpoint)
SDXL_models = {
"checkpoint": checkpoint_path,
"lora_path": lora_path
}
return (SDXL_models,)
class InitSDXLLoRATraining:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"SDXL_models": ("TRAIN_SDXL_MODELS",),
"dataset": ("JSON",),
"optimizer_settings": ("ARGS",),
"output_name": ("STRING", {"default": "SDXL_lora", "multiline": False}),
"output_dir": ("STRING", {"default": "SDXL_trainer_output", "multiline": False, "tooltip": "path to dataset, root is the 'ComfyUI' folder, with windows portable 'ComfyUI_windows_portable'"}),
"network_dim": ("INT", {"default": 16, "min": 1, "max": 100000, "step": 1, "tooltip": "network dim"}),
"network_alpha": ("FLOAT", {"default": 16, "min": 0.0, "max": 2048.0, "step": 0.01, "tooltip": "network alpha"}),
"learning_rate": ("FLOAT", {"default": 1e-6, "min": 0.0, "max": 10.0, "step": 0.0000001, "tooltip": "learning rate"}),
"max_train_steps": ("INT", {"default": 1500, "min": 1, "max": 100000, "step": 1, "tooltip": "max number of training steps"}),
"cache_latents": (["disk", "memory", "disabled"], {"tooltip": "caches text encoder outputs"}),
"cache_text_encoder_outputs": (["disk", "memory", "disabled"], {"tooltip": "caches text encoder outputs"}),
"highvram": ("BOOLEAN", {"default": False, "tooltip": "memory mode"}),
"blocks_to_swap": ("INT", {"default": 0, "min": 0, "max": 100, "step": 1, "tooltip": "option for memory use reduction. The maximum number of blocks that can be swapped is 36 for SDXL.5L and 22 for SDXL.5M"}),
"fp8_base": ("BOOLEAN", {"default": False, "tooltip": "use fp8 for base model"}),
"gradient_dtype": (["fp32", "fp16", "bf16"], {"default": "fp32", "tooltip": "the actual dtype training uses"}),
"save_dtype": (["fp32", "fp16", "bf16", "fp8_e4m3fn", "fp8_e5m2"], {"default": "fp16", "tooltip": "the dtype to save checkpoints as"}),
"attention_mode": (["sdpa", "xformers", "disabled"], {"default": "sdpa", "tooltip": "memory efficient attention mode"}),
"train_text_encoder": (['disabled', 'clip_l',], {"default": 'disabled', "tooltip": "also train the selected text encoders using specified dtype, T5 can not be trained without clip_l"}),
"clip_l_lr": ("FLOAT", {"default": 0, "min": 0.0, "max": 10.0, "step": 0.000001, "tooltip": "text encoder learning rate"}),
"clip_g_lr": ("FLOAT", {"default": 0, "min": 0.0, "max": 10.0, "step": 0.000001, "tooltip": "text encoder learning rate"}),
"sample_prompts_pos": ("STRING", {"multiline": True, "default": "illustration of a kitten | photograph of a turtle", "tooltip": "validation sample prompts, for multiple prompts, separate by `|`"}),
"sample_prompts_neg": ("STRING", {"multiline": True, "default": "", "tooltip": "validation sample prompts, for multiple prompts, separate by `|`"}),
"gradient_checkpointing": (["enabled", "disabled"], {"default": "enabled", "tooltip": "use gradient checkpointing"}),
},
"optional": {
"additional_args": ("STRING", {"multiline": True, "default": "", "tooltip": "additional args to pass to the training command"}),
"resume_args": ("ARGS", {"default": "", "tooltip": "resume args to pass to the training command"}),
"block_args": ("ARGS", {"default": "", "tooltip": "limit the blocks used in the LoRA"}),
"loss_args": ("ARGS", {"default": "", "tooltip": "loss args"}),
"network_config": ("NETWORK_CONFIG", {"tooltip": "additional network config"}),
},
"hidden": {
"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"
},
}
RETURN_TYPES = ("NETWORKTRAINER", "INT", "KOHYA_ARGS",)
RETURN_NAMES = ("network_trainer", "epochs_count", "args",)
FUNCTION = "init_training"
CATEGORY = "FluxTrainer/SDXL"
def init_training(self, SDXL_models, dataset, optimizer_settings, sample_prompts_pos, sample_prompts_neg, output_name, attention_mode,
gradient_dtype, save_dtype, additional_args=None, resume_args=None, train_text_encoder='disabled',
gradient_checkpointing="enabled", prompt=None, extra_pnginfo=None, clip_l_lr=0, clip_g_lr=0, loss_args=None, network_config=None, **kwargs):
mm.soft_empty_cache()
output_dir = os.path.abspath(kwargs.get("output_dir"))
os.makedirs(output_dir, exist_ok=True)
total, used, free = shutil.disk_usage(output_dir)
required_free_space = 2 * (2**30)
if free <= required_free_space:
raise ValueError(f"Insufficient disk space. Required: {required_free_space/2**30}GB. Available: {free/2**30}GB")
dataset_config = dataset["datasets"]
dataset_toml = toml.dumps(json.loads(dataset_config))
parser = train_network_setup_parser()
#sdxl_train_util.add_sdxl_training_arguments(parser)
if additional_args is not None:
print(f"additional_args: {additional_args}")
args, _ = parser.parse_known_args(args=shlex.split(additional_args))
else:
args, _ = parser.parse_known_args()
if kwargs.get("cache_latents") == "memory":
kwargs["cache_latents"] = True
kwargs["cache_latents_to_disk"] = False
elif kwargs.get("cache_latents") == "disk":
kwargs["cache_latents"] = True
kwargs["cache_latents_to_disk"] = True
kwargs["caption_dropout_rate"] = 0.0
kwargs["shuffle_caption"] = False
kwargs["token_warmup_step"] = 0.0
kwargs["caption_tag_dropout_rate"] = 0.0
else:
kwargs["cache_latents"] = False
kwargs["cache_latents_to_disk"] = False
if kwargs.get("cache_text_encoder_outputs") == "memory":
kwargs["cache_text_encoder_outputs"] = True
kwargs["cache_text_encoder_outputs_to_disk"] = False
elif kwargs.get("cache_text_encoder_outputs") == "disk":
kwargs["cache_text_encoder_outputs"] = True
kwargs["cache_text_encoder_outputs_to_disk"] = True
else:
kwargs["cache_text_encoder_outputs"] = False
kwargs["cache_text_encoder_outputs_to_disk"] = False
if '|' in sample_prompts_pos:
positive_prompts = sample_prompts_pos.split('|')
else:
positive_prompts = [sample_prompts_pos]
if '|' in sample_prompts_neg:
negative_prompts = sample_prompts_neg.split('|')
else:
negative_prompts = [sample_prompts_neg]
config_dict = {
"sample_prompts": positive_prompts,
"negative_prompts": negative_prompts,
"save_precision": save_dtype,
"mixed_precision": "bf16",
"num_cpu_threads_per_process": 1,
"pretrained_model_name_or_path": SDXL_models["checkpoint"],
"save_model_as": "safetensors",
"persistent_data_loader_workers": False,
"max_data_loader_n_workers": 0,
"seed": 42,
"network_module": ".networks.lora" if network_config is None else network_config["network_module"],
"dataset_config": dataset_toml,
"output_name": f"{output_name}_rank{kwargs.get('network_dim')}_{save_dtype}",
"loss_type": "l2",
"alpha_mask": dataset["alpha_mask"],
"network_train_unet_only": True if train_text_encoder == 'disabled' else False,
"disable_mmap_load_safetensors": False,
"network_args": None if network_config is None else network_config["network_args"],
}
attention_settings = {
"sdpa": {"mem_eff_attn": True, "xformers": False, "spda": True},
"xformers": {"mem_eff_attn": True, "xformers": True, "spda": False}
}
config_dict.update(attention_settings.get(attention_mode, {}))
gradient_dtype_settings = {
"fp16": {"full_fp16": True, "full_bf16": False, "mixed_precision": "fp16"},
"bf16": {"full_bf16": True, "full_fp16": False, "mixed_precision": "bf16"}
}
config_dict.update(gradient_dtype_settings.get(gradient_dtype, {}))
if train_text_encoder != 'disabled':
config_dict["text_encoder_lr"] = [clip_l_lr, clip_g_lr]
#network args
additional_network_args = []
# Handle network_args in args Namespace
if hasattr(args, 'network_args') and isinstance(args.network_args, list):
args.network_args.extend(additional_network_args)
else:
setattr(args, 'network_args', additional_network_args)
if gradient_checkpointing == "disabled":
config_dict["gradient_checkpointing"] = False
elif gradient_checkpointing == "enabled_with_cpu_offloading":
config_dict["gradient_checkpointing"] = True
config_dict["cpu_offload_checkpointing"] = True
else:
config_dict["gradient_checkpointing"] = True
if SDXL_models["lora_path"]:
config_dict["network_weights"] = SDXL_models["lora_path"]
config_dict.update(kwargs)
config_dict.update(optimizer_settings)
if loss_args:
config_dict.update(loss_args)
if resume_args:
config_dict.update(resume_args)
for key, value in config_dict.items():
setattr(args, key, value)
saved_args_file_path = os.path.join(output_dir, f"{output_name}_args.json")
with open(saved_args_file_path, 'w') as f:
json.dump(vars(args), f, indent=4)
#workflow saving
metadata = {}
if extra_pnginfo is not None:
metadata.update(extra_pnginfo["workflow"])
saved_workflow_file_path = os.path.join(output_dir, f"{output_name}_workflow.json")
with open(saved_workflow_file_path, 'w') as f:
json.dump(metadata, f, indent=4)
#pass args to kohya and initialize trainer
with torch.inference_mode(False):
network_trainer = SdxlNetworkTrainer()
training_loop = network_trainer.init_train(args)
epochs_count = network_trainer.num_train_epochs
trainer = {
"network_trainer": network_trainer,
"training_loop": training_loop,
}
return (trainer, epochs_count, args)
class SDXLTrainLoop:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"network_trainer": ("NETWORKTRAINER",),
"steps": ("INT", {"default": 1, "min": 1, "max": 10000, "step": 1, "tooltip": "the step point in training to validate/save"}),
},
}
RETURN_TYPES = ("NETWORKTRAINER", "INT",)
RETURN_NAMES = ("network_trainer", "steps",)
FUNCTION = "train"
CATEGORY = "FluxTrainer/SDXL"
def train(self, network_trainer, steps):
with torch.inference_mode(False):
training_loop = network_trainer["training_loop"]
network_trainer = network_trainer["network_trainer"]
initial_global_step = network_trainer.global_step
target_global_step = network_trainer.global_step + steps
comfy_pbar = comfy.utils.ProgressBar(steps)
network_trainer.comfy_pbar = comfy_pbar
network_trainer.optimizer_train_fn()
while network_trainer.global_step < target_global_step:
steps_done = training_loop(
break_at_steps = target_global_step,
epoch = network_trainer.current_epoch.value,
)
# Also break if the global steps have reached the max train steps
if network_trainer.global_step >= network_trainer.args.max_train_steps:
break
trainer = {
"network_trainer": network_trainer,
"training_loop": training_loop,
}
return (trainer, network_trainer.global_step)
class SDXLTrainLoRASave:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"network_trainer": ("NETWORKTRAINER",),
"save_state": ("BOOLEAN", {"default": False, "tooltip": "save the whole model state as well"}),
"copy_to_comfy_lora_folder": ("BOOLEAN", {"default": False, "tooltip": "copy the lora model to the comfy lora folder"}),
},
}
RETURN_TYPES = ("NETWORKTRAINER", "STRING", "INT",)
RETURN_NAMES = ("network_trainer","lora_path", "steps",)
FUNCTION = "save"
CATEGORY = "FluxTrainer/SDXL"
def save(self, network_trainer, save_state, copy_to_comfy_lora_folder):
import shutil
with torch.inference_mode(False):
trainer = network_trainer["network_trainer"]
global_step = trainer.global_step
ckpt_name = train_util.get_step_ckpt_name(trainer.args, "." + trainer.args.save_model_as, global_step)
trainer.save_model(ckpt_name, trainer.accelerator.unwrap_model(trainer.network), global_step, trainer.current_epoch.value + 1)
remove_step_no = train_util.get_remove_step_no(trainer.args, global_step)
if remove_step_no is not None:
remove_ckpt_name = train_util.get_step_ckpt_name(trainer.args, "." + trainer.args.save_model_as, remove_step_no)
trainer.remove_model(remove_ckpt_name)
if save_state:
train_util.save_and_remove_state_stepwise(trainer.args, trainer.accelerator, global_step)
lora_path = os.path.join(trainer.args.output_dir, ckpt_name)
if copy_to_comfy_lora_folder:
destination_dir = os.path.join(folder_paths.models_dir, "loras", "flux_trainer")
os.makedirs(destination_dir, exist_ok=True)
shutil.copy(lora_path, os.path.join(destination_dir, ckpt_name))
return (network_trainer, lora_path, global_step)
class SDXLTrainEnd:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"network_trainer": ("NETWORKTRAINER",),
"save_state": ("BOOLEAN", {"default": True}),
},
}
RETURN_TYPES = ("STRING", "STRING", "STRING",)
RETURN_NAMES = ("lora_name", "metadata", "lora_path",)
FUNCTION = "endtrain"
CATEGORY = "FluxTrainer/SDXL"
OUTPUT_NODE = True
def endtrain(self, network_trainer, save_state):
with torch.inference_mode(False):
training_loop = network_trainer["training_loop"]
network_trainer = network_trainer["network_trainer"]
network_trainer.metadata["ss_epoch"] = str(network_trainer.num_train_epochs)
network_trainer.metadata["ss_training_finished_at"] = str(time.time())
network = network_trainer.accelerator.unwrap_model(network_trainer.network)
network_trainer.accelerator.end_training()
network_trainer.optimizer_eval_fn()
if save_state:
train_util.save_state_on_train_end(network_trainer.args, network_trainer.accelerator)
ckpt_name = train_util.get_last_ckpt_name(network_trainer.args, "." + network_trainer.args.save_model_as)
network_trainer.save_model(ckpt_name, network, network_trainer.global_step, network_trainer.num_train_epochs, force_sync_upload=True)
logger.info("model saved.")
final_lora_name = str(network_trainer.args.output_name)
final_lora_path = os.path.join(network_trainer.args.output_dir, ckpt_name)
# metadata
metadata = json.dumps(network_trainer.metadata, indent=2)
training_loop = None
network_trainer = None
mm.soft_empty_cache()
return (final_lora_name, metadata, final_lora_path)
class SDXLTrainValidationSettings:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"steps": ("INT", {"default": 20, "min": 1, "max": 256, "step": 1, "tooltip": "sampling steps"}),
"width": ("INT", {"default": 1024, "min": 64, "max": 4096, "step": 8, "tooltip": "image width"}),
"height": ("INT", {"default": 1024, "min": 64, "max": 4096, "step": 8, "tooltip": "image height"}),
"guidance_scale": ("FLOAT", {"default": 7.5, "min": 1.0, "max": 32.0, "step": 0.05, "tooltip": "guidance scale"}),
"sampler": (["ddim", "ddpm", "pndm", "lms", "euler", "euler_a", "dpmsolver", "dpmsingle", "heun", "dpm_2", "dpm_2_a",], {"default": "dpm_2", "tooltip": "sampler"}),
"seed": ("INT", {"default": 42,"min": 0, "max": 0xffffffffffffffff, "step": 1}),
},
}
RETURN_TYPES = ("VALSETTINGS", )
RETURN_NAMES = ("validation_settings", )
FUNCTION = "set"
CATEGORY = "FluxTrainer/SDXL"
def set(self, **kwargs):
validation_settings = kwargs
print(validation_settings)
return (validation_settings,)
class SDXLTrainValidate:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"network_trainer": ("NETWORKTRAINER",),
},
"optional": {
"validation_settings": ("VALSETTINGS",),
}
}
RETURN_TYPES = ("NETWORKTRAINER", "IMAGE",)
RETURN_NAMES = ("network_trainer", "validation_images",)
FUNCTION = "validate"
CATEGORY = "FluxTrainer/SDXL"
def validate(self, network_trainer, validation_settings=None):
training_loop = network_trainer["training_loop"]
network_trainer = network_trainer["network_trainer"]
params = (
network_trainer.accelerator,
network_trainer.args,
network_trainer.current_epoch.value,
network_trainer.global_step,
network_trainer.accelerator.device,
network_trainer.vae,
network_trainer.tokenizers,
network_trainer.text_encoder,
network_trainer.unet,
validation_settings,
)
network_trainer.optimizer_eval_fn()
with torch.inference_mode(False):
image_tensors = network_trainer.sample_images(*params)
trainer = {
"network_trainer": network_trainer,
"training_loop": training_loop,
}
return (trainer, (0.5 * (image_tensors + 1.0)).cpu().float(),)
NODE_CLASS_MAPPINGS = {
"SDXLModelSelect": SDXLModelSelect,
"InitSDXLLoRATraining": InitSDXLLoRATraining,
"SDXLTrainValidationSettings": SDXLTrainValidationSettings,
"SDXLTrainValidate": SDXLTrainValidate,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"SDXLModelSelect": "SDXL Model Select",
"InitSDXLLoRATraining": "Init SDXL LoRA Training",
"SDXLTrainValidationSettings": "SDXL Train Validation Settings",
"SDXLTrainValidate": "SDXL Train Validate",
}
-485
View File
@@ -1,485 +0,0 @@
import math
import torch
from statistics import mean, harmonic_mean, geometric_mean
class CoreOptimiser(torch.optim.Optimizer):
def __init__(self, params, lr=1.0,
betas=(0.9, 0.99), beta3=None, beta4=0,
weight_decay=0.0,
use_bias_correction=False,
d0=1e-6, d_coef=1.0,
prodigy_steps=0,
warmup_steps=0,
eps=1e-8,
split_groups=True,
split_groups_mean="harmonic_mean",
factored=True,
fused_back_pass=False,
use_stableadamw=True,
use_muon_pp=False,
use_cautious=False,
use_adopt=False,
stochastic_rounding=True):
if not 0.0 < d0:
raise ValueError("Invalid d0 value: {}".format(d0))
if not 0.0 < lr:
raise ValueError("Invalid learning rate: {}".format(lr))
if eps is not None and not 0.0 < eps:
raise ValueError("Invalid epsilon value: {}".format(eps))
if not 0.0 <= betas[0] < 1.0:
raise ValueError("Invalid beta parameter at index 0: {}".format(betas[0]))
if not 0.0 <= betas[1] < 1.0:
raise ValueError("Invalid beta parameter at index 1: {}".format(betas[1]))
if beta3 is not None and not 0.0 <= beta3 < 1.0:
raise ValueError("Invalid beta3 parameter: {}".format(beta3))
if beta4 is not None and not 0.0 <= beta4 < 1.0:
raise ValueError("Invalid beta4 parameter: {}".format(beta4))
if split_groups_mean not in {None, "mean", "harmonic_mean", "geometric_mean"}:
raise ValueError(f"Invalid value for split_groups_mean: '{split_groups_mean}'. Must be one of {None, 'mean', 'harmonic_mean', 'geometric_mean'}")
if use_adopt and use_muon_pp:
print(f"[{self.__class__.__name__}] Muon and ADOPT cannot be used at the same time. Muon has been disabled.")
use_muon_pp = False
defaults = dict(lr=lr, betas=betas, beta3=beta3, beta4=beta4,
eps=eps,
weight_decay=weight_decay,
d=d0, d0=d0, d_coef=d_coef,
k=1, train_mode=True,
weight_sum=0,
prodigy_steps=prodigy_steps,
warmup_steps=warmup_steps,
use_bias_correction=use_bias_correction,
d_numerator=0.0,
factored=factored,
use_stableadamw=use_stableadamw,
use_muon_pp=use_muon_pp,
use_cautious=use_cautious,
use_adopt=use_adopt,
stochastic_rounding=stochastic_rounding)
super().__init__(params, defaults)
self.d0 = d0
if split_groups and len(self.param_groups) == 1:
print(f"[{self.__class__.__name__}] Optimiser contains single param_group -- 'split_groups' has been disabled.")
split_groups = False
self.split_groups = split_groups
self.split_groups_mean = split_groups_mean
# Properties for fused backward pass.
self.groups_to_process = None
self.shared_d = None
self.fused_back_pass = fused_back_pass
# Use tensors to keep everything on device during parameter loop.
for group in (self.param_groups if self.split_groups else self.param_groups[:1]):
p = group['params'][0]
group['running_d_numerator'] = torch.tensor(0.0, dtype=torch.float32, device=p.device)
group['running_d_denom'] = torch.tensor(0.0, dtype=torch.float32, device=p.device)
@torch.no_grad()
def eval(self):
pass
@torch.no_grad()
def train(self):
pass
@property
def supports_memory_efficient_fp16(self):
return False
@property
def supports_flat_params(self):
return True
def supports_fused_back_pass(self):
return True
@torch.no_grad()
def get_sliced_tensor(self, tensor, slice_p=11):
return tensor.ravel()[::slice_p]
@torch.no_grad()
def get_running_values_for_group(self, group):
if not self.split_groups:
group = self.param_groups[0]
return group['running_d_numerator'], group['running_d_denom']
@torch.no_grad()
def get_d_mean(self, groups, mode):
if mode is None:
return None
elif mode == "harmonic_mean":
return harmonic_mean(group['d'] for group in groups)
elif mode == "geometric_mean":
return geometric_mean(group['d'] for group in groups)
elif mode == "mean":
return mean(group['d'] for group in groups)
raise ValueError(f"Invalid value for split_groups_mean: '{mode}'. Must be one of {None, 'mean', 'harmonic_mean', 'geometric_mean'}")
# From: https://github.com/KellerJordan/Muon/blob/master/muon.py
@torch.no_grad()
def newton_schulz_(self, G, steps=6, eps=1e-7):
# Inline reshaping step within the method itself.
original_shape = None
if len(G.shape) > 2:
original_shape = G.shape
G = G.view(G.size(0), -1)
a, b, c = (3.4445, -4.7750, 2.0315)
X = G.bfloat16()
X /= (X.norm() + eps) # ensure top singular value <= 1
if G.size(0) > G.size(1):
X = X.T
for _ in range(steps):
A = X @ X.T
B = b * A + c * A @ A
X = a * X + B @ X
if G.size(0) > G.size(1):
X = X.T
if X is not G:
G.copy_(X)
del X
if original_shape is not None:
G = G.view(*original_shape)
return G
# Implementation by Nerogar. From: https://github.com/pytorch/pytorch/issues/120376#issuecomment-1974828905
def copy_stochastic_(self, target, source):
# create a random 16 bit integer
result = torch.randint_like(
source,
dtype=torch.int32,
low=0,
high=(1 << 16),
)
# add the random number to the lower 16 bit of the mantissa
result.add_(source.view(dtype=torch.int32))
# mask off the lower 16 bit of the mantissa
result.bitwise_and_(-65536) # -65536 = FFFF0000 as a signed int32
# copy the higher 16 bit into the target tensor
target.copy_(result.view(dtype=torch.float32))
# Modified Adafactor factorisation implementation by Ross Wightman
# https://github.com/huggingface/pytorch-image-models/pull/2320
@torch.no_grad()
def factored_dims(self,
shape,
factored,
min_dim_size_to_factor):
r"""Whether to use a factored second moment estimator.
This function returns a tuple with the two largest axes to reduce over.
If all dimensions have size < min_dim_size_to_factor, return None.
Args:
shape: an input shape
factored: whether to use factored second-moment estimator for > 2d vars.
min_dim_size_to_factor: only factor accumulator if all array dimensions are greater than this size.
Returns:
None or a tuple of ints
"""
if not factored or len(shape) < 2:
return None
if all(dim < min_dim_size_to_factor for dim in shape):
return None
sorted_dims = sorted(((x, i) for i, x in enumerate(shape)))
return int(sorted_dims[-2][1]), int(sorted_dims[-1][1])
@torch.no_grad()
def initialise_state(self, p, factored, use_muon_pp):
raise Exception("Not implemented!")
@torch.no_grad()
def initialise_state_internal(self, p, factored, use_muon_pp):
state = self.state[p]
needs_init = len(state) == 0
if needs_init:
grad = p.grad
dtype = torch.bfloat16 if p.dtype == torch.float32 else p.dtype
sliced_data = self.get_sliced_tensor(p)
# NOTE: We don't initialise z/exp_avg here -- subclass needs to do that.
state['muon'] = use_muon_pp and len(grad.shape) >= 2 and grad.size(0) < 10000
if not state['muon']:
factored_dims = self.factored_dims(
grad.shape,
factored=factored,
min_dim_size_to_factor=32
)
if factored_dims is not None:
dc, dr = factored_dims
row_shape = list(p.grad.shape)
row_shape[dr] = 1
col_shape = list(p.grad.shape)
col_shape[dc] = 1
reduce_dc = dc - 1 if dc > dr else dc
# Store reduction variables so we don't have to recalculate each step.
# Always store second moment low ranks in fp32 to avoid precision issues. Memory difference
# between bf16/fp16 and fp32 is negligible here.
state["exp_avg_sq"] = [torch.zeros(row_shape, dtype=torch.float32, device=p.device).detach(),
torch.zeros(col_shape, dtype=torch.float32, device=p.device).detach(),
dr, dc, reduce_dc]
else:
state['exp_avg_sq'] = torch.zeros_like(p, memory_format=torch.preserve_format).detach()
# If the initial weights are zero, don't bother storing them.
if p.count_nonzero() > 0:
state['p0'] = sliced_data.to(dtype=dtype, memory_format=torch.preserve_format, copy=True).detach()
else:
state['p0'] = torch.tensor(0.0, dtype=dtype, device=p.device)
state['s'] = torch.zeros_like(sliced_data, memory_format=torch.preserve_format, dtype=dtype).detach()
return state, needs_init
@torch.no_grad()
def update_d_and_reset(self, group):
k = group['k']
prodigy_steps = group['prodigy_steps']
if prodigy_steps > 0 and k >= prodigy_steps:
return
beta1, beta2 = group['betas']
beta3, beta4 = group['beta3'], group['beta4']
if beta3 is None:
beta3 = beta2 ** 0.5
if beta4 is None:
beta4 = beta1 ** 0.5
d = group['d']
d0 = group['d0']
d_coef = group['d_coef']
running_d_numerator, running_d_denom = self.get_running_values_for_group(group)
d_numerator = group['d_numerator']
d_numerator *= beta3
d_numerator_item = running_d_numerator.item()
d_denom_item = running_d_denom.item()
# Prevent the accumulation of negative values in the numerator in early training.
# We still allow negative updates once progress starts being made, as this is
# important for regulating the adaptive stepsize.
if d_numerator_item > 0 or d > d0:
d_numerator = max(0, d_numerator + d_numerator_item)
if d_denom_item > 0:
d_hat = max(math.atan2(d_coef * d_numerator, d_denom_item), d)
d = d * beta4 + d_hat * (1 - beta4) if beta4 > 0 else d_hat
group['d'] = d
group['d_numerator'] = d_numerator
running_d_numerator.zero_()
running_d_denom.zero_()
def on_start_step(self, group):
if self.groups_to_process is None:
# Optimiser hasn't run yet, so initialise.
self.groups_to_process = {i: len(group['params']) for i, group in enumerate(self.param_groups)}
elif len(self.groups_to_process) == 0:
# Start of new optimiser run, so grab updated d.
self.groups_to_process = {i: len(group['params']) for i, group in enumerate(self.param_groups)}
if not self.split_groups:
# When groups aren't split, calculate d for the first group,
# then copy to all other groups.
self.update_d_and_reset(group)
for g in self.param_groups:
g['d'] = group['d']
self.shared_d = self.get_d_mean(self.param_groups, self.split_groups_mean) if self.split_groups else None
def on_end_step(self, group):
group_index = self.param_groups.index(group)
# Decrement params processed so far.
self.groups_to_process[group_index] -= 1
# End of param loop for group, update calculations.
if self.groups_to_process[group_index] == 0:
k = group['k']
prodigy_steps = group['prodigy_steps']
if prodigy_steps > 0 and k == prodigy_steps:
print(f"[{self.__class__.__name__}] Prodigy stepsize adaptation disabled after {k} steps for param_group {group_index}.")
self.groups_to_process.pop(group_index)
if self.split_groups: # When groups are split, calculate per-group d.
self.update_d_and_reset(group)
group['k'] = k + 1
return True
return False
def get_dlr(self, group):
lr = group['lr']
k = group['k']
warmup_steps = group['warmup_steps']
d = group['d']
dlr = (self.shared_d if self.split_groups and self.shared_d else d) * lr
# Apply warmup separate to the denom and numerator updates.
if k < warmup_steps:
dlr *= k / warmup_steps
return dlr
def update_prodigy(self, state, group, grad, data, dlr):
k = group['k']
prodigy_steps = group['prodigy_steps']
if prodigy_steps <= 0 or k < prodigy_steps:
d, d0 = group['d'], group['d0']
beta3 = group['beta3']
if beta3 is None:
beta3 = group['betas'][1] ** 0.5
sliced_grad = self.get_sliced_tensor(grad)
sliced_data = self.get_sliced_tensor(data)
running_d_numerator, running_d_denom = self.get_running_values_for_group(group)
s = state['s']
x0_minus = state['p0'] - sliced_data
running_d_numerator.add_(torch.dot(sliced_grad, x0_minus), alpha=(d / d0) * dlr)
del x0_minus
s.mul_(beta3).add_(sliced_grad, alpha=(d / d0) * dlr)
running_d_denom.add_(s.abs().sum())
elif 's' in state: # Free the memory used by Prodigy, as we no longer need it.
del state['s']
del state['p0']
def get_update(self, num, denom, group):
d = group['d']
if group['eps'] is None:
# Adam-atan2. Use atan2 rather than epsilon and division
# for parameter updates (https://arxiv.org/abs/2407.05872).
# Has the nice property of "clipping" the gradient as well.
update = num.mul_(d).atan2_(denom)
else:
# Assume eps as already been added.
update = num.div_(denom).mul_(d)
return update
def get_denom(self, state, group):
exp_avg_sq = state['exp_avg_sq']
eps = group['eps']
# Adam EMA updates
if isinstance(exp_avg_sq, list):
row_var, col_var, _, _, reduce_dc = exp_avg_sq
row_col_mean = row_var.mean(dim=reduce_dc, keepdim=True).add_(1e-30)
row_factor = row_var.div(row_col_mean).sqrt_()
col_factor = col_var.sqrt()
denom = row_factor * col_factor
else:
denom = exp_avg_sq.sqrt()
if eps is not None:
denom.add_(group['d'] * eps)
return denom
def update_first_moment(self, exp_avg, group, grad):
d = group['d']
beta1, _ = group['betas']
exp_avg.mul_(beta1).add_(grad, value=d * (1 - beta1))
return exp_avg
def update_second_moment(self, state, group, grad, beta2, return_denom=True):
d = group['d']
exp_avg_sq = state['exp_avg_sq']
# Adafactor / PaLM beta2 decay. Clip beta2 as per Scaling ViT paper.
if group['use_bias_correction']:
beta2 = min(1 - group['k'] ** -0.8, beta2)
one_minus_beta2_d = d * d * (1 - beta2)
# Adam EMA updates
if isinstance(exp_avg_sq, list):
row_var, col_var, dr, dc, _ = exp_avg_sq
row_var.mul_(beta2).add_(
grad.norm(dim=dr, keepdim=True).square_().div_(grad.shape[dr]),
alpha=one_minus_beta2_d)
col_var.mul_(beta2).add_(
grad.norm(dim=dc, keepdim=True).square_().div_(grad.shape[dc]),
alpha=one_minus_beta2_d)
else:
exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=one_minus_beta2_d)
return self.get_denom(state, group) if return_denom else None
def rms_(self, tensor, rms_min):
if rms_min is not None:
rms = tensor.norm().div(tensor.numel() ** 0.5).add(rms_min)
tensor.div_(rms)
return tensor
# "Cautious Optimizer (C-Optim): Improving Training with One Line of Code"
# https://github.com/kyleliang919/c-optim
def cautious_(self, update, grad, reuse_grad):
if reuse_grad:
mask = grad.mul_(update) > 0
else:
mask = grad.mul(update) > 0
mask_scale = mask.numel() / mask.sum().add(1)
update.mul_(mask).mul_(mask_scale)
del mask
return update
@torch.no_grad()
def step_param(self, p, group):
raise Exception("Not implemented!")
@torch.no_grad()
def step_parameter(self, p, group, i):
self.step_param(p, group)
@torch.no_grad()
def step(self, closure=None):
if self.fused_back_pass:
return
"""Performs a single optimisation step.
Arguments:
closure (callable, optional): A closure that reevaluates the model
and returns the loss.
"""
loss = None
if closure is not None:
with torch.enable_grad():
loss = closure()
for param_group in self.param_groups:
for p in param_group["params"]:
self.step_param(p, param_group)
return loss
@@ -1,263 +0,0 @@
# source: @LoganBooker - https://github.com/LoganBooker/prodigy-plus-schedule-free
import torch
from .core_optimiser import CoreOptimiser
class ProdigyPlusScheduleFree(CoreOptimiser):
r"""
An optimiser based on Prodigy that includes schedule-free logic. Has additional improvements in the form of optional StableAdamW
gradient scaling and Adam-atan2 updates, per parameter group adaptation, lower memory utilisation, fused back pass support and
tweaks to mitigate uncontrolled LR growth.
Based on code from:
https://github.com/facebookresearch/schedule_free
https://github.com/konstmish/prodigy
Incorporates improvements from these pull requests (credit to https://github.com/dxqbYD and https://github.com/sangoi-exe):
https://github.com/konstmish/prodigy/pull/23
https://github.com/konstmish/prodigy/pull/22
https://github.com/konstmish/prodigy/pull/20
As with the reference implementation of schedule-free, a constant scheduler should be used, along with the appropriate
calls to `train()` and `eval()`. See the schedule-free documentation for more details: https://github.com/facebookresearch/schedule_free
If you do use another scheduler, linear or cosine is preferred, as a restarting scheduler can confuse Prodigy's adaptation logic.
Leave `lr` set to 1 unless you encounter instability. Do not use with gradient clipping, as this can hamper the
ability for the optimiser to predict stepsizes. Gradient clipping/normalisation is already handled in the following configurations:
1) `use_stableadamw=True,eps=1e8` (or any reasonable positive epsilon)
2) `eps=None` (Adam-atan2, scale invariant and can mess with Prodigy's stepsize calculations in some scenarios)
A new parameter, `beta4`, allows `d` to be updated via a moving average, rather than being immediately updated. This can help
smooth out learning rate adjustments. Values of 0.9-0.99 are recommended if trying out the feature. If set to None, the
square root of `beta1` is used, while a setting of 0 (the default) disables the feature.
By default, `split_groups` is set to `True`, so each parameter group will have its own adaptation values. So if you're training
different networks together, they won't contaminate each other's learning rates. The disadvantage of this approach is that some
networks can take a long time to reach a good learning rate when trained alongside others (for example, SDXL's Unet).
It's recommended to use a higher `d0` (1e-5, 5e-5, 1e-4) so these networks don't get stuck at a low learning rate.
For Prodigy's reference behaviour, which lumps all parameter groups together, set `split_groups` to `False`.
In some scenarios, it can be advantageous to freeze Prodigy's adaptive stepsize after a certain number of steps. This
can be controlled via the `prodigy_steps` settings.
Arguments:
params (iterable):
Iterable of parameters to optimize or dicts defining parameter groups.
lr (float):
Learning rate adjustment parameter. Increases or decreases the Prodigy learning rate.
betas (Tuple[float, float], optional):
Coefficients used for computing running averages of gradient and its square
(default: (0.9, 0.99))
eps (float):
Term added to the denominator outside of the root operation to improve numerical stability. If set to None,
Adam-atan2 is used instead. This removes the need for epsilon tuning, but may not work well in all situations.
(default: 1e-8).
beta3 (float):
Coefficient for computing the Prodigy stepsize using running averages.
If set to None, uses the value of square root of beta2 (default: None).
beta4 (float):
Coefficient for updating the learning rate from Prodigy's adaptive stepsize. Smooths out spikes in learning rate adjustments.
If set to None, beta1 is used instead. (default 0, which disables smoothing and uses original Prodigy behaviour).
weight_decay (float):
Decoupled weight decay. Value is multiplied by the adaptive learning rate.
(default: 0).
use_bias_correction (boolean):
Turn on Adafactor-style bias correction, which scales beta2 directly. (default False).
d0 (float):
Initial estimate for Prodigy (default 1e-6).
d_coef (float):
Coefficient in the expression for the estimate of d (default 1.0). Values such as 0.5 and 2.0 typically work as well.
Changing this parameter is the preferred way to tune the method.
prodigy_steps (int):
Freeze Prodigy stepsize adjustments after a certain optimiser step.
(default 0)
warmup_steps (int):
Enables a linear learning rate warmup (default 0). Use this over the warmup settings of your LR scheduler.
split_groups (boolean):
Track individual adaptation values for each parameter group. For example, if training
a text encoder beside a Unet. Note this can have a significant impact on training dynamics.
Set to False for original Prodigy behaviour, where all groups share the same values.
(default True)
split_groups_mean (str: None, "mean", "harmonic_mean", "geometric_mean"):
When split_groups is True, use specified mean of learning rates for all groups. This favours
a more conservative LR. Calculation remains per-group. If split_groups is False, this value has no effect.
Set to None to have each group use its own learning rate calculation.
(default "harmonic_mean")
factored (boolean):
Use factored approximation of the second moment, similar to Adafactor. Reduces memory usage. Disable
if training results in NaNs or the learning rate fails to grow.
(default True)
fused_back_pass (boolean):
Stops the optimiser from running the normal step method. Set to True if using fused backward pass.
(default False)
use_stableadamw (boolean):
Scales parameter updates by the root-mean-square of the normalised gradient, in essence identical to
Adafactor's gradient scaling. Set to False if the adaptive learning rate never improves.
(default True)
use_muon_pp (boolean):
Experimental. Perform orthogonalisation post-processing on 2D+ parameter updates ala Shampoo/SOAP/Muon.
(https://github.com/KellerJordan/Muon/blob/master/muon.py). Not suitable for all training scenarios.
May not work well with small batch sizes or finetuning. (default False)
use_cautious (boolean):
Experimental. Perform "cautious" updates, as proposed in https://arxiv.org/pdf/2411.16085. Modifies
the update to isolate and boost values that align with the current gradient.
(default False)
use_adopt (boolean):
Experimental. Performs a modified step where the second moment is updated after the parameter update,
so as not to include the current gradient in the denominator. This is a partial implementation of ADOPT
(https://arxiv.org/abs/2411.02853), as we don't have a first moment to use for the update.
(default False)
stochastic_rounding (boolean):
Use stochastic rounding for bfloat16 weights (https://github.com/pytorch/pytorch/issues/120376). Brings
bfloat16 training performance close to that of float32.
(default True)
"""
def __init__(self, params, lr=1.0,
betas=(0.9, 0.99), beta3=None, beta4=0,
weight_decay=0.0,
use_bias_correction=False,
d0=1e-6, d_coef=1.0,
prodigy_steps=0,
warmup_steps=0,
eps=1e-8,
split_groups=True,
split_groups_mean="harmonic_mean",
factored=True,
fused_back_pass=False,
use_stableadamw=True,
use_muon_pp=False,
use_cautious=False,
use_adopt=False,
stochastic_rounding=True):
super().__init__(params=params, lr=lr, betas=betas, beta3=beta3, beta4=beta4,
weight_decay=weight_decay, use_bias_correction=use_bias_correction,
d0=d0, d_coef=d_coef, prodigy_steps=prodigy_steps,
warmup_steps=warmup_steps, eps=eps, split_groups=split_groups,
split_groups_mean=split_groups_mean, factored=factored,
fused_back_pass=fused_back_pass, use_stableadamw=use_stableadamw,
use_muon_pp=use_muon_pp, use_cautious=use_cautious, use_adopt=use_adopt,
stochastic_rounding=stochastic_rounding)
@torch.no_grad()
def eval(self):
for group in self.param_groups:
if not group['train_mode']:
continue
beta1, _ = group['betas']
for p in group['params']:
z = self.state[p].get('z')
if z is not None:
# Set p to x
p.lerp_(end=z.to(device=p.device), weight=1 - 1 / beta1)
group['train_mode'] = False
@torch.no_grad()
def train(self):
for group in self.param_groups:
if group['train_mode']:
continue
beta1, _ = group['betas']
for p in group['params']:
z = self.state[p].get('z')
if z is not None:
# Set p to y
p.lerp_(end=z.to(device=p.device), weight=1 - beta1)
group['train_mode'] = True
@torch.no_grad()
def initialise_state(self, p, factored, use_muon_pp):
state, needs_init = self.initialise_state_internal(p, factored, use_muon_pp)
if needs_init:
state['z'] = p.detach().clone(memory_format=torch.preserve_format)
return state
@torch.no_grad()
def update_params(self, y, z, update, dlr, group):
# Weight decay.
weight_decay = group['weight_decay']
if weight_decay != 0:
update.add_(y, alpha=weight_decay)
weight = dlr ** 2
weight_sum = group['weight_sum'] + weight
ckp1 = weight / weight_sum if weight_sum else 0
y.lerp_(end=z, weight=ckp1)
y.add_(update, alpha=dlr * (group['betas'][0] * (1 - ckp1) - 1))
z.sub_(update, alpha=dlr)
return weight_sum
@torch.no_grad()
def step_param(self, p, group):
if not group['train_mode']:
raise Exception("Not in train mode!")
self.on_start_step(group)
weight_sum = group['weight_sum']
if p.grad is not None:
grad = p.grad
state = self.initialise_state(p, group['factored'], group['use_muon_pp'])
use_adopt = group['use_adopt']
if use_adopt and group['k'] == 1:
self.update_second_moment(state, group, grad.float(), 0, return_denom=False)
else:
dlr = self.get_dlr(group)
rms_min = 1.0 if group['use_stableadamw'] else None
y, z = p, state['z']
self.update_prodigy(state, group, grad, z, dlr)
grad_mask = grad.clone() if group['use_cautious'] else None
if state['muon']:
# newton_schulz_ casts to bf16 internally, so do float cast afterwards.
update = self.newton_schulz_(grad).float()
rms_min = 1e-30
else:
grad = grad.float()
_, beta2 = group['betas']
if use_adopt:
denom = self.get_denom(state, group)
self.update_second_moment(state, group, grad, beta2, return_denom=False)
else:
denom = self.update_second_moment(state, group, grad, beta2)
update = self.get_update(grad, denom, group)
del denom
if group['eps'] is None:
rms_min = None
self.rms_(update, rms_min)
if grad_mask is not None:
self.cautious_(update, grad_mask, reuse_grad=True)
if group['stochastic_rounding'] and y.dtype == z.dtype == torch.bfloat16:
y_fp32, z_fp32 = y.float(), z.float()
weight_sum = self.update_params(y_fp32, z_fp32, update, dlr, group)
self.copy_stochastic_(y, y_fp32)
self.copy_stochastic_(z, z_fp32)
del y_fp32, z_fp32
else:
weight_sum = self.update_params(y, z, update, dlr, group)
del update
if self.on_end_step(group):
group['weight_sum'] = weight_sum
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "comfyui-fluxtrainer"
description = "Currently supports LoRA training, and untested full finetune with code from kohya's scripts: [a/https://github.com/kohya-ss/sd-scripts](https://github.com/kohya-ss/sd-scripts)"
version = "1.0.1"
version = "1.0.2"
license = {file = "LICENSE"}
dependencies = ["accelerate>=0.33.0", "numpy<=1.26.4", "transformers>=4.44.0", "diffusers>=0.25.0", "ftfy>=6.1.1", "opencv-python>=4.7.0.68", "einops>=0.7.0", "bitsandbytes>=0.43.3", "prodigyopt>=1.0", "lion-pytorch>=0.0.6", "safetensors>=0.4.2", "altair>=4.2.2", "toml>=0.10.2", "voluptuous>=0.13.1", "huggingface-hub>=0.24.5", "# for Image utils", "imagesize>=1.4.1", "rich>=13.7.0", "came_pytorch", "matplotlib", "# for T5XXL tokenizer (SD3/FLUX)", "sentencepiece>=0.2.0"]
+2 -1
View File
@@ -21,4 +21,5 @@ matplotlib
# for T5XXL tokenizer (SD3/FLUX)
sentencepiece>=0.2.0
protobuf
schedulefree>=1.2.7
schedulefree>=1.2.7
prodigy-plus-schedule-free>=1.9.0
+228
View File
@@ -0,0 +1,228 @@
import argparse
from typing import List, Optional
import torch
from accelerate import Accelerator
from .library.device_utils import init_ipex, clean_memory_on_device
init_ipex()
from .library import sdxl_model_util, sdxl_train_util, strategy_base, strategy_sd, strategy_sdxl, train_util
from . import train_network
from .library.utils import setup_logging
setup_logging()
import logging
logger = logging.getLogger(__name__)
class SdxlNetworkTrainer(train_network.NetworkTrainer):
def __init__(self):
super().__init__()
self.vae_scale_factor = sdxl_model_util.VAE_SCALE_FACTOR
self.is_sdxl = True
def assert_extra_args(self, args, train_dataset_group):
super().assert_extra_args(args, train_dataset_group)
sdxl_train_util.verify_sdxl_training_args(args)
if args.cache_text_encoder_outputs:
assert (
train_dataset_group.is_text_encoder_output_cacheable()
), "when caching Text Encoder output, either caption_dropout_rate, shuffle_caption, token_warmup_step or caption_tag_dropout_rate cannot be used / Text Encoderの出力をキャッシュするときはcaption_dropout_rate, shuffle_caption, token_warmup_step, caption_tag_dropout_rateは使えません"
assert (
args.network_train_unet_only or not args.cache_text_encoder_outputs
), "network for Text Encoder cannot be trained with caching Text Encoder outputs / Text Encoderの出力をキャッシュしながらText Encoderのネットワークを学習することはできません"
train_dataset_group.verify_bucket_reso_steps(32)
def load_target_model(self, args, weight_dtype, accelerator):
(
load_stable_diffusion_format,
text_encoder1,
text_encoder2,
vae,
unet,
logit_scale,
ckpt_info,
) = sdxl_train_util.load_target_model(args, accelerator, sdxl_model_util.MODEL_VERSION_SDXL_BASE_V1_0, weight_dtype)
self.load_stable_diffusion_format = load_stable_diffusion_format
self.logit_scale = logit_scale
self.ckpt_info = ckpt_info
# モデルに xformers とか memory efficient attention を組み込む
train_util.replace_unet_modules(unet, args.mem_eff_attn, args.xformers, args.sdpa)
if torch.__version__ >= "2.0.0": # PyTorch 2.0.0 以上対応のxformersなら以下が使える
vae.set_use_memory_efficient_attention_xformers(args.xformers)
return sdxl_model_util.MODEL_VERSION_SDXL_BASE_V1_0, [text_encoder1, text_encoder2], vae, unet
def get_tokenize_strategy(self, args):
return strategy_sdxl.SdxlTokenizeStrategy(args.max_token_length, args.tokenizer_cache_dir)
def get_tokenizers(self, tokenize_strategy: strategy_sdxl.SdxlTokenizeStrategy):
return [tokenize_strategy.tokenizer1, tokenize_strategy.tokenizer2]
def get_latents_caching_strategy(self, args):
latents_caching_strategy = strategy_sd.SdSdxlLatentsCachingStrategy(
False, args.cache_latents_to_disk, args.vae_batch_size, args.skip_cache_check
)
return latents_caching_strategy
def get_text_encoding_strategy(self, args):
return strategy_sdxl.SdxlTextEncodingStrategy()
def get_models_for_text_encoding(self, args, accelerator, text_encoders):
return text_encoders + [accelerator.unwrap_model(text_encoders[-1])]
def get_text_encoder_outputs_caching_strategy(self, args):
if args.cache_text_encoder_outputs:
return strategy_sdxl.SdxlTextEncoderOutputsCachingStrategy(
args.cache_text_encoder_outputs_to_disk, None, args.skip_cache_check, is_weighted=args.weighted_captions
)
else:
return None
def cache_text_encoder_outputs_if_needed(
self, args, accelerator: Accelerator, unet, vae, text_encoders, dataset: train_util.DatasetGroup, weight_dtype
):
if args.cache_text_encoder_outputs:
if not args.lowram:
# メモリ消費を減らす
logger.info("move vae and unet to cpu to save memory")
org_vae_device = vae.device
org_unet_device = unet.device
vae.to("cpu")
unet.to("cpu")
clean_memory_on_device(accelerator.device)
# When TE is not be trained, it will not be prepared so we need to use explicit autocast
text_encoders[0].to(accelerator.device, dtype=weight_dtype)
text_encoders[1].to(accelerator.device, dtype=weight_dtype)
with accelerator.autocast():
dataset.new_cache_text_encoder_outputs(text_encoders + [accelerator.unwrap_model(text_encoders[-1])], accelerator)
accelerator.wait_for_everyone()
text_encoders[0].to("cpu", dtype=torch.float32) # Text Encoder doesn't work with fp16 on CPU
text_encoders[1].to("cpu", dtype=torch.float32)
clean_memory_on_device(accelerator.device)
if not args.lowram:
logger.info("move vae and unet back to original device")
vae.to(org_vae_device)
unet.to(org_unet_device)
else:
# Text Encoderから毎回出力を取得するので、GPUに乗せておく
text_encoders[0].to(accelerator.device, dtype=weight_dtype)
text_encoders[1].to(accelerator.device, dtype=weight_dtype)
def get_text_cond(self, args, accelerator, batch, tokenizers, text_encoders, weight_dtype):
if "text_encoder_outputs1_list" not in batch or batch["text_encoder_outputs1_list"] is None:
input_ids1 = batch["input_ids"]
input_ids2 = batch["input_ids2"]
with torch.enable_grad():
# Get the text embedding for conditioning
# TODO support weighted captions
# if args.weighted_captions:
# encoder_hidden_states = get_weighted_text_embeddings(
# tokenizer,
# text_encoder,
# batch["captions"],
# accelerator.device,
# args.max_token_length // 75 if args.max_token_length else 1,
# clip_skip=args.clip_skip,
# )
# else:
input_ids1 = input_ids1.to(accelerator.device)
input_ids2 = input_ids2.to(accelerator.device)
encoder_hidden_states1, encoder_hidden_states2, pool2 = train_util.get_hidden_states_sdxl(
args.max_token_length,
input_ids1,
input_ids2,
tokenizers[0],
tokenizers[1],
text_encoders[0],
text_encoders[1],
None if not args.full_fp16 else weight_dtype,
accelerator=accelerator,
)
else:
encoder_hidden_states1 = batch["text_encoder_outputs1_list"].to(accelerator.device).to(weight_dtype)
encoder_hidden_states2 = batch["text_encoder_outputs2_list"].to(accelerator.device).to(weight_dtype)
pool2 = batch["text_encoder_pool2_list"].to(accelerator.device).to(weight_dtype)
# # verify that the text encoder outputs are correct
# ehs1, ehs2, p2 = train_util.get_hidden_states_sdxl(
# args.max_token_length,
# batch["input_ids"].to(text_encoders[0].device),
# batch["input_ids2"].to(text_encoders[0].device),
# tokenizers[0],
# tokenizers[1],
# text_encoders[0],
# text_encoders[1],
# None if not args.full_fp16 else weight_dtype,
# )
# b_size = encoder_hidden_states1.shape[0]
# assert ((encoder_hidden_states1.to("cpu") - ehs1.to(dtype=weight_dtype)).abs().max() > 1e-2).sum() <= b_size * 2
# assert ((encoder_hidden_states2.to("cpu") - ehs2.to(dtype=weight_dtype)).abs().max() > 1e-2).sum() <= b_size * 2
# assert ((pool2.to("cpu") - p2.to(dtype=weight_dtype)).abs().max() > 1e-2).sum() <= b_size * 2
# logger.info("text encoder outputs verified")
return encoder_hidden_states1, encoder_hidden_states2, pool2
def call_unet(
self,
args,
accelerator,
unet,
noisy_latents,
timesteps,
text_conds,
batch,
weight_dtype,
indices: Optional[List[int]] = None,
):
noisy_latents = noisy_latents.to(weight_dtype) # TODO check why noisy_latents is not weight_dtype
# get size embeddings
orig_size = batch["original_sizes_hw"]
crop_size = batch["crop_top_lefts"]
target_size = batch["target_sizes_hw"]
embs = sdxl_train_util.get_size_embeddings(orig_size, crop_size, target_size, accelerator.device).to(weight_dtype)
# concat embeddings
encoder_hidden_states1, encoder_hidden_states2, pool2 = text_conds
vector_embedding = torch.cat([pool2, embs], dim=1).to(weight_dtype)
text_embedding = torch.cat([encoder_hidden_states1, encoder_hidden_states2], dim=2).to(weight_dtype)
if indices is not None and len(indices) > 0:
noisy_latents = noisy_latents[indices]
timesteps = timesteps[indices]
text_embedding = text_embedding[indices]
vector_embedding = vector_embedding[indices]
noise_pred = unet(noisy_latents, timesteps, text_embedding, vector_embedding)
return noise_pred
def sample_images(self, accelerator, args, epoch, global_step, device, vae, tokenizer, text_encoder, unet, validation_settings=None):
image_tensors = sdxl_train_util.sample_images(accelerator, args, epoch, global_step, device, vae, tokenizer, text_encoder, unet, validation_settings)
return image_tensors
def setup_parser() -> argparse.ArgumentParser:
parser = train_network.setup_parser()
sdxl_train_util.add_sdxl_training_arguments(parser)
return parser
if __name__ == "__main__":
parser = setup_parser()
args = parser.parse_args()
train_util.verify_command_line_training_args(args)
args = train_util.read_config_from_file(args, parser)
trainer = SdxlNetworkTrainer()
trainer.train(args)