Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
09fef403d1 | ||
|
|
639b3e80ba | ||
|
|
74611324dc | ||
|
|
f7025638fa | ||
|
|
6a91611a2b | ||
|
|
f6af45a169 | ||
|
|
5f254225c7 | ||
|
|
580bd8bb06 | ||
|
|
4343f2060a | ||
|
|
998968f5ff | ||
|
|
30cea9e372 | ||
|
|
136697a655 |
@@ -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
|
||||
|
||||

|
||||
|
||||
|
||||
@@ -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"]
|
||||
+611
-578
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
Binary file not shown.
|
Before Width: | Height: | Size: 2.5 MiB |
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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"
|
||||
)
|
||||
|
||||
@@ -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
@@ -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)
|
||||
@@ -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
@@ -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)
|
||||
@@ -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
@@ -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
|
||||
|
||||
|
||||
@@ -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.
|
||||
@@ -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
|
||||
@@ -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},
|
||||
},
|
||||
},
|
||||
}
|
||||
@@ -0,0 +1,9 @@
|
||||
from .general import (
|
||||
rebuild_tucker,
|
||||
factorization,
|
||||
power2factorization,
|
||||
FUNC_LIST,
|
||||
tucker_weight,
|
||||
tucker_weight_from_conv,
|
||||
apply_dora_scale,
|
||||
)
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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))
|
||||
@@ -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)
|
||||
@@ -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")
|
||||
@@ -0,0 +1,5 @@
|
||||
def product(xs: list[int | float]):
|
||||
res = 1
|
||||
for x in xs:
|
||||
res *= x
|
||||
return res
|
||||
@@ -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.")
|
||||
@@ -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
|
||||
@@ -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."
|
||||
)
|
||||
@@ -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
|
||||
@@ -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
@@ -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()
|
||||
|
||||
@@ -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
@@ -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",
|
||||
}
|
||||
@@ -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
@@ -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
@@ -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
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user