Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
09fef403d1 | ||
|
|
639b3e80ba | ||
|
|
74611324dc | ||
|
|
f7025638fa | ||
|
|
6a91611a2b | ||
|
|
f6af45a169 | ||
|
|
5f254225c7 | ||
|
|
580bd8bb06 | ||
|
|
4343f2060a | ||
|
|
998968f5ff | ||
|
|
30cea9e372 | ||
|
|
136697a655 | ||
|
|
a15cfb181a | ||
|
|
688482c0b4 | ||
|
|
3ba4c14e53 | ||
|
|
b8fe91abdd | ||
|
|
381b0a4bb1 | ||
|
|
971e1cf553 | ||
|
|
d8ae9b0be1 | ||
|
|
05a71c595b | ||
|
|
a348bae0a2 | ||
|
|
8bfe1f668b | ||
|
|
f05d2c3c96 | ||
|
|
7932a63b86 | ||
|
|
959d434d71 | ||
|
|
9df7868dc7 | ||
|
|
4ad054cce5 | ||
|
|
f9d3824c8a | ||
|
|
8c4cf7ffbc | ||
|
|
de56a60f28 | ||
|
|
a549df730f | ||
|
|
0419754233 | ||
|
|
3274cd0574 | ||
|
|
608073d244 | ||
|
|
94ae756b5e | ||
|
|
753f617f26 | ||
|
|
47c62fd8c9 | ||
|
|
09a8d7ffca | ||
|
|
c3aa4ea889 | ||
|
|
b52849ad9a | ||
|
|
4cf65f66e2 | ||
|
|
1d2e4c225f | ||
|
|
2c41f44ef0 | ||
|
|
54d8a4b60d | ||
|
|
c61fcf09cf | ||
|
|
66a5fa96c9 | ||
|
|
284bb6f112 | ||
|
|
95de0cbcc9 | ||
|
|
280f2557b9 | ||
|
|
299a994197 | ||
|
|
7f486c8916 | ||
|
|
712ffb7a1b | ||
|
|
1c70cce963 | ||
|
|
cd8e756103 | ||
|
|
c6de455381 | ||
|
|
98ca556b23 | ||
|
|
0a720115de | ||
|
|
77587e0193 | ||
|
|
0d1d93931f | ||
|
|
ba58d55a6c | ||
|
|
201891ef88 | ||
|
|
a4d2f7545b | ||
|
|
721dfc94f3 | ||
|
|
38d3d9ddf6 | ||
|
|
638076f337 | ||
|
|
372787dfb6 |
@@ -0,0 +1,2 @@
|
||||
github: [kijai]
|
||||
custom: ["https://www.paypal.me/kijaidesign"]
|
||||
@@ -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,3 +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"]
|
||||
+620
-621
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: 5.1 MiB |
+1
-1
@@ -175,7 +175,7 @@ def train(args):
|
||||
vae.requires_grad_(False)
|
||||
vae.eval()
|
||||
|
||||
train_dataset_group.new_cache_latents(vae, accelerator.is_main_process)
|
||||
train_dataset_group.new_cache_latents(vae, accelerator)
|
||||
|
||||
vae.to("cpu")
|
||||
clean_memory_on_device(accelerator.device)
|
||||
|
||||
@@ -1,809 +0,0 @@
|
||||
# training with captions
|
||||
|
||||
# Swap blocks between CPU and GPU:
|
||||
# This implementation is inspired by and based on the work of 2kpr.
|
||||
# Many thanks to 2kpr for the original concept and implementation of memory-efficient offloading.
|
||||
# The original idea has been adapted and extended to fit the current project's needs.
|
||||
|
||||
# Key features:
|
||||
# - CPU offloading during forward and backward passes
|
||||
# - Use of fused optimizer and grad_hook for efficient gradient processing
|
||||
# - Per-block fused optimizer instances
|
||||
|
||||
import argparse
|
||||
import copy
|
||||
import math
|
||||
import os
|
||||
from multiprocessing import Value
|
||||
from typing import List
|
||||
import toml
|
||||
|
||||
from tqdm import tqdm
|
||||
|
||||
import torch
|
||||
from .library.device_utils import init_ipex, clean_memory_on_device
|
||||
|
||||
init_ipex()
|
||||
|
||||
from accelerate.utils import set_seed
|
||||
from .library import deepspeed_utils, flux_train_utils, flux_utils, strategy_base, strategy_flux
|
||||
from .library.sd3_train_utils import load_prompts, FlowMatchEulerDiscreteScheduler
|
||||
|
||||
from .library import train_util as train_util
|
||||
|
||||
from .library.utils import setup_logging, add_logging_arguments
|
||||
|
||||
setup_logging()
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
from .library import config_util as config_util
|
||||
|
||||
from .library.config_util import (
|
||||
ConfigSanitizer,
|
||||
BlueprintGenerator,
|
||||
)
|
||||
from .library.custom_train_functions import apply_masked_loss, add_custom_train_arguments
|
||||
|
||||
|
||||
def train(args):
|
||||
train_util.verify_training_args(args)
|
||||
train_util.prepare_dataset_args(args, True)
|
||||
# sdxl_train_util.verify_sdxl_training_args(args)
|
||||
deepspeed_utils.prepare_deepspeed_args(args)
|
||||
setup_logging(args, reset=True)
|
||||
|
||||
# assert (
|
||||
# not args.weighted_captions
|
||||
# ), "weighted_captions is not supported currently / weighted_captionsは現在サポートされていません"
|
||||
if args.cache_text_encoder_outputs_to_disk and not args.cache_text_encoder_outputs:
|
||||
logger.warning(
|
||||
"cache_text_encoder_outputs_to_disk is enabled, so cache_text_encoder_outputs is also enabled / cache_text_encoder_outputs_to_diskが有効になっているため、cache_text_encoder_outputsも有効になります"
|
||||
)
|
||||
args.cache_text_encoder_outputs = True
|
||||
|
||||
if args.cpu_offload_checkpointing and not args.gradient_checkpointing:
|
||||
logger.warning(
|
||||
"cpu_offload_checkpointing is enabled, so gradient_checkpointing is also enabled / cpu_offload_checkpointingが有効になっているため、gradient_checkpointingも有効になります"
|
||||
)
|
||||
args.gradient_checkpointing = True
|
||||
|
||||
cache_latents = args.cache_latents
|
||||
use_dreambooth_method = args.in_json is None
|
||||
|
||||
if args.seed is not None:
|
||||
set_seed(args.seed) # 乱数系列を初期化する
|
||||
|
||||
# prepare caching strategy: this must be set before preparing dataset. because dataset may use this strategy for initialization.
|
||||
if args.cache_latents:
|
||||
latents_caching_strategy = strategy_flux.FluxLatentsCachingStrategy(
|
||||
args.cache_latents_to_disk, args.vae_batch_size, args.skip_latents_validity_check
|
||||
)
|
||||
strategy_base.LatentsCachingStrategy.set_strategy(latents_caching_strategy)
|
||||
|
||||
# データセットを準備する
|
||||
if args.dataset_class is None:
|
||||
blueprint_generator = BlueprintGenerator(ConfigSanitizer(True, True, args.masked_loss, True))
|
||||
if args.dataset_config is not None:
|
||||
logger.info(f"Load dataset config from {args.dataset_config}")
|
||||
user_config = config_util.load_user_config(args.dataset_config)
|
||||
ignored = ["train_data_dir", "in_json"]
|
||||
if any(getattr(args, attr) is not None for attr in ignored):
|
||||
logger.warning(
|
||||
"ignore following options because config file is found: {0} / 設定ファイルが利用されるため以下のオプションは無視されます: {0}".format(
|
||||
", ".join(ignored)
|
||||
)
|
||||
)
|
||||
else:
|
||||
if use_dreambooth_method:
|
||||
logger.info("Using DreamBooth method.")
|
||||
user_config = {
|
||||
"datasets": [
|
||||
{
|
||||
"subsets": config_util.generate_dreambooth_subsets_config_by_subdirs(
|
||||
args.train_data_dir, args.reg_data_dir
|
||||
)
|
||||
}
|
||||
]
|
||||
}
|
||||
else:
|
||||
logger.info("Training with captions.")
|
||||
user_config = {
|
||||
"datasets": [
|
||||
{
|
||||
"subsets": [
|
||||
{
|
||||
"image_dir": args.train_data_dir,
|
||||
"metadata_file": args.in_json,
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
blueprint = blueprint_generator.generate(user_config, args)
|
||||
train_dataset_group = config_util.generate_dataset_group_by_blueprint(blueprint.dataset_group)
|
||||
else:
|
||||
train_dataset_group = train_util.load_arbitrary_dataset(args)
|
||||
|
||||
current_epoch = Value("i", 0)
|
||||
current_step = Value("i", 0)
|
||||
ds_for_collator = train_dataset_group if args.max_data_loader_n_workers == 0 else None
|
||||
collator = train_util.collator_class(current_epoch, current_step, ds_for_collator)
|
||||
|
||||
train_dataset_group.verify_bucket_reso_steps(16) # TODO これでいいか確認
|
||||
|
||||
if args.debug_dataset:
|
||||
if args.cache_text_encoder_outputs:
|
||||
strategy_base.TextEncoderOutputsCachingStrategy.set_strategy(
|
||||
strategy_flux.FluxTextEncoderOutputsCachingStrategy(
|
||||
args.cache_text_encoder_outputs_to_disk, args.text_encoder_batch_size, False, False
|
||||
)
|
||||
)
|
||||
train_dataset_group.set_current_strategies()
|
||||
train_util.debug_dataset(train_dataset_group, True)
|
||||
return
|
||||
if len(train_dataset_group) == 0:
|
||||
logger.error(
|
||||
"No data found. Please verify the metadata file and train_data_dir option. / 画像がありません。メタデータおよびtrain_data_dirオプションを確認してください。"
|
||||
)
|
||||
return
|
||||
|
||||
if cache_latents:
|
||||
assert (
|
||||
train_dataset_group.is_latent_cacheable()
|
||||
), "when caching latents, either color_aug or random_crop cannot be used / latentをキャッシュするときはcolor_augとrandom_cropは使えません"
|
||||
|
||||
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は使えません"
|
||||
|
||||
# acceleratorを準備する
|
||||
logger.info("prepare accelerator")
|
||||
accelerator = train_util.prepare_accelerator(args)
|
||||
|
||||
# mixed precisionに対応した型を用意しておき適宜castする
|
||||
weight_dtype, save_dtype = train_util.prepare_dtype(args)
|
||||
|
||||
# モデルを読み込む
|
||||
name = "schnell" if "schnell" in args.pretrained_model_name_or_path else "dev"
|
||||
|
||||
# load VAE for caching latents
|
||||
ae = None
|
||||
if cache_latents:
|
||||
ae = flux_utils.load_ae(name, args.ae, weight_dtype, "cpu")
|
||||
ae.to(accelerator.device, dtype=weight_dtype)
|
||||
ae.requires_grad_(False)
|
||||
ae.eval()
|
||||
|
||||
train_dataset_group.new_cache_latents(ae, accelerator.is_main_process)
|
||||
|
||||
ae.to("cpu") # if no sampling, vae can be deleted
|
||||
clean_memory_on_device(accelerator.device)
|
||||
|
||||
accelerator.wait_for_everyone()
|
||||
|
||||
# prepare tokenize strategy
|
||||
if args.t5xxl_max_token_length is None:
|
||||
if name == "schnell":
|
||||
t5xxl_max_token_length = 256
|
||||
else:
|
||||
t5xxl_max_token_length = 512
|
||||
else:
|
||||
t5xxl_max_token_length = args.t5xxl_max_token_length
|
||||
|
||||
flux_tokenize_strategy = strategy_flux.FluxTokenizeStrategy(t5xxl_max_token_length)
|
||||
strategy_base.TokenizeStrategy.set_strategy(flux_tokenize_strategy)
|
||||
|
||||
# load clip_l, t5xxl for caching text encoder outputs
|
||||
clip_l = flux_utils.load_clip_l(args.clip_l, weight_dtype, "cpu")
|
||||
t5xxl = flux_utils.load_t5xxl(args.t5xxl, weight_dtype, "cpu")
|
||||
clip_l.eval()
|
||||
t5xxl.eval()
|
||||
clip_l.requires_grad_(False)
|
||||
t5xxl.requires_grad_(False)
|
||||
|
||||
text_encoding_strategy = strategy_flux.FluxTextEncodingStrategy(args.apply_t5_attn_mask)
|
||||
strategy_base.TextEncodingStrategy.set_strategy(text_encoding_strategy)
|
||||
|
||||
# cache text encoder outputs
|
||||
sample_prompts_te_outputs = None
|
||||
if args.cache_text_encoder_outputs:
|
||||
# Text Encodes are eval and no grad here
|
||||
clip_l.to(accelerator.device)
|
||||
t5xxl.to(accelerator.device)
|
||||
|
||||
text_encoder_caching_strategy = strategy_flux.FluxTextEncoderOutputsCachingStrategy(
|
||||
args.cache_text_encoder_outputs_to_disk, args.text_encoder_batch_size, False, False, args.apply_t5_attn_mask
|
||||
)
|
||||
strategy_base.TextEncoderOutputsCachingStrategy.set_strategy(text_encoder_caching_strategy)
|
||||
|
||||
with accelerator.autocast():
|
||||
train_dataset_group.new_cache_text_encoder_outputs([clip_l, t5xxl], accelerator.is_main_process)
|
||||
|
||||
# cache sample prompt's embeddings to free text encoder's memory
|
||||
if args.sample_prompts is not None:
|
||||
logger.info(f"cache Text Encoder outputs for sample prompt: {args.sample_prompts}")
|
||||
|
||||
tokenize_strategy: strategy_flux.FluxTokenizeStrategy = strategy_base.TokenizeStrategy.get_strategy()
|
||||
text_encoding_strategy: strategy_flux.FluxTextEncodingStrategy = strategy_base.TextEncodingStrategy.get_strategy()
|
||||
|
||||
prompts = load_prompts(args.sample_prompts)
|
||||
sample_prompts_te_outputs = {} # key: prompt, value: text encoder outputs
|
||||
with accelerator.autocast(), torch.no_grad():
|
||||
for prompt_dict in prompts:
|
||||
for p in [prompt_dict.get("prompt", ""), prompt_dict.get("negative_prompt", "")]:
|
||||
if p not in sample_prompts_te_outputs:
|
||||
logger.info(f"cache Text Encoder outputs for prompt: {p}")
|
||||
tokens_and_masks = tokenize_strategy.tokenize(p)
|
||||
sample_prompts_te_outputs[p] = text_encoding_strategy.encode_tokens(
|
||||
tokenize_strategy, [clip_l, t5xxl], tokens_and_masks, args.apply_t5_attn_mask
|
||||
)
|
||||
|
||||
accelerator.wait_for_everyone()
|
||||
|
||||
# now we can delete Text Encoders to free memory
|
||||
clip_l = None
|
||||
t5xxl = None
|
||||
clean_memory_on_device(accelerator.device)
|
||||
|
||||
# load FLUX
|
||||
# if we load to cpu, flux.to(fp8) takes a long time
|
||||
flux = flux_utils.load_flow_model(name, args.pretrained_model_name_or_path, weight_dtype, "cpu")
|
||||
|
||||
if args.gradient_checkpointing:
|
||||
flux.enable_gradient_checkpointing(args.cpu_offload_checkpointing)
|
||||
|
||||
flux.requires_grad_(True)
|
||||
|
||||
if args.double_blocks_to_swap is not None or args.single_blocks_to_swap is not None:
|
||||
# Swap blocks between CPU and GPU to reduce memory usage, in forward and backward passes.
|
||||
# This idea is based on 2kpr's great work. Thank you!
|
||||
logger.info(
|
||||
f"enable block swap: double_blocks_to_swap={args.double_blocks_to_swap}, single_blocks_to_swap={args.single_blocks_to_swap}"
|
||||
)
|
||||
flux.enable_block_swap(args.double_blocks_to_swap, args.single_blocks_to_swap)
|
||||
|
||||
if not cache_latents:
|
||||
# load VAE here if not cached
|
||||
ae = flux_utils.load_ae(name, args.ae, weight_dtype, "cpu")
|
||||
ae.requires_grad_(False)
|
||||
ae.eval()
|
||||
ae.to(accelerator.device, dtype=weight_dtype)
|
||||
|
||||
training_models = []
|
||||
params_to_optimize = []
|
||||
training_models.append(flux)
|
||||
params_to_optimize.append({"params": list(flux.parameters()), "lr": args.learning_rate})
|
||||
|
||||
# calculate number of trainable parameters
|
||||
n_params = 0
|
||||
for group in params_to_optimize:
|
||||
for p in group["params"]:
|
||||
n_params += p.numel()
|
||||
|
||||
accelerator.print(f"number of trainable parameters: {n_params}")
|
||||
|
||||
# 学習に必要なクラスを準備する
|
||||
accelerator.print("prepare optimizer, data loader etc.")
|
||||
|
||||
if args.blockwise_fused_optimizers:
|
||||
# fused backward pass: https://pytorch.org/tutorials/intermediate/optimizer_step_in_backward_tutorial.html
|
||||
# Instead of creating an optimizer for all parameters as in the tutorial, we create an optimizer for each block of parameters.
|
||||
# This balances memory usage and management complexity.
|
||||
|
||||
# split params into groups. currently different learning rates are not supported
|
||||
grouped_params = []
|
||||
param_group = {}
|
||||
for group in params_to_optimize:
|
||||
named_parameters = list(flux.named_parameters())
|
||||
assert len(named_parameters) == len(group["params"]), "number of parameters does not match"
|
||||
for p, np in zip(group["params"], named_parameters):
|
||||
# determine target layer and block index for each parameter
|
||||
block_type = "other" # double, single or other
|
||||
if np[0].startswith("double_blocks"):
|
||||
block_idx = int(np[0].split(".")[1])
|
||||
block_type = "double"
|
||||
elif np[0].startswith("single_blocks"):
|
||||
block_idx = int(np[0].split(".")[1])
|
||||
block_type = "single"
|
||||
else:
|
||||
block_idx = -1
|
||||
|
||||
param_group_key = (block_type, block_idx)
|
||||
if param_group_key not in param_group:
|
||||
param_group[param_group_key] = []
|
||||
param_group[param_group_key].append(p)
|
||||
|
||||
block_types_and_indices = []
|
||||
for param_group_key, param_group in param_group.items():
|
||||
block_types_and_indices.append(param_group_key)
|
||||
grouped_params.append({"params": param_group, "lr": args.learning_rate})
|
||||
|
||||
num_params = 0
|
||||
for p in param_group:
|
||||
num_params += p.numel()
|
||||
accelerator.print(f"block {param_group_key}: {num_params} parameters")
|
||||
|
||||
# prepare optimizers for each group
|
||||
optimizers = []
|
||||
for group in grouped_params:
|
||||
_, _, optimizer = train_util.get_optimizer(args, trainable_params=[group])
|
||||
optimizers.append(optimizer)
|
||||
optimizer = optimizers[0] # avoid error in the following code
|
||||
|
||||
logger.info(f"using {len(optimizers)} optimizers for blockwise fused optimizers")
|
||||
|
||||
else:
|
||||
_, _, optimizer = train_util.get_optimizer(args, trainable_params=params_to_optimize)
|
||||
|
||||
# prepare dataloader
|
||||
# strategies are set here because they cannot be referenced in another process. Copy them with the dataset
|
||||
# some strategies can be None
|
||||
train_dataset_group.set_current_strategies()
|
||||
|
||||
# DataLoaderのプロセス数:0 は persistent_workers が使えないので注意
|
||||
n_workers = min(args.max_data_loader_n_workers, os.cpu_count()) # cpu_count or max_data_loader_n_workers
|
||||
train_dataloader = torch.utils.data.DataLoader(
|
||||
train_dataset_group,
|
||||
batch_size=1,
|
||||
shuffle=True,
|
||||
collate_fn=collator,
|
||||
num_workers=n_workers,
|
||||
persistent_workers=args.persistent_data_loader_workers,
|
||||
)
|
||||
|
||||
# 学習ステップ数を計算する
|
||||
if args.max_train_epochs is not None:
|
||||
args.max_train_steps = args.max_train_epochs * math.ceil(
|
||||
len(train_dataloader) / accelerator.num_processes / args.gradient_accumulation_steps
|
||||
)
|
||||
accelerator.print(
|
||||
f"override steps. steps for {args.max_train_epochs} epochs is / 指定エポックまでのステップ数: {args.max_train_steps}"
|
||||
)
|
||||
|
||||
# データセット側にも学習ステップを送信
|
||||
train_dataset_group.set_max_train_steps(args.max_train_steps)
|
||||
|
||||
# lr schedulerを用意する
|
||||
if args.blockwise_fused_optimizers:
|
||||
# prepare lr schedulers for each optimizer
|
||||
lr_schedulers = [train_util.get_scheduler_fix(args, optimizer, accelerator.num_processes) for optimizer in optimizers]
|
||||
lr_scheduler = lr_schedulers[0] # avoid error in the following code
|
||||
else:
|
||||
lr_scheduler = train_util.get_scheduler_fix(args, optimizer, accelerator.num_processes)
|
||||
|
||||
# 実験的機能:勾配も含めたfp16/bf16学習を行う モデル全体をfp16/bf16にする
|
||||
if args.full_fp16:
|
||||
assert (
|
||||
args.mixed_precision == "fp16"
|
||||
), "full_fp16 requires mixed precision='fp16' / full_fp16を使う場合はmixed_precision='fp16'を指定してください。"
|
||||
accelerator.print("enable full fp16 training.")
|
||||
flux.to(weight_dtype)
|
||||
if clip_l is not None:
|
||||
clip_l.to(weight_dtype)
|
||||
t5xxl.to(weight_dtype) # TODO check works with fp16 or not
|
||||
elif args.full_bf16:
|
||||
assert (
|
||||
args.mixed_precision == "bf16"
|
||||
), "full_bf16 requires mixed precision='bf16' / full_bf16を使う場合はmixed_precision='bf16'を指定してください。"
|
||||
accelerator.print("enable full bf16 training.")
|
||||
flux.to(weight_dtype)
|
||||
if clip_l is not None:
|
||||
clip_l.to(weight_dtype)
|
||||
t5xxl.to(weight_dtype)
|
||||
|
||||
# if we don't cache text encoder outputs, move them to device
|
||||
if not args.cache_text_encoder_outputs:
|
||||
clip_l.to(accelerator.device)
|
||||
t5xxl.to(accelerator.device)
|
||||
|
||||
clean_memory_on_device(accelerator.device)
|
||||
|
||||
if args.deepspeed:
|
||||
ds_model = deepspeed_utils.prepare_deepspeed_model(args, mmdit=flux)
|
||||
# most of ZeRO stage uses optimizer partitioning, so we have to prepare optimizer and ds_model at the same time. # pull/1139#issuecomment-1986790007
|
||||
ds_model, optimizer, train_dataloader, lr_scheduler = accelerator.prepare(
|
||||
ds_model, optimizer, train_dataloader, lr_scheduler
|
||||
)
|
||||
training_models = [ds_model]
|
||||
|
||||
else:
|
||||
# acceleratorがなんかよろしくやってくれるらしい
|
||||
flux = accelerator.prepare(flux)
|
||||
optimizer, train_dataloader, lr_scheduler = accelerator.prepare(optimizer, train_dataloader, lr_scheduler)
|
||||
|
||||
# 実験的機能:勾配も含めたfp16学習を行う PyTorchにパッチを当ててfp16でのgrad scaleを有効にする
|
||||
if args.full_fp16:
|
||||
# During deepseed training, accelerate not handles fp16/bf16|mixed precision directly via scaler. Let deepspeed engine do.
|
||||
# -> But we think it's ok to patch accelerator even if deepspeed is enabled.
|
||||
train_util.patch_accelerator_for_fp16_training(accelerator)
|
||||
|
||||
# resumeする
|
||||
train_util.resume_from_local_or_hf_if_specified(accelerator, args)
|
||||
|
||||
if args.fused_backward_pass:
|
||||
# use fused optimizer for backward pass: other optimizers will be supported in the future
|
||||
import library.adafactor_fused
|
||||
|
||||
library.adafactor_fused.patch_adafactor_fused(optimizer)
|
||||
for param_group in optimizer.param_groups:
|
||||
for parameter in param_group["params"]:
|
||||
if parameter.requires_grad:
|
||||
|
||||
def __grad_hook(tensor: torch.Tensor, param_group=param_group):
|
||||
if accelerator.sync_gradients and args.max_grad_norm != 0.0:
|
||||
accelerator.clip_grad_norm_(tensor, args.max_grad_norm)
|
||||
optimizer.step_param(tensor, param_group)
|
||||
tensor.grad = None
|
||||
|
||||
parameter.register_post_accumulate_grad_hook(__grad_hook)
|
||||
|
||||
elif args.blockwise_fused_optimizers:
|
||||
# prepare for additional optimizers and lr schedulers
|
||||
for i in range(1, len(optimizers)):
|
||||
optimizers[i] = accelerator.prepare(optimizers[i])
|
||||
lr_schedulers[i] = accelerator.prepare(lr_schedulers[i])
|
||||
|
||||
# counters are used to determine when to step the optimizer
|
||||
global optimizer_hooked_count
|
||||
global num_parameters_per_group
|
||||
global parameter_optimizer_map
|
||||
|
||||
optimizer_hooked_count = {}
|
||||
num_parameters_per_group = [0] * len(optimizers)
|
||||
parameter_optimizer_map = {}
|
||||
|
||||
double_blocks_to_swap = args.double_blocks_to_swap
|
||||
single_blocks_to_swap = args.single_blocks_to_swap
|
||||
num_double_blocks = len(flux.double_blocks)
|
||||
num_single_blocks = len(flux.single_blocks)
|
||||
|
||||
for opt_idx, optimizer in enumerate(optimizers):
|
||||
for param_group in optimizer.param_groups:
|
||||
for parameter in param_group["params"]:
|
||||
if parameter.requires_grad:
|
||||
block_type, block_idx = block_types_and_indices[opt_idx]
|
||||
|
||||
def create_optimizer_hook(btype, bidx):
|
||||
def optimizer_hook(parameter: torch.Tensor):
|
||||
# print(f"optimizer_hook: {btype}, {bidx}")
|
||||
if accelerator.sync_gradients and args.max_grad_norm != 0.0:
|
||||
accelerator.clip_grad_norm_(parameter, args.max_grad_norm)
|
||||
|
||||
i = parameter_optimizer_map[parameter]
|
||||
optimizer_hooked_count[i] += 1
|
||||
if optimizer_hooked_count[i] == num_parameters_per_group[i]:
|
||||
optimizers[i].step()
|
||||
optimizers[i].zero_grad(set_to_none=True)
|
||||
|
||||
# swap blocks if necessary
|
||||
if btype == "double" and double_blocks_to_swap:
|
||||
if bidx >= num_double_blocks - double_blocks_to_swap:
|
||||
bidx_cuda = double_blocks_to_swap - (num_double_blocks - bidx)
|
||||
flux.double_blocks[bidx].to("cpu")
|
||||
flux.double_blocks[bidx_cuda].to(accelerator.device)
|
||||
# print(f"Move double block {bidx} to cpu and {bidx_cuda} to device")
|
||||
elif btype == "single" and single_blocks_to_swap:
|
||||
if bidx >= num_single_blocks - single_blocks_to_swap:
|
||||
bidx_cuda = single_blocks_to_swap - (num_single_blocks - bidx)
|
||||
flux.single_blocks[bidx].to("cpu")
|
||||
flux.single_blocks[bidx_cuda].to(accelerator.device)
|
||||
# print(f"Move single block {bidx} to cpu and {bidx_cuda} to device")
|
||||
|
||||
return optimizer_hook
|
||||
|
||||
parameter.register_post_accumulate_grad_hook(create_optimizer_hook(block_type, block_idx))
|
||||
parameter_optimizer_map[parameter] = opt_idx
|
||||
num_parameters_per_group[opt_idx] += 1
|
||||
|
||||
# epoch数を計算する
|
||||
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
|
||||
num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
|
||||
if (args.save_n_epoch_ratio is not None) and (args.save_n_epoch_ratio > 0):
|
||||
args.save_every_n_epochs = math.floor(num_train_epochs / args.save_n_epoch_ratio) or 1
|
||||
|
||||
# 学習する
|
||||
# total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps
|
||||
accelerator.print("running training / 学習開始")
|
||||
accelerator.print(f" num examples / サンプル数: {train_dataset_group.num_train_images}")
|
||||
accelerator.print(f" num batches per epoch / 1epochのバッチ数: {len(train_dataloader)}")
|
||||
accelerator.print(f" num epochs / epoch数: {num_train_epochs}")
|
||||
accelerator.print(
|
||||
f" batch size per device / バッチサイズ: {', '.join([str(d.batch_size) for d in train_dataset_group.datasets])}"
|
||||
)
|
||||
# accelerator.print(
|
||||
# f" total train batch size (with parallel & distributed & accumulation) / 総バッチサイズ(並列学習、勾配合計含む): {total_batch_size}"
|
||||
# )
|
||||
accelerator.print(f" gradient accumulation steps / 勾配を合計するステップ数 = {args.gradient_accumulation_steps}")
|
||||
accelerator.print(f" total optimization steps / 学習ステップ数: {args.max_train_steps}")
|
||||
|
||||
progress_bar = tqdm(range(args.max_train_steps), smoothing=0, disable=not accelerator.is_local_main_process, desc="steps")
|
||||
global_step = 0
|
||||
|
||||
noise_scheduler = FlowMatchEulerDiscreteScheduler(num_train_timesteps=1000, shift=args.discrete_flow_shift)
|
||||
noise_scheduler_copy = copy.deepcopy(noise_scheduler)
|
||||
|
||||
if accelerator.is_main_process:
|
||||
init_kwargs = {}
|
||||
if args.wandb_run_name:
|
||||
init_kwargs["wandb"] = {"name": args.wandb_run_name}
|
||||
if args.log_tracker_config is not None:
|
||||
init_kwargs = toml.load(args.log_tracker_config)
|
||||
accelerator.init_trackers(
|
||||
"finetuning" if args.log_tracker_name is None else args.log_tracker_name,
|
||||
config=train_util.get_sanitized_config_or_none(args),
|
||||
init_kwargs=init_kwargs,
|
||||
)
|
||||
|
||||
if args.double_blocks_to_swap is not None or args.single_blocks_to_swap is not None:
|
||||
flux.prepare_block_swap_before_forward()
|
||||
|
||||
# For --sample_at_first
|
||||
#flux_train_utils.sample_images(accelerator, args, 0, global_step, flux, ae, [clip_l, t5xxl], sample_prompts_te_outputs)
|
||||
|
||||
loss_recorder = train_util.LossRecorder()
|
||||
epoch = 0 # avoid error when max_train_steps is 0
|
||||
for epoch in range(num_train_epochs):
|
||||
accelerator.print(f"\nepoch {epoch+1}/{num_train_epochs}")
|
||||
current_epoch.value = epoch + 1
|
||||
|
||||
for m in training_models:
|
||||
m.train()
|
||||
|
||||
for step, batch in enumerate(train_dataloader):
|
||||
current_step.value = global_step
|
||||
|
||||
if args.blockwise_fused_optimizers:
|
||||
optimizer_hooked_count = {i: 0 for i in range(len(optimizers))} # reset counter for each step
|
||||
|
||||
with accelerator.accumulate(*training_models):
|
||||
if "latents" in batch and batch["latents"] is not None:
|
||||
latents = batch["latents"].to(accelerator.device, dtype=weight_dtype)
|
||||
else:
|
||||
with torch.no_grad():
|
||||
# encode images to latents. images are [-1, 1]
|
||||
latents = ae.encode(batch["images"])
|
||||
|
||||
# NaNが含まれていれば警告を表示し0に置き換える
|
||||
if torch.any(torch.isnan(latents)):
|
||||
accelerator.print("NaN found in latents, replacing with zeros")
|
||||
latents = torch.nan_to_num(latents, 0, out=latents)
|
||||
|
||||
text_encoder_outputs_list = batch.get("text_encoder_outputs_list", None)
|
||||
if text_encoder_outputs_list is not None:
|
||||
text_encoder_conds = text_encoder_outputs_list
|
||||
else:
|
||||
# not cached or training, so get from text encoders
|
||||
tokens_and_masks = batch["input_ids_list"]
|
||||
with torch.no_grad():
|
||||
input_ids = [ids.to(accelerator.device) for ids in batch["input_ids_list"]]
|
||||
text_encoder_conds = text_encoding_strategy.encode_tokens(
|
||||
tokenize_strategy, [clip_l, t5xxl], input_ids, args.apply_t5_attn_mask
|
||||
)
|
||||
if args.full_fp16:
|
||||
text_encoder_conds = [c.to(weight_dtype) for c in text_encoder_conds]
|
||||
|
||||
# TODO support some features for noise implemented in get_noise_noisy_latents_and_timesteps
|
||||
|
||||
# Sample noise that we'll add to the latents
|
||||
noise = torch.randn_like(latents)
|
||||
bsz = latents.shape[0]
|
||||
|
||||
# get noisy model input and timesteps
|
||||
noisy_model_input, timesteps, sigmas = flux_train_utils.get_noisy_model_input_and_timesteps(
|
||||
args, noise_scheduler, latents, noise, accelerator.device, weight_dtype
|
||||
)
|
||||
|
||||
# pack latents and get img_ids
|
||||
packed_noisy_model_input = flux_utils.pack_latents(noisy_model_input) # b, c, h*2, w*2 -> b, h*w, c*4
|
||||
packed_latent_height, packed_latent_width = noisy_model_input.shape[2] // 2, noisy_model_input.shape[3] // 2
|
||||
img_ids = flux_utils.prepare_img_ids(bsz, packed_latent_height, packed_latent_width).to(device=accelerator.device)
|
||||
|
||||
# get guidance
|
||||
guidance_vec = torch.full((bsz,), args.guidance_scale, device=accelerator.device)
|
||||
|
||||
# call model
|
||||
l_pooled, t5_out, txt_ids = text_encoder_conds
|
||||
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(
|
||||
img=packed_noisy_model_input,
|
||||
img_ids=img_ids,
|
||||
txt=t5_out,
|
||||
txt_ids=txt_ids,
|
||||
y=l_pooled,
|
||||
timesteps=timesteps / 1000,
|
||||
guidance=guidance_vec,
|
||||
)
|
||||
|
||||
# unpack latents
|
||||
model_pred = flux_utils.unpack_latents(model_pred, packed_latent_height, packed_latent_width)
|
||||
|
||||
# apply model prediction type
|
||||
model_pred, weighting = flux_train_utils.apply_model_prediction_type(args, model_pred, noisy_model_input, sigmas)
|
||||
|
||||
# flow matching loss: this is different from SD3
|
||||
target = noise - latents
|
||||
|
||||
# calculate loss
|
||||
loss = train_util.conditional_loss(
|
||||
model_pred.float(), target.float(), reduction="none", loss_type=args.loss_type, huber_c=None
|
||||
)
|
||||
if weighting is not None:
|
||||
loss = loss * weighting
|
||||
if args.masked_loss or ("alpha_masks" in batch and batch["alpha_masks"] is not None):
|
||||
loss = apply_masked_loss(loss, batch)
|
||||
loss = loss.mean([1, 2, 3])
|
||||
|
||||
loss_weights = batch["loss_weights"] # 各sampleごとのweight
|
||||
loss = loss * loss_weights
|
||||
loss = loss.mean()
|
||||
|
||||
# backward
|
||||
accelerator.backward(loss)
|
||||
|
||||
if not (args.fused_backward_pass or args.blockwise_fused_optimizers):
|
||||
if accelerator.sync_gradients and args.max_grad_norm != 0.0:
|
||||
params_to_clip = []
|
||||
for m in training_models:
|
||||
params_to_clip.extend(m.parameters())
|
||||
accelerator.clip_grad_norm_(params_to_clip, args.max_grad_norm)
|
||||
|
||||
optimizer.step()
|
||||
lr_scheduler.step()
|
||||
optimizer.zero_grad(set_to_none=True)
|
||||
else:
|
||||
# optimizer.step() and optimizer.zero_grad() are called in the optimizer hook
|
||||
lr_scheduler.step()
|
||||
if args.blockwise_fused_optimizers:
|
||||
for i in range(1, len(optimizers)):
|
||||
lr_schedulers[i].step()
|
||||
|
||||
# Checks if the accelerator has performed an optimization step behind the scenes
|
||||
if accelerator.sync_gradients:
|
||||
progress_bar.update(1)
|
||||
global_step += 1
|
||||
|
||||
flux_train_utils.sample_images(
|
||||
accelerator, args, None, global_step, flux, ae, [clip_l, t5xxl], sample_prompts_te_outputs
|
||||
)
|
||||
|
||||
# 指定ステップごとにモデルを保存
|
||||
if args.save_every_n_steps is not None and global_step % args.save_every_n_steps == 0:
|
||||
accelerator.wait_for_everyone()
|
||||
if accelerator.is_main_process:
|
||||
flux_train_utils.save_flux_model_on_epoch_end_or_stepwise(
|
||||
args,
|
||||
False,
|
||||
accelerator,
|
||||
save_dtype,
|
||||
epoch,
|
||||
num_train_epochs,
|
||||
global_step,
|
||||
accelerator.unwrap_model(flux),
|
||||
)
|
||||
|
||||
current_loss = loss.detach().item() # 平均なのでbatch sizeは関係ないはず
|
||||
if args.logging_dir is not None:
|
||||
logs = {"loss": current_loss}
|
||||
train_util.append_lr_to_logs(logs, lr_scheduler, args.optimizer_type, including_unet=True)
|
||||
|
||||
accelerator.log(logs, step=global_step)
|
||||
|
||||
loss_recorder.add(epoch=epoch, step=step, loss=current_loss)
|
||||
avr_loss: float = loss_recorder.moving_average
|
||||
logs = {"avr_loss": avr_loss} # , "lr": lr_scheduler.get_last_lr()[0]}
|
||||
progress_bar.set_postfix(**logs)
|
||||
|
||||
if global_step >= args.max_train_steps:
|
||||
break
|
||||
|
||||
if args.logging_dir is not None:
|
||||
logs = {"loss/epoch": loss_recorder.moving_average}
|
||||
accelerator.log(logs, step=epoch + 1)
|
||||
|
||||
accelerator.wait_for_everyone()
|
||||
|
||||
if args.save_every_n_epochs is not None:
|
||||
if accelerator.is_main_process:
|
||||
flux_train_utils.save_flux_model_on_epoch_end_or_stepwise(
|
||||
args,
|
||||
True,
|
||||
accelerator,
|
||||
save_dtype,
|
||||
epoch,
|
||||
num_train_epochs,
|
||||
global_step,
|
||||
accelerator.unwrap_model(flux),
|
||||
)
|
||||
|
||||
flux_train_utils.sample_images(
|
||||
accelerator, args, epoch + 1, global_step, flux, ae, [clip_l, t5xxl], sample_prompts_te_outputs
|
||||
)
|
||||
|
||||
is_main_process = accelerator.is_main_process
|
||||
# if is_main_process:
|
||||
flux = accelerator.unwrap_model(flux)
|
||||
|
||||
accelerator.end_training()
|
||||
|
||||
if args.save_state or args.save_state_on_train_end:
|
||||
train_util.save_state_on_train_end(args, accelerator)
|
||||
|
||||
del accelerator # この後メモリを使うのでこれは消す
|
||||
|
||||
if is_main_process:
|
||||
flux_train_utils.save_flux_model_on_train_end(args, save_dtype, epoch, global_step, flux)
|
||||
logger.info("model saved.")
|
||||
|
||||
|
||||
def setup_parser() -> argparse.ArgumentParser:
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
add_logging_arguments(parser)
|
||||
train_util.add_sd_models_arguments(parser) # TODO split this
|
||||
train_util.add_dataset_arguments(parser, True, True, True)
|
||||
train_util.add_training_arguments(parser, False)
|
||||
train_util.add_masked_loss_arguments(parser)
|
||||
deepspeed_utils.add_deepspeed_arguments(parser)
|
||||
train_util.add_sd_saving_arguments(parser)
|
||||
train_util.add_optimizer_arguments(parser)
|
||||
config_util.add_config_arguments(parser)
|
||||
add_custom_train_arguments(parser) # TODO remove this from here
|
||||
flux_train_utils.add_flux_train_arguments(parser)
|
||||
|
||||
parser.add_argument(
|
||||
"--fused_optimizer_groups",
|
||||
type=int,
|
||||
default=None,
|
||||
help="**this option is not working** will be removed in the future / このオプションは動作しません。将来削除されます",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--blockwise_fused_optimizers",
|
||||
action="store_true",
|
||||
help="enable blockwise optimizers for fused backward pass and optimizer step / fused backward passとoptimizer step のためブロック単位のoptimizerを有効にする",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--skip_latents_validity_check",
|
||||
action="store_true",
|
||||
help="skip latents validity check / latentsの正当性チェックをスキップする",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--double_blocks_to_swap",
|
||||
type=int,
|
||||
default=None,
|
||||
help="[EXPERIMENTAL] "
|
||||
"Sets the number of 'double_blocks' (~640MB) to swap during the forward and backward passes."
|
||||
"Increasing this number lowers the overall VRAM used during training at the expense of training speed (s/it)."
|
||||
" / 順伝播および逆伝播中にスワップする'変換ブロック'(約640MB)の数を設定します。"
|
||||
"この数を増やすと、トレーニング中のVRAM使用量が減りますが、トレーニング速度(s/it)も低下します。",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--single_blocks_to_swap",
|
||||
type=int,
|
||||
default=None,
|
||||
help="[EXPERIMENTAL] "
|
||||
"Sets the number of 'single_blocks' (~320MB) to swap during the forward and backward passes."
|
||||
"Increasing this number lowers the overall VRAM used during training at the expense of training speed (s/it)."
|
||||
" / 順伝播および逆伝播中にスワップする'変換ブロック'(約320MB)の数を設定します。"
|
||||
"この数を増やすと、トレーニング中のVRAM使用量が減りますが、トレーニング速度(s/it)も低下します。",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--cpu_offload_checkpointing",
|
||||
action="store_true",
|
||||
help="[EXPERIMENTAL] enable offloading of tensors to CPU during checkpointing / チェックポイント時にテンソルをCPUにオフロードする",
|
||||
)
|
||||
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)
|
||||
|
||||
train(args)
|
||||
+120
-243
@@ -15,7 +15,6 @@ import copy
|
||||
import math
|
||||
import os
|
||||
from multiprocessing import Value
|
||||
from typing import List
|
||||
import toml
|
||||
|
||||
from tqdm import tqdm
|
||||
@@ -27,7 +26,7 @@ init_ipex()
|
||||
|
||||
from accelerate.utils import set_seed
|
||||
from .library import deepspeed_utils, flux_train_utils, flux_utils, strategy_base, strategy_flux
|
||||
from .library.sd3_train_utils import load_prompts, FlowMatchEulerDiscreteScheduler
|
||||
from .library.sd3_train_utils import FlowMatchEulerDiscreteScheduler
|
||||
|
||||
from .library import train_util as train_util
|
||||
|
||||
@@ -50,6 +49,12 @@ from .library.custom_train_functions import apply_masked_loss, add_custom_train_
|
||||
class FluxTrainer:
|
||||
def __init__(self):
|
||||
self.sample_prompts_te_outputs = None
|
||||
|
||||
def sample_images(self, epoch, global_step, validation_settings):
|
||||
image_tensors = flux_train_utils.sample_images(
|
||||
self.accelerator, self.args, epoch, global_step, self.unet, self.vae, self.text_encoder, self.sample_prompts_te_outputs, validation_settings)
|
||||
return image_tensors
|
||||
|
||||
def init_train(self, args):
|
||||
train_util.verify_training_args(args)
|
||||
train_util.prepare_dataset_args(args, True)
|
||||
@@ -57,9 +62,10 @@ class FluxTrainer:
|
||||
deepspeed_utils.prepare_deepspeed_args(args)
|
||||
setup_logging(args, reset=True)
|
||||
|
||||
# assert (
|
||||
# not args.weighted_captions
|
||||
# ), "weighted_captions is not supported currently / weighted_captionsは現在サポートされていません"
|
||||
# temporary: backward compatibility for deprecated options. remove in the future
|
||||
if not args.skip_cache_check:
|
||||
args.skip_cache_check = args.skip_latents_validity_check
|
||||
|
||||
if args.cache_text_encoder_outputs_to_disk and not args.cache_text_encoder_outputs:
|
||||
logger.warning(
|
||||
"cache_text_encoder_outputs_to_disk is enabled, so cache_text_encoder_outputs is also enabled / cache_text_encoder_outputs_to_diskが有効になっているため、cache_text_encoder_outputsも有効になります"
|
||||
@@ -72,11 +78,17 @@ class FluxTrainer:
|
||||
)
|
||||
args.gradient_checkpointing = True
|
||||
|
||||
assert (
|
||||
args.blocks_to_swap is None or args.blocks_to_swap == 0
|
||||
) or not args.cpu_offload_checkpointing, (
|
||||
"blocks_to_swap is not supported with cpu_offload_checkpointing / blocks_to_swapはcpu_offload_checkpointingと併用できません"
|
||||
)
|
||||
|
||||
cache_latents = args.cache_latents
|
||||
use_dreambooth_method = args.in_json is None
|
||||
|
||||
if args.seed is not None:
|
||||
set_seed(args.seed) # 乱数系列を初期化する
|
||||
set_seed(args.seed)
|
||||
|
||||
# prepare caching strategy: this must be set before preparing dataset. because dataset may use this strategy for initialization.
|
||||
if args.cache_latents:
|
||||
@@ -85,7 +97,7 @@ class FluxTrainer:
|
||||
)
|
||||
strategy_base.LatentsCachingStrategy.set_strategy(latents_caching_strategy)
|
||||
|
||||
# データセットを準備する
|
||||
# Prepare the dataset
|
||||
if args.dataset_class is None:
|
||||
blueprint_generator = BlueprintGenerator(ConfigSanitizer(True, True, args.masked_loss, True))
|
||||
if args.dataset_config is not None:
|
||||
@@ -137,13 +149,19 @@ class FluxTrainer:
|
||||
|
||||
train_dataset_group.verify_bucket_reso_steps(16) # TODO これでいいか確認
|
||||
|
||||
_, is_schnell, _, _ = flux_utils.analyze_checkpoint_state(args.pretrained_model_name_or_path)
|
||||
if args.debug_dataset:
|
||||
if args.cache_text_encoder_outputs:
|
||||
strategy_base.TextEncoderOutputsCachingStrategy.set_strategy(
|
||||
strategy_flux.FluxTextEncoderOutputsCachingStrategy(
|
||||
args.cache_text_encoder_outputs_to_disk, args.text_encoder_batch_size, False, False
|
||||
args.cache_text_encoder_outputs_to_disk, args.text_encoder_batch_size, args.skip_cache_check, False
|
||||
)
|
||||
)
|
||||
t5xxl_max_token_length = (
|
||||
args.t5xxl_max_token_length if args.t5xxl_max_token_length is not None else (256 if is_schnell else 512)
|
||||
)
|
||||
strategy_base.TokenizeStrategy.set_strategy(strategy_flux.FluxTokenizeStrategy(t5xxl_max_token_length))
|
||||
|
||||
train_dataset_group.set_current_strategies()
|
||||
train_util.debug_dataset(train_dataset_group, True)
|
||||
return
|
||||
@@ -170,18 +188,15 @@ class FluxTrainer:
|
||||
# mixed precisionに対応した型を用意しておき適宜castする
|
||||
weight_dtype, save_dtype = train_util.prepare_dtype(args)
|
||||
|
||||
# モデルを読み込む
|
||||
name = "schnell" if "schnell" in args.pretrained_model_name_or_path else "dev"
|
||||
|
||||
# load VAE for caching latents
|
||||
ae = None
|
||||
if cache_latents:
|
||||
ae = flux_utils.load_ae(name, args.ae, weight_dtype, "cpu")
|
||||
ae = flux_utils.load_ae(args.ae, weight_dtype, "cpu", args.disable_mmap_load_safetensors)
|
||||
ae.to(accelerator.device, dtype=weight_dtype)
|
||||
ae.requires_grad_(False)
|
||||
ae.eval()
|
||||
|
||||
train_dataset_group.new_cache_latents(ae, accelerator.is_main_process)
|
||||
train_dataset_group.new_cache_latents(ae, accelerator)
|
||||
|
||||
ae.to("cpu") # if no sampling, vae can be deleted
|
||||
clean_memory_on_device(accelerator.device)
|
||||
@@ -190,7 +205,7 @@ class FluxTrainer:
|
||||
|
||||
# prepare tokenize strategy
|
||||
if args.t5xxl_max_token_length is None:
|
||||
if name == "schnell":
|
||||
if is_schnell:
|
||||
t5xxl_max_token_length = 256
|
||||
else:
|
||||
t5xxl_max_token_length = 512
|
||||
@@ -201,8 +216,8 @@ class FluxTrainer:
|
||||
strategy_base.TokenizeStrategy.set_strategy(flux_tokenize_strategy)
|
||||
|
||||
# load clip_l, t5xxl for caching text encoder outputs
|
||||
clip_l = flux_utils.load_clip_l(args.clip_l, weight_dtype, "cpu")
|
||||
t5xxl = flux_utils.load_t5xxl(args.t5xxl, weight_dtype, "cpu")
|
||||
clip_l = flux_utils.load_clip_l(args.clip_l, weight_dtype, "cpu", args.disable_mmap_load_safetensors)
|
||||
t5xxl = flux_utils.load_t5xxl(args.t5xxl, weight_dtype, "cpu", args.disable_mmap_load_safetensors)
|
||||
clip_l.eval()
|
||||
t5xxl.eval()
|
||||
clip_l.requires_grad_(False)
|
||||
@@ -215,8 +230,8 @@ class FluxTrainer:
|
||||
sample_prompts_te_outputs = None
|
||||
if args.cache_text_encoder_outputs:
|
||||
# Text Encodes are eval and no grad here
|
||||
clip_l.to(accelerator.device, dtype=weight_dtype)
|
||||
t5xxl.to(accelerator.device, dtype=weight_dtype)
|
||||
clip_l.to(accelerator.device)
|
||||
t5xxl.to(accelerator.device)
|
||||
|
||||
text_encoder_caching_strategy = strategy_flux.FluxTextEncoderOutputsCachingStrategy(
|
||||
args.cache_text_encoder_outputs_to_disk, args.text_encoder_batch_size, False, False, args.apply_t5_attn_mask
|
||||
@@ -224,13 +239,12 @@ class FluxTrainer:
|
||||
strategy_base.TextEncoderOutputsCachingStrategy.set_strategy(text_encoder_caching_strategy)
|
||||
|
||||
with accelerator.autocast():
|
||||
train_dataset_group.new_cache_text_encoder_outputs([clip_l, t5xxl], accelerator.is_main_process)
|
||||
train_dataset_group.new_cache_text_encoder_outputs([clip_l, t5xxl], accelerator)
|
||||
|
||||
# cache sample prompt's embeddings to free text encoder's memory
|
||||
if args.sample_prompts is not None:
|
||||
logger.info(f"cache Text Encoder outputs for sample prompt: {args.sample_prompts}")
|
||||
|
||||
tokenize_strategy: strategy_flux.FluxTokenizeStrategy = strategy_base.TokenizeStrategy.get_strategy()
|
||||
text_encoding_strategy: strategy_flux.FluxTextEncodingStrategy = strategy_base.TextEncodingStrategy.get_strategy()
|
||||
|
||||
prompts = []
|
||||
@@ -259,9 +273,9 @@ class FluxTrainer:
|
||||
for p in [prompt_dict.get("prompt", ""), prompt_dict.get("negative_prompt", "")]:
|
||||
if p not in sample_prompts_te_outputs:
|
||||
logger.info(f"cache Text Encoder outputs for prompt: {p}")
|
||||
tokens_and_masks = tokenize_strategy.tokenize(p)
|
||||
tokens_and_masks = flux_tokenize_strategy.tokenize(p)
|
||||
sample_prompts_te_outputs[p] = text_encoding_strategy.encode_tokens(
|
||||
tokenize_strategy, [clip_l, t5xxl], tokens_and_masks, args.apply_t5_attn_mask
|
||||
flux_tokenize_strategy, [clip_l, t5xxl], tokens_and_masks, args.apply_t5_attn_mask
|
||||
)
|
||||
self.sample_prompts_te_outputs = sample_prompts_te_outputs
|
||||
accelerator.wait_for_everyone()
|
||||
@@ -272,26 +286,43 @@ class FluxTrainer:
|
||||
clean_memory_on_device(accelerator.device)
|
||||
|
||||
# load FLUX
|
||||
# if we load to cpu, flux.to(fp8) takes a long time
|
||||
flux = flux_utils.load_flow_model(name, args.pretrained_model_name_or_path, weight_dtype, "cpu")
|
||||
_, flux = flux_utils.load_flow_model(
|
||||
args.pretrained_model_name_or_path, weight_dtype, "cpu", args.disable_mmap_load_safetensors
|
||||
)
|
||||
|
||||
if args.gradient_checkpointing:
|
||||
flux.enable_gradient_checkpointing(args.cpu_offload_checkpointing)
|
||||
flux.enable_gradient_checkpointing(cpu_offload=args.cpu_offload_checkpointing)
|
||||
|
||||
flux.requires_grad_(True)
|
||||
|
||||
is_swapping_blocks = args.double_blocks_to_swap is not None or args.single_blocks_to_swap is not None
|
||||
if is_swapping_blocks:
|
||||
# block swap
|
||||
|
||||
# backward compatibility
|
||||
if args.blocks_to_swap is None:
|
||||
blocks_to_swap = args.double_blocks_to_swap or 0
|
||||
if args.single_blocks_to_swap is not None:
|
||||
blocks_to_swap += args.single_blocks_to_swap // 2
|
||||
if blocks_to_swap > 0:
|
||||
logger.warning(
|
||||
"double_blocks_to_swap and single_blocks_to_swap are deprecated. Use blocks_to_swap instead."
|
||||
" / double_blocks_to_swapとsingle_blocks_to_swapは非推奨です。blocks_to_swapを使ってください。"
|
||||
)
|
||||
logger.info(
|
||||
f"double_blocks_to_swap={args.double_blocks_to_swap} and single_blocks_to_swap={args.single_blocks_to_swap} are converted to blocks_to_swap={blocks_to_swap}."
|
||||
)
|
||||
args.blocks_to_swap = blocks_to_swap
|
||||
del blocks_to_swap
|
||||
|
||||
self.is_swapping_blocks = args.blocks_to_swap is not None and args.blocks_to_swap > 0
|
||||
if self.is_swapping_blocks:
|
||||
# Swap blocks between CPU and GPU to reduce memory usage, in forward and backward passes.
|
||||
# This idea is based on 2kpr's great work. Thank you!
|
||||
logger.info(
|
||||
f"enable block swap: double_blocks_to_swap={args.double_blocks_to_swap}, single_blocks_to_swap={args.single_blocks_to_swap}"
|
||||
)
|
||||
flux.enable_block_swap(args.double_blocks_to_swap, args.single_blocks_to_swap)
|
||||
logger.info(f"enable block swap: blocks_to_swap={args.blocks_to_swap}")
|
||||
flux.enable_block_swap(args.blocks_to_swap, accelerator.device)
|
||||
|
||||
if not cache_latents:
|
||||
# load VAE here if not cached
|
||||
ae = flux_utils.load_ae(name, args.ae, weight_dtype, "cpu")
|
||||
ae = flux_utils.load_ae(args.ae, weight_dtype, "cpu")
|
||||
ae.requires_grad_(False)
|
||||
ae.eval()
|
||||
ae.to(accelerator.device, dtype=weight_dtype)
|
||||
@@ -330,15 +361,15 @@ class FluxTrainer:
|
||||
# determine target layer and block index for each parameter
|
||||
block_type = "other" # double, single or other
|
||||
if np[0].startswith("double_blocks"):
|
||||
block_idx = int(np[0].split(".")[1])
|
||||
block_index = int(np[0].split(".")[1])
|
||||
block_type = "double"
|
||||
elif np[0].startswith("single_blocks"):
|
||||
block_idx = int(np[0].split(".")[1])
|
||||
block_index = int(np[0].split(".")[1])
|
||||
block_type = "single"
|
||||
else:
|
||||
block_idx = -1
|
||||
block_index = -1
|
||||
|
||||
param_group_key = (block_type, block_idx)
|
||||
param_group_key = (block_type, block_index)
|
||||
if param_group_key not in param_group:
|
||||
param_group[param_group_key] = []
|
||||
param_group[param_group_key].append(p)
|
||||
@@ -362,8 +393,13 @@ class FluxTrainer:
|
||||
|
||||
logger.info(f"using {len(optimizers)} optimizers for blockwise fused optimizers")
|
||||
|
||||
if train_util.is_schedulefree_optimizer(optimizers[0], args):
|
||||
raise ValueError("Schedule-free optimizer is not supported with blockwise fused optimizers")
|
||||
self.optimizer_train_fn = lambda: None # dummy function
|
||||
self.optimizer_eval_fn = lambda: None # dummy function
|
||||
else:
|
||||
_, _, optimizer = train_util.get_optimizer(args, trainable_params=params_to_optimize)
|
||||
self.optimizer_train_fn, self.optimizer_eval_fn = train_util.get_optimizer_train_eval_fn(optimizer, args)
|
||||
|
||||
# prepare dataloader
|
||||
# strategies are set here because they cannot be referenced in another process. Copy them with the dataset
|
||||
@@ -439,9 +475,9 @@ class FluxTrainer:
|
||||
else:
|
||||
# accelerator does some magic
|
||||
# if we doesn't swap blocks, we can move the model to device
|
||||
flux = accelerator.prepare(flux, device_placement=[not is_swapping_blocks])
|
||||
if is_swapping_blocks:
|
||||
flux.move_to_device_except_swap_blocks(accelerator.device) # reduce peak memory usage
|
||||
flux = accelerator.prepare(flux, device_placement=[not self.is_swapping_blocks])
|
||||
if self.is_swapping_blocks:
|
||||
accelerator.unwrap_model(flux).move_to_device_except_swap_blocks(accelerator.device) # reduce peak memory usage
|
||||
optimizer, train_dataloader, lr_scheduler = accelerator.prepare(optimizer, train_dataloader, lr_scheduler)
|
||||
|
||||
# 実験的機能:勾配も含めたfp16学習を行う PyTorchにパッチを当ててfp16でのgrad scaleを有効にする
|
||||
@@ -458,88 +494,21 @@ class FluxTrainer:
|
||||
from .library import adafactor_fused
|
||||
|
||||
adafactor_fused.patch_adafactor_fused(optimizer)
|
||||
double_blocks_to_swap = args.double_blocks_to_swap
|
||||
single_blocks_to_swap = args.single_blocks_to_swap
|
||||
num_double_blocks = len(flux.double_blocks)
|
||||
num_single_blocks = len(flux.single_blocks)
|
||||
handled_double_block_indices = set()
|
||||
handled_single_block_indices = set()
|
||||
|
||||
for param_group, param_name_group in zip(optimizer.param_groups, param_names):
|
||||
for parameter, param_name in zip(param_group["params"], param_name_group):
|
||||
if parameter.requires_grad:
|
||||
grad_hook = None
|
||||
|
||||
if double_blocks_to_swap:
|
||||
if param_name.startswith("double_blocks"):
|
||||
block_idx = int(param_name.split(".")[1])
|
||||
if (
|
||||
block_idx not in handled_double_block_indices
|
||||
and block_idx >= (num_double_blocks - double_blocks_to_swap) - 1
|
||||
and block_idx < num_double_blocks - 1
|
||||
):
|
||||
# swap next (already backpropagated) block
|
||||
handled_double_block_indices.add(block_idx)
|
||||
block_idx_cpu = block_idx + 1
|
||||
block_idx_cuda = double_blocks_to_swap - (num_double_blocks - block_idx_cpu)
|
||||
|
||||
# create swap hook
|
||||
def create_double_swap_grad_hook(bidx, bidx_cuda):
|
||||
def __grad_hook(tensor: torch.Tensor):
|
||||
def create_grad_hook(p_name, p_group):
|
||||
def grad_hook(tensor: torch.Tensor):
|
||||
if accelerator.sync_gradients and args.max_grad_norm != 0.0:
|
||||
accelerator.clip_grad_norm_(tensor, args.max_grad_norm)
|
||||
optimizer.step_param(tensor, param_group)
|
||||
optimizer.step_param(tensor, p_group)
|
||||
tensor.grad = None
|
||||
|
||||
# swap blocks if necessary
|
||||
flux.double_blocks[bidx].to("cpu")
|
||||
flux.double_blocks[bidx_cuda].to(accelerator.device)
|
||||
# print(f"Move double block {bidx} to cpu and {bidx_cuda} to device")
|
||||
return grad_hook
|
||||
|
||||
return __grad_hook
|
||||
|
||||
grad_hook = create_double_swap_grad_hook(block_idx_cpu, block_idx_cuda)
|
||||
if single_blocks_to_swap:
|
||||
if param_name.startswith("single_blocks"):
|
||||
block_idx = int(param_name.split(".")[1])
|
||||
if (
|
||||
block_idx not in handled_single_block_indices
|
||||
and block_idx >= (num_single_blocks - single_blocks_to_swap) - 1
|
||||
and block_idx < num_single_blocks - 1
|
||||
):
|
||||
handled_single_block_indices.add(block_idx)
|
||||
block_idx_cpu = block_idx + 1
|
||||
block_idx_cuda = single_blocks_to_swap - (num_single_blocks - block_idx_cpu)
|
||||
# print(param_name, block_idx_cpu, block_idx_cuda)
|
||||
|
||||
# create swap hook
|
||||
def create_single_swap_grad_hook(bidx, bidx_cuda):
|
||||
def __grad_hook(tensor: torch.Tensor):
|
||||
if accelerator.sync_gradients and args.max_grad_norm != 0.0:
|
||||
accelerator.clip_grad_norm_(tensor, args.max_grad_norm)
|
||||
optimizer.step_param(tensor, param_group)
|
||||
tensor.grad = None
|
||||
|
||||
# swap blocks if necessary
|
||||
flux.single_blocks[bidx].to("cpu")
|
||||
flux.single_blocks[bidx_cuda].to(accelerator.device)
|
||||
# print(f"Move single block {bidx} to cpu and {bidx_cuda} to device")
|
||||
|
||||
return __grad_hook
|
||||
|
||||
grad_hook = create_single_swap_grad_hook(block_idx_cpu, block_idx_cuda)
|
||||
|
||||
if grad_hook is None:
|
||||
|
||||
def __grad_hook(tensor: torch.Tensor, param_group=param_group):
|
||||
if accelerator.sync_gradients and args.max_grad_norm != 0.0:
|
||||
accelerator.clip_grad_norm_(tensor, args.max_grad_norm)
|
||||
optimizer.step_param(tensor, param_group)
|
||||
tensor.grad = None
|
||||
|
||||
grad_hook = __grad_hook
|
||||
|
||||
parameter.register_post_accumulate_grad_hook(grad_hook)
|
||||
parameter.register_post_accumulate_grad_hook(create_grad_hook(param_name, param_group))
|
||||
|
||||
elif args.blockwise_fused_optimizers:
|
||||
# prepare for additional optimizers and lr schedulers
|
||||
@@ -556,20 +525,12 @@ class FluxTrainer:
|
||||
num_parameters_per_group = [0] * len(optimizers)
|
||||
parameter_optimizer_map = {}
|
||||
|
||||
double_blocks_to_swap = args.double_blocks_to_swap
|
||||
single_blocks_to_swap = args.single_blocks_to_swap
|
||||
num_double_blocks = len(flux.double_blocks)
|
||||
num_single_blocks = len(flux.single_blocks)
|
||||
|
||||
for opt_idx, optimizer in enumerate(optimizers):
|
||||
for param_group in optimizer.param_groups:
|
||||
for parameter in param_group["params"]:
|
||||
if parameter.requires_grad:
|
||||
block_type, block_idx = block_types_and_indices[opt_idx]
|
||||
|
||||
def create_optimizer_hook(btype, bidx):
|
||||
def optimizer_hook(parameter: torch.Tensor):
|
||||
# print(f"optimizer_hook: {btype}, {bidx}")
|
||||
def grad_hook(parameter: torch.Tensor):
|
||||
if accelerator.sync_gradients and args.max_grad_norm != 0.0:
|
||||
accelerator.clip_grad_norm_(parameter, args.max_grad_norm)
|
||||
|
||||
@@ -579,23 +540,7 @@ class FluxTrainer:
|
||||
optimizers[i].step()
|
||||
optimizers[i].zero_grad(set_to_none=True)
|
||||
|
||||
# swap blocks if necessary
|
||||
if btype == "double" and double_blocks_to_swap:
|
||||
if bidx >= num_double_blocks - double_blocks_to_swap:
|
||||
bidx_cuda = double_blocks_to_swap - (num_double_blocks - bidx)
|
||||
flux.double_blocks[bidx].to("cpu")
|
||||
flux.double_blocks[bidx_cuda].to(accelerator.device)
|
||||
# print(f"Move double block {bidx} to cpu and {bidx_cuda} to device")
|
||||
elif btype == "single" and single_blocks_to_swap:
|
||||
if bidx >= num_single_blocks - single_blocks_to_swap:
|
||||
bidx_cuda = single_blocks_to_swap - (num_single_blocks - bidx)
|
||||
flux.single_blocks[bidx].to("cpu")
|
||||
flux.single_blocks[bidx_cuda].to(accelerator.device)
|
||||
# print(f"Move single block {bidx} to cpu and {bidx_cuda} to device")
|
||||
|
||||
return optimizer_hook
|
||||
|
||||
parameter.register_post_accumulate_grad_hook(create_optimizer_hook(block_type, block_idx))
|
||||
parameter.register_post_accumulate_grad_hook(grad_hook)
|
||||
parameter_optimizer_map[parameter] = opt_idx
|
||||
num_parameters_per_group[opt_idx] += 1
|
||||
|
||||
@@ -607,18 +552,18 @@ class FluxTrainer:
|
||||
|
||||
# 学習する
|
||||
# total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps
|
||||
accelerator.print("running training / 学習開始")
|
||||
accelerator.print(f" num examples / サンプル数: {train_dataset_group.num_train_images}")
|
||||
accelerator.print(f" num batches per epoch / 1epochのバッチ数: {len(train_dataloader)}")
|
||||
accelerator.print(f" num epochs / epoch数: {num_train_epochs}")
|
||||
accelerator.print("running training")
|
||||
accelerator.print(f" num examples: {train_dataset_group.num_train_images}")
|
||||
accelerator.print(f" num batches per epoch: {len(train_dataloader)}")
|
||||
accelerator.print(f" num epochs: {num_train_epochs}")
|
||||
accelerator.print(
|
||||
f" batch size per device / バッチサイズ: {', '.join([str(d.batch_size) for d in train_dataset_group.datasets])}"
|
||||
f" batch size per device: {', '.join([str(d.batch_size) for d in train_dataset_group.datasets])}"
|
||||
)
|
||||
# accelerator.print(
|
||||
# f" total train batch size (with parallel & distributed & accumulation) / 総バッチサイズ(並列学習、勾配合計含む): {total_batch_size}"
|
||||
# )
|
||||
accelerator.print(f" gradient accumulation steps / 勾配を合計するステップ数 = {args.gradient_accumulation_steps}")
|
||||
accelerator.print(f" total optimization steps / 学習ステップ数: {args.max_train_steps}")
|
||||
accelerator.print(f" gradient accumulation steps = {args.gradient_accumulation_steps}")
|
||||
accelerator.print(f" total optimization steps: {args.max_train_steps}")
|
||||
|
||||
progress_bar = tqdm(range(args.max_train_steps), smoothing=0, disable=not accelerator.is_local_main_process, desc="steps")
|
||||
self.global_step = 0
|
||||
@@ -638,13 +583,13 @@ class FluxTrainer:
|
||||
init_kwargs=init_kwargs,
|
||||
)
|
||||
|
||||
if is_swapping_blocks:
|
||||
flux.prepare_block_swap_before_forward()
|
||||
if self.is_swapping_blocks:
|
||||
accelerator.unwrap_model(flux).prepare_block_swap_before_forward()
|
||||
|
||||
# For --sample_at_first
|
||||
#flux_train_utils.sample_images(accelerator, args, 0, global_step, flux, ae, [clip_l, t5xxl], sample_prompts_te_outputs)
|
||||
|
||||
loss_recorder = train_util.LossRecorder()
|
||||
self.loss_recorder = train_util.LossRecorder()
|
||||
epoch = 0 # avoid error when max_train_steps is 0
|
||||
|
||||
self.tokens_and_masks = tokens_and_masks
|
||||
@@ -655,7 +600,6 @@ class FluxTrainer:
|
||||
self.unet = flux
|
||||
self.vae = ae
|
||||
self.text_encoder = [clip_l, t5xxl]
|
||||
self.loss_recorder = loss_recorder
|
||||
self.save_dtype = save_dtype
|
||||
|
||||
def training_loop(break_at_steps, epoch):
|
||||
@@ -680,7 +624,7 @@ class FluxTrainer:
|
||||
else:
|
||||
with torch.no_grad():
|
||||
# encode images to latents. images are [-1, 1]
|
||||
latents = ae.encode(batch["images"])
|
||||
latents = ae.encode(batch["images"].to(ae.dtype)).to(accelerator.device, dtype=weight_dtype)
|
||||
|
||||
# NaNが含まれていれば警告を表示し0に置き換える
|
||||
if torch.any(torch.isnan(latents)):
|
||||
@@ -692,11 +636,11 @@ class FluxTrainer:
|
||||
text_encoder_conds = text_encoder_outputs_list
|
||||
else:
|
||||
# not cached or training, so get from text encoders
|
||||
self.tokens_and_masks = batch["input_ids_list"]
|
||||
tokens_and_masks = batch["input_ids_list"]
|
||||
with torch.no_grad():
|
||||
input_ids = [ids.to(accelerator.device) for ids in batch["input_ids_list"]]
|
||||
text_encoder_conds = text_encoding_strategy.encode_tokens(
|
||||
tokenize_strategy, [clip_l, t5xxl], input_ids, args.apply_t5_attn_mask
|
||||
flux_tokenize_strategy, [clip_l, t5xxl], input_ids, args.apply_t5_attn_mask
|
||||
)
|
||||
if args.full_fp16:
|
||||
text_encoder_conds = [c.to(weight_dtype) for c in text_encoder_conds]
|
||||
@@ -717,13 +661,17 @@ class FluxTrainer:
|
||||
packed_latent_height, packed_latent_width = noisy_model_input.shape[2] // 2, noisy_model_input.shape[3] // 2
|
||||
img_ids = flux_utils.prepare_img_ids(bsz, packed_latent_height, packed_latent_width).to(device=accelerator.device)
|
||||
|
||||
# get guidance
|
||||
guidance_vec = torch.full((bsz,), args.guidance_scale, device=accelerator.device)
|
||||
# get guidance: ensure args.guidance_scale is float
|
||||
guidance_vec = torch.full((bsz,), float(args.guidance_scale), device=accelerator.device)
|
||||
|
||||
# call model
|
||||
l_pooled, t5_out, txt_ids, t5_attn_mask = text_encoder_conds
|
||||
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(
|
||||
@@ -734,11 +682,15 @@ class FluxTrainer:
|
||||
y=l_pooled,
|
||||
timesteps=timesteps / 1000,
|
||||
guidance=guidance_vec,
|
||||
txt_attention_mask=t5_attn_mask,
|
||||
)
|
||||
|
||||
# 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)
|
||||
|
||||
@@ -746,9 +698,8 @@ class FluxTrainer:
|
||||
target = noise - latents
|
||||
|
||||
# calculate loss
|
||||
loss = train_util.conditional_loss(
|
||||
model_pred.float(), target.float(), reduction="none", loss_type=args.loss_type, huber_c=None
|
||||
)
|
||||
huber_c = train_util.get_huber_threshold_if_needed(args, timesteps, noise_scheduler)
|
||||
loss = train_util.conditional_loss(model_pred.float(), target.float(), args.loss_type, "none", huber_c)
|
||||
if weighting is not None:
|
||||
loss = loss * weighting
|
||||
if args.masked_loss or ("alpha_masks" in batch and batch["alpha_masks"] is not None):
|
||||
@@ -784,34 +735,16 @@ class FluxTrainer:
|
||||
progress_bar.update(1)
|
||||
self.global_step += 1
|
||||
|
||||
# flux_train_utils.sample_images(
|
||||
# accelerator, args, None, global_step, flux, ae, [clip_l, t5xxl], sample_prompts_te_outputs
|
||||
# )
|
||||
|
||||
# # 指定ステップごとにモデルを保存
|
||||
# if args.save_every_n_steps is not None and global_step % args.save_every_n_steps == 0:
|
||||
# accelerator.wait_for_everyone()
|
||||
# if accelerator.is_main_process:
|
||||
# flux_train_utils.save_flux_model_on_epoch_end_or_stepwise(
|
||||
# args,
|
||||
# False,
|
||||
# accelerator,
|
||||
# save_dtype,
|
||||
# epoch,
|
||||
# num_train_epochs,
|
||||
# global_step,
|
||||
# accelerator.unwrap_model(flux),
|
||||
# )
|
||||
|
||||
current_loss = loss.detach().item() # 平均なのでbatch sizeは関係ないはず
|
||||
if args.logging_dir is not None:
|
||||
if len(accelerator.trackers) > 0:
|
||||
logs = {"loss": current_loss}
|
||||
train_util.append_lr_to_logs(logs, lr_scheduler, args.optimizer_type, including_unet=True)
|
||||
|
||||
accelerator.log(logs, step=self.global_step)
|
||||
|
||||
loss_recorder.add(epoch=epoch, step=step, loss=current_loss, global_step=self.global_step)
|
||||
avr_loss: float = loss_recorder.moving_average
|
||||
self.loss_recorder.add(epoch=epoch, step=step, loss=current_loss, global_step=self.global_step)
|
||||
avr_loss: float = self.loss_recorder.moving_average
|
||||
logs = {"avr_loss": avr_loss} # , "lr": lr_scheduler.get_last_lr()[0]}
|
||||
progress_bar.set_postfix(**logs)
|
||||
|
||||
@@ -819,46 +752,12 @@ class FluxTrainer:
|
||||
break
|
||||
steps_done += 1
|
||||
|
||||
if args.logging_dir is not None:
|
||||
logs = {"loss/epoch": loss_recorder.moving_average}
|
||||
if len(accelerator.trackers) > 0:
|
||||
logs = {"loss/epoch": self.loss_recorder.moving_average}
|
||||
accelerator.log(logs, step=epoch + 1)
|
||||
return steps_done
|
||||
|
||||
return training_loop
|
||||
#accelerator.wait_for_everyone()
|
||||
|
||||
# if args.save_every_n_epochs is not None:
|
||||
# if accelerator.is_main_process:
|
||||
# flux_train_utils.save_flux_model_on_epoch_end_or_stepwise(
|
||||
# args,
|
||||
# True,
|
||||
# accelerator,
|
||||
# save_dtype,
|
||||
# epoch,
|
||||
# num_train_epochs,
|
||||
# global_step,
|
||||
# accelerator.unwrap_model(flux),
|
||||
# )
|
||||
|
||||
# flux_train_utils.sample_images(
|
||||
# accelerator, args, epoch + 1, global_step, flux, ae, [clip_l, t5xxl], sample_prompts_te_outputs
|
||||
# )
|
||||
|
||||
# is_main_process = accelerator.is_main_process
|
||||
# # if is_main_process:
|
||||
# flux = accelerator.unwrap_model(flux)
|
||||
|
||||
# accelerator.end_training()
|
||||
|
||||
# if args.save_state or args.save_state_on_train_end:
|
||||
# train_util.save_state_on_train_end(args, accelerator)
|
||||
|
||||
# del accelerator # この後メモリを使うのでこれは消す
|
||||
|
||||
# if is_main_process:
|
||||
# flux_train_utils.save_flux_model_on_train_end(args, save_dtype, epoch, global_step, flux)
|
||||
# logger.info("model saved.")
|
||||
|
||||
|
||||
def setup_parser() -> argparse.ArgumentParser:
|
||||
parser = argparse.ArgumentParser()
|
||||
@@ -873,8 +772,15 @@ def setup_parser() -> argparse.ArgumentParser:
|
||||
train_util.add_optimizer_arguments(parser)
|
||||
config_util.add_config_arguments(parser)
|
||||
add_custom_train_arguments(parser) # TODO remove this from here
|
||||
train_util.add_dit_training_arguments(parser)
|
||||
flux_train_utils.add_flux_train_arguments(parser)
|
||||
|
||||
parser.add_argument(
|
||||
"--mem_eff_save",
|
||||
action="store_true",
|
||||
help="[EXPERIMENTAL] use memory efficient custom model saving method / メモリ効率の良い独自のモデル保存方法を使う",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--fused_optimizer_groups",
|
||||
type=int,
|
||||
@@ -891,39 +797,10 @@ def setup_parser() -> argparse.ArgumentParser:
|
||||
action="store_true",
|
||||
help="skip latents validity check / latentsの正当性チェックをスキップする",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--double_blocks_to_swap",
|
||||
type=int,
|
||||
default=None,
|
||||
help="[EXPERIMENTAL] "
|
||||
"Sets the number of 'double_blocks' (~640MB) to swap during the forward and backward passes."
|
||||
"Increasing this number lowers the overall VRAM used during training at the expense of training speed (s/it)."
|
||||
" / 順伝播および逆伝播中にスワップする'変換ブロック'(約640MB)の数を設定します。"
|
||||
"この数を増やすと、トレーニング中のVRAM使用量が減りますが、トレーニング速度(s/it)も低下します。",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--single_blocks_to_swap",
|
||||
type=int,
|
||||
default=None,
|
||||
help="[EXPERIMENTAL] "
|
||||
"Sets the number of 'single_blocks' (~320MB) to swap during the forward and backward passes."
|
||||
"Increasing this number lowers the overall VRAM used during training at the expense of training speed (s/it)."
|
||||
" / 順伝播および逆伝播中にスワップする'変換ブロック'(約320MB)の数を設定します。"
|
||||
"この数を増やすと、トレーニング中のVRAM使用量が減りますが、トレーニング速度(s/it)も低下します。",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--cpu_offload_checkpointing",
|
||||
action="store_true",
|
||||
help="[EXPERIMENTAL] enable offloading of tensors to CPU during checkpointing / チェックポイント時にテンソルをCPUにオフロードする",
|
||||
)
|
||||
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)
|
||||
|
||||
# train(args)
|
||||
|
||||
+218
-154
@@ -1,10 +1,10 @@
|
||||
import torch
|
||||
import copy
|
||||
import math
|
||||
from typing import Any
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
import argparse
|
||||
from .library import flux_models, flux_train_utils, flux_utils, sd3_train_utils, strategy_base, strategy_flux, train_util
|
||||
from .train_network import NetworkTrainer, clean_memory_on_device
|
||||
from .train_network import NetworkTrainer, clean_memory_on_device, setup_parser
|
||||
|
||||
from accelerate import Accelerator
|
||||
|
||||
@@ -17,93 +17,91 @@ class FluxNetworkTrainer(NetworkTrainer):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.sample_prompts_te_outputs = None
|
||||
self.is_schnell: Optional[bool] = None
|
||||
self.is_swapping_blocks: bool = False
|
||||
|
||||
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.fp8_base_unet:
|
||||
args.fp8_base = True # if fp8_base_unet is enabled, fp8_base is also enabled for FLUX.1
|
||||
|
||||
if args.cache_text_encoder_outputs_to_disk and not args.cache_text_encoder_outputs:
|
||||
logger.warning(
|
||||
"cache_text_encoder_outputs_to_disk is enabled, so cache_text_encoder_outputs is also enabled / cache_text_encoder_outputs_to_diskが有効になっているため、cache_text_encoder_outputsも有効になります"
|
||||
)
|
||||
args.cache_text_encoder_outputs = True
|
||||
|
||||
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"
|
||||
), "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"
|
||||
if not args.network_train_unet_only:
|
||||
logger.info(
|
||||
"network for CLIP-L only will be trained. T5XXL will not be trained / CLIP-Lのネットワークのみが学習されます。T5XXLは学習されません"
|
||||
)
|
||||
# prepare CLIP-L/T5XXL training flags
|
||||
self.train_clip_l = not args.network_train_unet_only
|
||||
self.train_t5xxl = False # default is False even if args.network_train_unet_only is False
|
||||
|
||||
if args.max_token_length is not None:
|
||||
logger.warning("max_token_length is not used in Flux training")
|
||||
logger.warning("max_token_length is not used in Flux training / max_token_lengthはFluxのトレーニングでは使用されません")
|
||||
|
||||
assert (
|
||||
args.blocks_to_swap is None or args.blocks_to_swap == 0
|
||||
) or not args.cpu_offload_checkpointing, "blocks_to_swap is not supported with cpu_offload_checkpointing / blocks_to_swapはcpu_offload_checkpointingと併用できません"
|
||||
|
||||
train_dataset_group.verify_bucket_reso_steps(32) # TODO check this
|
||||
|
||||
def get_flux_model_name(self, args):
|
||||
return "schnell" if "schnell" in args.pretrained_model_name_or_path else "dev"
|
||||
|
||||
def load_target_model(self, args, weight_dtype, accelerator):
|
||||
# currently offload to cpu for some models
|
||||
name = self.get_flux_model_name(args)
|
||||
# if we load to cpu, flux.to(fp8) takes a long time
|
||||
model = flux_utils.load_flow_model(name, args.pretrained_model_name_or_path, weight_dtype, "cpu")
|
||||
|
||||
if args.split_mode:
|
||||
model = self.prepare_split_model(model, weight_dtype, accelerator, args)
|
||||
# if the file is fp8 and we are using fp8_base, we can load it as is (fp8)
|
||||
loading_dtype = None if args.fp8_base else weight_dtype
|
||||
|
||||
clip_l = flux_utils.load_clip_l(args.clip_l, weight_dtype, "cpu")
|
||||
# if we load to cpu, flux.to(fp8) takes a long time, so we should load to gpu in future
|
||||
self.is_schnell, model = flux_utils.load_flow_model(
|
||||
args.pretrained_model_name_or_path, loading_dtype, "cpu", disable_mmap=args.disable_mmap_load_safetensors
|
||||
)
|
||||
if args.fp8_base:
|
||||
# check dtype of model
|
||||
if model.dtype == torch.float8_e4m3fnuz or model.dtype == torch.float8_e5m2fnuz:
|
||||
raise ValueError(f"Unsupported fp8 model dtype: {model.dtype}")
|
||||
elif model.dtype == torch.float8_e4m3fn or model.dtype == torch.float8_e5m2:
|
||||
logger.info(f"Loaded {model.dtype} FLUX model")
|
||||
|
||||
self.is_swapping_blocks = args.blocks_to_swap is not None and args.blocks_to_swap > 0
|
||||
if self.is_swapping_blocks:
|
||||
# Swap blocks between CPU and GPU to reduce memory usage, in forward and backward passes.
|
||||
logger.info(f"enable block swap: blocks_to_swap={args.blocks_to_swap}")
|
||||
model.enable_block_swap(args.blocks_to_swap, accelerator.device)
|
||||
|
||||
clip_l = flux_utils.load_clip_l(args.clip_l, weight_dtype, "cpu", disable_mmap=args.disable_mmap_load_safetensors)
|
||||
clip_l.eval()
|
||||
|
||||
# loading t5xxl to cpu takes a long time, so we should load to gpu in future
|
||||
t5xxl = flux_utils.load_t5xxl(args.t5xxl, weight_dtype, "cpu")
|
||||
t5xxl.eval()
|
||||
# if the file is fp8 and we are using fp8_base (not unet), we can load it as is (fp8)
|
||||
if args.fp8_base and not args.fp8_base_unet:
|
||||
loading_dtype = None # as is
|
||||
else:
|
||||
loading_dtype = weight_dtype
|
||||
|
||||
ae = flux_utils.load_ae(name, args.ae, weight_dtype, "cpu")
|
||||
# loading t5xxl to cpu takes a long time, so we should load to gpu in future
|
||||
t5xxl = flux_utils.load_t5xxl(args.t5xxl, loading_dtype, "cpu", disable_mmap=args.disable_mmap_load_safetensors)
|
||||
t5xxl.eval()
|
||||
if args.fp8_base and not args.fp8_base_unet:
|
||||
# check dtype of model
|
||||
if t5xxl.dtype == torch.float8_e4m3fnuz or t5xxl.dtype == torch.float8_e5m2 or t5xxl.dtype == torch.float8_e5m2fnuz:
|
||||
raise ValueError(f"Unsupported fp8 model dtype: {t5xxl.dtype}")
|
||||
elif t5xxl.dtype == torch.float8_e4m3fn:
|
||||
logger.info("Loaded fp8 T5XXL model")
|
||||
|
||||
ae = flux_utils.load_ae(args.ae, weight_dtype, "cpu", disable_mmap=args.disable_mmap_load_safetensors)
|
||||
|
||||
return flux_utils.MODEL_VERSION_FLUX_V1, [clip_l, t5xxl], ae, model
|
||||
|
||||
def prepare_split_model(self, model, weight_dtype, accelerator, args):
|
||||
from accelerate import init_empty_weights
|
||||
|
||||
logger.info("prepare split model")
|
||||
with init_empty_weights():
|
||||
flux_upper = flux_models.FluxUpper(model.params)
|
||||
flux_lower = flux_models.FluxLower(model.params)
|
||||
sd = model.state_dict()
|
||||
|
||||
# lower (trainable)
|
||||
logger.info("load state dict for lower")
|
||||
flux_lower.load_state_dict(sd, strict=False, assign=True)
|
||||
flux_lower.to(dtype=weight_dtype)
|
||||
|
||||
# upper (frozen)
|
||||
logger.info("load state dict for upper")
|
||||
flux_upper.load_state_dict(sd, strict=False, assign=True)
|
||||
|
||||
logger.info("prepare upper model")
|
||||
target_dtype = torch.float8_e4m3fn if args.fp8_base else weight_dtype
|
||||
flux_upper.to(accelerator.device, dtype=target_dtype)
|
||||
flux_upper.eval()
|
||||
|
||||
if args.fp8_base:
|
||||
# this is required to run on fp8
|
||||
flux_upper = accelerator.prepare(flux_upper)
|
||||
|
||||
flux_upper.to("cpu")
|
||||
|
||||
self.flux_upper = flux_upper
|
||||
del model # we don't need model anymore
|
||||
clean_memory_on_device(accelerator.device)
|
||||
|
||||
logger.info("split model prepared")
|
||||
|
||||
return flux_lower
|
||||
|
||||
def get_tokenize_strategy(self, args):
|
||||
name = self.get_flux_model_name(args)
|
||||
_, is_schnell, _, _ = flux_utils.analyze_checkpoint_state(args.pretrained_model_name_or_path)
|
||||
|
||||
if args.t5xxl_max_token_length is None:
|
||||
if name == "schnell":
|
||||
if is_schnell:
|
||||
t5xxl_max_token_length = 256
|
||||
else:
|
||||
t5xxl_max_token_length = 512
|
||||
@@ -123,25 +121,35 @@ class FluxNetworkTrainer(NetworkTrainer):
|
||||
def get_text_encoding_strategy(self, args):
|
||||
return strategy_flux.FluxTextEncodingStrategy(apply_t5_attn_mask=args.apply_t5_attn_mask)
|
||||
|
||||
def post_process_network(self, args, accelerator, network, text_encoders, unet):
|
||||
# check t5xxl is trained or not
|
||||
self.train_t5xxl = network.train_t5xxl
|
||||
|
||||
if self.train_t5xxl and args.cache_text_encoder_outputs:
|
||||
raise ValueError(
|
||||
"T5XXL is trained, so cache_text_encoder_outputs cannot be used / T5XXL学習時はcache_text_encoder_outputsは使用できません"
|
||||
)
|
||||
|
||||
def get_models_for_text_encoding(self, args, accelerator, text_encoders):
|
||||
if args.cache_text_encoder_outputs:
|
||||
if self.is_train_text_encoder(args):
|
||||
if self.train_clip_l and not self.train_t5xxl:
|
||||
return text_encoders[0:1] # only CLIP-L is needed for encoding because T5XXL is cached
|
||||
else:
|
||||
return text_encoders # ignored
|
||||
return None # no text encoders are needed for encoding because both are cached
|
||||
else:
|
||||
return text_encoders # both CLIP-L and T5XXL are needed for encoding
|
||||
|
||||
def get_text_encoders_train_flags(self, args, text_encoders):
|
||||
return [True, False] if self.is_train_text_encoder(args) else [False, False]
|
||||
return [self.train_clip_l, self.train_t5xxl]
|
||||
|
||||
def get_text_encoder_outputs_caching_strategy(self, args):
|
||||
if args.cache_text_encoder_outputs:
|
||||
# if the text encoders is trained, we need tokenization, so is_partial is True
|
||||
return strategy_flux.FluxTextEncoderOutputsCachingStrategy(
|
||||
args.cache_text_encoder_outputs_to_disk,
|
||||
None,
|
||||
False,
|
||||
is_partial=self.is_train_text_encoder(args),
|
||||
args.text_encoder_batch_size,
|
||||
args.skip_cache_check,
|
||||
is_partial=self.train_clip_l or self.train_t5xxl,
|
||||
apply_t5_attn_mask=args.apply_t5_attn_mask,
|
||||
)
|
||||
else:
|
||||
@@ -162,13 +170,20 @@ class FluxNetworkTrainer(NetworkTrainer):
|
||||
|
||||
# When TE is not be trained, it will not be prepared so we need to use explicit autocast
|
||||
logger.info("move text encoders to gpu")
|
||||
text_encoders[0].to(accelerator.device, dtype=weight_dtype)
|
||||
text_encoders[1].to(accelerator.device, dtype=weight_dtype)
|
||||
text_encoders[0].to(accelerator.device, dtype=weight_dtype) # always not fp8
|
||||
text_encoders[1].to(accelerator.device)
|
||||
|
||||
if text_encoders[1].dtype == torch.float8_e4m3fn:
|
||||
# if we load fp8 weights, the model is already fp8, so we use it as is
|
||||
self.prepare_text_encoder_fp8(1, text_encoders[1], text_encoders[1].dtype, weight_dtype)
|
||||
else:
|
||||
# otherwise, we need to convert it to target dtype
|
||||
text_encoders[1].to(weight_dtype)
|
||||
|
||||
with accelerator.autocast():
|
||||
dataset.new_cache_text_encoder_outputs(text_encoders, accelerator.is_main_process)
|
||||
dataset.new_cache_text_encoder_outputs(text_encoders, accelerator)
|
||||
|
||||
# cache sample prompts
|
||||
|
||||
if args.sample_prompts is not None:
|
||||
logger.info(f"cache Text Encoder outputs for sample prompt: {args.sample_prompts}")
|
||||
|
||||
@@ -206,8 +221,10 @@ class FluxNetworkTrainer(NetworkTrainer):
|
||||
tokenize_strategy, text_encoders, tokens_and_masks, args.apply_t5_attn_mask
|
||||
)
|
||||
self.sample_prompts_te_outputs = sample_prompts_te_outputs
|
||||
|
||||
accelerator.wait_for_everyone()
|
||||
|
||||
# move back to cpu
|
||||
if not self.is_train_text_encoder(args):
|
||||
logger.info("move CLIP-L back to cpu")
|
||||
text_encoders[0].to("cpu")
|
||||
@@ -222,33 +239,14 @@ class FluxNetworkTrainer(NetworkTrainer):
|
||||
else:
|
||||
# Text Encoder
|
||||
text_encoders[0].to(accelerator.device, dtype=weight_dtype)
|
||||
text_encoders[1].to(accelerator.device, dtype=weight_dtype)
|
||||
text_encoders[1].to(accelerator.device)
|
||||
|
||||
def sample_images_split_mode(self, accelerator, args, epoch, global_step, flux, ae, text_encoder, sample_prompts_te_outputs, validation_settings):
|
||||
def sample_images(self, epoch, global_step, validation_settings):
|
||||
text_encoders = self.get_models_for_text_encoding(self.args, self.accelerator, self.text_encoder)
|
||||
|
||||
class FluxUpperLowerWrapper(torch.nn.Module):
|
||||
def __init__(self, flux_upper: flux_models.FluxUpper, flux_lower: flux_models.FluxLower, device: torch.device):
|
||||
super().__init__()
|
||||
self.flux_upper = flux_upper
|
||||
self.flux_lower = flux_lower
|
||||
self.target_device = device
|
||||
|
||||
def forward(self, img, img_ids, txt, txt_ids, timesteps, y, guidance=None, txt_attention_mask=None):
|
||||
self.flux_lower.to("cpu")
|
||||
clean_memory_on_device(self.target_device)
|
||||
self.flux_upper.to(self.target_device)
|
||||
img, txt, vec, pe = self.flux_upper(img, img_ids, txt, txt_ids, timesteps, y, guidance, txt_attention_mask)
|
||||
self.flux_upper.to("cpu")
|
||||
clean_memory_on_device(self.target_device)
|
||||
self.flux_lower.to(self.target_device)
|
||||
return self.flux_lower(img, txt, vec, pe, txt_attention_mask)
|
||||
|
||||
wrapper = FluxUpperLowerWrapper(self.flux_upper, flux, accelerator.device)
|
||||
clean_memory_on_device(accelerator.device)
|
||||
image_tensors = flux_train_utils.sample_images(
|
||||
accelerator, args, epoch, global_step, wrapper, ae, text_encoder, sample_prompts_te_outputs, validation_settings
|
||||
)
|
||||
clean_memory_on_device(accelerator.device)
|
||||
self.accelerator, self.args, epoch, global_step, self.unet, self.vae, text_encoders, self.sample_prompts_te_outputs, validation_settings)
|
||||
clean_memory_on_device(self.accelerator.device)
|
||||
return image_tensors
|
||||
|
||||
def get_noise_scheduler(self, args: argparse.Namespace, device: torch.device) -> Any:
|
||||
@@ -256,9 +254,6 @@ class FluxNetworkTrainer(NetworkTrainer):
|
||||
self.noise_scheduler_copy = copy.deepcopy(noise_scheduler)
|
||||
return noise_scheduler
|
||||
|
||||
def is_text_encoder_not_needed_for_training(self, args):
|
||||
return args.cache_text_encoder_outputs and not self.is_train_text_encoder(args)
|
||||
|
||||
def encode_images_to_latents(self, args, accelerator, vae, images):
|
||||
return vae.encode(images)
|
||||
|
||||
@@ -278,55 +273,6 @@ class FluxNetworkTrainer(NetworkTrainer):
|
||||
weight_dtype,
|
||||
train_unet,
|
||||
):
|
||||
# copy from sd3_train.py and modified
|
||||
|
||||
def get_sigmas(timesteps, n_dim=4, dtype=torch.float32):
|
||||
sigmas = self.noise_scheduler_copy.sigmas.to(device=accelerator.device, dtype=dtype)
|
||||
schedule_timesteps = self.noise_scheduler_copy.timesteps.to(accelerator.device)
|
||||
timesteps = timesteps.to(accelerator.device)
|
||||
step_indices = [(schedule_timesteps == t).nonzero().item() for t in timesteps]
|
||||
|
||||
sigma = sigmas[step_indices].flatten()
|
||||
while len(sigma.shape) < n_dim:
|
||||
sigma = sigma.unsqueeze(-1)
|
||||
return sigma
|
||||
|
||||
def compute_density_for_timestep_sampling(
|
||||
weighting_scheme: str, batch_size: int, logit_mean: float = None, logit_std: float = None, mode_scale: float = None
|
||||
):
|
||||
"""Compute the density for sampling the timesteps when doing SD3 training.
|
||||
|
||||
Courtesy: This was contributed by Rafie Walker in https://github.com/huggingface/diffusers/pull/8528.
|
||||
|
||||
SD3 paper reference: https://arxiv.org/abs/2403.03206v1.
|
||||
"""
|
||||
if weighting_scheme == "logit_normal":
|
||||
# See 3.1 in the SD3 paper ($rf/lognorm(0.00,1.00)$).
|
||||
u = torch.normal(mean=logit_mean, std=logit_std, size=(batch_size,), device="cpu")
|
||||
u = torch.nn.functional.sigmoid(u)
|
||||
elif weighting_scheme == "mode":
|
||||
u = torch.rand(size=(batch_size,), device="cpu")
|
||||
u = 1 - u - mode_scale * (torch.cos(math.pi * u / 2) ** 2 - 1 + u)
|
||||
else:
|
||||
u = torch.rand(size=(batch_size,), device="cpu")
|
||||
return u
|
||||
|
||||
def compute_loss_weighting_for_sd3(weighting_scheme: str, sigmas=None):
|
||||
"""Computes loss weighting scheme for SD3 training.
|
||||
|
||||
Courtesy: This was contributed by Rafie Walker in https://github.com/huggingface/diffusers/pull/8528.
|
||||
|
||||
SD3 paper reference: https://arxiv.org/abs/2403.03206v1.
|
||||
"""
|
||||
if weighting_scheme == "sigma_sqrt":
|
||||
weighting = (sigmas**-2.0).float()
|
||||
elif weighting_scheme == "cosmap":
|
||||
bot = 1 - 2 * sigmas + 2 * sigmas**2
|
||||
weighting = 2 / (math.pi * bot)
|
||||
else:
|
||||
weighting = torch.ones_like(sigmas)
|
||||
return weighting
|
||||
|
||||
# Sample noise that we'll add to the latents
|
||||
noise = torch.randn_like(latents)
|
||||
bsz = latents.shape[0]
|
||||
@@ -342,13 +288,14 @@ class FluxNetworkTrainer(NetworkTrainer):
|
||||
img_ids = flux_utils.prepare_img_ids(bsz, packed_latent_height, packed_latent_width).to(device=accelerator.device)
|
||||
|
||||
# get guidance
|
||||
guidance_vec = torch.full((bsz,), args.guidance_scale, device=accelerator.device)
|
||||
# ensure guidance_scale in args is float
|
||||
guidance_vec = torch.full((bsz,), float(args.guidance_scale), device=accelerator.device)
|
||||
|
||||
# ensure the hidden state will require grad
|
||||
if args.gradient_checkpointing:
|
||||
noisy_model_input.requires_grad_(True)
|
||||
for t in text_encoder_conds:
|
||||
if t.dtype.is_floating_point:
|
||||
if t is not None and t.dtype.is_floating_point:
|
||||
t.requires_grad_(True)
|
||||
img_ids.requires_grad_(True)
|
||||
guidance_vec.requires_grad_(True)
|
||||
@@ -358,12 +305,12 @@ class FluxNetworkTrainer(NetworkTrainer):
|
||||
if not args.apply_t5_attn_mask:
|
||||
t5_attn_mask = None
|
||||
|
||||
if not args.split_mode:
|
||||
def call_dit(img, img_ids, t5_out, txt_ids, l_pooled, timesteps, guidance_vec, t5_attn_mask):
|
||||
# normal forward
|
||||
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 = unet(
|
||||
img=packed_noisy_model_input,
|
||||
img=img,
|
||||
img_ids=img_ids,
|
||||
txt=t5_out,
|
||||
txt_ids=txt_ids,
|
||||
@@ -372,6 +319,7 @@ class FluxNetworkTrainer(NetworkTrainer):
|
||||
guidance=guidance_vec,
|
||||
txt_attention_mask=t5_attn_mask,
|
||||
)
|
||||
"""
|
||||
else:
|
||||
# split forward to reduce memory usage
|
||||
assert network.train_blocks == "single", "train_blocks must be single for split mode"
|
||||
@@ -405,17 +353,69 @@ class FluxNetworkTrainer(NetworkTrainer):
|
||||
vec.requires_grad_(True)
|
||||
pe.requires_grad_(True)
|
||||
model_pred = unet(img=intermediate_img, txt=intermediate_txt, vec=vec, pe=pe, txt_attention_mask=t5_attn_mask)
|
||||
"""
|
||||
|
||||
return model_pred
|
||||
|
||||
if args.bypass_flux_guidance:
|
||||
flux_utils.bypass_flux_guidance(unet)
|
||||
|
||||
model_pred = call_dit(
|
||||
img=packed_noisy_model_input,
|
||||
img_ids=img_ids,
|
||||
t5_out=t5_out,
|
||||
txt_ids=txt_ids,
|
||||
l_pooled=l_pooled,
|
||||
timesteps=timesteps,
|
||||
guidance_vec=guidance_vec,
|
||||
t5_attn_mask=t5_attn_mask,
|
||||
)
|
||||
|
||||
# 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)
|
||||
|
||||
# flow matching loss: this is different from SD3
|
||||
target = noise - latents
|
||||
|
||||
return model_pred, target, timesteps, None, weighting
|
||||
# differential output preservation
|
||||
if "custom_attributes" in batch:
|
||||
diff_output_pr_indices = []
|
||||
for i, custom_attributes in enumerate(batch["custom_attributes"]):
|
||||
if "diff_output_preservation" in custom_attributes and custom_attributes["diff_output_preservation"]:
|
||||
diff_output_pr_indices.append(i)
|
||||
|
||||
if len(diff_output_pr_indices) > 0:
|
||||
network.set_multiplier(0.0)
|
||||
unet.prepare_block_swap_before_forward()
|
||||
with torch.no_grad():
|
||||
model_pred_prior = call_dit(
|
||||
img=packed_noisy_model_input[diff_output_pr_indices],
|
||||
img_ids=img_ids[diff_output_pr_indices],
|
||||
t5_out=t5_out[diff_output_pr_indices],
|
||||
txt_ids=txt_ids[diff_output_pr_indices],
|
||||
l_pooled=l_pooled[diff_output_pr_indices],
|
||||
timesteps=timesteps[diff_output_pr_indices],
|
||||
guidance_vec=guidance_vec[diff_output_pr_indices] if guidance_vec is not None else None,
|
||||
t5_attn_mask=t5_attn_mask[diff_output_pr_indices] if t5_attn_mask is not None else None,
|
||||
)
|
||||
network.set_multiplier(1.0) # may be overwritten by "network_multipliers" in the next step
|
||||
|
||||
model_pred_prior = flux_utils.unpack_latents(model_pred_prior, packed_latent_height, packed_latent_width)
|
||||
model_pred_prior, _ = flux_train_utils.apply_model_prediction_type(
|
||||
args,
|
||||
model_pred_prior,
|
||||
noisy_model_input[diff_output_pr_indices],
|
||||
sigmas[diff_output_pr_indices] if sigmas is not None else None,
|
||||
)
|
||||
target[diff_output_pr_indices] = model_pred_prior.to(target.dtype)
|
||||
|
||||
return model_pred, target, timesteps, weighting
|
||||
|
||||
def post_process_loss(self, loss, args, timesteps, noise_scheduler):
|
||||
return loss
|
||||
@@ -434,3 +434,67 @@ class FluxNetworkTrainer(NetworkTrainer):
|
||||
metadata["ss_sigmoid_scale"] = args.sigmoid_scale
|
||||
metadata["ss_model_prediction_type"] = args.model_prediction_type
|
||||
metadata["ss_discrete_flow_shift"] = args.discrete_flow_shift
|
||||
|
||||
def is_text_encoder_not_needed_for_training(self, args):
|
||||
return args.cache_text_encoder_outputs and not self.is_train_text_encoder(args)
|
||||
|
||||
def prepare_text_encoder_grad_ckpt_workaround(self, index, text_encoder):
|
||||
if index == 0: # CLIP-L
|
||||
return super().prepare_text_encoder_grad_ckpt_workaround(index, text_encoder)
|
||||
else: # T5XXL
|
||||
text_encoder.encoder.embed_tokens.requires_grad_(True)
|
||||
|
||||
def prepare_text_encoder_fp8(self, index, text_encoder, te_weight_dtype, weight_dtype):
|
||||
if index == 0: # CLIP-L
|
||||
logger.info(f"prepare CLIP-L for fp8: set to {te_weight_dtype}, set embeddings to {weight_dtype}")
|
||||
text_encoder.to(te_weight_dtype) # fp8
|
||||
text_encoder.text_model.embeddings.to(dtype=weight_dtype)
|
||||
else: # T5XXL
|
||||
|
||||
def prepare_fp8(text_encoder, target_dtype):
|
||||
def forward_hook(module):
|
||||
def forward(hidden_states):
|
||||
hidden_gelu = module.act(module.wi_0(hidden_states))
|
||||
hidden_linear = module.wi_1(hidden_states)
|
||||
hidden_states = hidden_gelu * hidden_linear
|
||||
hidden_states = module.dropout(hidden_states)
|
||||
|
||||
hidden_states = module.wo(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
return forward
|
||||
|
||||
for module in text_encoder.modules():
|
||||
if module.__class__.__name__ in ["T5LayerNorm", "Embedding"]:
|
||||
# print("set", module.__class__.__name__, "to", target_dtype)
|
||||
module.to(target_dtype)
|
||||
if module.__class__.__name__ in ["T5DenseGatedActDense"]:
|
||||
# print("set", module.__class__.__name__, "hooks")
|
||||
module.forward = forward_hook(module)
|
||||
|
||||
if flux_utils.get_t5xxl_actual_dtype(text_encoder) == torch.float8_e4m3fn and text_encoder.dtype == weight_dtype:
|
||||
logger.info(f"T5XXL already prepared for fp8")
|
||||
else:
|
||||
logger.info(f"prepare T5XXL for fp8: set to {te_weight_dtype}, set embeddings to {weight_dtype}, add hooks")
|
||||
text_encoder.to(te_weight_dtype) # fp8
|
||||
prepare_fp8(text_encoder, weight_dtype)
|
||||
|
||||
def prepare_unet_with_accelerator(
|
||||
self, args: argparse.Namespace, accelerator: Accelerator, unet: torch.nn.Module
|
||||
) -> torch.nn.Module:
|
||||
if not self.is_swapping_blocks:
|
||||
return super().prepare_unet_with_accelerator(args, accelerator, unet)
|
||||
|
||||
# if we doesn't swap blocks, we can move the model to device
|
||||
flux: flux_models.Flux = unet
|
||||
flux = accelerator.prepare(flux, device_placement=[not self.is_swapping_blocks])
|
||||
accelerator.unwrap_model(flux).move_to_device_except_swap_blocks(accelerator.device) # reduce peak memory usage
|
||||
accelerator.unwrap_model(flux).prepare_block_swap_before_forward()
|
||||
|
||||
return flux
|
||||
|
||||
|
||||
def setup_parser() -> argparse.ArgumentParser:
|
||||
parser = setup_parser()
|
||||
train_util.add_dit_training_arguments(parser)
|
||||
flux_train_utils.add_flux_train_arguments(parser)
|
||||
|
||||
+9
-12
@@ -10,13 +10,7 @@ import json
|
||||
from pathlib import Path
|
||||
|
||||
# from toolz import curry
|
||||
from typing import (
|
||||
List,
|
||||
Optional,
|
||||
Sequence,
|
||||
Tuple,
|
||||
Union,
|
||||
)
|
||||
from typing import Dict, List, Optional, Sequence, Tuple, Union
|
||||
|
||||
import toml
|
||||
import voluptuous
|
||||
@@ -28,7 +22,7 @@ from voluptuous import (
|
||||
Required,
|
||||
Schema,
|
||||
)
|
||||
from transformers import CLIPTokenizer
|
||||
|
||||
|
||||
from . import train_util
|
||||
from .train_util import (
|
||||
@@ -78,6 +72,7 @@ class BaseSubsetParams:
|
||||
caption_tag_dropout_rate: float = 0.0
|
||||
token_warmup_min: int = 1
|
||||
token_warmup_step: float = 0
|
||||
custom_attributes: Optional[Dict[str, Any]] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -197,6 +192,7 @@ class ConfigSanitizer:
|
||||
"token_warmup_step": Any(float, int),
|
||||
"caption_prefix": str,
|
||||
"caption_suffix": str,
|
||||
"custom_attributes": dict,
|
||||
}
|
||||
# DO means DropOut
|
||||
DO_SUBSET_ASCENDABLE_SCHEMA = {
|
||||
@@ -530,7 +526,7 @@ def generate_dataset_group_by_blueprint(dataset_group_blueprint: DatasetGroupBlu
|
||||
secondary_separator: {subset.secondary_separator}
|
||||
enable_wildcard: {subset.enable_wildcard}
|
||||
caption_dropout_rate: {subset.caption_dropout_rate}
|
||||
caption_dropout_every_n_epoches: {subset.caption_dropout_every_n_epochs}
|
||||
caption_dropout_every_n_epochs: {subset.caption_dropout_every_n_epochs}
|
||||
caption_tag_dropout_rate: {subset.caption_tag_dropout_rate}
|
||||
caption_prefix: {subset.caption_prefix}
|
||||
caption_suffix: {subset.caption_suffix}
|
||||
@@ -538,9 +534,10 @@ def generate_dataset_group_by_blueprint(dataset_group_blueprint: DatasetGroupBlu
|
||||
flip_aug: {subset.flip_aug}
|
||||
face_crop_aug_range: {subset.face_crop_aug_range}
|
||||
random_crop: {subset.random_crop}
|
||||
token_warmup_min: {subset.token_warmup_min},
|
||||
token_warmup_step: {subset.token_warmup_step},
|
||||
alpha_mask: {subset.alpha_mask},
|
||||
token_warmup_min: {subset.token_warmup_min}
|
||||
token_warmup_step: {subset.token_warmup_step}
|
||||
alpha_mask: {subset.alpha_mask}
|
||||
custom_attributes: {subset.custom_attributes}
|
||||
"""
|
||||
),
|
||||
" ",
|
||||
|
||||
@@ -0,0 +1,227 @@
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
import time
|
||||
from typing import Optional
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from .device_utils import clean_memory_on_device
|
||||
|
||||
|
||||
def synchronize_device(device: torch.device):
|
||||
if device.type == "cuda":
|
||||
torch.cuda.synchronize()
|
||||
elif device.type == "xpu":
|
||||
torch.xpu.synchronize()
|
||||
elif device.type == "mps":
|
||||
torch.mps.synchronize()
|
||||
|
||||
|
||||
def swap_weight_devices_cuda(device: torch.device, layer_to_cpu: nn.Module, layer_to_cuda: nn.Module):
|
||||
assert layer_to_cpu.__class__ == layer_to_cuda.__class__
|
||||
|
||||
weight_swap_jobs = []
|
||||
|
||||
# This is not working for all cases (e.g. SD3), so we need to find the corresponding modules
|
||||
# for module_to_cpu, module_to_cuda in zip(layer_to_cpu.modules(), layer_to_cuda.modules()):
|
||||
# print(module_to_cpu.__class__, module_to_cuda.__class__)
|
||||
# if hasattr(module_to_cpu, "weight") and module_to_cpu.weight is not None:
|
||||
# weight_swap_jobs.append((module_to_cpu, module_to_cuda, module_to_cpu.weight.data, module_to_cuda.weight.data))
|
||||
|
||||
modules_to_cpu = {k: v for k, v in layer_to_cpu.named_modules()}
|
||||
for module_to_cuda_name, module_to_cuda in layer_to_cuda.named_modules():
|
||||
if hasattr(module_to_cuda, "weight") and module_to_cuda.weight is not None:
|
||||
module_to_cpu = modules_to_cpu.get(module_to_cuda_name, None)
|
||||
if module_to_cpu is not None and module_to_cpu.weight.shape == module_to_cuda.weight.shape:
|
||||
weight_swap_jobs.append((module_to_cpu, module_to_cuda, module_to_cpu.weight.data, module_to_cuda.weight.data))
|
||||
else:
|
||||
if module_to_cuda.weight.data.device.type != device.type:
|
||||
# print(
|
||||
# f"Module {module_to_cuda_name} not found in CPU model or shape mismatch, so not swapping and moving to device"
|
||||
# )
|
||||
module_to_cuda.weight.data = module_to_cuda.weight.data.to(device)
|
||||
|
||||
torch.cuda.current_stream().synchronize() # this prevents the illegal loss value
|
||||
|
||||
stream = torch.cuda.Stream()
|
||||
with torch.cuda.stream(stream):
|
||||
# cuda to cpu
|
||||
for module_to_cpu, module_to_cuda, cuda_data_view, cpu_data_view in weight_swap_jobs:
|
||||
cuda_data_view.record_stream(stream)
|
||||
module_to_cpu.weight.data = cuda_data_view.data.to("cpu", non_blocking=True)
|
||||
|
||||
stream.synchronize()
|
||||
|
||||
# cpu to cuda
|
||||
for module_to_cpu, module_to_cuda, cuda_data_view, cpu_data_view in weight_swap_jobs:
|
||||
cuda_data_view.copy_(module_to_cuda.weight.data, non_blocking=True)
|
||||
module_to_cuda.weight.data = cuda_data_view
|
||||
|
||||
stream.synchronize()
|
||||
torch.cuda.current_stream().synchronize() # this prevents the illegal loss value
|
||||
|
||||
|
||||
def swap_weight_devices_no_cuda(device: torch.device, layer_to_cpu: nn.Module, layer_to_cuda: nn.Module):
|
||||
"""
|
||||
not tested
|
||||
"""
|
||||
assert layer_to_cpu.__class__ == layer_to_cuda.__class__
|
||||
|
||||
weight_swap_jobs = []
|
||||
for module_to_cpu, module_to_cuda in zip(layer_to_cpu.modules(), layer_to_cuda.modules()):
|
||||
if hasattr(module_to_cpu, "weight") and module_to_cpu.weight is not None:
|
||||
weight_swap_jobs.append((module_to_cpu, module_to_cuda, module_to_cpu.weight.data, module_to_cuda.weight.data))
|
||||
|
||||
# device to cpu
|
||||
for module_to_cpu, module_to_cuda, cuda_data_view, cpu_data_view in weight_swap_jobs:
|
||||
module_to_cpu.weight.data = cuda_data_view.data.to("cpu", non_blocking=True)
|
||||
|
||||
synchronize_device()
|
||||
|
||||
# cpu to device
|
||||
for module_to_cpu, module_to_cuda, cuda_data_view, cpu_data_view in weight_swap_jobs:
|
||||
cuda_data_view.copy_(module_to_cuda.weight.data, non_blocking=True)
|
||||
module_to_cuda.weight.data = cuda_data_view
|
||||
|
||||
synchronize_device()
|
||||
|
||||
|
||||
def weighs_to_device(layer: nn.Module, device: torch.device):
|
||||
for module in layer.modules():
|
||||
if hasattr(module, "weight") and module.weight is not None:
|
||||
module.weight.data = module.weight.data.to(device, non_blocking=True)
|
||||
|
||||
|
||||
class Offloader:
|
||||
"""
|
||||
common offloading class
|
||||
"""
|
||||
|
||||
def __init__(self, num_blocks: int, blocks_to_swap: int, device: torch.device, debug: bool = False):
|
||||
self.num_blocks = num_blocks
|
||||
self.blocks_to_swap = blocks_to_swap
|
||||
self.device = device
|
||||
self.debug = debug
|
||||
|
||||
self.thread_pool = ThreadPoolExecutor(max_workers=1)
|
||||
self.futures = {}
|
||||
self.cuda_available = device.type == "cuda"
|
||||
|
||||
def swap_weight_devices(self, block_to_cpu: nn.Module, block_to_cuda: nn.Module):
|
||||
if self.cuda_available:
|
||||
swap_weight_devices_cuda(self.device, block_to_cpu, block_to_cuda)
|
||||
else:
|
||||
swap_weight_devices_no_cuda(self.device, block_to_cpu, block_to_cuda)
|
||||
|
||||
def _submit_move_blocks(self, blocks, block_idx_to_cpu, block_idx_to_cuda):
|
||||
def move_blocks(bidx_to_cpu, block_to_cpu, bidx_to_cuda, block_to_cuda):
|
||||
if self.debug:
|
||||
start_time = time.perf_counter()
|
||||
print(f"Move block {bidx_to_cpu} to CPU and block {bidx_to_cuda} to {'CUDA' if self.cuda_available else 'device'}")
|
||||
|
||||
self.swap_weight_devices(block_to_cpu, block_to_cuda)
|
||||
|
||||
if self.debug:
|
||||
print(f"Moved blocks {bidx_to_cpu} and {bidx_to_cuda} in {time.perf_counter()-start_time:.2f}s")
|
||||
return bidx_to_cpu, bidx_to_cuda # , event
|
||||
|
||||
block_to_cpu = blocks[block_idx_to_cpu]
|
||||
block_to_cuda = blocks[block_idx_to_cuda]
|
||||
|
||||
self.futures[block_idx_to_cuda] = self.thread_pool.submit(
|
||||
move_blocks, block_idx_to_cpu, block_to_cpu, block_idx_to_cuda, block_to_cuda
|
||||
)
|
||||
|
||||
def _wait_blocks_move(self, block_idx):
|
||||
if block_idx not in self.futures:
|
||||
return
|
||||
|
||||
if self.debug:
|
||||
print(f"Wait for block {block_idx}")
|
||||
start_time = time.perf_counter()
|
||||
|
||||
future = self.futures.pop(block_idx)
|
||||
_, bidx_to_cuda = future.result()
|
||||
|
||||
assert block_idx == bidx_to_cuda, f"Block index mismatch: {block_idx} != {bidx_to_cuda}"
|
||||
|
||||
if self.debug:
|
||||
print(f"Waited for block {block_idx}: {time.perf_counter()-start_time:.2f}s")
|
||||
|
||||
|
||||
class ModelOffloader(Offloader):
|
||||
"""
|
||||
supports forward offloading
|
||||
"""
|
||||
|
||||
def __init__(self, blocks: list[nn.Module], num_blocks: int, blocks_to_swap: int, device: torch.device, debug: bool = False):
|
||||
super().__init__(num_blocks, blocks_to_swap, device, debug)
|
||||
|
||||
# register backward hooks
|
||||
self.remove_handles = []
|
||||
for i, block in enumerate(blocks):
|
||||
hook = self.create_backward_hook(blocks, i)
|
||||
if hook is not None:
|
||||
handle = block.register_full_backward_hook(hook)
|
||||
self.remove_handles.append(handle)
|
||||
|
||||
def __del__(self):
|
||||
for handle in self.remove_handles:
|
||||
handle.remove()
|
||||
|
||||
def create_backward_hook(self, blocks: list[nn.Module], block_index: int) -> Optional[callable]:
|
||||
# -1 for 0-based index
|
||||
num_blocks_propagated = self.num_blocks - block_index - 1
|
||||
swapping = num_blocks_propagated > 0 and num_blocks_propagated <= self.blocks_to_swap
|
||||
waiting = block_index > 0 and block_index <= self.blocks_to_swap
|
||||
|
||||
if not swapping and not waiting:
|
||||
return None
|
||||
|
||||
# create hook
|
||||
block_idx_to_cpu = self.num_blocks - num_blocks_propagated
|
||||
block_idx_to_cuda = self.blocks_to_swap - num_blocks_propagated
|
||||
block_idx_to_wait = block_index - 1
|
||||
|
||||
def backward_hook(module, grad_input, grad_output):
|
||||
if self.debug:
|
||||
print(f"Backward hook for block {block_index}")
|
||||
|
||||
if swapping:
|
||||
self._submit_move_blocks(blocks, block_idx_to_cpu, block_idx_to_cuda)
|
||||
if waiting:
|
||||
self._wait_blocks_move(block_idx_to_wait)
|
||||
return None
|
||||
|
||||
return backward_hook
|
||||
|
||||
def prepare_block_devices_before_forward(self, blocks: list[nn.Module]):
|
||||
if self.blocks_to_swap is None or self.blocks_to_swap == 0:
|
||||
return
|
||||
|
||||
if self.debug:
|
||||
print("Prepare block devices before forward")
|
||||
|
||||
for b in blocks[0 : self.num_blocks - self.blocks_to_swap]:
|
||||
b.to(self.device)
|
||||
weighs_to_device(b, self.device) # make sure weights are on device
|
||||
|
||||
for b in blocks[self.num_blocks - self.blocks_to_swap :]:
|
||||
b.to(self.device) # move block to device first
|
||||
weighs_to_device(b, "cpu") # make sure weights are on cpu
|
||||
|
||||
synchronize_device(self.device)
|
||||
clean_memory_on_device(self.device)
|
||||
|
||||
def wait_for_block(self, block_idx: int):
|
||||
if self.blocks_to_swap is None or self.blocks_to_swap == 0:
|
||||
return
|
||||
self._wait_blocks_move(block_idx)
|
||||
|
||||
def submit_move_blocks(self, blocks: list[nn.Module], block_idx: int):
|
||||
if self.blocks_to_swap is None or self.blocks_to_swap == 0:
|
||||
return
|
||||
if block_idx >= self.blocks_to_swap:
|
||||
return
|
||||
block_idx_to_cpu = block_idx
|
||||
block_idx_to_cuda = self.num_blocks - self.blocks_to_swap + block_idx
|
||||
self._submit_move_blocks(blocks, block_idx_to_cpu, block_idx_to_cuda)
|
||||
+47
-246
@@ -1,15 +1,15 @@
|
||||
# copy from FLUX repo: https://github.com/black-forest-labs/flux
|
||||
# license: Apache-2.0 License
|
||||
|
||||
|
||||
from dataclasses import dataclass
|
||||
import math
|
||||
from typing import Optional
|
||||
import torch
|
||||
from typing import Dict, List, Optional, Union
|
||||
|
||||
from .device_utils import init_ipex, clean_memory_on_device
|
||||
from .device_utils import init_ipex
|
||||
from .custom_offloading_utils import ModelOffloader
|
||||
init_ipex()
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
from torch import Tensor, nn
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
@@ -916,8 +916,12 @@ class Flux(nn.Module):
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
self.cpu_offload_checkpointing = False
|
||||
self.double_blocks_to_swap = None
|
||||
self.single_blocks_to_swap = None
|
||||
self.blocks_to_swap = None
|
||||
|
||||
self.offloader_double = None
|
||||
self.offloader_single = None
|
||||
self.num_double_blocks = len(self.double_blocks)
|
||||
self.num_single_blocks = len(self.single_blocks)
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
@@ -955,39 +959,45 @@ class Flux(nn.Module):
|
||||
|
||||
print("FLUX: Gradient checkpointing disabled.")
|
||||
|
||||
def enable_block_swap(self, double_blocks: Optional[int], single_blocks: Optional[int]):
|
||||
self.double_blocks_to_swap = double_blocks
|
||||
self.single_blocks_to_swap = single_blocks
|
||||
def enable_block_swap(self, num_blocks: int, device: torch.device):
|
||||
self.blocks_to_swap = num_blocks
|
||||
double_blocks_to_swap = num_blocks // 2
|
||||
single_blocks_to_swap = (num_blocks - double_blocks_to_swap) * 2
|
||||
|
||||
assert double_blocks_to_swap <= self.num_double_blocks - 2 and single_blocks_to_swap <= self.num_single_blocks - 2, (
|
||||
f"Cannot swap more than {self.num_double_blocks - 2} double blocks and {self.num_single_blocks - 2} single blocks. "
|
||||
f"Requested {double_blocks_to_swap} double blocks and {single_blocks_to_swap} single blocks."
|
||||
)
|
||||
|
||||
self.offloader_double = ModelOffloader(
|
||||
self.double_blocks, self.num_double_blocks, double_blocks_to_swap, device # , debug=True
|
||||
)
|
||||
self.offloader_single = ModelOffloader(
|
||||
self.single_blocks, self.num_single_blocks, single_blocks_to_swap, device # , debug=True
|
||||
)
|
||||
print(
|
||||
f"FLUX: Block swap enabled. Swapping {num_blocks} blocks, double blocks: {double_blocks_to_swap}, single blocks: {single_blocks_to_swap}."
|
||||
)
|
||||
|
||||
def move_to_device_except_swap_blocks(self, device: torch.device):
|
||||
# assume model is on cpu
|
||||
if self.double_blocks_to_swap:
|
||||
# assume model is on cpu. do not move blocks to device to reduce temporary memory usage
|
||||
if self.blocks_to_swap:
|
||||
save_double_blocks = self.double_blocks
|
||||
self.double_blocks = None
|
||||
if self.single_blocks_to_swap:
|
||||
save_single_blocks = self.single_blocks
|
||||
self.double_blocks = None
|
||||
self.single_blocks = None
|
||||
|
||||
self.to(device)
|
||||
|
||||
if self.double_blocks_to_swap:
|
||||
if self.blocks_to_swap:
|
||||
self.double_blocks = save_double_blocks
|
||||
if self.single_blocks_to_swap:
|
||||
self.single_blocks = save_single_blocks
|
||||
|
||||
def prepare_block_swap_before_forward(self):
|
||||
# move last n blocks to cpu: they are on cuda
|
||||
if self.double_blocks_to_swap:
|
||||
for i in range(len(self.double_blocks) - self.double_blocks_to_swap):
|
||||
self.double_blocks[i].to(self.device)
|
||||
for i in range(len(self.double_blocks) - self.double_blocks_to_swap, len(self.double_blocks)):
|
||||
self.double_blocks[i].to("cpu") # , non_blocking=True)
|
||||
if self.single_blocks_to_swap:
|
||||
for i in range(len(self.single_blocks) - self.single_blocks_to_swap):
|
||||
self.single_blocks[i].to(self.device)
|
||||
for i in range(len(self.single_blocks) - self.single_blocks_to_swap, len(self.single_blocks)):
|
||||
self.single_blocks[i].to("cpu") # , non_blocking=True)
|
||||
clean_memory_on_device(self.device)
|
||||
if self.blocks_to_swap is None or self.blocks_to_swap == 0:
|
||||
return
|
||||
self.offloader_double.prepare_block_devices_before_forward(self.double_blocks)
|
||||
self.offloader_single.prepare_block_devices_before_forward(self.single_blocks)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -1016,69 +1026,28 @@ class Flux(nn.Module):
|
||||
ids = torch.cat((txt_ids, img_ids), dim=1)
|
||||
pe = self.pe_embedder(ids)
|
||||
|
||||
if not self.double_blocks_to_swap:
|
||||
if not self.blocks_to_swap:
|
||||
for block in self.double_blocks:
|
||||
img, txt = block(img=img, txt=txt, vec=vec, pe=pe, txt_attention_mask=txt_attention_mask)
|
||||
else:
|
||||
# make sure first n blocks are on cuda, and last n blocks are on cpu at beginning
|
||||
for block_idx in range(self.double_blocks_to_swap):
|
||||
block = self.double_blocks[len(self.double_blocks) - self.double_blocks_to_swap + block_idx]
|
||||
if block.parameters().__next__().device.type != "cpu":
|
||||
block.to("cpu") # , non_blocking=True)
|
||||
# print(f"Moved double block {len(self.double_blocks) - self.double_blocks_to_swap + block_idx} to cpu.")
|
||||
|
||||
block = self.double_blocks[block_idx]
|
||||
if block.parameters().__next__().device.type == "cpu":
|
||||
block.to(self.device)
|
||||
# print(f"Moved double block {block_idx} to cuda.")
|
||||
|
||||
to_cpu_block_index = 0
|
||||
for block_idx, block in enumerate(self.double_blocks):
|
||||
# move last n blocks to cuda: they are on cpu, and move first n blocks to cpu: they are on cuda
|
||||
moving = block_idx >= len(self.double_blocks) - self.double_blocks_to_swap
|
||||
if moving:
|
||||
block.to(self.device) # move to cuda
|
||||
# print(f"Moved double block {block_idx} to cuda.")
|
||||
|
||||
img, txt = block(img=img, txt=txt, vec=vec, pe=pe, txt_attention_mask=txt_attention_mask)
|
||||
|
||||
if moving:
|
||||
self.double_blocks[to_cpu_block_index].to("cpu") # , non_blocking=True)
|
||||
# print(f"Moved double block {to_cpu_block_index} to cpu.")
|
||||
to_cpu_block_index += 1
|
||||
|
||||
img = torch.cat((txt, img), 1)
|
||||
|
||||
if not self.single_blocks_to_swap:
|
||||
for block in self.single_blocks:
|
||||
img = block(img, vec=vec, pe=pe, txt_attention_mask=txt_attention_mask)
|
||||
else:
|
||||
# make sure first n blocks are on cuda, and last n blocks are on cpu at beginning
|
||||
for block_idx in range(self.single_blocks_to_swap):
|
||||
block = self.single_blocks[len(self.single_blocks) - self.single_blocks_to_swap + block_idx]
|
||||
if block.parameters().__next__().device.type != "cpu":
|
||||
block.to("cpu") # , non_blocking=True)
|
||||
# print(f"Moved single block {len(self.single_blocks) - self.single_blocks_to_swap + block_idx} to cpu.")
|
||||
for block_idx, block in enumerate(self.double_blocks):
|
||||
self.offloader_double.wait_for_block(block_idx)
|
||||
|
||||
block = self.single_blocks[block_idx]
|
||||
if block.parameters().__next__().device.type == "cpu":
|
||||
block.to(self.device)
|
||||
# print(f"Moved single block {block_idx} to cuda.")
|
||||
img, txt = block(img=img, txt=txt, vec=vec, pe=pe, txt_attention_mask=txt_attention_mask)
|
||||
|
||||
self.offloader_double.submit_move_blocks(self.double_blocks, block_idx)
|
||||
|
||||
img = torch.cat((txt, img), 1)
|
||||
|
||||
to_cpu_block_index = 0
|
||||
for block_idx, block in enumerate(self.single_blocks):
|
||||
# move last n blocks to cuda: they are on cpu, and move first n blocks to cpu: they are on cuda
|
||||
moving = block_idx >= len(self.single_blocks) - self.single_blocks_to_swap
|
||||
if moving:
|
||||
block.to(self.device) # move to cuda
|
||||
# print(f"Moved single block {block_idx} to cuda.")
|
||||
self.offloader_single.wait_for_block(block_idx)
|
||||
|
||||
img = block(img, vec=vec, pe=pe, txt_attention_mask=txt_attention_mask)
|
||||
|
||||
if moving:
|
||||
self.single_blocks[to_cpu_block_index].to("cpu") # , non_blocking=True)
|
||||
# print(f"Moved single block {to_cpu_block_index} to cpu.")
|
||||
to_cpu_block_index += 1
|
||||
self.offloader_single.submit_move_blocks(self.single_blocks, block_idx)
|
||||
|
||||
img = img[:, txt.shape[1] :, ...]
|
||||
|
||||
@@ -1087,173 +1056,5 @@ class Flux(nn.Module):
|
||||
vec = vec.to(self.device)
|
||||
|
||||
img = self.final_layer(img, vec) # (N, T, patch_size ** 2 * out_channels)
|
||||
return img
|
||||
|
||||
|
||||
class FluxUpper(nn.Module):
|
||||
"""
|
||||
Transformer model for flow matching on sequences.
|
||||
"""
|
||||
|
||||
def __init__(self, params: FluxParams):
|
||||
super().__init__()
|
||||
|
||||
self.params = params
|
||||
self.in_channels = params.in_channels
|
||||
self.out_channels = self.in_channels
|
||||
if params.hidden_size % params.num_heads != 0:
|
||||
raise ValueError(f"Hidden size {params.hidden_size} must be divisible by num_heads {params.num_heads}")
|
||||
pe_dim = params.hidden_size // params.num_heads
|
||||
if sum(params.axes_dim) != pe_dim:
|
||||
raise ValueError(f"Got {params.axes_dim} but expected positional dim {pe_dim}")
|
||||
self.hidden_size = params.hidden_size
|
||||
self.num_heads = params.num_heads
|
||||
self.pe_embedder = EmbedND(dim=pe_dim, theta=params.theta, axes_dim=params.axes_dim)
|
||||
self.img_in = nn.Linear(self.in_channels, self.hidden_size, bias=True)
|
||||
self.time_in = MLPEmbedder(in_dim=256, hidden_dim=self.hidden_size)
|
||||
self.vector_in = MLPEmbedder(params.vec_in_dim, self.hidden_size)
|
||||
self.guidance_in = MLPEmbedder(in_dim=256, hidden_dim=self.hidden_size) if params.guidance_embed else nn.Identity()
|
||||
self.txt_in = nn.Linear(params.context_in_dim, self.hidden_size)
|
||||
|
||||
self.double_blocks = nn.ModuleList(
|
||||
[
|
||||
DoubleStreamBlock(
|
||||
self.hidden_size,
|
||||
self.num_heads,
|
||||
mlp_ratio=params.mlp_ratio,
|
||||
qkv_bias=params.qkv_bias,
|
||||
)
|
||||
for _ in range(params.depth)
|
||||
]
|
||||
)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
return next(self.parameters()).device
|
||||
|
||||
@property
|
||||
def dtype(self):
|
||||
return next(self.parameters()).dtype
|
||||
|
||||
def enable_gradient_checkpointing(self):
|
||||
self.gradient_checkpointing = True
|
||||
|
||||
self.time_in.enable_gradient_checkpointing()
|
||||
self.vector_in.enable_gradient_checkpointing()
|
||||
if self.guidance_in.__class__ != nn.Identity:
|
||||
self.guidance_in.enable_gradient_checkpointing()
|
||||
|
||||
for block in self.double_blocks:
|
||||
block.enable_gradient_checkpointing()
|
||||
|
||||
print("FLUX: Gradient checkpointing enabled.")
|
||||
|
||||
def disable_gradient_checkpointing(self):
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
self.time_in.disable_gradient_checkpointing()
|
||||
self.vector_in.disable_gradient_checkpointing()
|
||||
if self.guidance_in.__class__ != nn.Identity:
|
||||
self.guidance_in.disable_gradient_checkpointing()
|
||||
|
||||
for block in self.double_blocks:
|
||||
block.disable_gradient_checkpointing()
|
||||
|
||||
print("FLUX: Gradient checkpointing disabled.")
|
||||
|
||||
def forward(
|
||||
self,
|
||||
img: Tensor,
|
||||
img_ids: Tensor,
|
||||
txt: Tensor,
|
||||
txt_ids: Tensor,
|
||||
timesteps: Tensor,
|
||||
y: Tensor,
|
||||
guidance: Tensor | None = None,
|
||||
txt_attention_mask: Tensor | None = None,
|
||||
) -> Tensor:
|
||||
if img.ndim != 3 or txt.ndim != 3:
|
||||
raise ValueError("Input img and txt tensors must have 3 dimensions.")
|
||||
|
||||
# running on sequences img
|
||||
img = self.img_in(img)
|
||||
vec = self.time_in(timestep_embedding(timesteps, 256))
|
||||
if self.params.guidance_embed:
|
||||
if guidance is None:
|
||||
raise ValueError("Didn't get guidance strength for guidance distilled model.")
|
||||
vec = vec + self.guidance_in(timestep_embedding(guidance, 256))
|
||||
vec = vec + self.vector_in(y)
|
||||
txt = self.txt_in(txt)
|
||||
|
||||
ids = torch.cat((txt_ids, img_ids), dim=1)
|
||||
pe = self.pe_embedder(ids)
|
||||
|
||||
for block in self.double_blocks:
|
||||
img, txt = block(img=img, txt=txt, vec=vec, pe=pe, txt_attention_mask=txt_attention_mask)
|
||||
|
||||
return img, txt, vec, pe
|
||||
|
||||
|
||||
class FluxLower(nn.Module):
|
||||
"""
|
||||
Transformer model for flow matching on sequences.
|
||||
"""
|
||||
|
||||
def __init__(self, params: FluxParams):
|
||||
super().__init__()
|
||||
self.hidden_size = params.hidden_size
|
||||
self.num_heads = params.num_heads
|
||||
self.out_channels = params.in_channels
|
||||
|
||||
self.single_blocks = nn.ModuleList(
|
||||
[
|
||||
SingleStreamBlock(self.hidden_size, self.num_heads, mlp_ratio=params.mlp_ratio)
|
||||
for _ in range(params.depth_single_blocks)
|
||||
]
|
||||
)
|
||||
|
||||
self.final_layer = LastLayer(self.hidden_size, 1, self.out_channels)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
return next(self.parameters()).device
|
||||
|
||||
@property
|
||||
def dtype(self):
|
||||
return next(self.parameters()).dtype
|
||||
|
||||
def enable_gradient_checkpointing(self):
|
||||
self.gradient_checkpointing = True
|
||||
|
||||
for block in self.single_blocks:
|
||||
block.enable_gradient_checkpointing()
|
||||
|
||||
print("FLUX: Gradient checkpointing enabled.")
|
||||
|
||||
def disable_gradient_checkpointing(self):
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
for block in self.single_blocks:
|
||||
block.disable_gradient_checkpointing()
|
||||
|
||||
print("FLUX: Gradient checkpointing disabled.")
|
||||
|
||||
def forward(
|
||||
self,
|
||||
img: Tensor,
|
||||
txt: Tensor,
|
||||
vec: Tensor | None = None,
|
||||
pe: Tensor | None = None,
|
||||
txt_attention_mask: Tensor | None = None,
|
||||
) -> Tensor:
|
||||
img = torch.cat((txt, img), 1)
|
||||
for block in self.single_blocks:
|
||||
img = block(img, vec=vec, pe=pe, txt_attention_mask=txt_attention_mask)
|
||||
img = img[:, txt.shape[1] :, ...]
|
||||
|
||||
img = self.final_layer(img, vec) # (N, T, patch_size ** 2 * out_channels)
|
||||
return img
|
||||
+34
-47
@@ -15,7 +15,6 @@ from PIL import Image
|
||||
|
||||
from safetensors.torch import save_file
|
||||
from . import flux_models, flux_utils, strategy_base, train_util
|
||||
from .sd3_train_utils import load_prompts
|
||||
from .device_utils import init_ipex, clean_memory_on_device
|
||||
|
||||
init_ipex()
|
||||
@@ -83,7 +82,7 @@ def sample_images(
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
with torch.no_grad():
|
||||
with torch.no_grad(), accelerator.autocast():
|
||||
image_tensor_list = []
|
||||
for prompt_dict in prompts:
|
||||
image_tensor = sample_image_inference(
|
||||
@@ -180,13 +179,26 @@ def sample_image_inference(
|
||||
tokenize_strategy = strategy_base.TokenizeStrategy.get_strategy()
|
||||
encoding_strategy = strategy_base.TextEncodingStrategy.get_strategy()
|
||||
|
||||
text_encoder_conds = []
|
||||
if sample_prompts_te_outputs and prompt in sample_prompts_te_outputs:
|
||||
te_outputs = sample_prompts_te_outputs[prompt]
|
||||
else:
|
||||
text_encoder_conds = sample_prompts_te_outputs[prompt]
|
||||
print(f"Using cached text encoder outputs for prompt: {prompt}")
|
||||
if text_encoders is not None:
|
||||
print(f"Encoding prompt: {prompt}")
|
||||
tokens_and_masks = tokenize_strategy.tokenize(prompt)
|
||||
te_outputs = encoding_strategy.encode_tokens(tokenize_strategy, text_encoders, tokens_and_masks)
|
||||
# strategy has apply_t5_attn_mask option
|
||||
encoded_text_encoder_conds = encoding_strategy.encode_tokens(tokenize_strategy, text_encoders, tokens_and_masks)
|
||||
|
||||
l_pooled, t5_out, txt_ids, t5_attn_mask = te_outputs
|
||||
# if text_encoder_conds is not cached, use encoded_text_encoder_conds
|
||||
if len(text_encoder_conds) == 0:
|
||||
text_encoder_conds = encoded_text_encoder_conds
|
||||
else:
|
||||
# if encoded_text_encoder_conds is not None, update cached text_encoder_conds
|
||||
for i in range(len(encoded_text_encoder_conds)):
|
||||
if encoded_text_encoder_conds[i] is not None:
|
||||
text_encoder_conds[i] = encoded_text_encoder_conds[i]
|
||||
|
||||
l_pooled, t5_out, txt_ids, t5_attn_mask = text_encoder_conds
|
||||
|
||||
# sample image
|
||||
weight_dtype = ae.dtype # TOFO give dtype as argument
|
||||
@@ -293,6 +305,7 @@ def denoise(
|
||||
comfy_pbar = ProgressBar(total=len(timesteps))
|
||||
for t_curr, t_prev in zip(tqdm(timesteps[:-1]), timesteps[1:]):
|
||||
t_vec = torch.full((img.shape[0],), t_curr, dtype=img.dtype, device=img.device)
|
||||
model.prepare_block_swap_before_forward()
|
||||
pred = model(
|
||||
img=img,
|
||||
img_ids=img_ids,
|
||||
@@ -306,7 +319,7 @@ def denoise(
|
||||
|
||||
img = img + (t_prev - t_curr) * pred
|
||||
comfy_pbar.update(1)
|
||||
|
||||
model.prepare_block_swap_before_forward()
|
||||
return img
|
||||
|
||||
# endregion
|
||||
@@ -329,7 +342,9 @@ def compute_density_for_timestep_sampling(
|
||||
weighting_scheme: str, batch_size: int, logit_mean: float = None, logit_std: float = None, mode_scale: float = None
|
||||
):
|
||||
"""Compute the density for sampling the timesteps when doing SD3 training.
|
||||
|
||||
Courtesy: This was contributed by Rafie Walker in https://github.com/huggingface/diffusers/pull/8528.
|
||||
|
||||
SD3 paper reference: https://arxiv.org/abs/2403.03206v1.
|
||||
"""
|
||||
if weighting_scheme == "logit_normal":
|
||||
@@ -346,7 +361,9 @@ def compute_density_for_timestep_sampling(
|
||||
|
||||
def compute_loss_weighting_for_sd3(weighting_scheme: str, sigmas=None):
|
||||
"""Computes loss weighting scheme for SD3 training.
|
||||
|
||||
Courtesy: This was contributed by Rafie Walker in https://github.com/huggingface/diffusers/pull/8528.
|
||||
|
||||
SD3 paper reference: https://arxiv.org/abs/2403.03206v1.
|
||||
"""
|
||||
if weighting_scheme == "sigma_sqrt":
|
||||
@@ -372,6 +389,7 @@ def get_noisy_model_input_and_timesteps(
|
||||
t = torch.sigmoid(args.sigmoid_scale * torch.randn((bsz,), device=device))
|
||||
else:
|
||||
t = torch.rand((bsz,), device=device)
|
||||
|
||||
timesteps = t * 1000.0
|
||||
t = t.view(-1, 1, 1, 1)
|
||||
noisy_model_input = (1 - t) * latents + t * noise
|
||||
@@ -522,46 +540,9 @@ def add_flux_train_arguments(parser: argparse.ArgumentParser):
|
||||
parser.add_argument(
|
||||
"--apply_t5_attn_mask",
|
||||
action="store_true",
|
||||
help="apply attention mask (zero embs) to T5-XXL / T5-XXLにアテンションマスク(ゼロ埋め)を適用する",
|
||||
)
|
||||
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(
|
||||
"--text_encoder_batch_size",
|
||||
type=int,
|
||||
default=None,
|
||||
help="text encoder batch size (default: None, use dataset's batch size)"
|
||||
+ " / text encoderのバッチサイズ(デフォルト: None, データセットのバッチサイズを使用)",
|
||||
)
|
||||
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環境等でモデル読み込みを高速化できる",
|
||||
help="apply attention mask to T5-XXL encode and FLUX double blocks / T5-XXLエンコードとFLUXダブルブロックにアテンションマスクを適用する",
|
||||
)
|
||||
|
||||
# copy from Diffusers
|
||||
parser.add_argument(
|
||||
"--weighting_scheme",
|
||||
type=str,
|
||||
default="none",
|
||||
choices=["sigma_sqrt", "logit_normal", "mode", "cosmap", "none"],
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logit_mean", type=float, default=0.0, help="mean to use when using the `'logit_normal'` weighting scheme."
|
||||
)
|
||||
parser.add_argument("--logit_std", type=float, default=1.0, help="std to use when using the `'logit_normal'` weighting scheme.")
|
||||
parser.add_argument(
|
||||
"--mode_scale",
|
||||
type=float,
|
||||
default=1.29,
|
||||
help="Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--guidance_scale",
|
||||
type=float,
|
||||
@@ -571,9 +552,10 @@ def add_flux_train_arguments(parser: argparse.ArgumentParser):
|
||||
|
||||
parser.add_argument(
|
||||
"--timestep_sampling",
|
||||
choices=["sigma", "uniform", "sigmoid"],
|
||||
choices=["sigma", "uniform", "sigmoid", "shift", "flux_shift"],
|
||||
default="sigma",
|
||||
help="Method to sample timesteps: sigma-based, uniform random, or sigmoid of random normal. / タイムステップをサンプリングする方法:sigma、random uniform、またはrandom normalのsigmoid。",
|
||||
help="Method to sample timesteps: sigma-based, uniform random, sigmoid of random normal, shift of sigmoid and FLUX.1 shifting."
|
||||
" / タイムステップをサンプリングする方法:sigma、random uniform、random normalのsigmoid、sigmoidのシフト、FLUX.1のシフト。",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--sigmoid_scale",
|
||||
@@ -596,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"
|
||||
)
|
||||
|
||||
+277
-17
@@ -1,14 +1,16 @@
|
||||
from dataclasses import replace
|
||||
import json
|
||||
from typing import Union
|
||||
import os
|
||||
from typing import List, Optional, Tuple, Union
|
||||
import einops
|
||||
import torch
|
||||
|
||||
from safetensors.torch import load_file
|
||||
from safetensors import safe_open
|
||||
from accelerate import init_empty_weights
|
||||
from transformers import CLIPTextModel, CLIPConfig, T5EncoderModel, T5Config
|
||||
|
||||
from .flux_models import Flux, AutoEncoder, configs
|
||||
from .utils import setup_logging
|
||||
from .utils import setup_logging, load_safetensors
|
||||
|
||||
setup_logging()
|
||||
import logging
|
||||
@@ -16,35 +18,151 @@ import logging
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
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か、ブロック数を計算して返す。
|
||||
|
||||
Args:
|
||||
ckpt_path (str): チェックポイントファイルまたはディレクトリのパス。
|
||||
|
||||
Returns:
|
||||
Tuple[bool, bool, Tuple[int, int], List[str]]:
|
||||
- bool: Diffusersかどうかを示すフラグ。
|
||||
- bool: Schnellかどうかを示すフラグ。
|
||||
- Tuple[int, int]: ダブルブロックとシングルブロックの数。
|
||||
- List[str]: チェックポイントに含まれるキーのリスト。
|
||||
"""
|
||||
# check the state dict: Diffusers or BFL, dev or schnell, number of blocks
|
||||
logger.info(f"Checking the state dict: Diffusers or BFL, dev or schnell")
|
||||
|
||||
if os.path.isdir(ckpt_path): # if ckpt_path is a directory, it is Diffusers
|
||||
ckpt_path = os.path.join(ckpt_path, "transformer", "diffusion_pytorch_model-00001-of-00003.safetensors")
|
||||
if "00001-of-00003" in ckpt_path:
|
||||
ckpt_paths = [ckpt_path.replace("00001-of-00003", f"0000{i}-of-00003") for i in range(1, 4)]
|
||||
else:
|
||||
ckpt_paths = [ckpt_path]
|
||||
|
||||
keys = []
|
||||
for ckpt_path in ckpt_paths:
|
||||
with safe_open(ckpt_path, framework="pt") as f:
|
||||
keys.extend(f.keys())
|
||||
|
||||
if keys[0].startswith("model.diffusion_model."):
|
||||
keys = [key.replace("model.diffusion_model.", "") for key in keys]
|
||||
|
||||
is_diffusers = "transformer_blocks.0.attn.add_k_proj.bias" in keys
|
||||
is_schnell = not ("guidance_in.in_layer.bias" in keys or "time_text_embed.guidance_embedder.linear_1.bias" in keys)
|
||||
|
||||
# check number of double and single blocks
|
||||
if not is_diffusers:
|
||||
max_double_block_index = max(
|
||||
[int(key.split(".")[1]) for key in keys if key.startswith("double_blocks.") and key.endswith(".img_attn.proj.bias")]
|
||||
)
|
||||
max_single_block_index = max(
|
||||
[int(key.split(".")[1]) for key in keys if key.startswith("single_blocks.") and key.endswith(".modulation.lin.bias")]
|
||||
)
|
||||
else:
|
||||
max_double_block_index = max(
|
||||
[
|
||||
int(key.split(".")[1])
|
||||
for key in keys
|
||||
if key.startswith("transformer_blocks.") and key.endswith(".attn.add_k_proj.bias")
|
||||
]
|
||||
)
|
||||
max_single_block_index = max(
|
||||
[
|
||||
int(key.split(".")[1])
|
||||
for key in keys
|
||||
if key.startswith("single_transformer_blocks.") and key.endswith(".attn.to_k.bias")
|
||||
]
|
||||
)
|
||||
|
||||
num_double_blocks = max_double_block_index + 1
|
||||
num_single_blocks = max_single_block_index + 1
|
||||
|
||||
return is_diffusers, is_schnell, (num_double_blocks, num_single_blocks), ckpt_paths
|
||||
|
||||
|
||||
def load_flow_model(name: str, ckpt_path: str, dtype: torch.dtype, device: Union[str, torch.device]) -> Flux:
|
||||
logger.info(f"Building Flux model {name}")
|
||||
def load_flow_model(
|
||||
ckpt_path: str, dtype: Optional[torch.dtype], device: Union[str, torch.device], disable_mmap: bool = False
|
||||
) -> Tuple[bool, Flux]:
|
||||
is_diffusers, is_schnell, (num_double_blocks, num_single_blocks), ckpt_paths = analyze_checkpoint_state(ckpt_path)
|
||||
name = MODEL_NAME_DEV if not is_schnell else MODEL_NAME_SCHNELL
|
||||
|
||||
# build model
|
||||
logger.info(f"Building Flux model {name} from {'Diffusers' if is_diffusers else 'BFL'} checkpoint")
|
||||
with torch.device("meta"):
|
||||
model = Flux(configs[name].params).to(dtype)
|
||||
params = configs[name].params
|
||||
|
||||
# set the number of blocks
|
||||
if params.depth != num_double_blocks:
|
||||
logger.info(f"Setting the number of double blocks from {params.depth} to {num_double_blocks}")
|
||||
params = replace(params, depth=num_double_blocks)
|
||||
if params.depth_single_blocks != num_single_blocks:
|
||||
logger.info(f"Setting the number of single blocks from {params.depth_single_blocks} to {num_single_blocks}")
|
||||
params = replace(params, depth_single_blocks=num_single_blocks)
|
||||
|
||||
model = Flux(params)
|
||||
if dtype is not None:
|
||||
model = model.to(dtype)
|
||||
|
||||
# load_sft doesn't support torch.device
|
||||
logger.info(f"Loading state dict from {ckpt_path}")
|
||||
sd = load_file(ckpt_path, device=str(device))
|
||||
sd = {}
|
||||
for ckpt_path in ckpt_paths:
|
||||
sd.update(load_safetensors(ckpt_path, device=str(device), disable_mmap=disable_mmap, dtype=dtype))
|
||||
|
||||
# convert Diffusers to BFL
|
||||
if is_diffusers:
|
||||
logger.info("Converting Diffusers to BFL")
|
||||
sd = convert_diffusers_sd_to_bfl(sd, num_double_blocks, num_single_blocks)
|
||||
logger.info("Converted Diffusers to BFL")
|
||||
|
||||
for key in list(sd.keys()):
|
||||
new_key = key.replace("model.diffusion_model.", "")
|
||||
if new_key == key:
|
||||
break
|
||||
sd[new_key] = sd.pop(key)
|
||||
|
||||
info = model.load_state_dict(sd, strict=False, assign=True)
|
||||
logger.info(f"Loaded Flux: {info}")
|
||||
return model
|
||||
return is_schnell, model
|
||||
|
||||
|
||||
def load_ae(name: str, ckpt_path: str, dtype: torch.dtype, device: Union[str, torch.device]) -> AutoEncoder:
|
||||
def load_ae(
|
||||
ckpt_path: str, dtype: torch.dtype, device: Union[str, torch.device], disable_mmap: bool = False
|
||||
) -> AutoEncoder:
|
||||
logger.info("Building AutoEncoder")
|
||||
with torch.device("meta"):
|
||||
ae = AutoEncoder(configs[name].ae_params).to(dtype)
|
||||
# dev and schnell have the same AE params
|
||||
ae = AutoEncoder(configs[MODEL_NAME_DEV].ae_params).to(dtype)
|
||||
|
||||
logger.info(f"Loading state dict from {ckpt_path}")
|
||||
sd = load_file(ckpt_path, device=str(device))
|
||||
sd = load_safetensors(ckpt_path, device=str(device), disable_mmap=disable_mmap, dtype=dtype)
|
||||
info = ae.load_state_dict(sd, strict=False, assign=True)
|
||||
logger.info(f"Loaded AE: {info}")
|
||||
return ae
|
||||
|
||||
|
||||
def load_clip_l(ckpt_path: str, dtype: torch.dtype, device: Union[str, torch.device]) -> CLIPTextModel:
|
||||
logger.info("Building CLIP")
|
||||
def load_clip_l(
|
||||
ckpt_path: Optional[str],
|
||||
dtype: torch.dtype,
|
||||
device: Union[str, torch.device],
|
||||
disable_mmap: bool = False,
|
||||
state_dict: Optional[dict] = None,
|
||||
) -> CLIPTextModel:
|
||||
logger.info("Building CLIP-L")
|
||||
CLIPL_CONFIG = {
|
||||
"_name_or_path": "clip-vit-large-patch14/",
|
||||
"architectures": ["CLIPModel"],
|
||||
@@ -137,14 +255,23 @@ def load_clip_l(ckpt_path: str, dtype: torch.dtype, device: Union[str, torch.dev
|
||||
with init_empty_weights():
|
||||
clip = CLIPTextModel._from_config(config)
|
||||
|
||||
if state_dict is not None:
|
||||
sd = state_dict
|
||||
else:
|
||||
logger.info(f"Loading state dict from {ckpt_path}")
|
||||
sd = load_file(ckpt_path, device=str(device))
|
||||
sd = load_safetensors(ckpt_path, device=str(device), disable_mmap=disable_mmap, dtype=dtype)
|
||||
info = clip.load_state_dict(sd, strict=False, assign=True)
|
||||
logger.info(f"Loaded CLIP: {info}")
|
||||
logger.info(f"Loaded CLIP-L: {info}")
|
||||
return clip
|
||||
|
||||
|
||||
def load_t5xxl(ckpt_path: str, dtype: torch.dtype, device: Union[str, torch.device]) -> T5EncoderModel:
|
||||
def load_t5xxl(
|
||||
ckpt_path: str,
|
||||
dtype: Optional[torch.dtype],
|
||||
device: Union[str, torch.device],
|
||||
disable_mmap: bool = False,
|
||||
state_dict: Optional[dict] = None,
|
||||
) -> T5EncoderModel:
|
||||
T5_CONFIG_JSON = """
|
||||
{
|
||||
"architectures": [
|
||||
@@ -183,13 +310,21 @@ def load_t5xxl(ckpt_path: str, dtype: torch.dtype, device: Union[str, torch.devi
|
||||
with init_empty_weights():
|
||||
t5xxl = T5EncoderModel._from_config(config)
|
||||
|
||||
if state_dict is not None:
|
||||
sd = state_dict
|
||||
else:
|
||||
logger.info(f"Loading state dict from {ckpt_path}")
|
||||
sd = load_file(ckpt_path, device=str(device))
|
||||
sd = load_safetensors(ckpt_path, device=str(device), disable_mmap=disable_mmap, dtype=dtype)
|
||||
info = t5xxl.load_state_dict(sd, strict=False, assign=True)
|
||||
logger.info(f"Loaded T5xxl: {info}")
|
||||
return t5xxl
|
||||
|
||||
|
||||
def get_t5xxl_actual_dtype(t5xxl: T5EncoderModel) -> torch.dtype:
|
||||
# nn.Embedding is the first layer, but it could be casted to bfloat16 or float32
|
||||
return t5xxl.encoder.block[0].layer[0].SelfAttention.q.weight.dtype
|
||||
|
||||
|
||||
def prepare_img_ids(batch_size: int, packed_latent_height: int, packed_latent_width: int):
|
||||
img_ids = torch.zeros(packed_latent_height, packed_latent_width, 3)
|
||||
img_ids[..., 1] = img_ids[..., 1] + torch.arange(packed_latent_height)[:, None]
|
||||
@@ -212,3 +347,128 @@ def pack_latents(x: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
x = einops.rearrange(x, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=2, pw=2)
|
||||
return x
|
||||
|
||||
|
||||
# region Diffusers
|
||||
|
||||
NUM_DOUBLE_BLOCKS = 19
|
||||
NUM_SINGLE_BLOCKS = 38
|
||||
|
||||
BFL_TO_DIFFUSERS_MAP = {
|
||||
"time_in.in_layer.weight": ["time_text_embed.timestep_embedder.linear_1.weight"],
|
||||
"time_in.in_layer.bias": ["time_text_embed.timestep_embedder.linear_1.bias"],
|
||||
"time_in.out_layer.weight": ["time_text_embed.timestep_embedder.linear_2.weight"],
|
||||
"time_in.out_layer.bias": ["time_text_embed.timestep_embedder.linear_2.bias"],
|
||||
"vector_in.in_layer.weight": ["time_text_embed.text_embedder.linear_1.weight"],
|
||||
"vector_in.in_layer.bias": ["time_text_embed.text_embedder.linear_1.bias"],
|
||||
"vector_in.out_layer.weight": ["time_text_embed.text_embedder.linear_2.weight"],
|
||||
"vector_in.out_layer.bias": ["time_text_embed.text_embedder.linear_2.bias"],
|
||||
"guidance_in.in_layer.weight": ["time_text_embed.guidance_embedder.linear_1.weight"],
|
||||
"guidance_in.in_layer.bias": ["time_text_embed.guidance_embedder.linear_1.bias"],
|
||||
"guidance_in.out_layer.weight": ["time_text_embed.guidance_embedder.linear_2.weight"],
|
||||
"guidance_in.out_layer.bias": ["time_text_embed.guidance_embedder.linear_2.bias"],
|
||||
"txt_in.weight": ["context_embedder.weight"],
|
||||
"txt_in.bias": ["context_embedder.bias"],
|
||||
"img_in.weight": ["x_embedder.weight"],
|
||||
"img_in.bias": ["x_embedder.bias"],
|
||||
"double_blocks.().img_mod.lin.weight": ["norm1.linear.weight"],
|
||||
"double_blocks.().img_mod.lin.bias": ["norm1.linear.bias"],
|
||||
"double_blocks.().txt_mod.lin.weight": ["norm1_context.linear.weight"],
|
||||
"double_blocks.().txt_mod.lin.bias": ["norm1_context.linear.bias"],
|
||||
"double_blocks.().img_attn.qkv.weight": ["attn.to_q.weight", "attn.to_k.weight", "attn.to_v.weight"],
|
||||
"double_blocks.().img_attn.qkv.bias": ["attn.to_q.bias", "attn.to_k.bias", "attn.to_v.bias"],
|
||||
"double_blocks.().txt_attn.qkv.weight": ["attn.add_q_proj.weight", "attn.add_k_proj.weight", "attn.add_v_proj.weight"],
|
||||
"double_blocks.().txt_attn.qkv.bias": ["attn.add_q_proj.bias", "attn.add_k_proj.bias", "attn.add_v_proj.bias"],
|
||||
"double_blocks.().img_attn.norm.query_norm.scale": ["attn.norm_q.weight"],
|
||||
"double_blocks.().img_attn.norm.key_norm.scale": ["attn.norm_k.weight"],
|
||||
"double_blocks.().txt_attn.norm.query_norm.scale": ["attn.norm_added_q.weight"],
|
||||
"double_blocks.().txt_attn.norm.key_norm.scale": ["attn.norm_added_k.weight"],
|
||||
"double_blocks.().img_mlp.0.weight": ["ff.net.0.proj.weight"],
|
||||
"double_blocks.().img_mlp.0.bias": ["ff.net.0.proj.bias"],
|
||||
"double_blocks.().img_mlp.2.weight": ["ff.net.2.weight"],
|
||||
"double_blocks.().img_mlp.2.bias": ["ff.net.2.bias"],
|
||||
"double_blocks.().txt_mlp.0.weight": ["ff_context.net.0.proj.weight"],
|
||||
"double_blocks.().txt_mlp.0.bias": ["ff_context.net.0.proj.bias"],
|
||||
"double_blocks.().txt_mlp.2.weight": ["ff_context.net.2.weight"],
|
||||
"double_blocks.().txt_mlp.2.bias": ["ff_context.net.2.bias"],
|
||||
"double_blocks.().img_attn.proj.weight": ["attn.to_out.0.weight"],
|
||||
"double_blocks.().img_attn.proj.bias": ["attn.to_out.0.bias"],
|
||||
"double_blocks.().txt_attn.proj.weight": ["attn.to_add_out.weight"],
|
||||
"double_blocks.().txt_attn.proj.bias": ["attn.to_add_out.bias"],
|
||||
"single_blocks.().modulation.lin.weight": ["norm.linear.weight"],
|
||||
"single_blocks.().modulation.lin.bias": ["norm.linear.bias"],
|
||||
"single_blocks.().linear1.weight": ["attn.to_q.weight", "attn.to_k.weight", "attn.to_v.weight", "proj_mlp.weight"],
|
||||
"single_blocks.().linear1.bias": ["attn.to_q.bias", "attn.to_k.bias", "attn.to_v.bias", "proj_mlp.bias"],
|
||||
"single_blocks.().linear2.weight": ["proj_out.weight"],
|
||||
"single_blocks.().norm.query_norm.scale": ["attn.norm_q.weight"],
|
||||
"single_blocks.().norm.key_norm.scale": ["attn.norm_k.weight"],
|
||||
"single_blocks.().linear2.weight": ["proj_out.weight"],
|
||||
"single_blocks.().linear2.bias": ["proj_out.bias"],
|
||||
"final_layer.linear.weight": ["proj_out.weight"],
|
||||
"final_layer.linear.bias": ["proj_out.bias"],
|
||||
"final_layer.adaLN_modulation.1.weight": ["norm_out.linear.weight"],
|
||||
"final_layer.adaLN_modulation.1.bias": ["norm_out.linear.bias"],
|
||||
}
|
||||
|
||||
|
||||
def make_diffusers_to_bfl_map(num_double_blocks: int, num_single_blocks: int) -> dict[str, tuple[int, str]]:
|
||||
# make reverse map from diffusers map
|
||||
diffusers_to_bfl_map = {} # key: diffusers_key, value: (index, bfl_key)
|
||||
for b in range(num_double_blocks):
|
||||
for key, weights in BFL_TO_DIFFUSERS_MAP.items():
|
||||
if key.startswith("double_blocks."):
|
||||
block_prefix = f"transformer_blocks.{b}."
|
||||
for i, weight in enumerate(weights):
|
||||
diffusers_to_bfl_map[f"{block_prefix}{weight}"] = (i, key.replace("()", f"{b}"))
|
||||
for b in range(num_single_blocks):
|
||||
for key, weights in BFL_TO_DIFFUSERS_MAP.items():
|
||||
if key.startswith("single_blocks."):
|
||||
block_prefix = f"single_transformer_blocks.{b}."
|
||||
for i, weight in enumerate(weights):
|
||||
diffusers_to_bfl_map[f"{block_prefix}{weight}"] = (i, key.replace("()", f"{b}"))
|
||||
for key, weights in BFL_TO_DIFFUSERS_MAP.items():
|
||||
if not (key.startswith("double_blocks.") or key.startswith("single_blocks.")):
|
||||
for i, weight in enumerate(weights):
|
||||
diffusers_to_bfl_map[weight] = (i, key)
|
||||
return diffusers_to_bfl_map
|
||||
|
||||
|
||||
def convert_diffusers_sd_to_bfl(
|
||||
diffusers_sd: dict[str, torch.Tensor], num_double_blocks: int = NUM_DOUBLE_BLOCKS, num_single_blocks: int = NUM_SINGLE_BLOCKS
|
||||
) -> dict[str, torch.Tensor]:
|
||||
diffusers_to_bfl_map = make_diffusers_to_bfl_map(num_double_blocks, num_single_blocks)
|
||||
|
||||
# iterate over three safetensors files to reduce memory usage
|
||||
flux_sd = {}
|
||||
for diffusers_key, tensor in diffusers_sd.items():
|
||||
if diffusers_key in diffusers_to_bfl_map:
|
||||
index, bfl_key = diffusers_to_bfl_map[diffusers_key]
|
||||
if bfl_key not in flux_sd:
|
||||
flux_sd[bfl_key] = []
|
||||
flux_sd[bfl_key].append((index, tensor))
|
||||
else:
|
||||
logger.error(f"Error: Key not found in diffusers_to_bfl_map: {diffusers_key}")
|
||||
raise KeyError(f"Key not found in diffusers_to_bfl_map: {diffusers_key}")
|
||||
|
||||
# concat tensors if multiple tensors are mapped to a single key, sort by index
|
||||
for key, values in flux_sd.items():
|
||||
if len(values) == 1:
|
||||
flux_sd[key] = values[0][1]
|
||||
else:
|
||||
flux_sd[key] = torch.cat([value[1] for value in sorted(values, key=lambda x: x[0])])
|
||||
|
||||
# special case for final_layer.adaLN_modulation.1.weight and final_layer.adaLN_modulation.1.bias
|
||||
def swap_scale_shift(weight):
|
||||
shift, scale = weight.chunk(2, dim=0)
|
||||
new_weight = torch.cat([scale, shift], dim=0)
|
||||
return new_weight
|
||||
|
||||
if "final_layer.adaLN_modulation.1.weight" in flux_sd:
|
||||
flux_sd["final_layer.adaLN_modulation.1.weight"] = swap_scale_shift(flux_sd["final_layer.adaLN_modulation.1.weight"])
|
||||
if "final_layer.adaLN_modulation.1.bias" in flux_sd:
|
||||
flux_sd["final_layer.adaLN_modulation.1.bias"] = swap_scale_shift(flux_sd["final_layer.adaLN_modulation.1.bias"])
|
||||
|
||||
return flux_sd
|
||||
|
||||
|
||||
# endregion
|
||||
@@ -1,223 +0,0 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from diffusers.models.attention_processor import (
|
||||
Attention,
|
||||
AttnProcessor2_0,
|
||||
SlicedAttnProcessor,
|
||||
XFormersAttnProcessor
|
||||
)
|
||||
|
||||
try:
|
||||
import xformers.ops
|
||||
except:
|
||||
xformers = None
|
||||
|
||||
|
||||
loaded_networks = []
|
||||
|
||||
|
||||
def apply_single_hypernetwork(
|
||||
hypernetwork, hidden_states, encoder_hidden_states
|
||||
):
|
||||
context_k, context_v = hypernetwork.forward(hidden_states, encoder_hidden_states)
|
||||
return context_k, context_v
|
||||
|
||||
|
||||
def apply_hypernetworks(context_k, context_v, layer=None):
|
||||
if len(loaded_networks) == 0:
|
||||
return context_v, context_v
|
||||
for hypernetwork in loaded_networks:
|
||||
context_k, context_v = hypernetwork.forward(context_k, context_v)
|
||||
|
||||
context_k = context_k.to(dtype=context_k.dtype)
|
||||
context_v = context_v.to(dtype=context_k.dtype)
|
||||
|
||||
return context_k, context_v
|
||||
|
||||
|
||||
|
||||
def xformers_forward(
|
||||
self: XFormersAttnProcessor,
|
||||
attn: Attention,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor = None,
|
||||
attention_mask: torch.Tensor = None,
|
||||
):
|
||||
batch_size, sequence_length, _ = (
|
||||
hidden_states.shape
|
||||
if encoder_hidden_states is None
|
||||
else encoder_hidden_states.shape
|
||||
)
|
||||
|
||||
attention_mask = attn.prepare_attention_mask(
|
||||
attention_mask, sequence_length, batch_size
|
||||
)
|
||||
|
||||
query = attn.to_q(hidden_states)
|
||||
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
elif attn.norm_cross:
|
||||
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
|
||||
|
||||
context_k, context_v = apply_hypernetworks(hidden_states, encoder_hidden_states)
|
||||
|
||||
key = attn.to_k(context_k)
|
||||
value = attn.to_v(context_v)
|
||||
|
||||
query = attn.head_to_batch_dim(query).contiguous()
|
||||
key = attn.head_to_batch_dim(key).contiguous()
|
||||
value = attn.head_to_batch_dim(value).contiguous()
|
||||
|
||||
hidden_states = xformers.ops.memory_efficient_attention(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
attn_bias=attention_mask,
|
||||
op=self.attention_op,
|
||||
scale=attn.scale,
|
||||
)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
hidden_states = attn.batch_to_head_dim(hidden_states)
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
def sliced_attn_forward(
|
||||
self: SlicedAttnProcessor,
|
||||
attn: Attention,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor = None,
|
||||
attention_mask: torch.Tensor = None,
|
||||
):
|
||||
batch_size, sequence_length, _ = (
|
||||
hidden_states.shape
|
||||
if encoder_hidden_states is None
|
||||
else encoder_hidden_states.shape
|
||||
)
|
||||
attention_mask = attn.prepare_attention_mask(
|
||||
attention_mask, sequence_length, batch_size
|
||||
)
|
||||
|
||||
query = attn.to_q(hidden_states)
|
||||
dim = query.shape[-1]
|
||||
query = attn.head_to_batch_dim(query)
|
||||
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
elif attn.norm_cross:
|
||||
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
|
||||
|
||||
context_k, context_v = apply_hypernetworks(hidden_states, encoder_hidden_states)
|
||||
|
||||
key = attn.to_k(context_k)
|
||||
value = attn.to_v(context_v)
|
||||
key = attn.head_to_batch_dim(key)
|
||||
value = attn.head_to_batch_dim(value)
|
||||
|
||||
batch_size_attention, query_tokens, _ = query.shape
|
||||
hidden_states = torch.zeros(
|
||||
(batch_size_attention, query_tokens, dim // attn.heads),
|
||||
device=query.device,
|
||||
dtype=query.dtype,
|
||||
)
|
||||
|
||||
for i in range(batch_size_attention // self.slice_size):
|
||||
start_idx = i * self.slice_size
|
||||
end_idx = (i + 1) * self.slice_size
|
||||
|
||||
query_slice = query[start_idx:end_idx]
|
||||
key_slice = key[start_idx:end_idx]
|
||||
attn_mask_slice = (
|
||||
attention_mask[start_idx:end_idx] if attention_mask is not None else None
|
||||
)
|
||||
|
||||
attn_slice = attn.get_attention_scores(query_slice, key_slice, attn_mask_slice)
|
||||
|
||||
attn_slice = torch.bmm(attn_slice, value[start_idx:end_idx])
|
||||
|
||||
hidden_states[start_idx:end_idx] = attn_slice
|
||||
|
||||
hidden_states = attn.batch_to_head_dim(hidden_states)
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
def v2_0_forward(
|
||||
self: AttnProcessor2_0,
|
||||
attn: Attention,
|
||||
hidden_states,
|
||||
encoder_hidden_states=None,
|
||||
attention_mask=None,
|
||||
):
|
||||
batch_size, sequence_length, _ = (
|
||||
hidden_states.shape
|
||||
if encoder_hidden_states is None
|
||||
else encoder_hidden_states.shape
|
||||
)
|
||||
inner_dim = hidden_states.shape[-1]
|
||||
|
||||
if attention_mask is not None:
|
||||
attention_mask = attn.prepare_attention_mask(
|
||||
attention_mask, sequence_length, batch_size
|
||||
)
|
||||
# scaled_dot_product_attention expects attention_mask shape to be
|
||||
# (batch, heads, source_length, target_length)
|
||||
attention_mask = attention_mask.view(
|
||||
batch_size, attn.heads, -1, attention_mask.shape[-1]
|
||||
)
|
||||
|
||||
query = attn.to_q(hidden_states)
|
||||
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
elif attn.norm_cross:
|
||||
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
|
||||
|
||||
context_k, context_v = apply_hypernetworks(hidden_states, encoder_hidden_states)
|
||||
|
||||
key = attn.to_k(context_k)
|
||||
value = attn.to_v(context_v)
|
||||
|
||||
head_dim = inner_dim // attn.heads
|
||||
query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
|
||||
# the output of sdp = (batch, num_heads, seq_len, head_dim)
|
||||
# TODO: add support for attn.scale when we move to Torch 2.1
|
||||
hidden_states = F.scaled_dot_product_attention(
|
||||
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
|
||||
)
|
||||
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(
|
||||
batch_size, -1, attn.heads * head_dim
|
||||
)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
def replace_attentions_for_hypernetwork():
|
||||
import diffusers.models.attention_processor
|
||||
|
||||
diffusers.models.attention_processor.XFormersAttnProcessor.__call__ = (
|
||||
xformers_forward
|
||||
)
|
||||
diffusers.models.attention_processor.SlicedAttnProcessor.__call__ = (
|
||||
sliced_attn_forward
|
||||
)
|
||||
diffusers.models.attention_processor.AttnProcessor2_0.__call__ = v2_0_forward
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1327,6 +1327,14 @@ def make_bucket_resolutions(max_reso, min_size=256, max_size=1024, divisible=64)
|
||||
resos.add((width, height))
|
||||
resos.add((height, width))
|
||||
|
||||
# # make additional resos
|
||||
# if width >= height and width - divisible >= min_size:
|
||||
# resos.add((width - divisible, height))
|
||||
# resos.add((height, width - divisible))
|
||||
# if height >= width and height - divisible >= min_size:
|
||||
# resos.add((width, height - divisible))
|
||||
# resos.add((height - divisible, width))
|
||||
|
||||
width += divisible
|
||||
|
||||
resos = list(resos)
|
||||
|
||||
@@ -57,8 +57,8 @@ ARCH_SD_V1 = "stable-diffusion-v1"
|
||||
ARCH_SD_V2_512 = "stable-diffusion-v2-512"
|
||||
ARCH_SD_V2_768_V = "stable-diffusion-v2-768-v"
|
||||
ARCH_SD_XL_V1_BASE = "stable-diffusion-xl-v1-base"
|
||||
ARCH_SD3_M = "stable-diffusion-3-medium"
|
||||
ARCH_SD3_UNKNOWN = "stable-diffusion-3"
|
||||
ARCH_SD3_M = "stable-diffusion-3" # may be followed by "-m" or "-5-large" etc.
|
||||
# ARCH_SD3_UNKNOWN = "stable-diffusion-3"
|
||||
ARCH_FLUX_1_DEV = "flux-1-dev"
|
||||
ARCH_FLUX_1_UNKNOWN = "flux-1"
|
||||
|
||||
@@ -140,10 +140,7 @@ def build_metadata(
|
||||
if sdxl:
|
||||
arch = ARCH_SD_XL_V1_BASE
|
||||
elif sd3 is not None:
|
||||
if sd3 == "m":
|
||||
arch = ARCH_SD3_M
|
||||
else:
|
||||
arch = ARCH_SD3_UNKNOWN
|
||||
arch = ARCH_SD3_M + "-" + sd3
|
||||
elif flux is not None:
|
||||
if flux == "dev":
|
||||
arch = ARCH_FLUX_1_DEV
|
||||
|
||||
+357
-1015
File diff suppressed because it is too large
Load Diff
+268
-214
@@ -11,10 +11,11 @@ from safetensors.torch import save_file
|
||||
from accelerate import Accelerator, PartialState
|
||||
from tqdm import tqdm
|
||||
from PIL import Image
|
||||
from transformers import CLIPTextModelWithProjection, T5EncoderModel
|
||||
|
||||
from . import sd3_models, sd3_utils, strategy_base, train_util
|
||||
from .device_utils import init_ipex, clean_memory_on_device
|
||||
|
||||
from comfy.utils import ProgressBar
|
||||
init_ipex()
|
||||
|
||||
# from transformers import CLIPTokenizer
|
||||
@@ -28,57 +29,16 @@ import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
def load_target_model(
|
||||
model_type: str,
|
||||
args: argparse.Namespace,
|
||||
state_dict: dict,
|
||||
accelerator: Accelerator,
|
||||
attn_mode: str,
|
||||
model_dtype: Optional[torch.dtype],
|
||||
device: Optional[torch.device],
|
||||
) -> Union[
|
||||
sd3_models.MMDiT,
|
||||
Optional[sd3_models.SDClipModel],
|
||||
Optional[sd3_models.SDXLClipG],
|
||||
Optional[sd3_models.T5XXLModel],
|
||||
sd3_models.SDVAE,
|
||||
]:
|
||||
loading_device = device if device is not None else (accelerator.device if args.lowram else "cpu")
|
||||
|
||||
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}")
|
||||
|
||||
if model_type == "mmdit":
|
||||
model = sd3_utils.load_mmdit(state_dict, attn_mode, model_dtype, loading_device)
|
||||
elif model_type == "clip_l":
|
||||
model = sd3_utils.load_clip_l(state_dict, args.clip_l, attn_mode, model_dtype, loading_device)
|
||||
elif model_type == "clip_g":
|
||||
model = sd3_utils.load_clip_g(state_dict, args.clip_g, attn_mode, model_dtype, loading_device)
|
||||
elif model_type == "t5xxl":
|
||||
model = sd3_utils.load_t5xxl(state_dict, args.t5xxl, attn_mode, model_dtype, loading_device)
|
||||
elif model_type == "vae":
|
||||
model = sd3_utils.load_vae(state_dict, args.vae, model_dtype, loading_device)
|
||||
else:
|
||||
raise ValueError(f"Unknown model type: {model_type}")
|
||||
|
||||
# work on low-ram device: models are already loaded on accelerator.device, but we ensure they are on device
|
||||
if args.lowram:
|
||||
model = model.to(accelerator.device)
|
||||
|
||||
clean_memory_on_device(accelerator.device)
|
||||
accelerator.wait_for_everyone()
|
||||
|
||||
return model
|
||||
from . import sd3_models, sd3_utils, strategy_base, train_util
|
||||
|
||||
|
||||
def save_models(
|
||||
ckpt_path: str,
|
||||
mmdit: sd3_models.MMDiT,
|
||||
vae: sd3_models.SDVAE,
|
||||
clip_l: sd3_models.SDClipModel,
|
||||
clip_g: sd3_models.SDXLClipG,
|
||||
t5xxl: Optional[sd3_models.T5XXLModel],
|
||||
mmdit: Optional[sd3_models.MMDiT],
|
||||
vae: Optional[sd3_models.SDVAE],
|
||||
clip_l: Optional[CLIPTextModelWithProjection],
|
||||
clip_g: Optional[CLIPTextModelWithProjection],
|
||||
t5xxl: Optional[T5EncoderModel],
|
||||
sai_metadata: Optional[dict],
|
||||
save_dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
@@ -98,24 +58,42 @@ def save_models(
|
||||
update_sd("model.diffusion_model.", mmdit.state_dict())
|
||||
update_sd("first_stage_model.", vae.state_dict())
|
||||
|
||||
if clip_l is not None:
|
||||
update_sd("text_encoders.clip_l.", clip_l.state_dict())
|
||||
if clip_g is not None:
|
||||
update_sd("text_encoders.clip_g.", clip_g.state_dict())
|
||||
if t5xxl is not None:
|
||||
update_sd("text_encoders.t5xxl.", t5xxl.state_dict())
|
||||
# do not support unified checkpoint format for now
|
||||
# if clip_l is not None:
|
||||
# update_sd("text_encoders.clip_l.", clip_l.state_dict())
|
||||
# if clip_g is not None:
|
||||
# update_sd("text_encoders.clip_g.", clip_g.state_dict())
|
||||
# if t5xxl is not None:
|
||||
# update_sd("text_encoders.t5xxl.", t5xxl.state_dict())
|
||||
|
||||
save_file(state_dict, ckpt_path, metadata=sai_metadata)
|
||||
|
||||
if clip_l is not None:
|
||||
clip_l_path = ckpt_path.replace(".safetensors", "_clip_l.safetensors")
|
||||
save_file(clip_l.state_dict(), clip_l_path)
|
||||
if clip_g is not None:
|
||||
clip_g_path = ckpt_path.replace(".safetensors", "_clip_g.safetensors")
|
||||
save_file(clip_g.state_dict(), clip_g_path)
|
||||
if t5xxl is not None:
|
||||
t5xxl_path = ckpt_path.replace(".safetensors", "_t5xxl.safetensors")
|
||||
t5xxl_state_dict = t5xxl.state_dict()
|
||||
|
||||
# replace "shared.weight" with copy of it to avoid annoying shared tensor error on safetensors.save_file
|
||||
shared_weight = t5xxl_state_dict["shared.weight"]
|
||||
shared_weight_copy = shared_weight.detach().clone()
|
||||
t5xxl_state_dict["shared.weight"] = shared_weight_copy
|
||||
|
||||
save_file(t5xxl_state_dict, t5xxl_path)
|
||||
|
||||
|
||||
def save_sd3_model_on_train_end(
|
||||
args: argparse.Namespace,
|
||||
save_dtype: torch.dtype,
|
||||
epoch: int,
|
||||
global_step: int,
|
||||
clip_l: sd3_models.SDClipModel,
|
||||
clip_g: sd3_models.SDXLClipG,
|
||||
t5xxl: Optional[sd3_models.T5XXLModel],
|
||||
clip_l: Optional[CLIPTextModelWithProjection],
|
||||
clip_g: Optional[CLIPTextModelWithProjection],
|
||||
t5xxl: Optional[T5EncoderModel],
|
||||
mmdit: sd3_models.MMDiT,
|
||||
vae: sd3_models.SDVAE,
|
||||
):
|
||||
@@ -138,9 +116,9 @@ def save_sd3_model_on_epoch_end_or_stepwise(
|
||||
epoch: int,
|
||||
num_train_epochs: int,
|
||||
global_step: int,
|
||||
clip_l: sd3_models.SDClipModel,
|
||||
clip_g: sd3_models.SDXLClipG,
|
||||
t5xxl: Optional[sd3_models.T5XXLModel],
|
||||
clip_l: Optional[CLIPTextModelWithProjection],
|
||||
clip_g: Optional[CLIPTextModelWithProjection],
|
||||
t5xxl: Optional[T5EncoderModel],
|
||||
mmdit: sd3_models.MMDiT,
|
||||
vae: sd3_models.SDVAE,
|
||||
):
|
||||
@@ -165,27 +143,6 @@ def save_sd3_model_on_epoch_end_or_stepwise(
|
||||
|
||||
|
||||
def add_sd3_training_arguments(parser: argparse.ArgumentParser):
|
||||
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(
|
||||
"--text_encoder_batch_size",
|
||||
type=int,
|
||||
default=None,
|
||||
help="text encoder batch size (default: None, use dataset's batch size)"
|
||||
+ " / text encoderのバッチサイズ(デフォルト: None, データセットのバッチサイズを使用)",
|
||||
)
|
||||
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環境等でモデル読み込みを高速化できる",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--clip_l",
|
||||
type=str,
|
||||
@@ -205,41 +162,84 @@ def add_sd3_training_arguments(parser: argparse.ArgumentParser):
|
||||
help="T5-XXL model path. if not specified, use ckpt's state_dict / T5-XXLモデルのパス。指定しない場合はckptのstate_dictを使用",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--save_clip", action="store_true", help="save CLIP models to checkpoint / CLIPモデルをチェックポイントに保存する"
|
||||
"--save_clip",
|
||||
action="store_true",
|
||||
help="[DOES NOT WORK] unified checkpoint is not supported / 統合チェックポイントはまだサポートされていません",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--save_t5xxl", action="store_true", help="save T5-XXL model to checkpoint / T5-XXLモデルをチェックポイントに保存する"
|
||||
"--save_t5xxl",
|
||||
action="store_true",
|
||||
help="[DOES NOT WORK] unified checkpoint is not supported / 統合チェックポイントはまだサポートされていません",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--t5xxl_device",
|
||||
type=str,
|
||||
default=None,
|
||||
help="T5-XXL device. if not specified, use accelerator's device / T5-XXLデバイス。指定しない場合はacceleratorのデバイスを使用",
|
||||
help="[DOES NOT WORK] not supported yet. T5-XXL device. if not specified, use accelerator's device / T5-XXLデバイス。指定しない場合はacceleratorのデバイスを使用",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--t5xxl_dtype",
|
||||
type=str,
|
||||
default=None,
|
||||
help="T5-XXL dtype. if not specified, use default dtype (from mixed precision) / T5-XXL dtype。指定しない場合はデフォルトのdtype(mixed precisionから)を使用",
|
||||
help="[DOES NOT WORK] not supported yet. T5-XXL dtype. if not specified, use default dtype (from mixed precision) / T5-XXL dtype。指定しない場合はデフォルトのdtype(mixed precisionから)を使用",
|
||||
)
|
||||
|
||||
# copy from Diffusers
|
||||
parser.add_argument(
|
||||
"--weighting_scheme",
|
||||
type=str,
|
||||
default="logit_normal",
|
||||
choices=["sigma_sqrt", "logit_normal", "mode", "cosmap"],
|
||||
"--t5xxl_max_token_length",
|
||||
type=int,
|
||||
default=256,
|
||||
help="maximum token length for T5-XXL. 256 is the default value / T5-XXLの最大トークン長。デフォルトは256",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logit_mean", type=float, default=0.0, help="mean to use when using the `'logit_normal'` weighting scheme."
|
||||
"--apply_lg_attn_mask",
|
||||
action="store_true",
|
||||
help="apply attention mask (zero embs) to CLIP-L and G / CLIP-LとGにアテンションマスク(ゼロ埋め)を適用する",
|
||||
)
|
||||
parser.add_argument("--logit_std", type=float, default=1.0, help="std to use when using the `'logit_normal'` weighting scheme.")
|
||||
parser.add_argument(
|
||||
"--mode_scale",
|
||||
"--apply_t5_attn_mask",
|
||||
action="store_true",
|
||||
help="apply attention mask (zero embs) to T5-XXL / T5-XXLにアテンションマスク(ゼロ埋め)を適用する",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--clip_l_dropout_rate",
|
||||
type=float,
|
||||
default=1.29,
|
||||
help="Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.",
|
||||
default=0.0,
|
||||
help="Dropout rate for CLIP-L encoder, default is 0.0 / CLIP-Lエンコーダのドロップアウト率、デフォルトは0.0",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--clip_g_dropout_rate",
|
||||
type=float,
|
||||
default=0.0,
|
||||
help="Dropout rate for CLIP-G encoder, default is 0.0 / CLIP-Gエンコーダのドロップアウト率、デフォルトは0.0",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--t5_dropout_rate",
|
||||
type=float,
|
||||
default=0.0,
|
||||
help="Dropout rate for T5 encoder, default is 0.0 / T5エンコーダのドロップアウト率、デフォルトは0.0",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--pos_emb_random_crop_rate",
|
||||
type=float,
|
||||
default=0.0,
|
||||
help="Random crop rate for positional embeddings, default is 0.0. Only for SD3.5M"
|
||||
" / 位置埋め込みのランダムクロップ率、デフォルトは0.0。SD3.5M以外では予期しない動作になります",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--enable_scaled_pos_embed",
|
||||
action="store_true",
|
||||
help="Scale position embeddings for each resolution during multi-resolution training. Only for SD3.5M"
|
||||
" / 複数解像度学習時に解像度ごとに位置埋め込みをスケーリングする。SD3.5M以外では予期しない動作になります",
|
||||
)
|
||||
|
||||
# Dependencies of Diffusers noise sampler has been removed for clarity in training
|
||||
|
||||
parser.add_argument(
|
||||
"--training_shift",
|
||||
type=float,
|
||||
default=1.0,
|
||||
help="Discrete flow shift for training timestep distribution adjustment, applied in addition to the weighting scheme, default is 1.0. /タイムステップ分布のための離散フローシフト、重み付けスキームの上に適用される、デフォルトは1.0。",
|
||||
)
|
||||
|
||||
|
||||
@@ -280,7 +280,7 @@ def verify_sdxl_training_args(args: argparse.Namespace, supportTextEncoderCachin
|
||||
# temporary copied from sd3_minimal_inferece.py
|
||||
|
||||
|
||||
def get_sigmas(sampling: sd3_utils.ModelSamplingDiscreteFlow, steps):
|
||||
def get_all_sigmas(sampling: sd3_utils.ModelSamplingDiscreteFlow, steps):
|
||||
start = sampling.timestep(sampling.sigma_max)
|
||||
end = sampling.timestep(sampling.sigma_min)
|
||||
timesteps = torch.linspace(start, end, steps)
|
||||
@@ -316,6 +316,8 @@ def do_sample(
|
||||
# noise = get_noise(seed, latent).to(device)
|
||||
if seed is not None:
|
||||
generator = torch.manual_seed(seed)
|
||||
else:
|
||||
generator = None
|
||||
noise = (
|
||||
torch.randn(latent.size(), dtype=torch.float32, layout=latent.layout, generator=generator, device="cpu")
|
||||
.to(latent.dtype)
|
||||
@@ -324,7 +326,7 @@ def do_sample(
|
||||
|
||||
model_sampling = sd3_utils.ModelSamplingDiscreteFlow(shift=3.0) # 3.0 is for SD3
|
||||
|
||||
sigmas = get_sigmas(model_sampling, steps).to(device)
|
||||
sigmas = get_all_sigmas(model_sampling, steps).to(device)
|
||||
|
||||
noise_scaled = model_sampling.noise_scaling(sigmas[0], noise, latent, max_denoise(model_sampling, sigmas))
|
||||
|
||||
@@ -334,7 +336,8 @@ def do_sample(
|
||||
x = noise_scaled.to(device).to(dtype)
|
||||
# print(x.shape)
|
||||
|
||||
with torch.no_grad():
|
||||
# with torch.no_grad():
|
||||
comfy_pbar = ProgressBar(len(sigmas) - 1)
|
||||
for i in tqdm(range(len(sigmas) - 1)):
|
||||
sigma_hat = sigmas[i]
|
||||
|
||||
@@ -344,6 +347,7 @@ def do_sample(
|
||||
x_c_nc = torch.cat([x, x], dim=0)
|
||||
# print(x_c_nc.shape, timestep.shape, c_crossattn.shape, y.shape)
|
||||
|
||||
mmdit.prepare_block_swap_before_forward()
|
||||
model_output = mmdit(x_c_nc, timestep, context=c_crossattn, y=y)
|
||||
model_output = model_output.float()
|
||||
batched = model_sampling.calculate_denoised(sigma_hat, model_output, x)
|
||||
@@ -364,41 +368,12 @@ def do_sample(
|
||||
# Euler method
|
||||
x = x + d * dt
|
||||
x = x.to(dtype)
|
||||
comfy_pbar.update(1)
|
||||
|
||||
mmdit.prepare_block_swap_before_forward()
|
||||
return x
|
||||
|
||||
|
||||
def load_prompts(prompt_file: str) -> List[Dict]:
|
||||
# read prompts
|
||||
if prompt_file.endswith(".txt"):
|
||||
with open(prompt_file, "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 prompt_file.endswith(".toml"):
|
||||
with open(prompt_file, "r", encoding="utf-8") as f:
|
||||
data = toml.load(f)
|
||||
prompts = [dict(**data["prompt"], **subset) for subset in data["prompt"]["subset"]]
|
||||
elif prompt_file.endswith(".json"):
|
||||
with open(prompt_file, "r", encoding="utf-8") as f:
|
||||
prompts = json.load(f)
|
||||
|
||||
# preprocess prompts
|
||||
for i in range(len(prompts)):
|
||||
prompt_dict = prompts[i]
|
||||
if isinstance(prompt_dict, str):
|
||||
from library.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)
|
||||
|
||||
return prompts
|
||||
|
||||
|
||||
def sample_images(
|
||||
accelerator: Accelerator,
|
||||
args: argparse.Namespace,
|
||||
@@ -409,35 +384,35 @@ def sample_images(
|
||||
text_encoders,
|
||||
sample_prompts_te_outputs,
|
||||
prompt_replacement=None,
|
||||
validation_settings=None,
|
||||
):
|
||||
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
|
||||
|
||||
distributed_state = PartialState() # for multi gpu distributed inference. this is a singleton, so it's safe to use it here
|
||||
logger.info(f"generating sample images at step: {steps}")
|
||||
|
||||
# unwrap unet and text_encoder(s)
|
||||
mmdit = accelerator.unwrap_model(mmdit)
|
||||
text_encoders = [accelerator.unwrap_model(te) for te in text_encoders]
|
||||
text_encoders = None if text_encoders is None else [accelerator.unwrap_model(te) for te in text_encoders]
|
||||
# print([(te.parameters().__next__().device if te is not None else None) for te in text_encoders])
|
||||
|
||||
prompts = load_prompts(args.sample_prompts)
|
||||
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)
|
||||
|
||||
save_dir = args.output_dir + "/sample"
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
@@ -450,37 +425,10 @@ def sample_images(
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
org_vae_device = vae.device # will be on cpu
|
||||
vae.to(distributed_state.device) # distributed_state.device is same as accelerator.device
|
||||
|
||||
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():
|
||||
with torch.no_grad(), accelerator.autocast():
|
||||
image_tensor_list = []
|
||||
for prompt_dict in prompts:
|
||||
sample_image_inference(
|
||||
accelerator,
|
||||
args,
|
||||
mmdit,
|
||||
text_encoders,
|
||||
vae,
|
||||
save_dir,
|
||||
prompt_dict,
|
||||
epoch,
|
||||
steps,
|
||||
sample_prompts_te_outputs,
|
||||
prompt_replacement,
|
||||
)
|
||||
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(
|
||||
image_tensor = sample_image_inference(
|
||||
accelerator,
|
||||
args,
|
||||
mmdit,
|
||||
@@ -492,38 +440,49 @@ def sample_images(
|
||||
steps,
|
||||
sample_prompts_te_outputs,
|
||||
prompt_replacement,
|
||||
validation_settings
|
||||
)
|
||||
print(f"Sampled image shape: {image_tensor.shape}")
|
||||
image_tensor_list.append(image_tensor)
|
||||
|
||||
torch.set_rng_state(rng_state)
|
||||
if cuda_rng_state is not None:
|
||||
torch.cuda.set_rng_state(cuda_rng_state)
|
||||
|
||||
vae.to(org_vae_device)
|
||||
|
||||
clean_memory_on_device(accelerator.device)
|
||||
return torch.cat(image_tensor_list, dim=0)
|
||||
|
||||
|
||||
def sample_image_inference(
|
||||
accelerator: Accelerator,
|
||||
args: argparse.Namespace,
|
||||
mmdit: sd3_models.MMDiT,
|
||||
text_encoders: List[Union[sd3_models.SDClipModel, sd3_models.SDXLClipG, sd3_models.T5XXLModel]],
|
||||
text_encoders: List[Union[CLIPTextModelWithProjection, T5EncoderModel]],
|
||||
vae: sd3_models.SDVAE,
|
||||
save_dir,
|
||||
prompt_dict,
|
||||
epoch,
|
||||
steps,
|
||||
sample_prompts_te_outputs,
|
||||
prompt_replacement,
|
||||
validation_settings=None,
|
||||
prompt_replacement=None,
|
||||
|
||||
):
|
||||
assert isinstance(prompt_dict, dict)
|
||||
negative_prompt = prompt_dict.get("negative_prompt")
|
||||
if validation_settings is not None:
|
||||
sample_steps = validation_settings["steps"]
|
||||
width = validation_settings["width"]
|
||||
height = validation_settings["height"]
|
||||
scale = validation_settings["guidance_scale"]
|
||||
seed = validation_settings["seed"]
|
||||
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")
|
||||
negative_prompt = prompt_dict.get("negative_prompt")
|
||||
prompt: str = prompt_dict.get("prompt", "")
|
||||
# sampler_name: str = prompt_dict.get("sample_sampler", args.sample_sampler)
|
||||
|
||||
@@ -559,33 +518,50 @@ def sample_image_inference(
|
||||
tokenize_strategy = strategy_base.TokenizeStrategy.get_strategy()
|
||||
encoding_strategy = strategy_base.TextEncodingStrategy.get_strategy()
|
||||
|
||||
if sample_prompts_te_outputs and prompt in sample_prompts_te_outputs:
|
||||
te_outputs = sample_prompts_te_outputs[prompt]
|
||||
else:
|
||||
l_tokens, g_tokens, t5_tokens = tokenize_strategy.tokenize(prompt)
|
||||
te_outputs = encoding_strategy.encode_tokens(tokenize_strategy, text_encoders, [l_tokens, g_tokens, t5_tokens])
|
||||
def encode_prompt(prpt):
|
||||
text_encoder_conds = []
|
||||
if sample_prompts_te_outputs and prpt in sample_prompts_te_outputs:
|
||||
text_encoder_conds = sample_prompts_te_outputs[prpt]
|
||||
print(f"Using cached text encoder outputs for prompt: {prpt}")
|
||||
if text_encoders is not None:
|
||||
print(f"Encoding prompt: {prpt}")
|
||||
tokens_and_masks = tokenize_strategy.tokenize(prpt)
|
||||
# strategy has apply_t5_attn_mask option
|
||||
encoded_text_encoder_conds = encoding_strategy.encode_tokens(tokenize_strategy, text_encoders, tokens_and_masks)
|
||||
|
||||
lg_out, t5_out, pooled = te_outputs
|
||||
# if text_encoder_conds is not cached, use encoded_text_encoder_conds
|
||||
if len(text_encoder_conds) == 0:
|
||||
text_encoder_conds = encoded_text_encoder_conds
|
||||
else:
|
||||
# if encoded_text_encoder_conds is not None, update cached text_encoder_conds
|
||||
for i in range(len(encoded_text_encoder_conds)):
|
||||
if encoded_text_encoder_conds[i] is not None:
|
||||
text_encoder_conds[i] = encoded_text_encoder_conds[i]
|
||||
return text_encoder_conds
|
||||
|
||||
lg_out, t5_out, pooled, l_attn_mask, g_attn_mask, t5_attn_mask = encode_prompt(prompt)
|
||||
cond = encoding_strategy.concat_encodings(lg_out, t5_out, pooled)
|
||||
|
||||
# encode negative prompts
|
||||
if sample_prompts_te_outputs and negative_prompt in sample_prompts_te_outputs:
|
||||
neg_te_outputs = sample_prompts_te_outputs[negative_prompt]
|
||||
else:
|
||||
l_tokens, g_tokens, t5_tokens = tokenize_strategy.tokenize(negative_prompt)
|
||||
neg_te_outputs = encoding_strategy.encode_tokens(tokenize_strategy, text_encoders, [l_tokens, g_tokens, t5_tokens])
|
||||
|
||||
lg_out, t5_out, pooled = neg_te_outputs
|
||||
lg_out, t5_out, pooled, l_attn_mask, g_attn_mask, t5_attn_mask = encode_prompt(negative_prompt)
|
||||
neg_cond = encoding_strategy.concat_encodings(lg_out, t5_out, pooled)
|
||||
|
||||
# sample image
|
||||
latents = do_sample(height, width, seed, cond, neg_cond, mmdit, sample_steps, scale, mmdit.dtype, accelerator.device)
|
||||
latents = vae.process_out(latents.to(vae.device, dtype=vae.dtype))
|
||||
clean_memory_on_device(accelerator.device)
|
||||
with accelerator.autocast(), torch.no_grad():
|
||||
# mmdit may be fp8, so we need weight_dtype here. vae is always in that dtype.
|
||||
latents = do_sample(height, width, seed, cond, neg_cond, mmdit, sample_steps, scale, vae.dtype, accelerator.device)
|
||||
|
||||
# latent to image
|
||||
with torch.no_grad():
|
||||
image = vae.decode(latents)
|
||||
image = image.float()
|
||||
clean_memory_on_device(accelerator.device)
|
||||
org_vae_device = vae.device # will be on cpu
|
||||
vae.to(accelerator.device)
|
||||
latents = vae.process_out(latents.to(vae.device, dtype=vae.dtype))
|
||||
image_tensor = vae.decode(latents)
|
||||
vae.to(org_vae_device)
|
||||
clean_memory_on_device(accelerator.device)
|
||||
|
||||
image = image_tensor.float()
|
||||
image = torch.clamp((image + 1.0) / 2.0, min=0.0, max=1.0)[0]
|
||||
decoded_np = 255.0 * np.moveaxis(image.cpu().numpy(), 0, 2)
|
||||
decoded_np = decoded_np.astype(np.uint8)
|
||||
@@ -600,18 +576,16 @@ def sample_image_inference(
|
||||
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))
|
||||
return image_tensor
|
||||
|
||||
# wandb有効時のみログを送信
|
||||
try:
|
||||
wandb_tracker = accelerator.get_tracker("wandb")
|
||||
try:
|
||||
import wandb
|
||||
except ImportError: # 事前に一度確認するのでここはエラー出ないはず
|
||||
raise ImportError("No wandb / wandb がインストールされていないようです")
|
||||
# # send images to wandb if enabled
|
||||
# if "wandb" in [tracker.name for tracker in accelerator.trackers]:
|
||||
# wandb_tracker = accelerator.get_tracker("wandb")
|
||||
|
||||
wandb_tracker.log({f"sample_{i}": wandb.Image(image)})
|
||||
except: # wandb 無効時
|
||||
pass
|
||||
# import wandb
|
||||
|
||||
# # 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
|
||||
|
||||
|
||||
# region Diffusers
|
||||
@@ -881,4 +855,84 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
return self.config.num_train_timesteps
|
||||
|
||||
|
||||
def get_sigmas(noise_scheduler, timesteps, device, n_dim=4, dtype=torch.float32):
|
||||
sigmas = noise_scheduler.sigmas.to(device=device, dtype=dtype)
|
||||
schedule_timesteps = noise_scheduler.timesteps.to(device)
|
||||
timesteps = timesteps.to(device)
|
||||
step_indices = [(schedule_timesteps == t).nonzero().item() for t in timesteps]
|
||||
|
||||
sigma = sigmas[step_indices].flatten()
|
||||
while len(sigma.shape) < n_dim:
|
||||
sigma = sigma.unsqueeze(-1)
|
||||
return sigma
|
||||
|
||||
|
||||
def compute_density_for_timestep_sampling(
|
||||
weighting_scheme: str, batch_size: int, logit_mean: float = None, logit_std: float = None, mode_scale: float = None
|
||||
):
|
||||
"""Compute the density for sampling the timesteps when doing SD3 training.
|
||||
|
||||
Courtesy: This was contributed by Rafie Walker in https://github.com/huggingface/diffusers/pull/8528.
|
||||
|
||||
SD3 paper reference: https://arxiv.org/abs/2403.03206v1.
|
||||
"""
|
||||
if weighting_scheme == "logit_normal":
|
||||
# See 3.1 in the SD3 paper ($rf/lognorm(0.00,1.00)$).
|
||||
u = torch.normal(mean=logit_mean, std=logit_std, size=(batch_size,), device="cpu")
|
||||
u = torch.nn.functional.sigmoid(u)
|
||||
elif weighting_scheme == "mode":
|
||||
u = torch.rand(size=(batch_size,), device="cpu")
|
||||
u = 1 - u - mode_scale * (torch.cos(math.pi * u / 2) ** 2 - 1 + u)
|
||||
else:
|
||||
u = torch.rand(size=(batch_size,), device="cpu")
|
||||
return u
|
||||
|
||||
|
||||
def compute_loss_weighting_for_sd3(weighting_scheme: str, sigmas=None):
|
||||
"""Computes loss weighting scheme for SD3 training.
|
||||
|
||||
Courtesy: This was contributed by Rafie Walker in https://github.com/huggingface/diffusers/pull/8528.
|
||||
|
||||
SD3 paper reference: https://arxiv.org/abs/2403.03206v1.
|
||||
"""
|
||||
if weighting_scheme == "sigma_sqrt":
|
||||
weighting = (sigmas**-2.0).float()
|
||||
elif weighting_scheme == "cosmap":
|
||||
bot = 1 - 2 * sigmas + 2 * sigmas**2
|
||||
weighting = 2 / (math.pi * bot)
|
||||
else:
|
||||
weighting = torch.ones_like(sigmas)
|
||||
return weighting
|
||||
|
||||
|
||||
# endregion
|
||||
|
||||
|
||||
def get_noisy_model_input_and_timesteps(args, latents, noise, device, dtype) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
bsz = latents.shape[0]
|
||||
|
||||
# Sample a random timestep for each image
|
||||
# for weighting schemes where we sample timesteps non-uniformly
|
||||
u = compute_density_for_timestep_sampling(
|
||||
weighting_scheme=args.weighting_scheme,
|
||||
batch_size=bsz,
|
||||
logit_mean=args.logit_mean,
|
||||
logit_std=args.logit_std,
|
||||
mode_scale=args.mode_scale,
|
||||
)
|
||||
t_min = args.min_timestep if args.min_timestep is not None else 0
|
||||
t_max = args.max_timestep if args.max_timestep is not None else 1000
|
||||
shift = args.training_shift
|
||||
|
||||
# weighting shift, value >1 will shift distribution to noisy side (focus more on overall structure), value <1 will shift towards less-noisy side (focus more on details)
|
||||
u = (u * shift) / (1 + (shift - 1) * u)
|
||||
|
||||
indices = (u * (t_max - t_min) + t_min).long()
|
||||
timesteps = indices.to(device=device, dtype=dtype)
|
||||
|
||||
# sigmas according to flowmatching
|
||||
sigmas = timesteps / 1000
|
||||
sigmas = sigmas.view(-1, 1, 1, 1)
|
||||
noisy_model_input = sigmas * noise + (1.0 - sigmas) * latents
|
||||
|
||||
return noisy_model_input, timesteps, sigmas
|
||||
|
||||
+145
-379
@@ -1,10 +1,9 @@
|
||||
import math
|
||||
from typing import Dict, Optional, Union, List
|
||||
import re
|
||||
from typing import Dict, List, Optional, Union
|
||||
import torch
|
||||
import safetensors
|
||||
from safetensors.torch import load_file
|
||||
from accelerate import init_empty_weights
|
||||
from accelerate.utils.modeling import set_module_tensor_to_device
|
||||
from transformers import CLIPTextModel, CLIPTextModelWithProjection, CLIPConfig, CLIPTextConfig
|
||||
|
||||
from .utils import setup_logging
|
||||
|
||||
@@ -15,38 +14,64 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
from . import sd3_models
|
||||
|
||||
# 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())
|
||||
# region models
|
||||
|
||||
# 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>"
|
||||
# TODO remove dependency on flux_utils
|
||||
from .utils import load_safetensors
|
||||
from .flux_utils import load_t5xxl as flux_utils_load_t5xxl
|
||||
|
||||
# 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 analyze_state_dict_state(state_dict: Dict, prefix: str = ""):
|
||||
logger.info(f"Analyzing state dict state...")
|
||||
|
||||
def load_safetensors(path: str, dvc: Union[str, torch.device], disable_mmap: bool = False):
|
||||
if disable_mmap:
|
||||
return safetensors.torch.load(open(path, "rb").read())
|
||||
# analyze configs
|
||||
patch_size = state_dict[f"{prefix}x_embedder.proj.weight"].shape[2]
|
||||
depth = state_dict[f"{prefix}x_embedder.proj.weight"].shape[0] // 64
|
||||
num_patches = state_dict[f"{prefix}pos_embed"].shape[1]
|
||||
pos_embed_max_size = round(math.sqrt(num_patches))
|
||||
adm_in_channels = state_dict[f"{prefix}y_embedder.mlp.0.weight"].shape[1]
|
||||
context_shape = state_dict[f"{prefix}context_embedder.weight"].shape
|
||||
qk_norm = "rms" if f"{prefix}joint_blocks.0.context_block.attn.ln_k.weight" in state_dict.keys() else None
|
||||
|
||||
# x_block_self_attn_layers.append(int(key.split(".x_block.attn2.ln_k.weight")[0].split(".")[-1]))
|
||||
x_block_self_attn_layers = []
|
||||
re_attn = re.compile(r"\.(\d+)\.x_block\.attn2\.ln_k\.weight")
|
||||
for key in list(state_dict.keys()):
|
||||
m = re_attn.search(key)
|
||||
if m:
|
||||
x_block_self_attn_layers.append(int(m.group(1)))
|
||||
|
||||
context_embedder_in_features = context_shape[1]
|
||||
context_embedder_out_features = context_shape[0]
|
||||
|
||||
# only supports 3-5-large, medium or 3-medium
|
||||
if qk_norm is not None:
|
||||
if len(x_block_self_attn_layers) == 0:
|
||||
model_type = "3-5-large"
|
||||
else:
|
||||
try:
|
||||
return load_file(path, device=dvc)
|
||||
except:
|
||||
return load_file(path) # prevent device invalid Error
|
||||
model_type = "3-5-medium"
|
||||
else:
|
||||
model_type = "3-medium"
|
||||
|
||||
params = sd3_models.SD3Params(
|
||||
patch_size=patch_size,
|
||||
depth=depth,
|
||||
num_patches=num_patches,
|
||||
pos_embed_max_size=pos_embed_max_size,
|
||||
adm_in_channels=adm_in_channels,
|
||||
qk_norm=qk_norm,
|
||||
x_block_self_attn_layers=x_block_self_attn_layers,
|
||||
context_embedder_in_features=context_embedder_in_features,
|
||||
context_embedder_out_features=context_embedder_out_features,
|
||||
model_type=model_type,
|
||||
)
|
||||
logger.info(f"Analyzed state dict state: {params}")
|
||||
return params
|
||||
|
||||
|
||||
def load_mmdit(state_dict: Dict, attn_mode: str, dtype: Optional[Union[str, torch.dtype]], device: Union[str, torch.device]):
|
||||
def load_mmdit(
|
||||
state_dict: Dict, dtype: Optional[Union[str, torch.dtype]], device: Union[str, torch.device], attn_mode: str = "torch"
|
||||
) -> sd3_models.MMDiT:
|
||||
mmdit_sd = {}
|
||||
|
||||
mmdit_prefix = "model.diffusion_model."
|
||||
@@ -56,30 +81,25 @@ def load_mmdit(state_dict: Dict, attn_mode: str, dtype: Optional[Union[str, torc
|
||||
|
||||
# load MMDiT
|
||||
logger.info("Building MMDit")
|
||||
params = analyze_state_dict_state(mmdit_sd)
|
||||
with init_empty_weights():
|
||||
mmdit = sd3_models.create_mmdit_sd3_medium_configs(attn_mode)
|
||||
mmdit = sd3_models.create_sd3_mmdit(params, attn_mode)
|
||||
|
||||
logger.info("Loading state dict...")
|
||||
info = load_state_dict_on_device(mmdit, mmdit_sd, device, dtype)
|
||||
info = mmdit.load_state_dict(mmdit_sd, strict=False, assign=True)
|
||||
logger.info(f"Loaded MMDiT: {info}")
|
||||
return mmdit
|
||||
|
||||
|
||||
def load_clip_l(
|
||||
state_dict: Dict,
|
||||
clip_l_path: Optional[str],
|
||||
attn_mode: str,
|
||||
clip_dtype: Optional[Union[str, torch.dtype]],
|
||||
dtype: Optional[Union[str, torch.dtype]],
|
||||
device: Union[str, torch.device],
|
||||
disable_mmap: bool = False,
|
||||
state_dict: Optional[Dict] = None,
|
||||
):
|
||||
clip_l_sd = None
|
||||
if clip_l_path:
|
||||
logger.info(f"Loading clip_l from {clip_l_path}...")
|
||||
clip_l_sd = load_safetensors(clip_l_path, device, disable_mmap)
|
||||
for key in list(clip_l_sd.keys()):
|
||||
clip_l_sd["transformer." + key] = clip_l_sd.pop(key)
|
||||
else:
|
||||
if clip_l_path is None:
|
||||
if "text_encoders.clip_l.transformer.text_model.embeddings.position_embedding.weight" in state_dict:
|
||||
# found clip_l: remove prefix "text_encoders.clip_l."
|
||||
logger.info("clip_l is included in the checkpoint")
|
||||
@@ -88,34 +108,58 @@ def load_clip_l(
|
||||
for k in list(state_dict.keys()):
|
||||
if k.startswith(prefix):
|
||||
clip_l_sd[k[len(prefix) :]] = state_dict.pop(k)
|
||||
elif clip_l_path is None:
|
||||
logger.info("clip_l is not included in the checkpoint and clip_l_path is not provided")
|
||||
return None
|
||||
|
||||
# load clip_l
|
||||
logger.info("Building CLIP-L")
|
||||
config = 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():
|
||||
clip = CLIPTextModelWithProjection(config)
|
||||
|
||||
if clip_l_sd is None:
|
||||
clip_l = None
|
||||
else:
|
||||
logger.info("Building ClipL")
|
||||
clip_l = sd3_models.create_clip_l(device, clip_dtype, clip_l_sd)
|
||||
logger.info("Loading state dict...")
|
||||
info = clip_l.load_state_dict(clip_l_sd)
|
||||
logger.info(f"Loaded ClipL: {info}")
|
||||
clip_l.set_attn_mode(attn_mode)
|
||||
return clip_l
|
||||
logger.info(f"Loading state dict from {clip_l_path}")
|
||||
clip_l_sd = load_safetensors(clip_l_path, device=str(device), disable_mmap=disable_mmap, dtype=dtype)
|
||||
|
||||
if "text_projection.weight" not in clip_l_sd:
|
||||
logger.info("Adding text_projection.weight to clip_l_sd")
|
||||
clip_l_sd["text_projection.weight"] = torch.eye(768, dtype=dtype, device=device)
|
||||
|
||||
info = clip.load_state_dict(clip_l_sd, strict=False, assign=True)
|
||||
logger.info(f"Loaded CLIP-L: {info}")
|
||||
return clip
|
||||
|
||||
|
||||
def load_clip_g(
|
||||
state_dict: Dict,
|
||||
clip_g_path: Optional[str],
|
||||
attn_mode: str,
|
||||
clip_dtype: Optional[Union[str, torch.dtype]],
|
||||
dtype: Optional[Union[str, torch.dtype]],
|
||||
device: Union[str, torch.device],
|
||||
disable_mmap: bool = False,
|
||||
state_dict: Optional[Dict] = None,
|
||||
):
|
||||
clip_g_sd = None
|
||||
if clip_g_path:
|
||||
logger.info(f"Loading clip_g from {clip_g_path}...")
|
||||
clip_g_sd = load_safetensors(clip_g_path, device, disable_mmap)
|
||||
for key in list(clip_g_sd.keys()):
|
||||
clip_g_sd["transformer." + key] = clip_g_sd.pop(key)
|
||||
else:
|
||||
if state_dict is not None:
|
||||
if "text_encoders.clip_g.transformer.text_model.embeddings.position_embedding.weight" in state_dict:
|
||||
# found clip_g: remove prefix "text_encoders.clip_g."
|
||||
logger.info("clip_g is included in the checkpoint")
|
||||
@@ -124,34 +168,53 @@ def load_clip_g(
|
||||
for k in list(state_dict.keys()):
|
||||
if k.startswith(prefix):
|
||||
clip_g_sd[k[len(prefix) :]] = state_dict.pop(k)
|
||||
elif clip_g_path is None:
|
||||
logger.info("clip_g is not included in the checkpoint and clip_g_path is not provided")
|
||||
return None
|
||||
|
||||
# load clip_g
|
||||
logger.info("Building CLIP-G")
|
||||
config = 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():
|
||||
clip = CLIPTextModelWithProjection(config)
|
||||
|
||||
if clip_g_sd is None:
|
||||
clip_g = None
|
||||
else:
|
||||
logger.info("Building ClipG")
|
||||
clip_g = sd3_models.create_clip_g(device, clip_dtype, clip_g_sd)
|
||||
logger.info("Loading state dict...")
|
||||
info = clip_g.load_state_dict(clip_g_sd)
|
||||
logger.info(f"Loaded ClipG: {info}")
|
||||
clip_g.set_attn_mode(attn_mode)
|
||||
return clip_g
|
||||
logger.info(f"Loading state dict from {clip_g_path}")
|
||||
clip_g_sd = load_safetensors(clip_g_path, device=str(device), disable_mmap=disable_mmap, dtype=dtype)
|
||||
info = clip.load_state_dict(clip_g_sd, strict=False, assign=True)
|
||||
logger.info(f"Loaded CLIP-G: {info}")
|
||||
return clip
|
||||
|
||||
|
||||
def load_t5xxl(
|
||||
state_dict: Dict,
|
||||
t5xxl_path: Optional[str],
|
||||
attn_mode: str,
|
||||
dtype: Optional[Union[str, torch.dtype]],
|
||||
device: Union[str, torch.device],
|
||||
disable_mmap: bool = False,
|
||||
state_dict: Optional[Dict] = None,
|
||||
):
|
||||
t5xxl_sd = None
|
||||
if t5xxl_path:
|
||||
logger.info(f"Loading t5xxl from {t5xxl_path}...")
|
||||
t5xxl_sd = load_safetensors(t5xxl_path, device, disable_mmap)
|
||||
for key in list(t5xxl_sd.keys()):
|
||||
t5xxl_sd["transformer." + key] = t5xxl_sd.pop(key)
|
||||
else:
|
||||
if state_dict is not None:
|
||||
if "text_encoders.t5xxl.transformer.encoder.block.0.layer.0.SelfAttention.k.weight" in state_dict:
|
||||
# found t5xxl: remove prefix "text_encoders.t5xxl."
|
||||
logger.info("t5xxl is included in the checkpoint")
|
||||
@@ -160,29 +223,19 @@ def load_t5xxl(
|
||||
for k in list(state_dict.keys()):
|
||||
if k.startswith(prefix):
|
||||
t5xxl_sd[k[len(prefix) :]] = state_dict.pop(k)
|
||||
elif t5xxl_path is None:
|
||||
logger.info("t5xxl is not included in the checkpoint and t5xxl_path is not provided")
|
||||
return None
|
||||
|
||||
if t5xxl_sd is None:
|
||||
t5xxl = None
|
||||
else:
|
||||
logger.info("Building T5XXL")
|
||||
|
||||
# workaround for T5XXL model creation: create with fp16 takes too long TODO support virtual device
|
||||
t5xxl = sd3_models.create_t5xxl(device, torch.float32, t5xxl_sd)
|
||||
t5xxl.to(dtype=dtype)
|
||||
|
||||
logger.info("Loading state dict...")
|
||||
info = t5xxl.load_state_dict(t5xxl_sd)
|
||||
logger.info(f"Loaded T5XXL: {info}")
|
||||
t5xxl.set_attn_mode(attn_mode)
|
||||
return t5xxl
|
||||
return flux_utils_load_t5xxl(t5xxl_path, dtype, device, disable_mmap, state_dict=t5xxl_sd)
|
||||
|
||||
|
||||
def load_vae(
|
||||
state_dict: Dict,
|
||||
vae_path: Optional[str],
|
||||
vae_dtype: Optional[Union[str, torch.dtype]],
|
||||
device: Optional[Union[str, torch.device]],
|
||||
disable_mmap: bool = False,
|
||||
state_dict: Optional[Dict] = None,
|
||||
):
|
||||
vae_sd = {}
|
||||
if vae_path:
|
||||
@@ -197,299 +250,15 @@ def load_vae(
|
||||
vae_sd[k[len(vae_prefix) :]] = state_dict.pop(k)
|
||||
|
||||
logger.info("Building VAE")
|
||||
vae = sd3_models.SDVAE()
|
||||
vae = sd3_models.SDVAE(vae_dtype, device)
|
||||
logger.info("Loading state dict...")
|
||||
info = vae.load_state_dict(vae_sd)
|
||||
logger.info(f"Loaded VAE: {info}")
|
||||
vae.to(device=device, dtype=vae_dtype)
|
||||
vae.to(device=device, dtype=vae_dtype) # make sure it's in the right device and dtype
|
||||
return vae
|
||||
|
||||
|
||||
def load_models(
|
||||
ckpt_path: str,
|
||||
clip_l_path: str,
|
||||
clip_g_path: str,
|
||||
t5xxl_path: str,
|
||||
vae_path: str,
|
||||
attn_mode: str,
|
||||
device: Union[str, torch.device],
|
||||
weight_dtype: Optional[Union[str, torch.dtype]] = None,
|
||||
disable_mmap: bool = False,
|
||||
clip_dtype: Optional[Union[str, torch.dtype]] = None,
|
||||
t5xxl_device: Optional[Union[str, torch.device]] = None,
|
||||
t5xxl_dtype: Optional[Union[str, torch.dtype]] = None,
|
||||
vae_dtype: Optional[Union[str, torch.dtype]] = None,
|
||||
):
|
||||
"""
|
||||
Load SD3 models from checkpoint files.
|
||||
|
||||
Args:
|
||||
ckpt_path: Path to the SD3 checkpoint file.
|
||||
clip_l_path: Path to the clip_l checkpoint file.
|
||||
clip_g_path: Path to the clip_g checkpoint file.
|
||||
t5xxl_path: Path to the t5xxl checkpoint file.
|
||||
vae_path: Path to the VAE checkpoint file.
|
||||
attn_mode: Attention mode for MMDiT model.
|
||||
device: Device for MMDiT model.
|
||||
weight_dtype: Default dtype of weights for all models. This is weight dtype, so the model dtype may be different.
|
||||
disable_mmap: Disable memory mapping when loading state dict.
|
||||
clip_dtype: Dtype for Clip models, or None to use default dtype.
|
||||
t5xxl_device: Device for T5XXL model to load T5XXL in another device (eg. gpu). Default is None to use device.
|
||||
t5xxl_dtype: Dtype for T5XXL model, or None to use default dtype.
|
||||
vae_dtype: Dtype for VAE model, or None to use default dtype.
|
||||
|
||||
Returns:
|
||||
Tuple of MMDiT, ClipL, ClipG, T5XXL, and VAE models.
|
||||
"""
|
||||
|
||||
# In SD1/2 and SDXL, the model is created with empty weights and then loaded with state dict.
|
||||
# However, in SD3, Clip and T5XXL models are created with dtype, so we need to set dtype before loading state dict.
|
||||
# Therefore, we need clip_dtype and t5xxl_dtype.
|
||||
|
||||
def load_state_dict(path: str, dvc: Union[str, torch.device] = device):
|
||||
if disable_mmap:
|
||||
return safetensors.torch.load(open(path, "rb").read())
|
||||
else:
|
||||
try:
|
||||
return load_file(path, device=dvc)
|
||||
except:
|
||||
return load_file(path) # prevent device invalid Error
|
||||
|
||||
t5xxl_device = t5xxl_device or device
|
||||
clip_dtype = clip_dtype or weight_dtype or torch.float32
|
||||
t5xxl_dtype = t5xxl_dtype or weight_dtype or torch.float32
|
||||
vae_dtype = vae_dtype or weight_dtype or torch.float32
|
||||
|
||||
logger.info(f"Loading SD3 models from {ckpt_path}...")
|
||||
state_dict = load_state_dict(ckpt_path)
|
||||
|
||||
# load clip_l
|
||||
clip_l_sd = None
|
||||
if clip_l_path:
|
||||
logger.info(f"Loading clip_l from {clip_l_path}...")
|
||||
clip_l_sd = load_state_dict(clip_l_path)
|
||||
for key in list(clip_l_sd.keys()):
|
||||
clip_l_sd["transformer." + key] = clip_l_sd.pop(key)
|
||||
else:
|
||||
if "text_encoders.clip_l.transformer.text_model.embeddings.position_embedding.weight" in state_dict:
|
||||
# found clip_l: remove prefix "text_encoders.clip_l."
|
||||
logger.info("clip_l is included in the checkpoint")
|
||||
clip_l_sd = {}
|
||||
prefix = "text_encoders.clip_l."
|
||||
for k in list(state_dict.keys()):
|
||||
if k.startswith(prefix):
|
||||
clip_l_sd[k[len(prefix) :]] = state_dict.pop(k)
|
||||
|
||||
# load clip_g
|
||||
clip_g_sd = None
|
||||
if clip_g_path:
|
||||
logger.info(f"Loading clip_g from {clip_g_path}...")
|
||||
clip_g_sd = load_state_dict(clip_g_path)
|
||||
for key in list(clip_g_sd.keys()):
|
||||
clip_g_sd["transformer." + key] = clip_g_sd.pop(key)
|
||||
else:
|
||||
if "text_encoders.clip_g.transformer.text_model.embeddings.position_embedding.weight" in state_dict:
|
||||
# found clip_g: remove prefix "text_encoders.clip_g."
|
||||
logger.info("clip_g is included in the checkpoint")
|
||||
clip_g_sd = {}
|
||||
prefix = "text_encoders.clip_g."
|
||||
for k in list(state_dict.keys()):
|
||||
if k.startswith(prefix):
|
||||
clip_g_sd[k[len(prefix) :]] = state_dict.pop(k)
|
||||
|
||||
# load t5xxl
|
||||
t5xxl_sd = None
|
||||
if t5xxl_path:
|
||||
logger.info(f"Loading t5xxl from {t5xxl_path}...")
|
||||
t5xxl_sd = load_state_dict(t5xxl_path, t5xxl_device)
|
||||
for key in list(t5xxl_sd.keys()):
|
||||
t5xxl_sd["transformer." + key] = t5xxl_sd.pop(key)
|
||||
else:
|
||||
if "text_encoders.t5xxl.transformer.encoder.block.0.layer.0.SelfAttention.k.weight" in state_dict:
|
||||
# found t5xxl: remove prefix "text_encoders.t5xxl."
|
||||
logger.info("t5xxl is included in the checkpoint")
|
||||
t5xxl_sd = {}
|
||||
prefix = "text_encoders.t5xxl."
|
||||
for k in list(state_dict.keys()):
|
||||
if k.startswith(prefix):
|
||||
t5xxl_sd[k[len(prefix) :]] = state_dict.pop(k)
|
||||
|
||||
# MMDiT and VAE
|
||||
vae_sd = {}
|
||||
if vae_path:
|
||||
logger.info(f"Loading VAE from {vae_path}...")
|
||||
vae_sd = load_state_dict(vae_path)
|
||||
else:
|
||||
# remove prefix "first_stage_model."
|
||||
vae_sd = {}
|
||||
vae_prefix = "first_stage_model."
|
||||
for k in list(state_dict.keys()):
|
||||
if k.startswith(vae_prefix):
|
||||
vae_sd[k[len(vae_prefix) :]] = state_dict.pop(k)
|
||||
|
||||
mmdit_prefix = "model.diffusion_model."
|
||||
for k in list(state_dict.keys()):
|
||||
if k.startswith(mmdit_prefix):
|
||||
state_dict[k[len(mmdit_prefix) :]] = state_dict.pop(k)
|
||||
else:
|
||||
state_dict.pop(k) # remove other keys
|
||||
|
||||
# load MMDiT
|
||||
logger.info("Building MMDit")
|
||||
with init_empty_weights():
|
||||
mmdit = sd3_models.create_mmdit_sd3_medium_configs(attn_mode)
|
||||
|
||||
logger.info("Loading state dict...")
|
||||
info = load_state_dict_on_device(mmdit, state_dict, device, weight_dtype)
|
||||
logger.info(f"Loaded MMDiT: {info}")
|
||||
|
||||
# load ClipG and ClipL
|
||||
if clip_l_sd is None:
|
||||
clip_l = None
|
||||
else:
|
||||
logger.info("Building ClipL")
|
||||
clip_l = sd3_models.create_clip_l(device, clip_dtype, clip_l_sd)
|
||||
logger.info("Loading state dict...")
|
||||
info = clip_l.load_state_dict(clip_l_sd)
|
||||
logger.info(f"Loaded ClipL: {info}")
|
||||
clip_l.set_attn_mode(attn_mode)
|
||||
|
||||
if clip_g_sd is None:
|
||||
clip_g = None
|
||||
else:
|
||||
logger.info("Building ClipG")
|
||||
clip_g = sd3_models.create_clip_g(device, clip_dtype, clip_g_sd)
|
||||
logger.info("Loading state dict...")
|
||||
info = clip_g.load_state_dict(clip_g_sd)
|
||||
logger.info(f"Loaded ClipG: {info}")
|
||||
clip_g.set_attn_mode(attn_mode)
|
||||
|
||||
# load T5XXL
|
||||
if t5xxl_sd is None:
|
||||
t5xxl = None
|
||||
else:
|
||||
logger.info("Building T5XXL")
|
||||
t5xxl = sd3_models.create_t5xxl(t5xxl_device, t5xxl_dtype, t5xxl_sd)
|
||||
logger.info("Loading state dict...")
|
||||
info = t5xxl.load_state_dict(t5xxl_sd)
|
||||
logger.info(f"Loaded T5XXL: {info}")
|
||||
t5xxl.set_attn_mode(attn_mode)
|
||||
|
||||
# load VAE
|
||||
logger.info("Building VAE")
|
||||
vae = sd3_models.SDVAE()
|
||||
logger.info("Loading state dict...")
|
||||
info = vae.load_state_dict(vae_sd)
|
||||
logger.info(f"Loaded VAE: {info}")
|
||||
vae.to(device=device, dtype=vae_dtype)
|
||||
|
||||
return mmdit, clip_l, clip_g, t5xxl, vae
|
||||
|
||||
|
||||
# endregion
|
||||
# region utils
|
||||
|
||||
|
||||
def get_cond(
|
||||
prompt: str,
|
||||
tokenizer: sd3_models.SD3Tokenizer,
|
||||
clip_l: sd3_models.SDClipModel,
|
||||
clip_g: sd3_models.SDXLClipG,
|
||||
t5xxl: Optional[sd3_models.T5XXLModel] = None,
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
l_tokens, g_tokens, t5_tokens = tokenizer.tokenize_with_weights(prompt)
|
||||
print(t5_tokens)
|
||||
return get_cond_from_tokens(l_tokens, g_tokens, t5_tokens, clip_l, clip_g, t5xxl, device=device, dtype=dtype)
|
||||
|
||||
|
||||
def get_cond_from_tokens(
|
||||
l_tokens,
|
||||
g_tokens,
|
||||
t5_tokens,
|
||||
clip_l: sd3_models.SDClipModel,
|
||||
clip_g: sd3_models.SDXLClipG,
|
||||
t5xxl: Optional[sd3_models.T5XXLModel] = None,
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
l_out, l_pooled = clip_l.encode_token_weights(l_tokens)
|
||||
g_out, g_pooled = clip_g.encode_token_weights(g_tokens)
|
||||
lg_out = torch.cat([l_out, g_out], dim=-1)
|
||||
lg_out = torch.nn.functional.pad(lg_out, (0, 4096 - lg_out.shape[-1]))
|
||||
if device is not None:
|
||||
lg_out = lg_out.to(device=device)
|
||||
l_pooled = l_pooled.to(device=device)
|
||||
g_pooled = g_pooled.to(device=device)
|
||||
if dtype is not None:
|
||||
lg_out = lg_out.to(dtype=dtype)
|
||||
l_pooled = l_pooled.to(dtype=dtype)
|
||||
g_pooled = g_pooled.to(dtype=dtype)
|
||||
|
||||
# t5xxl may be in another device (eg. cpu)
|
||||
if t5_tokens is None:
|
||||
t5_out = torch.zeros((lg_out.shape[0], 77, 4096), device=lg_out.device, dtype=lg_out.dtype)
|
||||
else:
|
||||
t5_out, _ = t5xxl.encode_token_weights(t5_tokens) # t5_out is [1, 77, 4096], t5_pooled is None
|
||||
if device is not None:
|
||||
t5_out = t5_out.to(device=device)
|
||||
if dtype is not None:
|
||||
t5_out = t5_out.to(dtype=dtype)
|
||||
|
||||
# return torch.cat([lg_out, t5_out], dim=-2), torch.cat((l_pooled, g_pooled), dim=-1)
|
||||
return lg_out, t5_out, torch.cat((l_pooled, g_pooled), dim=-1)
|
||||
|
||||
|
||||
# used if other sd3 models is available
|
||||
r"""
|
||||
def get_sd3_configs(state_dict: Dict):
|
||||
# Important configuration values can be quickly determined by checking shapes in the source file
|
||||
# Some of these will vary between models (eg 2B vs 8B primarily differ in their depth, but also other details change)
|
||||
# prefix = "model.diffusion_model."
|
||||
prefix = ""
|
||||
|
||||
patch_size = state_dict[prefix + "x_embedder.proj.weight"].shape[2]
|
||||
depth = state_dict[prefix + "x_embedder.proj.weight"].shape[0] // 64
|
||||
num_patches = state_dict[prefix + "pos_embed"].shape[1]
|
||||
pos_embed_max_size = round(math.sqrt(num_patches))
|
||||
adm_in_channels = state_dict[prefix + "y_embedder.mlp.0.weight"].shape[1]
|
||||
context_shape = state_dict[prefix + "context_embedder.weight"].shape
|
||||
context_embedder_config = {
|
||||
"target": "torch.nn.Linear",
|
||||
"params": {"in_features": context_shape[1], "out_features": context_shape[0]},
|
||||
}
|
||||
return {
|
||||
"patch_size": patch_size,
|
||||
"depth": depth,
|
||||
"num_patches": num_patches,
|
||||
"pos_embed_max_size": pos_embed_max_size,
|
||||
"adm_in_channels": adm_in_channels,
|
||||
"context_embedder": context_embedder_config,
|
||||
}
|
||||
|
||||
|
||||
def create_mmdit_from_sd3_checkpoint(state_dict: Dict, attn_mode: str = "xformers"):
|
||||
""
|
||||
Doesn't load state dict.
|
||||
""
|
||||
sd3_configs = get_sd3_configs(state_dict)
|
||||
|
||||
mmdit = sd3_models.MMDiT(
|
||||
input_size=None,
|
||||
pos_embed_max_size=sd3_configs["pos_embed_max_size"],
|
||||
patch_size=sd3_configs["patch_size"],
|
||||
in_channels=16,
|
||||
adm_in_channels=sd3_configs["adm_in_channels"],
|
||||
depth=sd3_configs["depth"],
|
||||
mlp_ratio=4,
|
||||
qk_norm=None,
|
||||
num_patches=sd3_configs["num_patches"],
|
||||
context_size=4096,
|
||||
attn_mode=attn_mode,
|
||||
)
|
||||
return mmdit
|
||||
"""
|
||||
|
||||
|
||||
class ModelSamplingDiscreteFlow:
|
||||
@@ -525,6 +294,3 @@ class ModelSamplingDiscreteFlow:
|
||||
# assert max_denoise is False, "max_denoise not implemented"
|
||||
# max_denoise is always True, I'm not sure why it's there
|
||||
return sigma * noise + (1.0 - sigma) * latent_image
|
||||
|
||||
|
||||
# endregion
|
||||
|
||||
@@ -13,12 +13,20 @@ from tqdm import tqdm
|
||||
from transformers import CLIPFeatureExtractor, CLIPTextModel, CLIPTokenizer
|
||||
|
||||
from diffusers import SchedulerMixin, StableDiffusionPipeline
|
||||
from diffusers.models import AutoencoderKL, UNet2DConditionModel
|
||||
from diffusers.pipelines.stable_diffusion import StableDiffusionPipelineOutput, StableDiffusionSafetyChecker
|
||||
from diffusers.models import AutoencoderKL
|
||||
from diffusers.pipelines.stable_diffusion import StableDiffusionSafetyChecker
|
||||
from diffusers.utils import logging
|
||||
from PIL import Image
|
||||
|
||||
from library import sdxl_model_util, sdxl_train_util, train_util
|
||||
from . import (
|
||||
sdxl_model_util,
|
||||
sdxl_train_util,
|
||||
strategy_base,
|
||||
strategy_sdxl,
|
||||
train_util,
|
||||
sdxl_original_unet,
|
||||
sdxl_original_control_net,
|
||||
)
|
||||
|
||||
|
||||
try:
|
||||
@@ -40,6 +48,8 @@ except ImportError:
|
||||
"lanczos": PIL.Image.LANCZOS,
|
||||
"nearest": PIL.Image.NEAREST,
|
||||
}
|
||||
|
||||
from comfy.utils import ProgressBar
|
||||
# ------------------------------------------------------------------------------
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
@@ -537,7 +547,7 @@ class SdxlStableDiffusionLongPromptWeightingPipeline:
|
||||
vae: AutoencoderKL,
|
||||
text_encoder: List[CLIPTextModel],
|
||||
tokenizer: List[CLIPTokenizer],
|
||||
unet: UNet2DConditionModel,
|
||||
unet: Union[sdxl_original_unet.SdxlUNet2DConditionModel, sdxl_original_control_net.SdxlControlledUNet],
|
||||
scheduler: SchedulerMixin,
|
||||
# clip_skip: int,
|
||||
safety_checker: StableDiffusionSafetyChecker,
|
||||
@@ -594,74 +604,6 @@ class SdxlStableDiffusionLongPromptWeightingPipeline:
|
||||
return torch.device(module._hf_hook.execution_device)
|
||||
return self.device
|
||||
|
||||
def _encode_prompt(
|
||||
self,
|
||||
prompt,
|
||||
device,
|
||||
num_images_per_prompt,
|
||||
do_classifier_free_guidance,
|
||||
negative_prompt,
|
||||
max_embeddings_multiples,
|
||||
is_sdxl_text_encoder2,
|
||||
):
|
||||
r"""
|
||||
Encodes the prompt into text encoder hidden states.
|
||||
|
||||
Args:
|
||||
prompt (`str` or `list(int)`):
|
||||
prompt to be encoded
|
||||
device: (`torch.device`):
|
||||
torch device
|
||||
num_images_per_prompt (`int`):
|
||||
number of images that should be generated per prompt
|
||||
do_classifier_free_guidance (`bool`):
|
||||
whether to use classifier free guidance or not
|
||||
negative_prompt (`str` or `List[str]`):
|
||||
The prompt or prompts not to guide the image generation. Ignored when not using guidance (i.e., ignored
|
||||
if `guidance_scale` is less than `1`).
|
||||
max_embeddings_multiples (`int`, *optional*, defaults to `3`):
|
||||
The max multiple length of prompt embeddings compared to the max output length of text encoder.
|
||||
"""
|
||||
batch_size = len(prompt) if isinstance(prompt, list) else 1
|
||||
|
||||
if negative_prompt is None:
|
||||
negative_prompt = [""] * batch_size
|
||||
elif isinstance(negative_prompt, str):
|
||||
negative_prompt = [negative_prompt] * batch_size
|
||||
if batch_size != len(negative_prompt):
|
||||
raise ValueError(
|
||||
f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:"
|
||||
f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches"
|
||||
" the batch size of `prompt`."
|
||||
)
|
||||
|
||||
text_embeddings, text_pool, uncond_embeddings, uncond_pool = get_weighted_text_embeddings(
|
||||
pipe=self,
|
||||
prompt=prompt,
|
||||
uncond_prompt=negative_prompt if do_classifier_free_guidance else None,
|
||||
max_embeddings_multiples=max_embeddings_multiples,
|
||||
clip_skip=self.clip_skip,
|
||||
is_sdxl_text_encoder2=is_sdxl_text_encoder2,
|
||||
)
|
||||
bs_embed, seq_len, _ = text_embeddings.shape
|
||||
text_embeddings = text_embeddings.repeat(1, num_images_per_prompt, 1) # ??
|
||||
text_embeddings = text_embeddings.view(bs_embed * num_images_per_prompt, seq_len, -1)
|
||||
if text_pool is not None:
|
||||
text_pool = text_pool.repeat(1, num_images_per_prompt)
|
||||
text_pool = text_pool.view(bs_embed * num_images_per_prompt, -1)
|
||||
|
||||
if do_classifier_free_guidance:
|
||||
bs_embed, seq_len, _ = uncond_embeddings.shape
|
||||
uncond_embeddings = uncond_embeddings.repeat(1, num_images_per_prompt, 1)
|
||||
uncond_embeddings = uncond_embeddings.view(bs_embed * num_images_per_prompt, seq_len, -1)
|
||||
if uncond_pool is not None:
|
||||
uncond_pool = uncond_pool.repeat(1, num_images_per_prompt)
|
||||
uncond_pool = uncond_pool.view(bs_embed * num_images_per_prompt, -1)
|
||||
|
||||
return text_embeddings, text_pool, uncond_embeddings, uncond_pool
|
||||
|
||||
return text_embeddings, text_pool, None, None
|
||||
|
||||
def check_inputs(self, prompt, height, width, strength, callback_steps):
|
||||
if not isinstance(prompt, str) and not isinstance(prompt, list):
|
||||
raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
|
||||
@@ -711,11 +653,11 @@ class SdxlStableDiffusionLongPromptWeightingPipeline:
|
||||
# self.vae.set_use_memory_efficient_attention_xformers(False)
|
||||
# image = self.vae.decode(latents.to("cpu")).sample
|
||||
|
||||
image = self.vae.decode(latents.to(self.vae.dtype)).sample
|
||||
image = (image / 2 + 0.5).clamp(0, 1)
|
||||
tensors = self.vae.decode(latents.to(self.vae.dtype)).sample
|
||||
image = (tensors / 2 + 0.5).clamp(0, 1)
|
||||
# we always cast to float32 as this does not cause significant overhead and is compatible with bfloat16
|
||||
image = image.cpu().permute(0, 2, 3, 1).float().numpy()
|
||||
return image
|
||||
return image, tensors
|
||||
|
||||
def prepare_extra_step_kwargs(self, generator, eta):
|
||||
# prepare extra kwargs for the scheduler step, since not all schedulers have the same signature
|
||||
@@ -792,7 +734,7 @@ class SdxlStableDiffusionLongPromptWeightingPipeline:
|
||||
max_embeddings_multiples: Optional[int] = 3,
|
||||
output_type: Optional[str] = "pil",
|
||||
return_dict: bool = True,
|
||||
controlnet=None,
|
||||
controlnet: sdxl_original_control_net.SdxlControlNet = None,
|
||||
controlnet_image=None,
|
||||
callback: Optional[Callable[[int, int, torch.FloatTensor], None]] = None,
|
||||
is_cancelled_callback: Optional[Callable[[], bool]] = None,
|
||||
@@ -896,32 +838,24 @@ class SdxlStableDiffusionLongPromptWeightingPipeline:
|
||||
do_classifier_free_guidance = guidance_scale > 1.0
|
||||
|
||||
# 3. Encode input prompt
|
||||
# 実装を簡単にするためにtokenzer/text encoderを切り替えて二回呼び出す
|
||||
# To simplify the implementation, switch the tokenzer/text encoder and call it twice
|
||||
text_embeddings_list = []
|
||||
text_pool = None
|
||||
uncond_embeddings_list = []
|
||||
uncond_pool = None
|
||||
for i in range(len(self.tokenizers)):
|
||||
self.tokenizer = self.tokenizers[i]
|
||||
self.text_encoder = self.text_encoders[i]
|
||||
tokenize_strategy: strategy_sdxl.SdxlTokenizeStrategy = strategy_base.TokenizeStrategy.get_strategy()
|
||||
encoding_strategy: strategy_sdxl.SdxlTextEncodingStrategy = strategy_base.TextEncodingStrategy.get_strategy()
|
||||
|
||||
text_embeddings, tp1, uncond_embeddings, up1 = self._encode_prompt(
|
||||
prompt,
|
||||
device,
|
||||
num_images_per_prompt,
|
||||
do_classifier_free_guidance,
|
||||
negative_prompt,
|
||||
max_embeddings_multiples,
|
||||
is_sdxl_text_encoder2=i == 1,
|
||||
text_input_ids, text_weights = tokenize_strategy.tokenize_with_weights(prompt)
|
||||
hidden_states_1, hidden_states_2, text_pool = encoding_strategy.encode_tokens_with_weights(
|
||||
tokenize_strategy, self.text_encoders, text_input_ids, text_weights
|
||||
)
|
||||
text_embeddings_list.append(text_embeddings)
|
||||
uncond_embeddings_list.append(uncond_embeddings)
|
||||
text_embeddings = torch.cat([hidden_states_1, hidden_states_2], dim=-1)
|
||||
|
||||
if tp1 is not None:
|
||||
text_pool = tp1
|
||||
if up1 is not None:
|
||||
uncond_pool = up1
|
||||
if do_classifier_free_guidance:
|
||||
input_ids, weights = tokenize_strategy.tokenize_with_weights(negative_prompt or "")
|
||||
hidden_states_1, hidden_states_2, uncond_pool = encoding_strategy.encode_tokens_with_weights(
|
||||
tokenize_strategy, self.text_encoders, input_ids, weights
|
||||
)
|
||||
uncond_embeddings = torch.cat([hidden_states_1, hidden_states_2], dim=-1)
|
||||
else:
|
||||
uncond_embeddings = None
|
||||
uncond_pool = None
|
||||
|
||||
unet_dtype = self.unet.dtype
|
||||
dtype = unet_dtype
|
||||
@@ -970,45 +904,38 @@ class SdxlStableDiffusionLongPromptWeightingPipeline:
|
||||
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
|
||||
|
||||
# create size embs and concat embeddings for SDXL
|
||||
orig_size = torch.tensor([height, width]).repeat(batch_size * num_images_per_prompt, 1).to(dtype)
|
||||
orig_size = torch.tensor([height, width]).repeat(batch_size * num_images_per_prompt, 1).to(device, dtype)
|
||||
crop_size = torch.zeros_like(orig_size)
|
||||
target_size = orig_size
|
||||
embs = sdxl_train_util.get_size_embeddings(orig_size, crop_size, target_size, device).to(dtype)
|
||||
embs = sdxl_train_util.get_size_embeddings(orig_size, crop_size, target_size, device).to(device, dtype)
|
||||
|
||||
# make conditionings
|
||||
text_pool = text_pool.to(device, dtype)
|
||||
if do_classifier_free_guidance:
|
||||
text_embeddings = torch.cat(text_embeddings_list, dim=2)
|
||||
uncond_embeddings = torch.cat(uncond_embeddings_list, dim=2)
|
||||
text_embedding = torch.cat([uncond_embeddings, text_embeddings]).to(dtype)
|
||||
text_embedding = torch.cat([uncond_embeddings, text_embeddings]).to(device, dtype)
|
||||
|
||||
cond_vector = torch.cat([text_pool, embs], dim=1)
|
||||
uncond_vector = torch.cat([uncond_pool, embs], dim=1)
|
||||
vector_embedding = torch.cat([uncond_vector, cond_vector]).to(dtype)
|
||||
uncond_pool = uncond_pool.to(device, dtype)
|
||||
cond_vector = torch.cat([text_pool, embs], dim=1).to(dtype)
|
||||
uncond_vector = torch.cat([uncond_pool, embs], dim=1).to(dtype)
|
||||
vector_embedding = torch.cat([uncond_vector, cond_vector])
|
||||
else:
|
||||
text_embedding = torch.cat(text_embeddings_list, dim=2).to(dtype)
|
||||
vector_embedding = torch.cat([text_pool, embs], dim=1).to(dtype)
|
||||
text_embedding = text_embeddings.to(device, dtype)
|
||||
vector_embedding = torch.cat([text_pool, embs], dim=1)
|
||||
|
||||
# 8. Denoising loop
|
||||
comfy_pbar = ProgressBar(total=len(timesteps))
|
||||
for i, t in enumerate(self.progress_bar(timesteps)):
|
||||
# expand the latents if we are doing classifier free guidance
|
||||
latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents
|
||||
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
|
||||
|
||||
unet_additional_args = {}
|
||||
if controlnet is not None:
|
||||
down_block_res_samples, mid_block_res_sample = controlnet(
|
||||
latent_model_input,
|
||||
t,
|
||||
encoder_hidden_states=text_embeddings,
|
||||
controlnet_cond=controlnet_image,
|
||||
conditioning_scale=1.0,
|
||||
guess_mode=False,
|
||||
return_dict=False,
|
||||
)
|
||||
unet_additional_args["down_block_additional_residuals"] = down_block_res_samples
|
||||
unet_additional_args["mid_block_additional_residual"] = mid_block_res_sample
|
||||
# FIXME SD1 ControlNet is not working
|
||||
|
||||
# predict the noise residual
|
||||
if controlnet is not None:
|
||||
input_resi_add, mid_add = controlnet(latent_model_input, t, text_embedding, vector_embedding, controlnet_image)
|
||||
noise_pred = self.unet(latent_model_input, t, text_embedding, vector_embedding, input_resi_add, mid_add)
|
||||
else:
|
||||
noise_pred = self.unet(latent_model_input, t, text_embedding, vector_embedding)
|
||||
noise_pred = noise_pred.to(dtype) # U-Net changes dtype in LoRA training
|
||||
|
||||
@@ -1031,15 +958,16 @@ class SdxlStableDiffusionLongPromptWeightingPipeline:
|
||||
callback(i, t, latents)
|
||||
if is_cancelled_callback is not None and is_cancelled_callback():
|
||||
return None
|
||||
comfy_pbar.update(1)
|
||||
|
||||
self.unet.to(unet_dtype)
|
||||
return latents
|
||||
|
||||
def latents_to_image(self, latents):
|
||||
# 9. Post-processing
|
||||
image = self.decode_latents(latents.to(self.vae.dtype))
|
||||
image, tensors = self.decode_latents(latents.to(self.vae.dtype))
|
||||
image = self.numpy_to_pil(image)
|
||||
return image
|
||||
return image, tensors
|
||||
|
||||
# copy from pil_utils.py
|
||||
def numpy_to_pil(self, images: np.ndarray) -> Image.Image:
|
||||
|
||||
@@ -6,8 +6,8 @@ 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 library import model_util
|
||||
from library import sdxl_original_unet
|
||||
from . import model_util
|
||||
from . import sdxl_original_unet
|
||||
from .utils import setup_logging
|
||||
|
||||
setup_logging()
|
||||
|
||||
@@ -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")
|
||||
@@ -1156,9 +1156,9 @@ class InferSdxlUNet2DConditionModel:
|
||||
self.ds_timesteps_2 = ds_timesteps_2 if ds_timesteps_2 is not None else 1000
|
||||
self.ds_ratio = ds_ratio
|
||||
|
||||
def forward(self, x, timesteps=None, context=None, y=None, **kwargs):
|
||||
def forward(self, x, timesteps=None, context=None, y=None, input_resi_add=None, mid_add=None, **kwargs):
|
||||
r"""
|
||||
current implementation is a copy of `SdxlUNet2DConditionModel.forward()` with Deep Shrink.
|
||||
current implementation is a copy of `SdxlUNet2DConditionModel.forward()` with Deep Shrink and ControlNet.
|
||||
"""
|
||||
_self = self.delegate
|
||||
|
||||
@@ -1209,6 +1209,8 @@ class InferSdxlUNet2DConditionModel:
|
||||
hs.append(h)
|
||||
|
||||
h = call_module(_self.middle_block, h, emb, context)
|
||||
if mid_add is not None:
|
||||
h = h + mid_add
|
||||
|
||||
for module in _self.output_blocks:
|
||||
# Deep Shrink
|
||||
@@ -1217,7 +1219,11 @@ class InferSdxlUNet2DConditionModel:
|
||||
# print("upsample", h.shape, hs[-1].shape)
|
||||
h = resize_like(h, hs[-1])
|
||||
|
||||
h = torch.cat([h, hs.pop()], dim=1)
|
||||
resi = hs.pop()
|
||||
if input_resi_add is not None:
|
||||
resi = resi + input_resi_add.pop()
|
||||
|
||||
h = torch.cat([h, resi], dim=1)
|
||||
h = call_module(module, h, emb, context)
|
||||
|
||||
# Deep Shrink: in case of depth 0
|
||||
|
||||
@@ -4,7 +4,7 @@ import os
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
from .library.device_utils import init_ipex, clean_memory_on_device
|
||||
from .device_utils import init_ipex, clean_memory_on_device
|
||||
|
||||
init_ipex()
|
||||
|
||||
@@ -364,9 +364,9 @@ def verify_sdxl_training_args(args: argparse.Namespace, supportTextEncoderCachin
|
||||
# )
|
||||
# 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を有効にすることはできません"
|
||||
# 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:
|
||||
|
||||
+213
-5
@@ -1,11 +1,12 @@
|
||||
# base class for platform strategies. this file defines the interface for strategies
|
||||
|
||||
import os
|
||||
import re
|
||||
from typing import Any, List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from transformers import CLIPTokenizer
|
||||
from transformers import CLIPTokenizer, CLIPTextModel, CLIPTextModelWithProjection
|
||||
|
||||
|
||||
# TODO remove circular import by moving ImageInfo to a separate file
|
||||
@@ -22,8 +23,28 @@ logger = logging.getLogger(__name__)
|
||||
class TokenizeStrategy:
|
||||
_strategy = None # strategy instance: actual strategy class
|
||||
|
||||
_re_attention = re.compile(
|
||||
r"""\\\(|
|
||||
\\\)|
|
||||
\\\[|
|
||||
\\]|
|
||||
\\\\|
|
||||
\\|
|
||||
\(|
|
||||
\[|
|
||||
:([+-]?[.\d]+)\)|
|
||||
\)|
|
||||
]|
|
||||
[^\\()\[\]:]+|
|
||||
:
|
||||
""",
|
||||
re.X,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def set_strategy(cls, strategy):
|
||||
#if cls._strategy is not None:
|
||||
# raise RuntimeError(f"Internal error. {cls.__name__} strategy is already set")
|
||||
cls._strategy = strategy
|
||||
|
||||
@classmethod
|
||||
@@ -52,7 +73,154 @@ class TokenizeStrategy:
|
||||
def tokenize(self, text: Union[str, List[str]]) -> List[torch.Tensor]:
|
||||
raise NotImplementedError
|
||||
|
||||
def _get_input_ids(self, tokenizer: CLIPTokenizer, text: str, max_length: Optional[int] = None) -> torch.Tensor:
|
||||
def tokenize_with_weights(self, text: Union[str, List[str]]) -> Tuple[List[torch.Tensor], List[torch.Tensor]]:
|
||||
"""
|
||||
returns: [tokens1, tokens2, ...], [weights1, weights2, ...]
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def _get_weighted_input_ids(
|
||||
self, tokenizer: CLIPTokenizer, text: str, max_length: Optional[int] = None
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
max_length includes starting and ending tokens.
|
||||
"""
|
||||
|
||||
def parse_prompt_attention(text):
|
||||
"""
|
||||
Parses a string with attention tokens and returns a list of pairs: text and its associated weight.
|
||||
Accepted tokens are:
|
||||
(abc) - increases attention to abc by a multiplier of 1.1
|
||||
(abc:3.12) - increases attention to abc by a multiplier of 3.12
|
||||
[abc] - decreases attention to abc by a multiplier of 1.1
|
||||
\( - literal character '('
|
||||
\[ - literal character '['
|
||||
\) - literal character ')'
|
||||
\] - literal character ']'
|
||||
\\ - literal character '\'
|
||||
anything else - just text
|
||||
>>> parse_prompt_attention('normal text')
|
||||
[['normal text', 1.0]]
|
||||
>>> parse_prompt_attention('an (important) word')
|
||||
[['an ', 1.0], ['important', 1.1], [' word', 1.0]]
|
||||
>>> parse_prompt_attention('(unbalanced')
|
||||
[['unbalanced', 1.1]]
|
||||
>>> parse_prompt_attention('\(literal\]')
|
||||
[['(literal]', 1.0]]
|
||||
>>> parse_prompt_attention('(unnecessary)(parens)')
|
||||
[['unnecessaryparens', 1.1]]
|
||||
>>> parse_prompt_attention('a (((house:1.3)) [on] a (hill:0.5), sun, (((sky))).')
|
||||
[['a ', 1.0],
|
||||
['house', 1.5730000000000004],
|
||||
[' ', 1.1],
|
||||
['on', 1.0],
|
||||
[' a ', 1.1],
|
||||
['hill', 0.55],
|
||||
[', sun, ', 1.1],
|
||||
['sky', 1.4641000000000006],
|
||||
['.', 1.1]]
|
||||
"""
|
||||
|
||||
res = []
|
||||
round_brackets = []
|
||||
square_brackets = []
|
||||
|
||||
round_bracket_multiplier = 1.1
|
||||
square_bracket_multiplier = 1 / 1.1
|
||||
|
||||
def multiply_range(start_position, multiplier):
|
||||
for p in range(start_position, len(res)):
|
||||
res[p][1] *= multiplier
|
||||
|
||||
for m in TokenizeStrategy._re_attention.finditer(text):
|
||||
text = m.group(0)
|
||||
weight = m.group(1)
|
||||
|
||||
if text.startswith("\\"):
|
||||
res.append([text[1:], 1.0])
|
||||
elif text == "(":
|
||||
round_brackets.append(len(res))
|
||||
elif text == "[":
|
||||
square_brackets.append(len(res))
|
||||
elif weight is not None and len(round_brackets) > 0:
|
||||
multiply_range(round_brackets.pop(), float(weight))
|
||||
elif text == ")" and len(round_brackets) > 0:
|
||||
multiply_range(round_brackets.pop(), round_bracket_multiplier)
|
||||
elif text == "]" and len(square_brackets) > 0:
|
||||
multiply_range(square_brackets.pop(), square_bracket_multiplier)
|
||||
else:
|
||||
res.append([text, 1.0])
|
||||
|
||||
for pos in round_brackets:
|
||||
multiply_range(pos, round_bracket_multiplier)
|
||||
|
||||
for pos in square_brackets:
|
||||
multiply_range(pos, square_bracket_multiplier)
|
||||
|
||||
if len(res) == 0:
|
||||
res = [["", 1.0]]
|
||||
|
||||
# merge runs of identical weights
|
||||
i = 0
|
||||
while i + 1 < len(res):
|
||||
if res[i][1] == res[i + 1][1]:
|
||||
res[i][0] += res[i + 1][0]
|
||||
res.pop(i + 1)
|
||||
else:
|
||||
i += 1
|
||||
|
||||
return res
|
||||
|
||||
def get_prompts_with_weights(text: str, max_length: int):
|
||||
r"""
|
||||
Tokenize a list of prompts and return its tokens with weights of each token. max_length does not include starting and ending token.
|
||||
|
||||
No padding, starting or ending token is included.
|
||||
"""
|
||||
truncated = False
|
||||
|
||||
texts_and_weights = parse_prompt_attention(text)
|
||||
tokens = []
|
||||
weights = []
|
||||
for word, weight in texts_and_weights:
|
||||
# tokenize and discard the starting and the ending token
|
||||
token = tokenizer(word).input_ids[1:-1]
|
||||
tokens += token
|
||||
# copy the weight by length of token
|
||||
weights += [weight] * len(token)
|
||||
# stop if the text is too long (longer than truncation limit)
|
||||
if len(tokens) > max_length:
|
||||
truncated = True
|
||||
break
|
||||
# truncate
|
||||
if len(tokens) > max_length:
|
||||
truncated = True
|
||||
tokens = tokens[:max_length]
|
||||
weights = weights[:max_length]
|
||||
if truncated:
|
||||
logger.warning("Prompt was truncated. Try to shorten the prompt or increase max_embeddings_multiples")
|
||||
return tokens, weights
|
||||
|
||||
def pad_tokens_and_weights(tokens, weights, max_length, bos, eos, pad):
|
||||
r"""
|
||||
Pad the tokens (with starting and ending tokens) and weights (with 1.0) to max_length.
|
||||
"""
|
||||
tokens = [bos] + tokens + [eos] + [pad] * (max_length - 2 - len(tokens))
|
||||
weights = [1.0] + weights + [1.0] * (max_length - 1 - len(weights))
|
||||
return tokens, weights
|
||||
|
||||
if max_length is None:
|
||||
max_length = tokenizer.model_max_length
|
||||
|
||||
tokens, weights = get_prompts_with_weights(text, max_length - 2)
|
||||
tokens, weights = pad_tokens_and_weights(
|
||||
tokens, weights, max_length, tokenizer.bos_token_id, tokenizer.eos_token_id, tokenizer.pad_token_id
|
||||
)
|
||||
return torch.tensor(tokens).unsqueeze(0), torch.tensor(weights).unsqueeze(0)
|
||||
|
||||
def _get_input_ids(
|
||||
self, tokenizer: CLIPTokenizer, text: str, max_length: Optional[int] = None, weighted: bool = False
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
for SD1.5/2.0/SDXL
|
||||
TODO support batch input
|
||||
@@ -60,6 +228,9 @@ class TokenizeStrategy:
|
||||
if max_length is None:
|
||||
max_length = tokenizer.model_max_length - 2
|
||||
|
||||
if weighted:
|
||||
input_ids, weights = self._get_weighted_input_ids(tokenizer, text, max_length)
|
||||
else:
|
||||
input_ids = tokenizer(text, padding="max_length", truncation=True, max_length=max_length, return_tensors="pt").input_ids
|
||||
|
||||
if max_length > tokenizer.model_max_length:
|
||||
@@ -99,6 +270,17 @@ class TokenizeStrategy:
|
||||
iids_list.append(ids_chunk)
|
||||
|
||||
input_ids = torch.stack(iids_list) # 3,77
|
||||
|
||||
if weighted:
|
||||
weights = weights.squeeze(0)
|
||||
new_weights = torch.ones(input_ids.shape)
|
||||
for i in range(1, max_length - tokenizer.model_max_length + 2, tokenizer.model_max_length - 2):
|
||||
b = i // (tokenizer.model_max_length - 2)
|
||||
new_weights[b, 1 : 1 + tokenizer.model_max_length - 2] = weights[i : i + tokenizer.model_max_length - 2]
|
||||
weights = new_weights
|
||||
|
||||
if weighted:
|
||||
return input_ids, weights
|
||||
return input_ids
|
||||
|
||||
|
||||
@@ -125,17 +307,34 @@ class TextEncodingStrategy:
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def encode_tokens_with_weights(
|
||||
self, tokenize_strategy: TokenizeStrategy, models: List[Any], tokens: List[torch.Tensor], weights: List[torch.Tensor]
|
||||
) -> List[torch.Tensor]:
|
||||
"""
|
||||
Encode tokens into embeddings and outputs.
|
||||
:param tokens: list of token tensors for each TextModel
|
||||
:param weights: list of weight tensors for each TextModel
|
||||
:return: list of output embeddings for each architecture
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class TextEncoderOutputsCachingStrategy:
|
||||
_strategy = None # strategy instance: actual strategy class
|
||||
|
||||
def __init__(
|
||||
self, cache_to_disk: bool, batch_size: int, skip_disk_cache_validity_check: bool, is_partial: bool = False
|
||||
self,
|
||||
cache_to_disk: bool,
|
||||
batch_size: Optional[int],
|
||||
skip_disk_cache_validity_check: bool,
|
||||
is_partial: bool = False,
|
||||
is_weighted: bool = False,
|
||||
) -> None:
|
||||
self._cache_to_disk = cache_to_disk
|
||||
self._batch_size = batch_size
|
||||
self.skip_disk_cache_validity_check = skip_disk_cache_validity_check
|
||||
self._is_partial = is_partial
|
||||
self._is_weighted = is_weighted
|
||||
|
||||
@classmethod
|
||||
def set_strategy(cls, strategy):
|
||||
@@ -159,6 +358,10 @@ class TextEncoderOutputsCachingStrategy:
|
||||
def is_partial(self):
|
||||
return self._is_partial
|
||||
|
||||
@property
|
||||
def is_weighted(self):
|
||||
return self._is_weighted
|
||||
|
||||
def get_outputs_npz_path(self, image_abs_path: str) -> str:
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -202,9 +405,14 @@ class LatentsCachingStrategy:
|
||||
def batch_size(self):
|
||||
return self._batch_size
|
||||
|
||||
def get_image_size_from_disk_cache_path(self, absolute_path: str) -> Tuple[Optional[int], Optional[int]]:
|
||||
@property
|
||||
def cache_suffix(self):
|
||||
raise NotImplementedError
|
||||
|
||||
def get_image_size_from_disk_cache_path(self, absolute_path: str, npz_path: str) -> Tuple[Optional[int], Optional[int]]:
|
||||
w, h = os.path.splitext(npz_path)[0].split("_")[-2].split("x")
|
||||
return int(w), int(h)
|
||||
|
||||
def get_latents_npz_path(self, absolute_path: str, image_size: Tuple[int, int]) -> str:
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -310,7 +518,7 @@ class LatentsCachingStrategy:
|
||||
self, npz_path: str, bucket_reso: Tuple[int, int]
|
||||
) -> Tuple[Optional[np.ndarray], Optional[List[int]], Optional[List[int]], Optional[np.ndarray], Optional[np.ndarray]]:
|
||||
"""
|
||||
for SD/SDXL/SD3.0
|
||||
for SD/SDXL
|
||||
"""
|
||||
return self._default_load_latents_from_disk(None, npz_path, bucket_reso)
|
||||
|
||||
|
||||
@@ -6,6 +6,7 @@ import numpy as np
|
||||
from transformers import CLIPTokenizer, T5TokenizerFast
|
||||
|
||||
from . import train_util
|
||||
from .flux_utils import get_t5xxl_actual_dtype
|
||||
from .strategy_base import LatentsCachingStrategy, TextEncodingStrategy, TokenizeStrategy, TextEncoderOutputsCachingStrategy
|
||||
|
||||
from .utils import setup_logging
|
||||
@@ -99,6 +100,8 @@ class FluxTextEncoderOutputsCachingStrategy(TextEncoderOutputsCachingStrategy):
|
||||
super().__init__(cache_to_disk, batch_size, skip_disk_cache_validity_check, is_partial)
|
||||
self.apply_t5_attn_mask = apply_t5_attn_mask
|
||||
|
||||
self.warn_fp8_weights = False
|
||||
|
||||
def get_outputs_npz_path(self, image_abs_path: str) -> str:
|
||||
return os.path.splitext(image_abs_path)[0] + FluxTextEncoderOutputsCachingStrategy.FLUX_TEXT_ENCODER_OUTPUTS_NPZ_SUFFIX
|
||||
|
||||
@@ -138,12 +141,18 @@ class FluxTextEncoderOutputsCachingStrategy(TextEncoderOutputsCachingStrategy):
|
||||
txt_ids = data["txt_ids"]
|
||||
t5_attn_mask = data["t5_attn_mask"]
|
||||
# apply_t5_attn_mask should be same as self.apply_t5_attn_mask
|
||||
|
||||
return [l_pooled, t5_out, txt_ids, t5_attn_mask]
|
||||
|
||||
def cache_batch_outputs(
|
||||
self, tokenize_strategy: TokenizeStrategy, models: List[Any], text_encoding_strategy: TextEncodingStrategy, infos: List
|
||||
):
|
||||
if not self.warn_fp8_weights:
|
||||
if get_t5xxl_actual_dtype(models[1]) == torch.float8_e4m3fn:
|
||||
logger.warning(
|
||||
"T5 model is using fp8 weights for caching. This may affect the quality of the cached outputs."
|
||||
)
|
||||
self.warn_fp8_weights = True
|
||||
|
||||
flux_text_encoding_strategy: FluxTextEncodingStrategy = text_encoding_strategy
|
||||
captions = [info.caption for info in infos]
|
||||
|
||||
@@ -190,12 +199,9 @@ class FluxLatentsCachingStrategy(LatentsCachingStrategy):
|
||||
def __init__(self, cache_to_disk: bool, batch_size: int, skip_disk_cache_validity_check: bool) -> None:
|
||||
super().__init__(cache_to_disk, batch_size, skip_disk_cache_validity_check)
|
||||
|
||||
def get_image_size_from_disk_cache_path(self, absolute_path: str) -> Tuple[Optional[int], Optional[int]]:
|
||||
npz_file = glob.glob(os.path.splitext(absolute_path)[0] + "_*" + FluxLatentsCachingStrategy.FLUX_LATENTS_NPZ_SUFFIX)
|
||||
if len(npz_file) == 0:
|
||||
return None, None
|
||||
w, h = os.path.splitext(npz_file[0])[0].split("_")[-2].split("x")
|
||||
return int(w), int(h)
|
||||
@property
|
||||
def cache_suffix(self) -> str:
|
||||
return FluxLatentsCachingStrategy.FLUX_LATENTS_NPZ_SUFFIX
|
||||
|
||||
def get_latents_npz_path(self, absolute_path: str, image_size: Tuple[int, int]) -> str:
|
||||
return (
|
||||
@@ -205,7 +211,7 @@ class FluxLatentsCachingStrategy(LatentsCachingStrategy):
|
||||
)
|
||||
|
||||
def is_disk_cached_latents_expected(self, bucket_reso: Tuple[int, int], npz_path: str, flip_aug: bool, alpha_mask: bool):
|
||||
return self._default_is_disk_cached_latents_expected(8, bucket_reso, npz_path, flip_aug, alpha_mask, True)
|
||||
return self._default_is_disk_cached_latents_expected(8, bucket_reso, npz_path, flip_aug, alpha_mask, multi_resolution=True)
|
||||
|
||||
def load_latents_from_disk(
|
||||
self, npz_path: str, bucket_reso: Tuple[int, int]
|
||||
@@ -219,7 +225,7 @@ class FluxLatentsCachingStrategy(LatentsCachingStrategy):
|
||||
vae_dtype = vae.dtype
|
||||
|
||||
self._default_cache_batch_latents(
|
||||
encode_by_vae, vae_device, vae_dtype, image_infos, flip_aug, alpha_mask, random_crop, True
|
||||
encode_by_vae, vae_device, vae_dtype, image_infos, flip_aug, alpha_mask, random_crop, multi_resolution=True
|
||||
)
|
||||
|
||||
if not train_util.HIGH_VRAM:
|
||||
|
||||
+39
-7
@@ -40,6 +40,16 @@ class SdTokenizeStrategy(TokenizeStrategy):
|
||||
text = [text] if isinstance(text, str) else text
|
||||
return [torch.stack([self._get_input_ids(self.tokenizer, 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
|
||||
tokens_list = []
|
||||
weights_list = []
|
||||
for t in text:
|
||||
tokens, weights = self._get_input_ids(self.tokenizer, t, self.max_length, weighted=True)
|
||||
tokens_list.append(tokens)
|
||||
weights_list.append(weights)
|
||||
return [torch.stack(tokens_list, dim=0)], [torch.stack(weights_list, dim=0)]
|
||||
|
||||
|
||||
class SdTextEncodingStrategy(TextEncodingStrategy):
|
||||
def __init__(self, clip_skip: Optional[int] = None) -> None:
|
||||
@@ -58,6 +68,8 @@ class SdTextEncodingStrategy(TextEncodingStrategy):
|
||||
model_max_length = sd_tokenize_strategy.tokenizer.model_max_length
|
||||
tokens = tokens.reshape((-1, model_max_length)) # batch_size*3, 77
|
||||
|
||||
tokens = tokens.to(text_encoder.device)
|
||||
|
||||
if self.clip_skip is None:
|
||||
encoder_hidden_states = text_encoder(tokens)[0]
|
||||
else:
|
||||
@@ -93,6 +105,30 @@ class SdTextEncodingStrategy(TextEncodingStrategy):
|
||||
|
||||
return [encoder_hidden_states]
|
||||
|
||||
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]:
|
||||
encoder_hidden_states = self.encode_tokens(tokenize_strategy, models, tokens_list)[0]
|
||||
|
||||
weights = weights_list[0].to(encoder_hidden_states.device)
|
||||
|
||||
# apply weights
|
||||
if weights.shape[1] == 1: # no max_token_length
|
||||
# weights: ((b, 1, 77), (b, 1, 77)), hidden_states: (b, 77, 768), (b, 77, 768)
|
||||
encoder_hidden_states = encoder_hidden_states * weights.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 i in range(weights.shape[1]):
|
||||
encoder_hidden_states[:, i * 75 + 1 : i * 75 + 76] = encoder_hidden_states[:, i * 75 + 1 : i * 75 + 76] * weights[
|
||||
:, i, 1:-1
|
||||
].unsqueeze(-1)
|
||||
|
||||
return [encoder_hidden_states]
|
||||
|
||||
|
||||
class SdSdxlLatentsCachingStrategy(LatentsCachingStrategy):
|
||||
# sd and sdxl share the same strategy. we can make them separate, but the difference is only the suffix.
|
||||
@@ -109,13 +145,9 @@ class SdSdxlLatentsCachingStrategy(LatentsCachingStrategy):
|
||||
SdSdxlLatentsCachingStrategy.SD_LATENTS_NPZ_SUFFIX if sd else SdSdxlLatentsCachingStrategy.SDXL_LATENTS_NPZ_SUFFIX
|
||||
)
|
||||
|
||||
def get_image_size_from_disk_cache_path(self, absolute_path: str) -> Tuple[Optional[int], Optional[int]]:
|
||||
# does not include old npz
|
||||
npz_file = glob.glob(os.path.splitext(absolute_path)[0] + "_*" + self.suffix)
|
||||
if len(npz_file) == 0:
|
||||
return None, None
|
||||
w, h = os.path.splitext(npz_file[0])[0].split("_")[-2].split("x")
|
||||
return int(w), int(h)
|
||||
@property
|
||||
def cache_suffix(self) -> str:
|
||||
return self.suffix
|
||||
|
||||
def get_latents_npz_path(self, absolute_path: str, image_size: Tuple[int, int]) -> str:
|
||||
# support old .npz
|
||||
|
||||
+231
-64
@@ -3,10 +3,9 @@ import glob
|
||||
from typing import Any, List, Optional, Tuple, Union
|
||||
import torch
|
||||
import numpy as np
|
||||
from transformers import CLIPTokenizer, T5TokenizerFast
|
||||
from transformers import CLIPTokenizer, T5TokenizerFast, CLIPTextModel, CLIPTextModelWithProjection, T5EncoderModel
|
||||
|
||||
from . import sd3_utils, train_util
|
||||
from . import sd3_models
|
||||
from . import train_util
|
||||
from .strategy_base import LatentsCachingStrategy, TextEncodingStrategy, TokenizeStrategy, TextEncoderOutputsCachingStrategy
|
||||
|
||||
from .utils import setup_logging
|
||||
@@ -48,45 +47,200 @@ class Sd3TokenizeStrategy(TokenizeStrategy):
|
||||
|
||||
|
||||
class Sd3TextEncodingStrategy(TextEncodingStrategy):
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
def __init__(
|
||||
self,
|
||||
apply_lg_attn_mask: Optional[bool] = None,
|
||||
apply_t5_attn_mask: Optional[bool] = None,
|
||||
l_dropout_rate: float = 0.0,
|
||||
g_dropout_rate: float = 0.0,
|
||||
t5_dropout_rate: float = 0.0,
|
||||
) -> None:
|
||||
"""
|
||||
Args:
|
||||
apply_t5_attn_mask: Default value for apply_t5_attn_mask.
|
||||
"""
|
||||
self.apply_lg_attn_mask = apply_lg_attn_mask
|
||||
self.apply_t5_attn_mask = apply_t5_attn_mask
|
||||
self.l_dropout_rate = l_dropout_rate
|
||||
self.g_dropout_rate = g_dropout_rate
|
||||
self.t5_dropout_rate = t5_dropout_rate
|
||||
|
||||
def encode_tokens(
|
||||
self,
|
||||
tokenize_strategy: TokenizeStrategy,
|
||||
models: List[Any],
|
||||
tokens: List[torch.Tensor],
|
||||
apply_lg_attn_mask: bool = False,
|
||||
apply_t5_attn_mask: bool = False,
|
||||
apply_lg_attn_mask: Optional[bool] = False,
|
||||
apply_t5_attn_mask: Optional[bool] = False,
|
||||
enable_dropout: bool = True,
|
||||
) -> List[torch.Tensor]:
|
||||
"""
|
||||
returned embeddings are not masked
|
||||
"""
|
||||
clip_l, clip_g, t5xxl = models
|
||||
clip_l: Optional[CLIPTextModel]
|
||||
clip_g: Optional[CLIPTextModelWithProjection]
|
||||
t5xxl: Optional[T5EncoderModel]
|
||||
|
||||
l_tokens, g_tokens, t5_tokens = tokens[:3]
|
||||
l_attn_mask, g_attn_mask, t5_attn_mask = tokens[3:] if len(tokens) > 3 else [None, None, None]
|
||||
if l_tokens is None:
|
||||
if apply_lg_attn_mask is None:
|
||||
apply_lg_attn_mask = self.apply_lg_attn_mask
|
||||
if apply_t5_attn_mask is None:
|
||||
apply_t5_attn_mask = self.apply_t5_attn_mask
|
||||
|
||||
l_tokens, g_tokens, t5_tokens, l_attn_mask, g_attn_mask, t5_attn_mask = tokens
|
||||
|
||||
# dropout: if enable_dropout is False, dropout is not applied. dropout means zeroing out embeddings
|
||||
|
||||
if l_tokens is None or clip_l is None:
|
||||
assert g_tokens is None, "g_tokens must be None if l_tokens is None"
|
||||
lg_out = None
|
||||
lg_pooled = None
|
||||
l_attn_mask = None
|
||||
g_attn_mask = None
|
||||
else:
|
||||
assert g_tokens is not None, "g_tokens must not be None if l_tokens is not None"
|
||||
l_out, l_pooled = clip_l(l_tokens)
|
||||
g_out, g_pooled = clip_g(g_tokens)
|
||||
if apply_lg_attn_mask:
|
||||
l_out = l_out * l_attn_mask.to(l_out.device).unsqueeze(-1)
|
||||
g_out = g_out * g_attn_mask.to(g_out.device).unsqueeze(-1)
|
||||
|
||||
# drop some members of the batch: we do not call clip_l and clip_g for dropped members
|
||||
batch_size, l_seq_len = l_tokens.shape
|
||||
g_seq_len = g_tokens.shape[1]
|
||||
|
||||
non_drop_l_indices = []
|
||||
non_drop_g_indices = []
|
||||
for i in range(l_tokens.shape[0]):
|
||||
drop_l = enable_dropout and (self.l_dropout_rate > 0.0 and random.random() < self.l_dropout_rate)
|
||||
drop_g = enable_dropout and (self.g_dropout_rate > 0.0 and random.random() < self.g_dropout_rate)
|
||||
if not drop_l:
|
||||
non_drop_l_indices.append(i)
|
||||
if not drop_g:
|
||||
non_drop_g_indices.append(i)
|
||||
|
||||
# filter out dropped members
|
||||
if len(non_drop_l_indices) > 0 and len(non_drop_l_indices) < batch_size:
|
||||
l_tokens = l_tokens[non_drop_l_indices]
|
||||
l_attn_mask = l_attn_mask[non_drop_l_indices]
|
||||
if len(non_drop_g_indices) > 0 and len(non_drop_g_indices) < batch_size:
|
||||
g_tokens = g_tokens[non_drop_g_indices]
|
||||
g_attn_mask = g_attn_mask[non_drop_g_indices]
|
||||
|
||||
# call clip_l for non-dropped members
|
||||
if len(non_drop_l_indices) > 0:
|
||||
nd_l_attn_mask = l_attn_mask.to(clip_l.device)
|
||||
prompt_embeds = clip_l(
|
||||
l_tokens.to(clip_l.device), nd_l_attn_mask if apply_lg_attn_mask else None, output_hidden_states=True
|
||||
)
|
||||
nd_l_pooled = prompt_embeds[0]
|
||||
nd_l_out = prompt_embeds.hidden_states[-2]
|
||||
if len(non_drop_g_indices) > 0:
|
||||
nd_g_attn_mask = g_attn_mask.to(clip_g.device)
|
||||
prompt_embeds = clip_g(
|
||||
g_tokens.to(clip_g.device), nd_g_attn_mask if apply_lg_attn_mask else None, output_hidden_states=True
|
||||
)
|
||||
nd_g_pooled = prompt_embeds[0]
|
||||
nd_g_out = prompt_embeds.hidden_states[-2]
|
||||
|
||||
# fill in the dropped members
|
||||
if len(non_drop_l_indices) == batch_size:
|
||||
l_pooled = nd_l_pooled
|
||||
l_out = nd_l_out
|
||||
else:
|
||||
# model output is always float32 because of the models are wrapped with Accelerator
|
||||
l_pooled = torch.zeros((batch_size, 768), device=clip_l.device, dtype=torch.float32)
|
||||
l_out = torch.zeros((batch_size, l_seq_len, 768), device=clip_l.device, dtype=torch.float32)
|
||||
l_attn_mask = torch.zeros((batch_size, l_seq_len), device=clip_l.device, dtype=l_attn_mask.dtype)
|
||||
if len(non_drop_l_indices) > 0:
|
||||
l_pooled[non_drop_l_indices] = nd_l_pooled
|
||||
l_out[non_drop_l_indices] = nd_l_out
|
||||
l_attn_mask[non_drop_l_indices] = nd_l_attn_mask
|
||||
|
||||
if len(non_drop_g_indices) == batch_size:
|
||||
g_pooled = nd_g_pooled
|
||||
g_out = nd_g_out
|
||||
else:
|
||||
g_pooled = torch.zeros((batch_size, 1280), device=clip_g.device, dtype=torch.float32)
|
||||
g_out = torch.zeros((batch_size, g_seq_len, 1280), device=clip_g.device, dtype=torch.float32)
|
||||
g_attn_mask = torch.zeros((batch_size, g_seq_len), device=clip_g.device, dtype=g_attn_mask.dtype)
|
||||
if len(non_drop_g_indices) > 0:
|
||||
g_pooled[non_drop_g_indices] = nd_g_pooled
|
||||
g_out[non_drop_g_indices] = nd_g_out
|
||||
g_attn_mask[non_drop_g_indices] = nd_g_attn_mask
|
||||
|
||||
lg_pooled = torch.cat((l_pooled, g_pooled), dim=-1)
|
||||
lg_out = torch.cat([l_out, g_out], dim=-1)
|
||||
|
||||
if t5xxl is not None and t5_tokens is not None:
|
||||
t5_out, _ = t5xxl(t5_tokens) # t5_out is [1, max length, 4096]
|
||||
if apply_t5_attn_mask:
|
||||
t5_out = t5_out * t5_attn_mask.to(t5_out.device).unsqueeze(-1)
|
||||
else:
|
||||
if t5xxl is None or t5_tokens is None:
|
||||
t5_out = None
|
||||
t5_attn_mask = None
|
||||
else:
|
||||
# drop some members of the batch: we do not call t5xxl for dropped members
|
||||
batch_size, t5_seq_len = t5_tokens.shape
|
||||
non_drop_t5_indices = []
|
||||
for i in range(t5_tokens.shape[0]):
|
||||
drop_t5 = enable_dropout and (self.t5_dropout_rate > 0.0 and random.random() < self.t5_dropout_rate)
|
||||
if not drop_t5:
|
||||
non_drop_t5_indices.append(i)
|
||||
|
||||
lg_pooled = torch.cat((l_pooled, g_pooled), dim=-1) if l_tokens is not None else None
|
||||
return [lg_out, t5_out, lg_pooled]
|
||||
# filter out dropped members
|
||||
if len(non_drop_t5_indices) > 0 and len(non_drop_t5_indices) < batch_size:
|
||||
t5_tokens = t5_tokens[non_drop_t5_indices]
|
||||
t5_attn_mask = t5_attn_mask[non_drop_t5_indices]
|
||||
|
||||
# call t5xxl for non-dropped members
|
||||
if len(non_drop_t5_indices) > 0:
|
||||
nd_t5_attn_mask = t5_attn_mask.to(t5xxl.device)
|
||||
nd_t5_out, _ = t5xxl(
|
||||
t5_tokens.to(t5xxl.device),
|
||||
nd_t5_attn_mask if apply_t5_attn_mask else None,
|
||||
return_dict=False,
|
||||
output_hidden_states=True,
|
||||
)
|
||||
|
||||
# fill in the dropped members
|
||||
if len(non_drop_t5_indices) == batch_size:
|
||||
t5_out = nd_t5_out
|
||||
else:
|
||||
t5_out = torch.zeros((batch_size, t5_seq_len, 4096), device=t5xxl.device, dtype=torch.float32)
|
||||
t5_attn_mask = torch.zeros((batch_size, t5_seq_len), device=t5xxl.device, dtype=t5_attn_mask.dtype)
|
||||
if len(non_drop_t5_indices) > 0:
|
||||
t5_out[non_drop_t5_indices] = nd_t5_out
|
||||
t5_attn_mask[non_drop_t5_indices] = nd_t5_attn_mask
|
||||
|
||||
# masks are used for attention masking in transformer
|
||||
return [lg_out, t5_out, lg_pooled, l_attn_mask, g_attn_mask, t5_attn_mask]
|
||||
|
||||
def drop_cached_text_encoder_outputs(
|
||||
self,
|
||||
lg_out: torch.Tensor,
|
||||
t5_out: torch.Tensor,
|
||||
lg_pooled: torch.Tensor,
|
||||
l_attn_mask: torch.Tensor,
|
||||
g_attn_mask: torch.Tensor,
|
||||
t5_attn_mask: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
# dropout: if enable_dropout is True, dropout is not applied. dropout means zeroing out embeddings
|
||||
if lg_out is not None:
|
||||
for i in range(lg_out.shape[0]):
|
||||
drop_l = self.l_dropout_rate > 0.0 and random.random() < self.l_dropout_rate
|
||||
if drop_l:
|
||||
lg_out[i, :, :768] = torch.zeros_like(lg_out[i, :, :768])
|
||||
lg_pooled[i, :768] = torch.zeros_like(lg_pooled[i, :768])
|
||||
if l_attn_mask is not None:
|
||||
l_attn_mask[i] = torch.zeros_like(l_attn_mask[i])
|
||||
drop_g = self.g_dropout_rate > 0.0 and random.random() < self.g_dropout_rate
|
||||
if drop_g:
|
||||
lg_out[i, :, 768:] = torch.zeros_like(lg_out[i, :, 768:])
|
||||
lg_pooled[i, 768:] = torch.zeros_like(lg_pooled[i, 768:])
|
||||
if g_attn_mask is not None:
|
||||
g_attn_mask[i] = torch.zeros_like(g_attn_mask[i])
|
||||
|
||||
if t5_out is not None:
|
||||
for i in range(t5_out.shape[0]):
|
||||
drop_t5 = self.t5_dropout_rate > 0.0 and random.random() < self.t5_dropout_rate
|
||||
if drop_t5:
|
||||
t5_out[i] = torch.zeros_like(t5_out[i])
|
||||
if t5_attn_mask is not None:
|
||||
t5_attn_mask[i] = torch.zeros_like(t5_attn_mask[i])
|
||||
|
||||
return [lg_out, t5_out, lg_pooled, l_attn_mask, g_attn_mask, t5_attn_mask]
|
||||
|
||||
def concat_encodings(
|
||||
self, lg_out: torch.Tensor, t5_out: Optional[torch.Tensor], lg_pooled: torch.Tensor
|
||||
@@ -132,39 +286,38 @@ class Sd3TextEncoderOutputsCachingStrategy(TextEncoderOutputsCachingStrategy):
|
||||
return False
|
||||
if "clip_l_attn_mask" not in npz or "clip_g_attn_mask" not in npz: # necessary even if not used
|
||||
return False
|
||||
# t5xxl is optional
|
||||
if "apply_lg_attn_mask" not in npz:
|
||||
return False
|
||||
if "t5_out" not in npz:
|
||||
return False
|
||||
if "t5_attn_mask" not in npz:
|
||||
return False
|
||||
npz_apply_lg_attn_mask = npz["apply_lg_attn_mask"]
|
||||
if npz_apply_lg_attn_mask != self.apply_lg_attn_mask:
|
||||
return False
|
||||
if "apply_t5_attn_mask" not in npz:
|
||||
return False
|
||||
npz_apply_t5_attn_mask = npz["apply_t5_attn_mask"]
|
||||
if npz_apply_t5_attn_mask != self.apply_t5_attn_mask:
|
||||
return False
|
||||
except Exception as e:
|
||||
logger.error(f"Error loading file: {npz_path}")
|
||||
raise e
|
||||
|
||||
return True
|
||||
|
||||
def mask_lg_attn(self, lg_out: np.ndarray, l_attn_mask: np.ndarray, g_attn_mask: np.ndarray) -> np.ndarray:
|
||||
l_out = lg_out[..., :768]
|
||||
g_out = lg_out[..., 768:] # 1280
|
||||
l_out = l_out * np.expand_dims(l_attn_mask, -1) # l_out = l_out * l_attn_mask.
|
||||
g_out = g_out * np.expand_dims(g_attn_mask, -1) # g_out = g_out * g_attn_mask.
|
||||
return np.concatenate([l_out, g_out], axis=-1)
|
||||
|
||||
def mask_t5_attn(self, t5_out: np.ndarray, t5_attn_mask: np.ndarray) -> np.ndarray:
|
||||
return t5_out * np.expand_dims(t5_attn_mask, -1)
|
||||
|
||||
def load_outputs_npz(self, npz_path: str) -> List[np.ndarray]:
|
||||
data = np.load(npz_path)
|
||||
lg_out = data["lg_out"]
|
||||
lg_pooled = data["lg_pooled"]
|
||||
t5_out = data["t5_out"] if "t5_out" in data else None
|
||||
t5_out = data["t5_out"]
|
||||
|
||||
if self.apply_lg_attn_mask:
|
||||
l_attn_mask = data["clip_l_attn_mask"]
|
||||
g_attn_mask = data["clip_g_attn_mask"]
|
||||
lg_out = self.mask_lg_attn(lg_out, l_attn_mask, g_attn_mask)
|
||||
|
||||
if self.apply_t5_attn_mask and t5_out is not None:
|
||||
t5_attn_mask = data["t5_attn_mask"]
|
||||
t5_out = self.mask_t5_attn(t5_out, t5_attn_mask)
|
||||
|
||||
return [lg_out, t5_out, lg_pooled]
|
||||
# apply_t5_attn_mask and apply_lg_attn_mask are same as self.apply_t5_attn_mask and self.apply_lg_attn_mask
|
||||
return [lg_out, t5_out, lg_pooled, l_attn_mask, g_attn_mask, t5_attn_mask]
|
||||
|
||||
def cache_batch_outputs(
|
||||
self, tokenize_strategy: TokenizeStrategy, models: List[Any], text_encoding_strategy: TextEncodingStrategy, infos: List
|
||||
@@ -174,46 +327,56 @@ class Sd3TextEncoderOutputsCachingStrategy(TextEncoderOutputsCachingStrategy):
|
||||
|
||||
tokens_and_masks = tokenize_strategy.tokenize(captions)
|
||||
with torch.no_grad():
|
||||
lg_out, t5_out, lg_pooled = sd3_text_encoding_strategy.encode_tokens(
|
||||
tokenize_strategy, models, tokens_and_masks, self.apply_lg_attn_mask, self.apply_t5_attn_mask
|
||||
# always disable dropout during caching
|
||||
lg_out, t5_out, lg_pooled, l_attn_mask, g_attn_mask, t5_attn_mask = sd3_text_encoding_strategy.encode_tokens(
|
||||
tokenize_strategy,
|
||||
models,
|
||||
tokens_and_masks,
|
||||
apply_lg_attn_mask=self.apply_lg_attn_mask,
|
||||
apply_t5_attn_mask=self.apply_t5_attn_mask,
|
||||
enable_dropout=False,
|
||||
)
|
||||
|
||||
if lg_out.dtype == torch.bfloat16:
|
||||
lg_out = lg_out.float()
|
||||
if lg_pooled.dtype == torch.bfloat16:
|
||||
lg_pooled = lg_pooled.float()
|
||||
if t5_out is not None and t5_out.dtype == torch.bfloat16:
|
||||
if t5_out.dtype == torch.bfloat16:
|
||||
t5_out = t5_out.float()
|
||||
|
||||
lg_out = lg_out.cpu().numpy()
|
||||
lg_pooled = lg_pooled.cpu().numpy()
|
||||
if t5_out is not None:
|
||||
t5_out = t5_out.cpu().numpy()
|
||||
|
||||
l_attn_mask = tokens_and_masks[3].cpu().numpy()
|
||||
g_attn_mask = tokens_and_masks[4].cpu().numpy()
|
||||
t5_attn_mask = tokens_and_masks[5].cpu().numpy()
|
||||
|
||||
for i, info in enumerate(infos):
|
||||
lg_out_i = lg_out[i]
|
||||
t5_out_i = t5_out[i] if t5_out is not None else None
|
||||
t5_out_i = t5_out[i]
|
||||
lg_pooled_i = lg_pooled[i]
|
||||
l_attn_mask_i = l_attn_mask[i]
|
||||
g_attn_mask_i = g_attn_mask[i]
|
||||
t5_attn_mask_i = t5_attn_mask[i]
|
||||
apply_lg_attn_mask = self.apply_lg_attn_mask
|
||||
apply_t5_attn_mask = self.apply_t5_attn_mask
|
||||
|
||||
if self.cache_to_disk:
|
||||
clip_l_attn_mask, clip_g_attn_mask, t5_attn_mask = tokens_and_masks[3:6]
|
||||
clip_l_attn_mask_i = clip_l_attn_mask[i].cpu().numpy()
|
||||
clip_g_attn_mask_i = clip_g_attn_mask[i].cpu().numpy()
|
||||
t5_attn_mask_i = t5_attn_mask[i].cpu().numpy() if t5_attn_mask is not None else None # shouldn't be None
|
||||
kwargs = {}
|
||||
if t5_out is not None:
|
||||
kwargs["t5_out"] = t5_out_i
|
||||
np.savez(
|
||||
info.text_encoder_outputs_npz,
|
||||
lg_out=lg_out_i,
|
||||
lg_pooled=lg_pooled_i,
|
||||
clip_l_attn_mask=clip_l_attn_mask_i,
|
||||
clip_g_attn_mask=clip_g_attn_mask_i,
|
||||
t5_out=t5_out_i,
|
||||
clip_l_attn_mask=l_attn_mask_i,
|
||||
clip_g_attn_mask=g_attn_mask_i,
|
||||
t5_attn_mask=t5_attn_mask_i,
|
||||
**kwargs,
|
||||
apply_lg_attn_mask=apply_lg_attn_mask,
|
||||
apply_t5_attn_mask=apply_t5_attn_mask,
|
||||
)
|
||||
else:
|
||||
info.text_encoder_outputs = (lg_out_i, t5_out_i, lg_pooled_i)
|
||||
# it's fine that attn mask is not None. it's overwritten before calling the model if necessary
|
||||
info.text_encoder_outputs = (lg_out_i, t5_out_i, lg_pooled_i, l_attn_mask_i, g_attn_mask_i, t5_attn_mask_i)
|
||||
|
||||
|
||||
class Sd3LatentsCachingStrategy(LatentsCachingStrategy):
|
||||
@@ -222,12 +385,9 @@ class Sd3LatentsCachingStrategy(LatentsCachingStrategy):
|
||||
def __init__(self, cache_to_disk: bool, batch_size: int, skip_disk_cache_validity_check: bool) -> None:
|
||||
super().__init__(cache_to_disk, batch_size, skip_disk_cache_validity_check)
|
||||
|
||||
def get_image_size_from_disk_cache_path(self, absolute_path: str) -> Tuple[Optional[int], Optional[int]]:
|
||||
npz_file = glob.glob(os.path.splitext(absolute_path)[0] + "_*" + Sd3LatentsCachingStrategy.SD3_LATENTS_NPZ_SUFFIX)
|
||||
if len(npz_file) == 0:
|
||||
return None, None
|
||||
w, h = os.path.splitext(npz_file[0])[0].split("_")[-2].split("x")
|
||||
return int(w), int(h)
|
||||
@property
|
||||
def cache_suffix(self) -> str:
|
||||
return Sd3LatentsCachingStrategy.SD3_LATENTS_NPZ_SUFFIX
|
||||
|
||||
def get_latents_npz_path(self, absolute_path: str, image_size: Tuple[int, int]) -> str:
|
||||
return (
|
||||
@@ -237,7 +397,12 @@ class Sd3LatentsCachingStrategy(LatentsCachingStrategy):
|
||||
)
|
||||
|
||||
def is_disk_cached_latents_expected(self, bucket_reso: Tuple[int, int], npz_path: str, flip_aug: bool, alpha_mask: bool):
|
||||
return self._default_is_disk_cached_latents_expected(8, bucket_reso, npz_path, flip_aug, alpha_mask)
|
||||
return self._default_is_disk_cached_latents_expected(8, bucket_reso, npz_path, flip_aug, alpha_mask, multi_resolution=True)
|
||||
|
||||
def load_latents_from_disk(
|
||||
self, npz_path: str, bucket_reso: Tuple[int, int]
|
||||
) -> Tuple[Optional[np.ndarray], Optional[List[int]], Optional[List[int]], Optional[np.ndarray], Optional[np.ndarray]]:
|
||||
return self._default_load_latents_from_disk(8, npz_path, bucket_reso) # support multi-resolution
|
||||
|
||||
# TODO remove circular dependency for ImageInfo
|
||||
def cache_batch_latents(self, vae, image_infos: List, flip_aug: bool, alpha_mask: bool, random_crop: bool):
|
||||
@@ -245,7 +410,9 @@ class Sd3LatentsCachingStrategy(LatentsCachingStrategy):
|
||||
vae_device = vae.device
|
||||
vae_dtype = vae.dtype
|
||||
|
||||
self._default_cache_batch_latents(encode_by_vae, vae_device, vae_dtype, image_infos, flip_aug, alpha_mask, random_crop)
|
||||
self._default_cache_batch_latents(
|
||||
encode_by_vae, vae_device, vae_dtype, image_infos, flip_aug, alpha_mask, random_crop, multi_resolution=True
|
||||
)
|
||||
|
||||
if not train_util.HIGH_VRAM:
|
||||
train_util.clean_memory_on_device(vae.device)
|
||||
|
||||
@@ -37,6 +37,22 @@ class SdxlTokenizeStrategy(TokenizeStrategy):
|
||||
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:
|
||||
@@ -98,6 +114,9 @@ class SdxlTextEncodingStrategy(TextEncodingStrategy):
|
||||
):
|
||||
# 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
|
||||
@@ -155,7 +174,8 @@ class SdxlTextEncodingStrategy(TextEncodingStrategy):
|
||||
"""
|
||||
Args:
|
||||
tokenize_strategy: TokenizeStrategy
|
||||
models: List of models, [text_encoder1, text_encoder2, unwrapped text_encoder2 (optional)]
|
||||
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:
|
||||
@@ -172,14 +192,45 @@ class SdxlTextEncodingStrategy(TextEncodingStrategy):
|
||||
)
|
||||
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
|
||||
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)
|
||||
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
|
||||
@@ -215,11 +266,19 @@ class SdxlTextEncoderOutputsCachingStrategy(TextEncoderOutputsCachingStrategy):
|
||||
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:
|
||||
|
||||
+659
-231
File diff suppressed because it is too large
Load Diff
@@ -9,7 +9,28 @@ import struct
|
||||
from diffusers import EulerAncestralDiscreteScheduler
|
||||
import diffusers.schedulers.scheduling_euler_ancestral_discrete
|
||||
from diffusers.schedulers.scheduling_euler_ancestral_discrete import EulerAncestralDiscreteSchedulerOutput
|
||||
import cv2
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
from safetensors.torch import load_file
|
||||
|
||||
def load_safetensors(
|
||||
path: str, device: Union[str, torch.device], disable_mmap: bool = False, dtype: Optional[torch.dtype] = torch.float32
|
||||
):
|
||||
if disable_mmap:
|
||||
# return safetensors.torch.load(open(path, "rb").read())
|
||||
# use experimental loader
|
||||
#logger.info(f"Loading without mmap (experimental)")
|
||||
state_dict = {}
|
||||
with MemoryEfficientSafeOpen(path) as f:
|
||||
for key in f.keys():
|
||||
state_dict[key] = f.get_tensor(key).to(device, dtype=dtype)
|
||||
return state_dict
|
||||
else:
|
||||
try:
|
||||
return load_file(path, device=device)
|
||||
except:
|
||||
return load_file(path) # prevent device invalid Error
|
||||
|
||||
def fire_in_thread(f, *args, **kwargs):
|
||||
threading.Thread(target=f, args=args, kwargs=kwargs).start()
|
||||
@@ -81,6 +102,137 @@ def setup_logging(args=None, log_level=None, reset=False):
|
||||
logger.info(msg_init)
|
||||
|
||||
|
||||
def str_to_dtype(s: Optional[str], default_dtype: Optional[torch.dtype] = None) -> torch.dtype:
|
||||
"""
|
||||
Convert a string to a torch.dtype
|
||||
|
||||
Args:
|
||||
s: string representation of the dtype
|
||||
default_dtype: default dtype to return if s is None
|
||||
|
||||
Returns:
|
||||
torch.dtype: the corresponding torch.dtype
|
||||
|
||||
Raises:
|
||||
ValueError: if the dtype is not supported
|
||||
|
||||
Examples:
|
||||
>>> str_to_dtype("float32")
|
||||
torch.float32
|
||||
>>> str_to_dtype("fp32")
|
||||
torch.float32
|
||||
>>> str_to_dtype("float16")
|
||||
torch.float16
|
||||
>>> str_to_dtype("fp16")
|
||||
torch.float16
|
||||
>>> str_to_dtype("bfloat16")
|
||||
torch.bfloat16
|
||||
>>> str_to_dtype("bf16")
|
||||
torch.bfloat16
|
||||
>>> str_to_dtype("fp8")
|
||||
torch.float8_e4m3fn
|
||||
>>> str_to_dtype("fp8_e4m3fn")
|
||||
torch.float8_e4m3fn
|
||||
>>> str_to_dtype("fp8_e4m3fnuz")
|
||||
torch.float8_e4m3fnuz
|
||||
>>> str_to_dtype("fp8_e5m2")
|
||||
torch.float8_e5m2
|
||||
>>> str_to_dtype("fp8_e5m2fnuz")
|
||||
torch.float8_e5m2fnuz
|
||||
"""
|
||||
if s is None:
|
||||
return default_dtype
|
||||
if s in ["bf16", "bfloat16"]:
|
||||
return torch.bfloat16
|
||||
elif s in ["fp16", "float16"]:
|
||||
return torch.float16
|
||||
elif s in ["fp32", "float32", "float"]:
|
||||
return torch.float32
|
||||
elif s in ["fp8_e4m3fn", "e4m3fn", "float8_e4m3fn"]:
|
||||
return torch.float8_e4m3fn
|
||||
elif s in ["fp8_e4m3fnuz", "e4m3fnuz", "float8_e4m3fnuz"]:
|
||||
return torch.float8_e4m3fnuz
|
||||
elif s in ["fp8_e5m2", "e5m2", "float8_e5m2"]:
|
||||
return torch.float8_e5m2
|
||||
elif s in ["fp8_e5m2fnuz", "e5m2fnuz", "float8_e5m2fnuz"]:
|
||||
return torch.float8_e5m2fnuz
|
||||
elif s in ["fp8", "float8"]:
|
||||
return torch.float8_e4m3fn # default fp8
|
||||
else:
|
||||
raise ValueError(f"Unsupported dtype: {s}")
|
||||
|
||||
|
||||
def mem_eff_save_file(tensors: Dict[str, torch.Tensor], filename: str, metadata: Dict[str, Any] = None):
|
||||
"""
|
||||
memory efficient save file
|
||||
"""
|
||||
|
||||
_TYPES = {
|
||||
torch.float64: "F64",
|
||||
torch.float32: "F32",
|
||||
torch.float16: "F16",
|
||||
torch.bfloat16: "BF16",
|
||||
torch.int64: "I64",
|
||||
torch.int32: "I32",
|
||||
torch.int16: "I16",
|
||||
torch.int8: "I8",
|
||||
torch.uint8: "U8",
|
||||
torch.bool: "BOOL",
|
||||
getattr(torch, "float8_e5m2", None): "F8_E5M2",
|
||||
getattr(torch, "float8_e4m3fn", None): "F8_E4M3",
|
||||
}
|
||||
_ALIGN = 256
|
||||
|
||||
def validate_metadata(metadata: Dict[str, Any]) -> Dict[str, str]:
|
||||
validated = {}
|
||||
for key, value in metadata.items():
|
||||
if not isinstance(key, str):
|
||||
raise ValueError(f"Metadata key must be a string, got {type(key)}")
|
||||
if not isinstance(value, str):
|
||||
print(f"Warning: Metadata value for key '{key}' is not a string. Converting to string.")
|
||||
validated[key] = str(value)
|
||||
else:
|
||||
validated[key] = value
|
||||
return validated
|
||||
|
||||
print(f"Using memory efficient save file: {filename}")
|
||||
|
||||
header = {}
|
||||
offset = 0
|
||||
if metadata:
|
||||
header["__metadata__"] = validate_metadata(metadata)
|
||||
for k, v in tensors.items():
|
||||
if v.numel() == 0: # empty tensor
|
||||
header[k] = {"dtype": _TYPES[v.dtype], "shape": list(v.shape), "data_offsets": [offset, offset]}
|
||||
else:
|
||||
size = v.numel() * v.element_size()
|
||||
header[k] = {"dtype": _TYPES[v.dtype], "shape": list(v.shape), "data_offsets": [offset, offset + size]}
|
||||
offset += size
|
||||
|
||||
hjson = json.dumps(header).encode("utf-8")
|
||||
hjson += b" " * (-(len(hjson) + 8) % _ALIGN)
|
||||
|
||||
with open(filename, "wb") as f:
|
||||
f.write(struct.pack("<Q", len(hjson)))
|
||||
f.write(hjson)
|
||||
|
||||
for k, v in tensors.items():
|
||||
if v.numel() == 0:
|
||||
continue
|
||||
if v.is_cuda:
|
||||
# Direct GPU to disk save
|
||||
with torch.cuda.device(v.device):
|
||||
if v.dim() == 0: # if scalar, need to add a dimension to work with view
|
||||
v = v.unsqueeze(0)
|
||||
tensor_bytes = v.contiguous().view(torch.uint8)
|
||||
tensor_bytes.cpu().numpy().tofile(f)
|
||||
else:
|
||||
# CPU tensor save
|
||||
if v.dim() == 0: # if scalar, need to add a dimension to work with view
|
||||
v = v.unsqueeze(0)
|
||||
v.contiguous().view(torch.uint8).numpy().tofile(f)
|
||||
|
||||
|
||||
class MemoryEfficientSafeOpen:
|
||||
# does not support metadata loading
|
||||
def __init__(self, filename):
|
||||
@@ -169,6 +321,25 @@ class MemoryEfficientSafeOpen:
|
||||
# return byte_tensor.view(torch.uint8).to(torch.float16).reshape(shape)
|
||||
raise ValueError(f"Unsupported float8 type: {dtype_str} (upgrade PyTorch to support float8 types)")
|
||||
|
||||
def pil_resize(image, size, interpolation=Image.LANCZOS):
|
||||
has_alpha = image.shape[2] == 4 if len(image.shape) == 3 else False
|
||||
|
||||
if has_alpha:
|
||||
pil_image = Image.fromarray(cv2.cvtColor(image, cv2.COLOR_BGRA2RGBA))
|
||||
else:
|
||||
pil_image = Image.fromarray(cv2.cvtColor(image, cv2.COLOR_BGR2RGB))
|
||||
|
||||
resized_pil = pil_image.resize(size, interpolation)
|
||||
|
||||
# Convert back to cv2 format
|
||||
if has_alpha:
|
||||
resized_cv2 = cv2.cvtColor(np.array(resized_pil), cv2.COLOR_RGBA2BGRA)
|
||||
else:
|
||||
resized_cv2 = cv2.cvtColor(np.array(resized_pil), cv2.COLOR_RGB2BGR)
|
||||
|
||||
return resized_cv2
|
||||
|
||||
|
||||
# TODO make inf_utils.py
|
||||
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
+294
-60
@@ -38,6 +38,7 @@ class LoRAModule(torch.nn.Module):
|
||||
dropout=None,
|
||||
rank_dropout=None,
|
||||
module_dropout=None,
|
||||
split_dims: Optional[List[int]] = None,
|
||||
):
|
||||
"""if alpha == 0 or None, alpha is rank (no scaling)."""
|
||||
super().__init__()
|
||||
@@ -51,7 +52,9 @@ class LoRAModule(torch.nn.Module):
|
||||
out_dim = org_module.out_features
|
||||
|
||||
self.lora_dim = lora_dim
|
||||
self.split_dims = split_dims
|
||||
|
||||
if split_dims is None:
|
||||
if org_module.__class__.__name__ == "Conv2d":
|
||||
kernel_size = org_module.kernel_size
|
||||
stride = org_module.stride
|
||||
@@ -62,6 +65,22 @@ class LoRAModule(torch.nn.Module):
|
||||
self.lora_down = torch.nn.Linear(in_dim, self.lora_dim, bias=False)
|
||||
self.lora_up = torch.nn.Linear(self.lora_dim, out_dim, bias=False)
|
||||
|
||||
torch.nn.init.kaiming_uniform_(self.lora_down.weight, a=math.sqrt(5))
|
||||
torch.nn.init.zeros_(self.lora_up.weight)
|
||||
else:
|
||||
# conv2d not supported
|
||||
assert sum(split_dims) == out_dim, "sum of split_dims must be equal to out_dim"
|
||||
assert org_module.__class__.__name__ == "Linear", "split_dims is only supported for Linear"
|
||||
# print(f"split_dims: {split_dims}")
|
||||
self.lora_down = torch.nn.ModuleList(
|
||||
[torch.nn.Linear(in_dim, self.lora_dim, bias=False) for _ in range(len(split_dims))]
|
||||
)
|
||||
self.lora_up = torch.nn.ModuleList([torch.nn.Linear(self.lora_dim, split_dim, bias=False) for split_dim in split_dims])
|
||||
for lora_down in self.lora_down:
|
||||
torch.nn.init.kaiming_uniform_(lora_down.weight, a=math.sqrt(5))
|
||||
for lora_up in self.lora_up:
|
||||
torch.nn.init.zeros_(lora_up.weight)
|
||||
|
||||
if type(alpha) == torch.Tensor:
|
||||
alpha = alpha.detach().float().numpy() # without casting, bf16 causes error
|
||||
alpha = self.lora_dim if alpha is None or alpha == 0 else alpha
|
||||
@@ -69,9 +88,6 @@ class LoRAModule(torch.nn.Module):
|
||||
self.register_buffer("alpha", torch.tensor(alpha)) # 定数として扱える
|
||||
|
||||
# same as microsoft's
|
||||
torch.nn.init.kaiming_uniform_(self.lora_down.weight, a=math.sqrt(5))
|
||||
torch.nn.init.zeros_(self.lora_up.weight)
|
||||
|
||||
self.multiplier = multiplier
|
||||
self.org_module = org_module # remove in applying
|
||||
self.dropout = dropout
|
||||
@@ -91,6 +107,7 @@ class LoRAModule(torch.nn.Module):
|
||||
if torch.rand(1) < self.module_dropout:
|
||||
return org_forwarded
|
||||
|
||||
if self.split_dims is None:
|
||||
lx = self.lora_down(x)
|
||||
|
||||
# normal dropout
|
||||
@@ -115,6 +132,31 @@ class LoRAModule(torch.nn.Module):
|
||||
lx = self.lora_up(lx)
|
||||
|
||||
return org_forwarded + lx * self.multiplier * scale
|
||||
else:
|
||||
lxs = [lora_down(x) for lora_down in self.lora_down]
|
||||
|
||||
# normal dropout
|
||||
if self.dropout is not None and self.training:
|
||||
lxs = [torch.nn.functional.dropout(lx, p=self.dropout) for lx in lxs]
|
||||
|
||||
# rank dropout
|
||||
if self.rank_dropout is not None and self.training:
|
||||
masks = [torch.rand((lx.size(0), self.lora_dim), device=lx.device) > self.rank_dropout for lx in lxs]
|
||||
for i in range(len(lxs)):
|
||||
if len(lx.size()) == 3:
|
||||
masks[i] = masks[i].unsqueeze(1)
|
||||
elif len(lx.size()) == 4:
|
||||
masks[i] = masks[i].unsqueeze(-1).unsqueeze(-1)
|
||||
lxs[i] = lxs[i] * masks[i]
|
||||
|
||||
# scaling for rank dropout: treat as if the rank is changed
|
||||
scale = self.scale * (1.0 / (1.0 - self.rank_dropout)) # redundant for readability
|
||||
else:
|
||||
scale = self.scale
|
||||
|
||||
lxs = [lora_up(lx) for lora_up, lx in zip(self.lora_up, lxs)]
|
||||
|
||||
return org_forwarded + torch.cat(lxs, dim=-1) * self.multiplier * scale
|
||||
|
||||
|
||||
class LoRAInfModule(LoRAModule):
|
||||
@@ -151,9 +193,10 @@ class LoRAInfModule(LoRAModule):
|
||||
if device is None:
|
||||
device = org_device
|
||||
|
||||
if self.split_dims is None:
|
||||
# get up/down weight
|
||||
up_weight = sd["lora_up.weight"].to(torch.float).to(device)
|
||||
down_weight = sd["lora_down.weight"].to(torch.float).to(device)
|
||||
up_weight = sd["lora_up.weight"].to(torch.float).to(device)
|
||||
|
||||
# merge weight
|
||||
if len(weight.size()) == 2:
|
||||
@@ -176,6 +219,24 @@ class LoRAInfModule(LoRAModule):
|
||||
# set weight to org_module
|
||||
org_sd["weight"] = weight.to(dtype)
|
||||
self.org_module.load_state_dict(org_sd)
|
||||
else:
|
||||
# split_dims
|
||||
total_dims = sum(self.split_dims)
|
||||
for i in range(len(self.split_dims)):
|
||||
# get up/down weight
|
||||
down_weight = sd[f"lora_down.{i}.weight"].to(torch.float).to(device) # (rank, in_dim)
|
||||
up_weight = sd[f"lora_up.{i}.weight"].to(torch.float).to(device) # (split dim, rank)
|
||||
|
||||
# pad up_weight -> (total_dims, rank)
|
||||
padded_up_weight = torch.zeros((total_dims, up_weight.size(0)), device=device, dtype=torch.float)
|
||||
padded_up_weight[sum(self.split_dims[:i]) : sum(self.split_dims[: i + 1])] = up_weight
|
||||
|
||||
# merge weight
|
||||
weight = weight + self.multiplier * (up_weight @ down_weight) * self.scale
|
||||
|
||||
# set weight to org_module
|
||||
org_sd["weight"] = weight.to(dtype)
|
||||
self.org_module.load_state_dict(org_sd)
|
||||
|
||||
# 復元できるマージのため、このモジュールのweightを返す
|
||||
def get_weight(self, multiplier=None):
|
||||
@@ -210,7 +271,14 @@ class LoRAInfModule(LoRAModule):
|
||||
|
||||
def default_forward(self, x):
|
||||
# logger.info(f"default_forward {self.lora_name} {x.size()}")
|
||||
return self.org_forward(x) + self.lora_up(self.lora_down(x)) * self.multiplier * self.scale
|
||||
if self.split_dims is None:
|
||||
lx = self.lora_down(x)
|
||||
lx = self.lora_up(lx)
|
||||
return self.org_forward(x) + lx * self.multiplier * self.scale
|
||||
else:
|
||||
lxs = [lora_down(x) for lora_down in self.lora_down]
|
||||
lxs = [lora_up(lx) for lora_up, lx in zip(self.lora_up, lxs)]
|
||||
return self.org_forward(x) + torch.cat(lxs, dim=-1) * self.multiplier * self.scale
|
||||
|
||||
def forward(self, x):
|
||||
if not self.enabled:
|
||||
@@ -256,6 +324,20 @@ def create_network(
|
||||
if train_blocks is not None:
|
||||
assert train_blocks in ["all", "single", "double"], f"invalid train_blocks: {train_blocks}"
|
||||
|
||||
only_if_contains = kwargs.get("only_if_contains", None)
|
||||
if only_if_contains is not None:
|
||||
only_if_contains = [word.strip() for word in only_if_contains.split(',')]
|
||||
|
||||
# split qkv
|
||||
split_qkv = kwargs.get("split_qkv", False)
|
||||
if split_qkv is not None:
|
||||
split_qkv = True if split_qkv == "True" else False
|
||||
|
||||
# train T5XXL
|
||||
train_t5xxl = kwargs.get("train_t5xxl", False)
|
||||
if train_t5xxl is not None:
|
||||
train_t5xxl = True if train_t5xxl == "True" else False
|
||||
|
||||
# すごく引数が多いな ( ^ω^)・・・
|
||||
network = LoRANetwork(
|
||||
text_encoders,
|
||||
@@ -269,7 +351,10 @@ def create_network(
|
||||
conv_lora_dim=conv_dim,
|
||||
conv_alpha=conv_alpha,
|
||||
train_blocks=train_blocks,
|
||||
split_qkv=split_qkv,
|
||||
train_t5xxl=train_t5xxl,
|
||||
varbose=True,
|
||||
only_if_contains=only_if_contains
|
||||
)
|
||||
|
||||
loraplus_lr_ratio = kwargs.get("loraplus_lr_ratio", None)
|
||||
@@ -295,9 +380,10 @@ def create_network_from_weights(multiplier, file, ae, text_encoders, flux, weigh
|
||||
else:
|
||||
weights_sd = torch.load(file, map_location="cpu")
|
||||
|
||||
# get dim/alpha mapping
|
||||
# get dim/alpha mapping, and train t5xxl
|
||||
modules_dim = {}
|
||||
modules_alpha = {}
|
||||
train_t5xxl = None
|
||||
for key, value in weights_sd.items():
|
||||
if "." not in key:
|
||||
continue
|
||||
@@ -310,10 +396,41 @@ def create_network_from_weights(multiplier, file, ae, text_encoders, flux, weigh
|
||||
modules_dim[lora_name] = dim
|
||||
# logger.info(lora_name, value.size(), dim)
|
||||
|
||||
if train_t5xxl is None or train_t5xxl is False:
|
||||
train_t5xxl = "lora_te3" in lora_name
|
||||
|
||||
if train_t5xxl is None:
|
||||
train_t5xxl = False
|
||||
|
||||
# # split qkv
|
||||
# double_qkv_rank = None
|
||||
# single_qkv_rank = None
|
||||
# rank = None
|
||||
# for lora_name, dim in modules_dim.items():
|
||||
# if "double" in lora_name and "qkv" in lora_name:
|
||||
# double_qkv_rank = dim
|
||||
# elif "single" in lora_name and "linear1" in lora_name:
|
||||
# single_qkv_rank = dim
|
||||
# elif rank is None:
|
||||
# rank = dim
|
||||
# if double_qkv_rank is not None and single_qkv_rank is not None and rank is not None:
|
||||
# break
|
||||
# split_qkv = (double_qkv_rank is not None and double_qkv_rank != rank) or (
|
||||
# single_qkv_rank is not None and single_qkv_rank != rank
|
||||
# )
|
||||
split_qkv = False # split_qkv is not needed to care, because state_dict is qkv combined
|
||||
|
||||
module_class = LoRAInfModule if for_inference else LoRAModule
|
||||
|
||||
network = LoRANetwork(
|
||||
text_encoders, flux, multiplier=multiplier, modules_dim=modules_dim, modules_alpha=modules_alpha, module_class=module_class
|
||||
text_encoders,
|
||||
flux,
|
||||
multiplier=multiplier,
|
||||
modules_dim=modules_dim,
|
||||
modules_alpha=modules_alpha,
|
||||
module_class=module_class,
|
||||
split_qkv=split_qkv,
|
||||
train_t5xxl=train_t5xxl,
|
||||
)
|
||||
return network, weights_sd
|
||||
|
||||
@@ -322,10 +439,10 @@ class LoRANetwork(torch.nn.Module):
|
||||
# FLUX_TARGET_REPLACE_MODULE = ["DoubleStreamBlock", "SingleStreamBlock"]
|
||||
FLUX_TARGET_REPLACE_MODULE_DOUBLE = ["DoubleStreamBlock"]
|
||||
FLUX_TARGET_REPLACE_MODULE_SINGLE = ["SingleStreamBlock"]
|
||||
TEXT_ENCODER_TARGET_REPLACE_MODULE = ["CLIPAttention", "CLIPSdpaAttention", "CLIPMLP"]
|
||||
TEXT_ENCODER_TARGET_REPLACE_MODULE = ["CLIPAttention", "CLIPSdpaAttention", "CLIPMLP", "T5Attention", "T5DenseGatedActDense"]
|
||||
LORA_PREFIX_FLUX = "lora_unet" # make ComfyUI compatible
|
||||
LORA_PREFIX_TEXT_ENCODER_CLIP = "lora_te1"
|
||||
LORA_PREFIX_TEXT_ENCODER_T5 = "lora_te2"
|
||||
LORA_PREFIX_TEXT_ENCODER_T5 = "lora_te3" # make ComfyUI compatible
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -343,7 +460,10 @@ class LoRANetwork(torch.nn.Module):
|
||||
modules_dim: Optional[Dict[str, int]] = None,
|
||||
modules_alpha: Optional[Dict[str, int]] = None,
|
||||
train_blocks: Optional[str] = None,
|
||||
split_qkv: bool = False,
|
||||
train_t5xxl: bool = False,
|
||||
varbose: Optional[bool] = False,
|
||||
only_if_contains: Optional[List[str]] = None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.multiplier = multiplier
|
||||
@@ -356,11 +476,15 @@ class LoRANetwork(torch.nn.Module):
|
||||
self.rank_dropout = rank_dropout
|
||||
self.module_dropout = module_dropout
|
||||
self.train_blocks = train_blocks if train_blocks is not None else "all"
|
||||
self.split_qkv = split_qkv
|
||||
self.train_t5xxl = train_t5xxl
|
||||
|
||||
self.loraplus_lr_ratio = None
|
||||
self.loraplus_unet_lr_ratio = None
|
||||
self.loraplus_text_encoder_lr_ratio = None
|
||||
|
||||
self.only_if_contains = only_if_contains
|
||||
|
||||
if modules_dim is not None:
|
||||
logger.info(f"create LoRA network from weights")
|
||||
else:
|
||||
@@ -368,10 +492,18 @@ class LoRANetwork(torch.nn.Module):
|
||||
logger.info(
|
||||
f"neuron dropout: p={self.dropout}, rank dropout: p={self.rank_dropout}, module dropout: p={self.module_dropout}"
|
||||
)
|
||||
if self.conv_lora_dim is not None:
|
||||
logger.info(
|
||||
f"apply LoRA to Conv2d with kernel size (3,3). dim (rank): {self.conv_lora_dim}, alpha: {self.conv_alpha}"
|
||||
)
|
||||
# if self.conv_lora_dim is not None:
|
||||
# logger.info(
|
||||
# f"apply LoRA to Conv2d with kernel size (3,3). dim (rank): {self.conv_lora_dim}, alpha: {self.conv_alpha}"
|
||||
# )
|
||||
if self.split_qkv:
|
||||
logger.info(f"split qkv for LoRA")
|
||||
if self.train_blocks is not None:
|
||||
logger.info(f"train {self.train_blocks} blocks only")
|
||||
if train_t5xxl:
|
||||
logger.info(f"train T5XXL as well")
|
||||
|
||||
#self.only_if_contains = ["lora_unet_single_blocks_20_linear2"]
|
||||
|
||||
# create module instances
|
||||
def create_modules(
|
||||
@@ -395,6 +527,10 @@ class LoRANetwork(torch.nn.Module):
|
||||
if is_linear or is_conv2d:
|
||||
lora_name = prefix + "." + name + "." + child_name
|
||||
lora_name = lora_name.replace(".", "_")
|
||||
#lora_unet_single_blocks_20_linear2
|
||||
|
||||
if "unet" in lora_name and (self.only_if_contains is not None and not any(word in lora_name for word in self.only_if_contains)):
|
||||
continue
|
||||
|
||||
dim = None
|
||||
alpha = None
|
||||
@@ -419,6 +555,14 @@ class LoRANetwork(torch.nn.Module):
|
||||
skipped.append(lora_name)
|
||||
continue
|
||||
|
||||
# qkv split
|
||||
split_dims = None
|
||||
if is_flux and split_qkv:
|
||||
if "double" in lora_name and "qkv" in lora_name:
|
||||
split_dims = [3072] * 3
|
||||
elif "single" in lora_name and "linear1" in lora_name:
|
||||
split_dims = [3072] * 3 + [12288]
|
||||
|
||||
lora = module_class(
|
||||
lora_name,
|
||||
child_module,
|
||||
@@ -428,6 +572,7 @@ class LoRANetwork(torch.nn.Module):
|
||||
dropout=dropout,
|
||||
rank_dropout=rank_dropout,
|
||||
module_dropout=module_dropout,
|
||||
split_dims=split_dims,
|
||||
)
|
||||
loras.append(lora)
|
||||
return loras, skipped
|
||||
@@ -438,12 +583,15 @@ class LoRANetwork(torch.nn.Module):
|
||||
skipped_te = []
|
||||
for i, text_encoder in enumerate(text_encoders):
|
||||
index = i
|
||||
if not train_t5xxl and index > 0: # 0: CLIP, 1: T5XXL, so we skip T5XXL if train_t5xxl is False
|
||||
break
|
||||
|
||||
logger.info(f"create LoRA for Text Encoder {index+1}:")
|
||||
|
||||
text_encoder_loras, skipped = create_modules(False, index, text_encoder, LoRANetwork.TEXT_ENCODER_TARGET_REPLACE_MODULE)
|
||||
logger.info(f"create LoRA for Text Encoder {index+1}: {len(text_encoder_loras)} modules.")
|
||||
self.text_encoder_loras.extend(text_encoder_loras)
|
||||
skipped_te += skipped
|
||||
logger.info(f"create LoRA for Text Encoder: {len(self.text_encoder_loras)} modules.")
|
||||
|
||||
# create LoRA for U-Net
|
||||
if self.train_blocks == "all":
|
||||
@@ -456,6 +604,7 @@ class LoRANetwork(torch.nn.Module):
|
||||
self.unet_loras: List[Union[LoRAModule, LoRAInfModule]]
|
||||
self.unet_loras, skipped_un = create_modules(True, None, unet, target_replace_modules)
|
||||
logger.info(f"create LoRA for FLUX {self.train_blocks} blocks: {len(self.unet_loras)} modules.")
|
||||
#print(self.unet_loras)
|
||||
|
||||
skipped = skipped_te + skipped_un
|
||||
if varbose and len(skipped) > 0:
|
||||
@@ -491,6 +640,111 @@ class LoRANetwork(torch.nn.Module):
|
||||
info = self.load_state_dict(weights_sd, False)
|
||||
return info
|
||||
|
||||
def load_state_dict(self, state_dict, strict=True):
|
||||
# override to convert original weight to split qkv
|
||||
if not self.split_qkv:
|
||||
return super().load_state_dict(state_dict, strict)
|
||||
|
||||
# split qkv
|
||||
for key in list(state_dict.keys()):
|
||||
if "double" in key and "qkv" in key:
|
||||
split_dims = [3072] * 3
|
||||
elif "single" in key and "linear1" in key:
|
||||
split_dims = [3072] * 3 + [12288]
|
||||
else:
|
||||
continue
|
||||
|
||||
weight = state_dict[key]
|
||||
lora_name = key.split(".")[0]
|
||||
if "lora_down" in key and "weight" in key:
|
||||
# dense weight (rank*3, in_dim)
|
||||
split_weight = torch.chunk(weight, len(split_dims), dim=0)
|
||||
for i, split_w in enumerate(split_weight):
|
||||
state_dict[f"{lora_name}.lora_down.{i}.weight"] = split_w
|
||||
|
||||
del state_dict[key]
|
||||
# print(f"split {key}: {weight.shape} to {[w.shape for w in split_weight]}")
|
||||
elif "lora_up" in key and "weight" in key:
|
||||
# sparse weight (out_dim=sum(split_dims), rank*3)
|
||||
rank = weight.size(1) // len(split_dims)
|
||||
i = 0
|
||||
for j in range(len(split_dims)):
|
||||
state_dict[f"{lora_name}.lora_up.{j}.weight"] = weight[i : i + split_dims[j], j * rank : (j + 1) * rank]
|
||||
i += split_dims[j]
|
||||
del state_dict[key]
|
||||
|
||||
# # check is sparse
|
||||
# i = 0
|
||||
# is_zero = True
|
||||
# for j in range(len(split_dims)):
|
||||
# for k in range(len(split_dims)):
|
||||
# if j == k:
|
||||
# continue
|
||||
# is_zero = is_zero and torch.all(weight[i : i + split_dims[j], k * rank : (k + 1) * rank] == 0)
|
||||
# i += split_dims[j]
|
||||
# if not is_zero:
|
||||
# logger.warning(f"weight is not sparse: {key}")
|
||||
# else:
|
||||
# logger.info(f"weight is sparse: {key}")
|
||||
|
||||
# print(
|
||||
# f"split {key}: {weight.shape} to {[state_dict[k].shape for k in [f'{lora_name}.lora_up.{j}.weight' for j in range(len(split_dims))]]}"
|
||||
# )
|
||||
|
||||
# alpha is unchanged
|
||||
|
||||
return super().load_state_dict(state_dict, strict)
|
||||
|
||||
def state_dict(self, destination=None, prefix="", keep_vars=False):
|
||||
if not self.split_qkv:
|
||||
return super().state_dict(destination, prefix, keep_vars)
|
||||
|
||||
# merge qkv
|
||||
state_dict = super().state_dict(destination, prefix, keep_vars)
|
||||
new_state_dict = {}
|
||||
for key in list(state_dict.keys()):
|
||||
if "double" in key and "qkv" in key:
|
||||
split_dims = [3072] * 3
|
||||
elif "single" in key and "linear1" in key:
|
||||
split_dims = [3072] * 3 + [12288]
|
||||
else:
|
||||
new_state_dict[key] = state_dict[key]
|
||||
continue
|
||||
|
||||
if key not in state_dict:
|
||||
continue # already merged
|
||||
|
||||
lora_name = key.split(".")[0]
|
||||
|
||||
# (rank, in_dim) * 3
|
||||
down_weights = [state_dict.pop(f"{lora_name}.lora_down.{i}.weight") for i in range(len(split_dims))]
|
||||
# (split dim, rank) * 3
|
||||
up_weights = [state_dict.pop(f"{lora_name}.lora_up.{i}.weight") for i in range(len(split_dims))]
|
||||
|
||||
alpha = state_dict.pop(f"{lora_name}.alpha")
|
||||
|
||||
# merge down weight
|
||||
down_weight = torch.cat(down_weights, dim=0) # (rank, split_dim) * 3 -> (rank*3, sum of split_dim)
|
||||
|
||||
# merge up weight (sum of split_dim, rank*3)
|
||||
rank = up_weights[0].size(1)
|
||||
up_weight = torch.zeros((sum(split_dims), down_weight.size(0)), device=down_weight.device, dtype=down_weight.dtype)
|
||||
i = 0
|
||||
for j in range(len(split_dims)):
|
||||
up_weight[i : i + split_dims[j], j * rank : (j + 1) * rank] = up_weights[j]
|
||||
i += split_dims[j]
|
||||
|
||||
new_state_dict[f"{lora_name}.lora_down.weight"] = down_weight
|
||||
new_state_dict[f"{lora_name}.lora_up.weight"] = up_weight
|
||||
new_state_dict[f"{lora_name}.alpha"] = alpha
|
||||
|
||||
# print(
|
||||
# f"merged {lora_name}: {lora_name}, {[w.shape for w in down_weights]}, {[w.shape for w in up_weights]} to {down_weight.shape}, {up_weight.shape}"
|
||||
# )
|
||||
print(f"new key: {lora_name}.lora_down.weight, {lora_name}.lora_up.weight, {lora_name}.alpha")
|
||||
|
||||
return new_state_dict
|
||||
|
||||
def apply_to(self, text_encoders, flux, apply_text_encoder=True, apply_unet=True):
|
||||
if apply_text_encoder:
|
||||
logger.info(f"enable LoRA for text encoder: {len(self.text_encoder_loras)} modules")
|
||||
@@ -546,28 +800,26 @@ class LoRANetwork(torch.nn.Module):
|
||||
logger.info(f"LoRA+ UNet LR Ratio: {self.loraplus_unet_lr_ratio or self.loraplus_lr_ratio}")
|
||||
logger.info(f"LoRA+ Text Encoder LR Ratio: {self.loraplus_text_encoder_lr_ratio or self.loraplus_lr_ratio}")
|
||||
|
||||
# 二つのText Encoderに別々の学習率を設定できるようにするといいかも
|
||||
def prepare_optimizer_params(self, text_encoder_lr, unet_lr, default_lr):
|
||||
# TODO warn if optimizer is not compatible with LoRA+ (but it will cause error so we don't need to check it here?)
|
||||
# if (
|
||||
# self.loraplus_lr_ratio is not None
|
||||
# or self.loraplus_text_encoder_lr_ratio is not None
|
||||
# or self.loraplus_unet_lr_ratio is not None
|
||||
# ):
|
||||
# assert (
|
||||
# optimizer_type.lower() != "prodigy" and "dadapt" not in optimizer_type.lower()
|
||||
# ), "LoRA+ and Prodigy/DAdaptation is not supported / LoRA+とProdigy/DAdaptationの組み合わせはサポートされていません"
|
||||
def prepare_optimizer_params_with_multiple_te_lrs(self, text_encoder_lr, unet_lr, default_lr):
|
||||
# make sure text_encoder_lr as list of two elements
|
||||
# if float, use the same value for both text encoders
|
||||
if text_encoder_lr is None or (isinstance(text_encoder_lr, list) and len(text_encoder_lr) == 0):
|
||||
text_encoder_lr = [default_lr, default_lr]
|
||||
elif isinstance(text_encoder_lr, float) or isinstance(text_encoder_lr, int):
|
||||
text_encoder_lr = [float(text_encoder_lr), float(text_encoder_lr)]
|
||||
elif len(text_encoder_lr) == 1:
|
||||
text_encoder_lr = [text_encoder_lr[0], text_encoder_lr[0]]
|
||||
|
||||
self.requires_grad_(True)
|
||||
|
||||
all_params = []
|
||||
lr_descriptions = []
|
||||
|
||||
def assemble_params(loras, lr, ratio):
|
||||
def assemble_params(loras, lr, loraplus_ratio):
|
||||
param_groups = {"lora": {}, "plus": {}}
|
||||
for lora in loras:
|
||||
for name, param in lora.named_parameters():
|
||||
if ratio is not None and "lora_up" in name:
|
||||
if loraplus_ratio is not None and "lora_up" in name:
|
||||
param_groups["plus"][f"{lora.lora_name}.{name}"] = param
|
||||
else:
|
||||
param_groups["lora"][f"{lora.lora_name}.{name}"] = param
|
||||
@@ -582,7 +834,7 @@ class LoRANetwork(torch.nn.Module):
|
||||
|
||||
if lr is not None:
|
||||
if key == "plus":
|
||||
param_data["lr"] = lr * ratio
|
||||
param_data["lr"] = lr * loraplus_ratio
|
||||
else:
|
||||
param_data["lr"] = lr
|
||||
|
||||
@@ -596,41 +848,23 @@ class LoRANetwork(torch.nn.Module):
|
||||
return params, descriptions
|
||||
|
||||
if self.text_encoder_loras:
|
||||
params, descriptions = assemble_params(
|
||||
self.text_encoder_loras,
|
||||
text_encoder_lr if text_encoder_lr is not None else default_lr,
|
||||
self.loraplus_text_encoder_lr_ratio or self.loraplus_lr_ratio,
|
||||
)
|
||||
loraplus_lr_ratio = self.loraplus_text_encoder_lr_ratio or self.loraplus_lr_ratio
|
||||
|
||||
# split text encoder loras for te1 and te3
|
||||
te1_loras = [lora for lora in self.text_encoder_loras if lora.lora_name.startswith(self.LORA_PREFIX_TEXT_ENCODER_CLIP)]
|
||||
te3_loras = [lora for lora in self.text_encoder_loras if lora.lora_name.startswith(self.LORA_PREFIX_TEXT_ENCODER_T5)]
|
||||
if len(te1_loras) > 0:
|
||||
logger.info(f"Text Encoder 1 (CLIP-L): {len(te1_loras)} modules, LR {text_encoder_lr[0]}")
|
||||
params, descriptions = assemble_params(te1_loras, text_encoder_lr[0], loraplus_lr_ratio)
|
||||
all_params.extend(params)
|
||||
lr_descriptions.extend(["textencoder" + (" " + d if d else "") for d in descriptions])
|
||||
lr_descriptions.extend(["textencoder 1 " + (" " + d if d else "") for d in descriptions])
|
||||
if len(te3_loras) > 0:
|
||||
logger.info(f"Text Encoder 2 (T5XXL): {len(te3_loras)} modules, LR {text_encoder_lr[1]}")
|
||||
params, descriptions = assemble_params(te3_loras, text_encoder_lr[1], loraplus_lr_ratio)
|
||||
all_params.extend(params)
|
||||
lr_descriptions.extend(["textencoder 2 " + (" " + d if d else "") for d in descriptions])
|
||||
|
||||
if self.unet_loras:
|
||||
# if self.block_lr:
|
||||
# is_sdxl = False
|
||||
# for lora in self.unet_loras:
|
||||
# if "input_blocks" in lora.lora_name or "output_blocks" in lora.lora_name:
|
||||
# is_sdxl = True
|
||||
# break
|
||||
|
||||
# # 学習率のグラフをblockごとにしたいので、blockごとにloraを分類
|
||||
# block_idx_to_lora = {}
|
||||
# for lora in self.unet_loras:
|
||||
# idx = get_block_index(lora.lora_name, is_sdxl)
|
||||
# if idx not in block_idx_to_lora:
|
||||
# block_idx_to_lora[idx] = []
|
||||
# block_idx_to_lora[idx].append(lora)
|
||||
|
||||
# # blockごとにパラメータを設定する
|
||||
# for idx, block_loras in block_idx_to_lora.items():
|
||||
# params, descriptions = assemble_params(
|
||||
# block_loras,
|
||||
# (unet_lr if unet_lr is not None else default_lr) * self.get_lr_weight(idx),
|
||||
# self.loraplus_unet_lr_ratio or self.loraplus_lr_ratio,
|
||||
# )
|
||||
# all_params.extend(params)
|
||||
# lr_descriptions.extend([f"unet_block{idx}" + (" " + d if d else "") for d in descriptions])
|
||||
|
||||
# else:
|
||||
params, descriptions = assemble_params(
|
||||
self.unet_loras,
|
||||
unet_lr if unet_lr is not None else default_lr,
|
||||
|
||||
@@ -0,0 +1,837 @@
|
||||
# temporary minimum implementation of LoRA
|
||||
# SD3 doesn't have Conv2d, so we ignore it
|
||||
# TODO commonize with the original/SD3/FLUX implementation
|
||||
|
||||
# LoRA network module
|
||||
# reference:
|
||||
# https://github.com/microsoft/LoRA/blob/main/loralib/layers.py
|
||||
# https://github.com/cloneofsimo/lora/blob/master/lora_diffusion/lora.py
|
||||
|
||||
import os
|
||||
from typing import Dict, List, Optional, Tuple, Type, Union
|
||||
from transformers import CLIPTextModelWithProjection, T5EncoderModel
|
||||
import torch
|
||||
from ..library.utils import setup_logging
|
||||
|
||||
setup_logging()
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
from .lora_flux import LoRAModule, LoRAInfModule
|
||||
from ..library import sd3_models
|
||||
|
||||
|
||||
def create_network(
|
||||
multiplier: float,
|
||||
network_dim: Optional[int],
|
||||
network_alpha: Optional[float],
|
||||
vae: sd3_models.SDVAE,
|
||||
text_encoders: List[Union[CLIPTextModelWithProjection, T5EncoderModel]],
|
||||
mmdit,
|
||||
neuron_dropout: Optional[float] = None,
|
||||
**kwargs,
|
||||
):
|
||||
if network_dim is None:
|
||||
network_dim = 4 # default
|
||||
if network_alpha is None:
|
||||
network_alpha = 1.0
|
||||
|
||||
# extract dim/alpha for conv2d, and block dim
|
||||
conv_dim = kwargs.get("conv_dim", None)
|
||||
conv_alpha = kwargs.get("conv_alpha", None)
|
||||
if conv_dim is not None:
|
||||
conv_dim = int(conv_dim)
|
||||
if conv_alpha is None:
|
||||
conv_alpha = 1.0
|
||||
else:
|
||||
conv_alpha = float(conv_alpha)
|
||||
|
||||
# attn dim, mlp dim: only for DoubleStreamBlock. SingleStreamBlock is not supported because of combined qkv
|
||||
context_attn_dim = kwargs.get("context_attn_dim", None)
|
||||
context_mlp_dim = kwargs.get("context_mlp_dim", None)
|
||||
context_mod_dim = kwargs.get("context_mod_dim", None)
|
||||
x_attn_dim = kwargs.get("x_attn_dim", None)
|
||||
x_mlp_dim = kwargs.get("x_mlp_dim", None)
|
||||
x_mod_dim = kwargs.get("x_mod_dim", None)
|
||||
if context_attn_dim is not None:
|
||||
context_attn_dim = int(context_attn_dim)
|
||||
if context_mlp_dim is not None:
|
||||
context_mlp_dim = int(context_mlp_dim)
|
||||
if context_mod_dim is not None:
|
||||
context_mod_dim = int(context_mod_dim)
|
||||
if x_attn_dim is not None:
|
||||
x_attn_dim = int(x_attn_dim)
|
||||
if x_mlp_dim is not None:
|
||||
x_mlp_dim = int(x_mlp_dim)
|
||||
if x_mod_dim is not None:
|
||||
x_mod_dim = int(x_mod_dim)
|
||||
type_dims = [context_attn_dim, context_mlp_dim, context_mod_dim, x_attn_dim, x_mlp_dim, x_mod_dim]
|
||||
if all([d is None for d in type_dims]):
|
||||
type_dims = None
|
||||
|
||||
# emb_dims [context_embedder, t_embedder, x_embedder, y_embedder, final_mod, final_linear]
|
||||
emb_dims = kwargs.get("emb_dims", None)
|
||||
if emb_dims is not None:
|
||||
emb_dims = emb_dims.strip()
|
||||
if emb_dims.startswith("[") and emb_dims.endswith("]"):
|
||||
emb_dims = emb_dims[1:-1]
|
||||
emb_dims = [int(d) for d in emb_dims.split(",")] # is it better to use ast.literal_eval?
|
||||
assert len(emb_dims) == 6, f"invalid emb_dims: {emb_dims}, must be 6 dimensions (context, t, x, y, final_mod, final_linear)"
|
||||
|
||||
# double/single train blocks
|
||||
def parse_block_selection(selection: str, total_blocks: int) -> List[bool]:
|
||||
"""
|
||||
Parse a block selection string and return a list of booleans.
|
||||
|
||||
Args:
|
||||
selection (str): A string specifying which blocks to select.
|
||||
total_blocks (int): The total number of blocks available.
|
||||
|
||||
Returns:
|
||||
List[bool]: A list of booleans indicating which blocks are selected.
|
||||
"""
|
||||
if selection == "all":
|
||||
return [True] * total_blocks
|
||||
if selection == "none" or selection == "":
|
||||
return [False] * total_blocks
|
||||
|
||||
selected = [False] * total_blocks
|
||||
ranges = selection.split(",")
|
||||
|
||||
for r in ranges:
|
||||
if "-" in r:
|
||||
start, end = map(str.strip, r.split("-"))
|
||||
start = int(start)
|
||||
end = int(end)
|
||||
assert 0 <= start < total_blocks, f"invalid start index: {start}"
|
||||
assert 0 <= end < total_blocks, f"invalid end index: {end}"
|
||||
assert start <= end, f"invalid range: {start}-{end}"
|
||||
for i in range(start, end + 1):
|
||||
selected[i] = True
|
||||
else:
|
||||
index = int(r)
|
||||
assert 0 <= index < total_blocks, f"invalid index: {index}"
|
||||
selected[index] = True
|
||||
|
||||
return selected
|
||||
|
||||
train_block_indices = kwargs.get("train_block_indices", None)
|
||||
if train_block_indices is not None:
|
||||
train_block_indices = parse_block_selection(train_block_indices, 999) # 999 is a dummy number
|
||||
|
||||
# rank/module dropout
|
||||
rank_dropout = kwargs.get("rank_dropout", None)
|
||||
if rank_dropout is not None:
|
||||
rank_dropout = float(rank_dropout)
|
||||
module_dropout = kwargs.get("module_dropout", None)
|
||||
if module_dropout is not None:
|
||||
module_dropout = float(module_dropout)
|
||||
|
||||
# split qkv
|
||||
split_qkv = kwargs.get("split_qkv", False)
|
||||
if split_qkv is not None:
|
||||
split_qkv = True if split_qkv == "True" else False
|
||||
|
||||
# train T5XXL
|
||||
train_t5xxl = kwargs.get("train_t5xxl", False)
|
||||
if train_t5xxl is not None:
|
||||
train_t5xxl = True if train_t5xxl == "True" else False
|
||||
|
||||
# verbose
|
||||
verbose = kwargs.get("verbose", False)
|
||||
if verbose is not None:
|
||||
verbose = True if verbose == "True" else False
|
||||
|
||||
# すごく引数が多いな ( ^ω^)・・・
|
||||
network = LoRANetwork(
|
||||
text_encoders,
|
||||
mmdit,
|
||||
multiplier=multiplier,
|
||||
lora_dim=network_dim,
|
||||
alpha=network_alpha,
|
||||
dropout=neuron_dropout,
|
||||
rank_dropout=rank_dropout,
|
||||
module_dropout=module_dropout,
|
||||
conv_lora_dim=conv_dim,
|
||||
conv_alpha=conv_alpha,
|
||||
split_qkv=split_qkv,
|
||||
train_t5xxl=train_t5xxl,
|
||||
type_dims=type_dims,
|
||||
emb_dims=emb_dims,
|
||||
train_block_indices=train_block_indices,
|
||||
verbose=verbose,
|
||||
)
|
||||
|
||||
loraplus_lr_ratio = kwargs.get("loraplus_lr_ratio", None)
|
||||
loraplus_unet_lr_ratio = kwargs.get("loraplus_unet_lr_ratio", None)
|
||||
loraplus_text_encoder_lr_ratio = kwargs.get("loraplus_text_encoder_lr_ratio", None)
|
||||
loraplus_lr_ratio = float(loraplus_lr_ratio) if loraplus_lr_ratio is not None else None
|
||||
loraplus_unet_lr_ratio = float(loraplus_unet_lr_ratio) if loraplus_unet_lr_ratio is not None else None
|
||||
loraplus_text_encoder_lr_ratio = float(loraplus_text_encoder_lr_ratio) if loraplus_text_encoder_lr_ratio is not None else None
|
||||
if loraplus_lr_ratio is not None or loraplus_unet_lr_ratio is not None or loraplus_text_encoder_lr_ratio is not None:
|
||||
network.set_loraplus_lr_ratio(loraplus_lr_ratio, loraplus_unet_lr_ratio, loraplus_text_encoder_lr_ratio)
|
||||
|
||||
return network
|
||||
|
||||
|
||||
# Create network from weights for inference, weights are not loaded here (because can be merged)
|
||||
def create_network_from_weights(multiplier, file, ae, text_encoders, mmdit, weights_sd=None, for_inference=False, **kwargs):
|
||||
# if unet is an instance of SdxlUNet2DConditionModel or subclass, set is_sdxl to True
|
||||
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, and train t5xxl
|
||||
modules_dim = {}
|
||||
modules_alpha = {}
|
||||
train_t5xxl = None
|
||||
for key, value in weights_sd.items():
|
||||
if "." not in key:
|
||||
continue
|
||||
|
||||
lora_name = key.split(".")[0]
|
||||
if "alpha" in key:
|
||||
modules_alpha[lora_name] = value
|
||||
elif "lora_down" in key:
|
||||
dim = value.size()[0]
|
||||
modules_dim[lora_name] = dim
|
||||
# logger.info(lora_name, value.size(), dim)
|
||||
|
||||
if train_t5xxl is None or train_t5xxl is False:
|
||||
train_t5xxl = "lora_te3" in lora_name
|
||||
|
||||
if train_t5xxl is None:
|
||||
train_t5xxl = False
|
||||
|
||||
split_qkv = False # split_qkv is not needed to care, because state_dict is qkv combined
|
||||
|
||||
module_class = LoRAInfModule if for_inference else LoRAModule
|
||||
|
||||
network = LoRANetwork(
|
||||
text_encoders,
|
||||
mmdit,
|
||||
multiplier=multiplier,
|
||||
modules_dim=modules_dim,
|
||||
modules_alpha=modules_alpha,
|
||||
module_class=module_class,
|
||||
split_qkv=split_qkv,
|
||||
train_t5xxl=train_t5xxl,
|
||||
)
|
||||
return network, weights_sd
|
||||
|
||||
|
||||
class LoRANetwork(torch.nn.Module):
|
||||
SD3_TARGET_REPLACE_MODULE = ["SingleDiTBlock"]
|
||||
TEXT_ENCODER_TARGET_REPLACE_MODULE = ["CLIPAttention", "CLIPSdpaAttention", "CLIPMLP", "T5Attention", "T5DenseGatedActDense"]
|
||||
LORA_PREFIX_SD3 = "lora_unet" # make ComfyUI compatible
|
||||
LORA_PREFIX_TEXT_ENCODER_CLIP_L = "lora_te1"
|
||||
LORA_PREFIX_TEXT_ENCODER_CLIP_G = "lora_te2"
|
||||
LORA_PREFIX_TEXT_ENCODER_T5 = "lora_te3" # make ComfyUI compatible
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
text_encoders: List[Union[CLIPTextModelWithProjection, T5EncoderModel]],
|
||||
unet: sd3_models.MMDiT,
|
||||
multiplier: float = 1.0,
|
||||
lora_dim: int = 4,
|
||||
alpha: float = 1,
|
||||
dropout: Optional[float] = None,
|
||||
rank_dropout: Optional[float] = None,
|
||||
module_dropout: Optional[float] = None,
|
||||
conv_lora_dim: Optional[int] = None,
|
||||
conv_alpha: Optional[float] = None,
|
||||
module_class: Type[object] = LoRAModule,
|
||||
modules_dim: Optional[Dict[str, int]] = None,
|
||||
modules_alpha: Optional[Dict[str, int]] = None,
|
||||
split_qkv: bool = False,
|
||||
train_t5xxl: bool = False,
|
||||
type_dims: Optional[List[int]] = None,
|
||||
emb_dims: Optional[List[int]] = None,
|
||||
train_block_indices: Optional[List[bool]] = None,
|
||||
verbose: Optional[bool] = False,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.multiplier = multiplier
|
||||
|
||||
self.lora_dim = lora_dim
|
||||
self.alpha = alpha
|
||||
self.conv_lora_dim = conv_lora_dim
|
||||
self.conv_alpha = conv_alpha
|
||||
self.dropout = dropout
|
||||
self.rank_dropout = rank_dropout
|
||||
self.module_dropout = module_dropout
|
||||
self.split_qkv = split_qkv
|
||||
self.train_t5xxl = train_t5xxl
|
||||
|
||||
self.type_dims = type_dims
|
||||
self.emb_dims = emb_dims
|
||||
self.train_block_indices = train_block_indices
|
||||
|
||||
self.loraplus_lr_ratio = None
|
||||
self.loraplus_unet_lr_ratio = None
|
||||
self.loraplus_text_encoder_lr_ratio = None
|
||||
|
||||
if modules_dim is not None:
|
||||
logger.info(f"create LoRA network from weights")
|
||||
self.emb_dims = [0] * 6 # create emb_dims
|
||||
# verbose = True
|
||||
else:
|
||||
logger.info(f"create LoRA network. base dim (rank): {lora_dim}, alpha: {alpha}")
|
||||
logger.info(
|
||||
f"neuron dropout: p={self.dropout}, rank dropout: p={self.rank_dropout}, module dropout: p={self.module_dropout}"
|
||||
)
|
||||
# if self.conv_lora_dim is not None:
|
||||
# logger.info(
|
||||
# f"apply LoRA to Conv2d with kernel size (3,3). dim (rank): {self.conv_lora_dim}, alpha: {self.conv_alpha}"
|
||||
# )
|
||||
|
||||
qkv_dim = 0
|
||||
if self.split_qkv:
|
||||
logger.info(f"split qkv for LoRA")
|
||||
qkv_dim = unet.joint_blocks[0].context_block.attn.qkv.weight.size(0)
|
||||
if train_t5xxl:
|
||||
logger.info(f"train T5XXL as well")
|
||||
|
||||
# create module instances
|
||||
def create_modules(
|
||||
is_mmdit: bool,
|
||||
text_encoder_idx: Optional[int],
|
||||
root_module: torch.nn.Module,
|
||||
target_replace_modules: List[str],
|
||||
filter: Optional[str] = None,
|
||||
default_dim: Optional[int] = None,
|
||||
include_conv2d_if_filter: bool = False,
|
||||
) -> List[LoRAModule]:
|
||||
prefix = (
|
||||
self.LORA_PREFIX_SD3
|
||||
if is_mmdit
|
||||
else [self.LORA_PREFIX_TEXT_ENCODER_CLIP_L, self.LORA_PREFIX_TEXT_ENCODER_CLIP_G, self.LORA_PREFIX_TEXT_ENCODER_T5][
|
||||
text_encoder_idx
|
||||
]
|
||||
)
|
||||
|
||||
loras = []
|
||||
skipped = []
|
||||
for name, module in root_module.named_modules():
|
||||
if target_replace_modules is None or module.__class__.__name__ in target_replace_modules:
|
||||
if target_replace_modules is None: # dirty hack for all modules
|
||||
module = root_module # search all modules
|
||||
|
||||
for child_name, child_module in module.named_modules():
|
||||
is_linear = child_module.__class__.__name__ == "Linear"
|
||||
is_conv2d = child_module.__class__.__name__ == "Conv2d"
|
||||
is_conv2d_1x1 = is_conv2d and child_module.kernel_size == (1, 1)
|
||||
|
||||
if is_linear or is_conv2d:
|
||||
lora_name = prefix + "." + (name + "." if name else "") + child_name
|
||||
lora_name = lora_name.replace(".", "_")
|
||||
|
||||
force_incl_conv2d = False
|
||||
if filter is not None:
|
||||
if not filter in lora_name:
|
||||
continue
|
||||
force_incl_conv2d = include_conv2d_if_filter
|
||||
|
||||
dim = None
|
||||
alpha = None
|
||||
|
||||
if modules_dim is not None:
|
||||
# モジュール指定あり
|
||||
if lora_name in modules_dim:
|
||||
dim = modules_dim[lora_name]
|
||||
alpha = modules_alpha[lora_name]
|
||||
else:
|
||||
# 通常、すべて対象とする
|
||||
if is_linear or is_conv2d_1x1:
|
||||
dim = default_dim if default_dim is not None else self.lora_dim
|
||||
alpha = self.alpha
|
||||
|
||||
if is_mmdit and type_dims is not None:
|
||||
# type_dims = [context_attn_dim, context_mlp_dim, context_mod_dim, x_attn_dim, x_mlp_dim, x_mod_dim]
|
||||
identifier = [
|
||||
("context_block", "attn"),
|
||||
("context_block", "mlp"),
|
||||
("context_block", "adaLN_modulation"),
|
||||
("x_block", "attn"),
|
||||
("x_block", "mlp"),
|
||||
("x_block", "adaLN_modulation"),
|
||||
]
|
||||
for i, d in enumerate(type_dims):
|
||||
if d is not None and all([id in lora_name for id in identifier[i]]):
|
||||
dim = d # may be 0 for skip
|
||||
break
|
||||
|
||||
if is_mmdit and dim and self.train_block_indices is not None and "joint_blocks" in lora_name:
|
||||
# "lora_unet_joint_blocks_0_x_block_attn_proj..."
|
||||
block_index = int(lora_name.split("_")[4]) # bit dirty
|
||||
if self.train_block_indices is not None and not self.train_block_indices[block_index]:
|
||||
dim = 0
|
||||
|
||||
elif self.conv_lora_dim is not None:
|
||||
dim = self.conv_lora_dim
|
||||
alpha = self.conv_alpha
|
||||
elif force_incl_conv2d:
|
||||
# x_embedder
|
||||
dim = default_dim if default_dim is not None else self.lora_dim
|
||||
alpha = self.alpha
|
||||
|
||||
if dim is None or dim == 0:
|
||||
# skipした情報を出力
|
||||
if is_linear or is_conv2d_1x1 or (self.conv_lora_dim is not None):
|
||||
skipped.append(lora_name)
|
||||
continue
|
||||
|
||||
# qkv split
|
||||
split_dims = None
|
||||
if is_mmdit and split_qkv:
|
||||
if "joint_blocks" in lora_name and "qkv" in lora_name:
|
||||
split_dims = [qkv_dim // 3] * 3
|
||||
|
||||
lora = module_class(
|
||||
lora_name,
|
||||
child_module,
|
||||
self.multiplier,
|
||||
dim,
|
||||
alpha,
|
||||
dropout=dropout,
|
||||
rank_dropout=rank_dropout,
|
||||
module_dropout=module_dropout,
|
||||
split_dims=split_dims,
|
||||
)
|
||||
loras.append(lora)
|
||||
|
||||
if target_replace_modules is None:
|
||||
break # all modules are searched
|
||||
return loras, skipped
|
||||
|
||||
# create LoRA for text encoder
|
||||
# 毎回すべてのモジュールを作るのは無駄なので要検討
|
||||
self.text_encoder_loras: List[Union[LoRAModule, LoRAInfModule]] = []
|
||||
skipped_te = []
|
||||
for i, text_encoder in enumerate(text_encoders):
|
||||
index = i
|
||||
if not train_t5xxl and index >= 2: # 0: CLIP-L, 1: CLIP-G, 2: T5XXL, so we skip T5XXL if train_t5xxl is False
|
||||
break
|
||||
|
||||
logger.info(f"create LoRA for Text Encoder {index+1}:")
|
||||
|
||||
text_encoder_loras, skipped = create_modules(False, index, text_encoder, LoRANetwork.TEXT_ENCODER_TARGET_REPLACE_MODULE)
|
||||
logger.info(f"create LoRA for Text Encoder {index+1}: {len(text_encoder_loras)} modules.")
|
||||
self.text_encoder_loras.extend(text_encoder_loras)
|
||||
skipped_te += skipped
|
||||
|
||||
# create LoRA for U-Net
|
||||
self.unet_loras: List[Union[LoRAModule, LoRAInfModule]]
|
||||
self.unet_loras, skipped_un = create_modules(True, None, unet, LoRANetwork.SD3_TARGET_REPLACE_MODULE)
|
||||
|
||||
# emb_dims [context_embedder, t_embedder, x_embedder, y_embedder, final_mod, final_linear]
|
||||
if self.emb_dims:
|
||||
for filter, in_dim in zip(
|
||||
[
|
||||
"context_embedder",
|
||||
"_t_embedder", # don't use "t_embedder" because it's used in "context_embedder"
|
||||
"x_embedder",
|
||||
"y_embedder",
|
||||
"final_layer_adaLN_modulation",
|
||||
"final_layer_linear",
|
||||
],
|
||||
self.emb_dims,
|
||||
):
|
||||
# x_embedder is conv2d, so we need to include it
|
||||
loras, _ = create_modules(
|
||||
True, None, unet, None, filter=filter, default_dim=in_dim, include_conv2d_if_filter=filter == "x_embedder"
|
||||
)
|
||||
# if len(loras) > 0:
|
||||
# logger.info(f"create LoRA for {filter}: {len(loras)} modules.")
|
||||
self.unet_loras.extend(loras)
|
||||
|
||||
logger.info(f"create LoRA for SD3 MMDiT: {len(self.unet_loras)} modules.")
|
||||
if verbose:
|
||||
for lora in self.unet_loras:
|
||||
logger.info(f"\t{lora.lora_name:50} {lora.lora_dim}, {lora.alpha}")
|
||||
|
||||
skipped = skipped_te + skipped_un
|
||||
if verbose and len(skipped) > 0:
|
||||
logger.warning(
|
||||
f"because dim (rank) is 0, {len(skipped)} LoRA modules are skipped / dim (rank)が0の為、次の{len(skipped)}個のLoRAモジュールはスキップされます:"
|
||||
)
|
||||
for name in skipped:
|
||||
logger.info(f"\t{name}")
|
||||
|
||||
# assertion
|
||||
names = set()
|
||||
for lora in self.text_encoder_loras + self.unet_loras:
|
||||
assert lora.lora_name not in names, f"duplicated lora name: {lora.lora_name}"
|
||||
names.add(lora.lora_name)
|
||||
|
||||
def set_multiplier(self, multiplier):
|
||||
self.multiplier = multiplier
|
||||
for lora in self.text_encoder_loras + self.unet_loras:
|
||||
lora.multiplier = self.multiplier
|
||||
|
||||
def set_enabled(self, is_enabled):
|
||||
for lora in self.text_encoder_loras + self.unet_loras:
|
||||
lora.enabled = is_enabled
|
||||
|
||||
def load_weights(self, file):
|
||||
if os.path.splitext(file)[1] == ".safetensors":
|
||||
from safetensors.torch import load_file
|
||||
|
||||
weights_sd = load_file(file)
|
||||
else:
|
||||
weights_sd = torch.load(file, map_location="cpu")
|
||||
|
||||
info = self.load_state_dict(weights_sd, False)
|
||||
return info
|
||||
|
||||
def load_state_dict(self, state_dict, strict=True):
|
||||
# override to convert original weight to split qkv
|
||||
if not self.split_qkv:
|
||||
return super().load_state_dict(state_dict, strict)
|
||||
|
||||
# split qkv
|
||||
for key in list(state_dict.keys()):
|
||||
if not ("joint_blocks" in key and "qkv" in key):
|
||||
continue
|
||||
|
||||
weight = state_dict[key]
|
||||
lora_name = key.split(".")[0]
|
||||
if "lora_down" in key and "weight" in key:
|
||||
# dense weight (rank*3, in_dim)
|
||||
split_weight = torch.chunk(weight, 3, dim=0)
|
||||
for i, split_w in enumerate(split_weight):
|
||||
state_dict[f"{lora_name}.lora_down.{i}.weight"] = split_w
|
||||
|
||||
del state_dict[key]
|
||||
# print(f"split {key}: {weight.shape} to {[w.shape for w in split_weight]}")
|
||||
elif "lora_up" in key and "weight" in key:
|
||||
# sparse weight (out_dim=sum(split_dims), rank*3)
|
||||
rank = weight.size(1) // 3
|
||||
i = 0
|
||||
split_dim = weight.shape[0] // 3
|
||||
for j in range(3):
|
||||
state_dict[f"{lora_name}.lora_up.{j}.weight"] = weight[i : i + split_dim, j * rank : (j + 1) * rank]
|
||||
i += split_dim
|
||||
del state_dict[key]
|
||||
|
||||
# alpha is unchanged
|
||||
|
||||
return super().load_state_dict(state_dict, strict)
|
||||
|
||||
def state_dict(self, destination=None, prefix="", keep_vars=False):
|
||||
if not self.split_qkv:
|
||||
return super().state_dict(destination, prefix, keep_vars)
|
||||
|
||||
# merge qkv
|
||||
state_dict = super().state_dict(destination, prefix, keep_vars)
|
||||
new_state_dict = {}
|
||||
for key in list(state_dict.keys()):
|
||||
if not ("joint_blocks" in key and "qkv" in key):
|
||||
new_state_dict[key] = state_dict[key]
|
||||
continue
|
||||
|
||||
if key not in state_dict:
|
||||
continue # already merged
|
||||
|
||||
lora_name = key.split(".")[0]
|
||||
|
||||
# (rank, in_dim) * 3
|
||||
down_weights = [state_dict.pop(f"{lora_name}.lora_down.{i}.weight") for i in range(3)]
|
||||
# (split dim, rank) * 3
|
||||
up_weights = [state_dict.pop(f"{lora_name}.lora_up.{i}.weight") for i in range(3)]
|
||||
|
||||
alpha = state_dict.pop(f"{lora_name}.alpha")
|
||||
|
||||
# merge down weight
|
||||
down_weight = torch.cat(down_weights, dim=0) # (rank, split_dim) * 3 -> (rank*3, sum of split_dim)
|
||||
|
||||
# merge up weight (sum of split_dim, rank*3)
|
||||
split_dim, rank = up_weights[0].size()
|
||||
qkv_dim = split_dim * 3
|
||||
up_weight = torch.zeros((qkv_dim, down_weight.size(0)), device=down_weight.device, dtype=down_weight.dtype)
|
||||
i = 0
|
||||
for j in range(3):
|
||||
up_weight[i : i + split_dim, j * rank : (j + 1) * rank] = up_weights[j]
|
||||
i += split_dim
|
||||
|
||||
new_state_dict[f"{lora_name}.lora_down.weight"] = down_weight
|
||||
new_state_dict[f"{lora_name}.lora_up.weight"] = up_weight
|
||||
new_state_dict[f"{lora_name}.alpha"] = alpha
|
||||
|
||||
# print(
|
||||
# f"merged {lora_name}: {lora_name}, {[w.shape for w in down_weights]}, {[w.shape for w in up_weights]} to {down_weight.shape}, {up_weight.shape}"
|
||||
# )
|
||||
print(f"new key: {lora_name}.lora_down.weight, {lora_name}.lora_up.weight, {lora_name}.alpha")
|
||||
|
||||
return new_state_dict
|
||||
|
||||
def apply_to(self, text_encoders, mmdit, apply_text_encoder=True, apply_unet=True):
|
||||
if apply_text_encoder:
|
||||
logger.info(f"enable LoRA for text encoder: {len(self.text_encoder_loras)} modules")
|
||||
else:
|
||||
self.text_encoder_loras = []
|
||||
|
||||
if apply_unet:
|
||||
logger.info(f"enable LoRA for U-Net: {len(self.unet_loras)} modules")
|
||||
else:
|
||||
self.unet_loras = []
|
||||
|
||||
for lora in self.text_encoder_loras + self.unet_loras:
|
||||
lora.apply_to()
|
||||
self.add_module(lora.lora_name, lora)
|
||||
|
||||
# マージできるかどうかを返す
|
||||
def is_mergeable(self):
|
||||
return True
|
||||
|
||||
# TODO refactor to common function with apply_to
|
||||
def merge_to(self, text_encoders, mmdit, weights_sd, dtype=None, device=None):
|
||||
apply_text_encoder = apply_unet = False
|
||||
for key in weights_sd.keys():
|
||||
if (
|
||||
key.startswith(LoRANetwork.LORA_PREFIX_TEXT_ENCODER_CLIP_L)
|
||||
or key.startswith(LoRANetwork.LORA_PREFIX_TEXT_ENCODER_CLIP_G)
|
||||
or key.startswith(LoRANetwork.LORA_PREFIX_TEXT_ENCODER_T5)
|
||||
):
|
||||
apply_text_encoder = True
|
||||
elif key.startswith(LoRANetwork.LORA_PREFIX_SD3):
|
||||
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 = []
|
||||
|
||||
for lora in self.text_encoder_loras + self.unet_loras:
|
||||
sd_for_lora = {}
|
||||
for key in weights_sd.keys():
|
||||
if key.startswith(lora.lora_name):
|
||||
sd_for_lora[key[len(lora.lora_name) + 1 :]] = weights_sd[key]
|
||||
lora.merge_to(sd_for_lora, dtype, device)
|
||||
|
||||
logger.info(f"weights are merged")
|
||||
|
||||
def set_loraplus_lr_ratio(self, loraplus_lr_ratio, loraplus_unet_lr_ratio, loraplus_text_encoder_lr_ratio):
|
||||
self.loraplus_lr_ratio = loraplus_lr_ratio
|
||||
self.loraplus_unet_lr_ratio = loraplus_unet_lr_ratio
|
||||
self.loraplus_text_encoder_lr_ratio = loraplus_text_encoder_lr_ratio
|
||||
|
||||
logger.info(f"LoRA+ UNet LR Ratio: {self.loraplus_unet_lr_ratio or self.loraplus_lr_ratio}")
|
||||
logger.info(f"LoRA+ Text Encoder LR Ratio: {self.loraplus_text_encoder_lr_ratio or self.loraplus_lr_ratio}")
|
||||
|
||||
def prepare_optimizer_params_with_multiple_te_lrs(self, text_encoder_lr, unet_lr, default_lr):
|
||||
# make sure text_encoder_lr as list of three elements
|
||||
# if float, use the same value for all three
|
||||
if text_encoder_lr is None or (isinstance(text_encoder_lr, list) and len(text_encoder_lr) == 0):
|
||||
text_encoder_lr = [default_lr, default_lr, default_lr]
|
||||
elif isinstance(text_encoder_lr, float) or isinstance(text_encoder_lr, int):
|
||||
text_encoder_lr = [float(text_encoder_lr), float(text_encoder_lr), float(text_encoder_lr)]
|
||||
elif len(text_encoder_lr) == 1:
|
||||
text_encoder_lr = [text_encoder_lr[0], text_encoder_lr[0], text_encoder_lr[0]]
|
||||
elif len(text_encoder_lr) == 2:
|
||||
text_encoder_lr = [text_encoder_lr[0], text_encoder_lr[1], text_encoder_lr[1]]
|
||||
|
||||
self.requires_grad_(True)
|
||||
|
||||
all_params = []
|
||||
lr_descriptions = []
|
||||
|
||||
def assemble_params(loras, lr, loraplus_ratio):
|
||||
param_groups = {"lora": {}, "plus": {}}
|
||||
for lora in loras:
|
||||
for name, param in lora.named_parameters():
|
||||
if loraplus_ratio is not None and "lora_up" in name:
|
||||
param_groups["plus"][f"{lora.lora_name}.{name}"] = param
|
||||
else:
|
||||
param_groups["lora"][f"{lora.lora_name}.{name}"] = param
|
||||
|
||||
params = []
|
||||
descriptions = []
|
||||
for key in param_groups.keys():
|
||||
param_data = {"params": param_groups[key].values()}
|
||||
|
||||
if len(param_data["params"]) == 0:
|
||||
continue
|
||||
|
||||
if lr is not None:
|
||||
if key == "plus":
|
||||
param_data["lr"] = lr * loraplus_ratio
|
||||
else:
|
||||
param_data["lr"] = lr
|
||||
|
||||
if param_data.get("lr", None) == 0 or param_data.get("lr", None) is None:
|
||||
logger.info("NO LR skipping!")
|
||||
continue
|
||||
|
||||
params.append(param_data)
|
||||
descriptions.append("plus" if key == "plus" else "")
|
||||
|
||||
return params, descriptions
|
||||
|
||||
if self.text_encoder_loras:
|
||||
loraplus_lr_ratio = self.loraplus_text_encoder_lr_ratio or self.loraplus_lr_ratio
|
||||
|
||||
# split text encoder loras for te1 and te3
|
||||
te1_loras = [
|
||||
lora for lora in self.text_encoder_loras if lora.lora_name.startswith(self.LORA_PREFIX_TEXT_ENCODER_CLIP_L)
|
||||
]
|
||||
te2_loras = [
|
||||
lora for lora in self.text_encoder_loras if lora.lora_name.startswith(self.LORA_PREFIX_TEXT_ENCODER_CLIP_G)
|
||||
]
|
||||
te3_loras = [lora for lora in self.text_encoder_loras if lora.lora_name.startswith(self.LORA_PREFIX_TEXT_ENCODER_T5)]
|
||||
if len(te1_loras) > 0:
|
||||
logger.info(f"Text Encoder 1 (CLIP-L): {len(te1_loras)} modules, LR {text_encoder_lr[0]}")
|
||||
params, descriptions = assemble_params(te1_loras, text_encoder_lr[0], loraplus_lr_ratio)
|
||||
all_params.extend(params)
|
||||
lr_descriptions.extend(["textencoder 1 " + (" " + d if d else "") for d in descriptions])
|
||||
if len(te2_loras) > 0:
|
||||
logger.info(f"Text Encoder 2 (CLIP-G): {len(te2_loras)} modules, LR {text_encoder_lr[1]}")
|
||||
params, descriptions = assemble_params(te2_loras, text_encoder_lr[1], loraplus_lr_ratio)
|
||||
all_params.extend(params)
|
||||
lr_descriptions.extend(["textencoder 1 " + (" " + d if d else "") for d in descriptions])
|
||||
if len(te3_loras) > 0:
|
||||
logger.info(f"Text Encoder 3 (T5XXL): {len(te3_loras)} modules, LR {text_encoder_lr[2]}")
|
||||
params, descriptions = assemble_params(te3_loras, text_encoder_lr[2], loraplus_lr_ratio)
|
||||
all_params.extend(params)
|
||||
lr_descriptions.extend(["textencoder 3 " + (" " + d if d else "") for d in descriptions])
|
||||
|
||||
if self.unet_loras:
|
||||
params, descriptions = assemble_params(
|
||||
self.unet_loras,
|
||||
unet_lr if unet_lr is not None else default_lr,
|
||||
self.loraplus_unet_lr_ratio or self.loraplus_lr_ratio,
|
||||
)
|
||||
all_params.extend(params)
|
||||
lr_descriptions.extend(["unet" + (" " + d if d else "") for d in descriptions])
|
||||
|
||||
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 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
|
||||
from ..library import train_util
|
||||
|
||||
# Precalculate model hashes to save time on indexing
|
||||
if metadata is None:
|
||||
metadata = {}
|
||||
model_hash, legacy_hash = train_util.precalculate_safetensors_hashes(state_dict, metadata)
|
||||
metadata["sshs_model_hash"] = model_hash
|
||||
metadata["sshs_legacy_hash"] = legacy_hash
|
||||
|
||||
save_file(state_dict, file, metadata)
|
||||
else:
|
||||
torch.save(state_dict, file)
|
||||
|
||||
def backup_weights(self):
|
||||
# 重みのバックアップを行う
|
||||
loras: List[LoRAInfModule] = self.text_encoder_loras + self.unet_loras
|
||||
for lora in loras:
|
||||
org_module = lora.org_module_ref[0]
|
||||
if not hasattr(org_module, "_lora_org_weight"):
|
||||
sd = org_module.state_dict()
|
||||
org_module._lora_org_weight = sd["weight"].detach().clone()
|
||||
org_module._lora_restored = True
|
||||
|
||||
def restore_weights(self):
|
||||
# 重みのリストアを行う
|
||||
loras: List[LoRAInfModule] = self.text_encoder_loras + self.unet_loras
|
||||
for lora in loras:
|
||||
org_module = lora.org_module_ref[0]
|
||||
if not org_module._lora_restored:
|
||||
sd = org_module.state_dict()
|
||||
sd["weight"] = org_module._lora_org_weight
|
||||
org_module.load_state_dict(sd)
|
||||
org_module._lora_restored = True
|
||||
|
||||
def pre_calculation(self):
|
||||
# 事前計算を行う
|
||||
loras: List[LoRAInfModule] = self.text_encoder_loras + self.unet_loras
|
||||
for lora in loras:
|
||||
org_module = lora.org_module_ref[0]
|
||||
sd = org_module.state_dict()
|
||||
|
||||
org_weight = sd["weight"]
|
||||
lora_weight = lora.get_weight().to(org_weight.device, dtype=org_weight.dtype)
|
||||
sd["weight"] = org_weight + lora_weight
|
||||
assert sd["weight"].shape == org_weight.shape
|
||||
org_module.load_state_dict(sd)
|
||||
|
||||
org_module._lora_restored = False
|
||||
lora.enabled = False
|
||||
|
||||
def apply_max_norm_regularization(self, max_norm_value, device):
|
||||
downkeys = []
|
||||
upkeys = []
|
||||
alphakeys = []
|
||||
norms = []
|
||||
keys_scaled = 0
|
||||
|
||||
state_dict = self.state_dict()
|
||||
for key in state_dict.keys():
|
||||
if "lora_down" in key and "weight" in key:
|
||||
downkeys.append(key)
|
||||
upkeys.append(key.replace("lora_down", "lora_up"))
|
||||
alphakeys.append(key.replace("lora_down.weight", "alpha"))
|
||||
|
||||
for i in range(len(downkeys)):
|
||||
down = state_dict[downkeys[i]].to(device)
|
||||
up = state_dict[upkeys[i]].to(device)
|
||||
alpha = state_dict[alphakeys[i]].to(device)
|
||||
dim = down.shape[0]
|
||||
scale = alpha / dim
|
||||
|
||||
if up.shape[2:] == (1, 1) and down.shape[2:] == (1, 1):
|
||||
updown = (up.squeeze(2).squeeze(2) @ down.squeeze(2).squeeze(2)).unsqueeze(2).unsqueeze(3)
|
||||
elif up.shape[2:] == (3, 3) or down.shape[2:] == (3, 3):
|
||||
updown = torch.nn.functional.conv2d(down.permute(1, 0, 2, 3), up).permute(1, 0, 2, 3)
|
||||
else:
|
||||
updown = up @ down
|
||||
|
||||
updown *= scale
|
||||
|
||||
norm = updown.norm().clamp(min=max_norm_value / 2)
|
||||
desired = torch.clamp(norm, max=max_norm_value)
|
||||
ratio = desired.cpu() / norm.cpu()
|
||||
sqrt_ratio = ratio**0.5
|
||||
if ratio != 1:
|
||||
keys_scaled += 1
|
||||
state_dict[upkeys[i]] *= sqrt_ratio
|
||||
state_dict[downkeys[i]] *= sqrt_ratio
|
||||
scalednorm = updown.norm() * ratio
|
||||
norms.append(scalednorm.item())
|
||||
|
||||
return keys_scaled, sum(norms) / len(norms), max(norms)
|
||||
@@ -9,6 +9,8 @@ import toml
|
||||
import json
|
||||
import time
|
||||
import shutil
|
||||
import shlex
|
||||
|
||||
from pathlib import Path
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
@@ -40,6 +42,9 @@ class FluxTrainModelSelect:
|
||||
"clip_l": (folder_paths.get_filename_list("clip"), ),
|
||||
"t5": (folder_paths.get_filename_list("clip"), ),
|
||||
},
|
||||
"optional": {
|
||||
"lora_path": ("STRING",{"multiline": True, "forceInput": True, "default": "", "tooltip": "pre-trained LoRA path to load (network_weights)"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("TRAIN_FLUX_MODELS",)
|
||||
@@ -47,7 +52,7 @@ class FluxTrainModelSelect:
|
||||
FUNCTION = "loadmodel"
|
||||
CATEGORY = "FluxTrainer"
|
||||
|
||||
def loadmodel(self, transformer, vae, clip_l, t5):
|
||||
def loadmodel(self, transformer, vae, clip_l, t5, lora_path=""):
|
||||
|
||||
transformer_path = folder_paths.get_full_path("unet", transformer)
|
||||
vae_path = folder_paths.get_full_path("vae", vae)
|
||||
@@ -58,12 +63,20 @@ class FluxTrainModelSelect:
|
||||
"transformer": transformer_path,
|
||||
"vae": vae_path,
|
||||
"clip_l": clip_path,
|
||||
"t5": t5_path
|
||||
"t5": t5_path,
|
||||
"lora_path": lora_path
|
||||
}
|
||||
|
||||
return (flux_models,)
|
||||
|
||||
class TrainDatasetGeneralConfig:
|
||||
queue_counter = 0
|
||||
@classmethod
|
||||
def IS_CHANGED(s, reset_on_queue=False, **kwargs):
|
||||
if reset_on_queue:
|
||||
s.queue_counter += 1
|
||||
print(f"queue_counter: {s.queue_counter}")
|
||||
return s.queue_counter
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
@@ -73,6 +86,10 @@ class TrainDatasetGeneralConfig:
|
||||
"caption_dropout_rate": ("FLOAT",{"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01,"tooltip": "tag dropout rate"}),
|
||||
"alpha_mask": ("BOOLEAN",{"default": False, "tooltip": "use alpha channel as mask for training"}),
|
||||
},
|
||||
"optional": {
|
||||
"reset_on_queue": ("BOOLEAN",{"default": False, "tooltip": "Force refresh of everything for cleaner queueing"}),
|
||||
"caption_extension": ("STRING",{"default": ".txt", "tooltip": "extension for caption files"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("JSON",)
|
||||
@@ -80,12 +97,12 @@ class TrainDatasetGeneralConfig:
|
||||
FUNCTION = "create_config"
|
||||
CATEGORY = "FluxTrainer"
|
||||
|
||||
def create_config(self, shuffle_caption, caption_dropout_rate, color_aug, flip_aug, alpha_mask):
|
||||
def create_config(self, shuffle_caption, caption_dropout_rate, color_aug, flip_aug, alpha_mask, reset_on_queue=False, caption_extension=".txt"):
|
||||
|
||||
dataset = {
|
||||
"general": {
|
||||
"shuffle_caption": shuffle_caption,
|
||||
"caption_extension": ".txt",
|
||||
"caption_extension": caption_extension,
|
||||
"keep_tokens_separator": "|||",
|
||||
"caption_dropout_rate": caption_dropout_rate,
|
||||
"color_aug": color_aug,
|
||||
@@ -101,9 +118,37 @@ class TrainDatasetGeneralConfig:
|
||||
}
|
||||
return (dataset_config,)
|
||||
|
||||
class TrainDatasetRegularization:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"dataset_path": ("STRING",{"multiline": True, "default": "", "tooltip": "path to dataset, root is the 'ComfyUI' folder, with windows portable 'ComfyUI_windows_portable'"}),
|
||||
"class_tokens": ("STRING",{"multiline": True, "default": "", "tooltip": "aka trigger word, if specified, will be added to the start of each caption, if no captions exist, will be used on it's own"}),
|
||||
"num_repeats": ("INT", {"default": 1, "min": 1, "tooltip": "number of times to repeat dataset for an epoch"}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("JSON",)
|
||||
RETURN_NAMES = ("subset",)
|
||||
FUNCTION = "create_config"
|
||||
CATEGORY = "FluxTrainer"
|
||||
|
||||
def create_config(self, dataset_path, class_tokens, num_repeats):
|
||||
|
||||
reg_subset = {
|
||||
"image_dir": dataset_path,
|
||||
"class_tokens": class_tokens,
|
||||
"num_repeats": num_repeats,
|
||||
"is_reg": True
|
||||
}
|
||||
|
||||
return reg_subset,
|
||||
|
||||
class TrainDatasetAdd:
|
||||
def __init__(self):
|
||||
self.previous_dataset_signature = None
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
@@ -118,8 +163,10 @@ class TrainDatasetAdd:
|
||||
"num_repeats": ("INT", {"default": 1, "min": 1, "tooltip": "number of times to repeat dataset for an epoch"}),
|
||||
"min_bucket_reso": ("INT", {"default": 256, "min": 64, "max": 4096, "step": 8, "tooltip": "min bucket resolution"}),
|
||||
"max_bucket_reso": ("INT", {"default": 1024, "min": 64, "max": 4096, "step": 8, "tooltip": "max bucket resolution"}),
|
||||
|
||||
},
|
||||
"optional": {
|
||||
"regularization": ("JSON", {"tooltip": "reg data dir"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("JSON",)
|
||||
@@ -128,7 +175,7 @@ class TrainDatasetAdd:
|
||||
CATEGORY = "FluxTrainer"
|
||||
|
||||
def create_config(self, dataset_config, dataset_path, class_tokens, width, height, batch_size, num_repeats, enable_bucket,
|
||||
bucket_no_upscale, min_bucket_reso, max_bucket_reso):
|
||||
bucket_no_upscale, min_bucket_reso, max_bucket_reso, regularization=None):
|
||||
|
||||
new_dataset = {
|
||||
"resolution": (width, height),
|
||||
@@ -145,6 +192,8 @@ class TrainDatasetAdd:
|
||||
}
|
||||
]
|
||||
}
|
||||
if regularization is not None:
|
||||
new_dataset["subsets"].append(regularization)
|
||||
|
||||
# Generate a signature for the new dataset
|
||||
new_dataset_signature = self.generate_signature(new_dataset)
|
||||
@@ -179,7 +228,7 @@ class OptimizerConfig:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"optimizer_type": (["adamw8bit", "adamw","prodigy", "CAME"], {"default": "adamw8bit", "tooltip": "optimizer type"}),
|
||||
"optimizer_type": (["adamw8bit", "adamw","prodigy", "CAME", "Lion8bit", "Lion", "adamwschedulefree", "sgdschedulefree", "AdEMAMix8bit", "PagedAdEMAMix8bit", "ProdigyPlusScheduleFree"], {"default": "adamw8bit", "tooltip": "optimizer type"}),
|
||||
"max_grad_norm": ("FLOAT",{"default": 1.0, "min": 0.0, "tooltip": "gradient clipping"}),
|
||||
"lr_scheduler": (["constant", "cosine", "cosine_with_restarts", "polynomial", "constant_with_warmup"], {"default": "constant", "tooltip": "learning rate scheduler"}),
|
||||
"lr_warmup_steps": ("INT",{"default": 0, "min": 0, "tooltip": "learning rate warmup steps"}),
|
||||
@@ -197,7 +246,7 @@ class OptimizerConfig:
|
||||
|
||||
def create_config(self, min_snr_gamma, extra_optimizer_args, **kwargs):
|
||||
kwargs["min_snr_gamma"] = min_snr_gamma if min_snr_gamma != 0.0 else None
|
||||
kwargs["optimizer_args"] = [arg.strip() for arg in extra_optimizer_args.strip().split(',') if arg.strip()]
|
||||
kwargs["optimizer_args"] = [arg.strip() for arg in extra_optimizer_args.strip().split('|') if arg.strip()]
|
||||
return (kwargs,)
|
||||
|
||||
class OptimizerConfigAdafactor:
|
||||
@@ -225,7 +274,7 @@ class OptimizerConfigAdafactor:
|
||||
|
||||
def create_config(self, relative_step, scale_parameter, warmup_init, clip_threshold, min_snr_gamma, extra_optimizer_args, **kwargs):
|
||||
kwargs["optimizer_type"] = "adafactor"
|
||||
extra_args = [arg.strip() for arg in extra_optimizer_args.strip().split(',') if arg.strip()]
|
||||
extra_args = [arg.strip() for arg in extra_optimizer_args.strip().split('|') if arg.strip()]
|
||||
node_args = [
|
||||
f"relative_step={relative_step}",
|
||||
f"scale_parameter={scale_parameter}",
|
||||
@@ -237,6 +286,25 @@ class OptimizerConfigAdafactor:
|
||||
|
||||
return (kwargs,)
|
||||
|
||||
class FluxTrainerLossConfig:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"loss_type": (["l2", "huber","smooth_l1"], {"default": "huber", "tooltip": "The type of loss function to use"}),
|
||||
"huber_schedule": (["snr", "exponential", "constant"], {"default": "exponential", "tooltip": "The scheduling method for Huber loss (constant, exponential, or SNR-based). Only used when loss_type is 'huber' or 'smooth_l1'. default is snr"}),
|
||||
"huber_c": ("FLOAT",{"default": 0.25, "min": 0.0, "step": 0.01, "tooltip": "The Huber loss decay parameter. Only used if one of the huber loss modes (huber or smooth l1) is selected with loss_type. default is 0.1"}),
|
||||
"huber_scale": ("FLOAT",{"default": 1.75, "min": 0.0, "step": 0.01, "tooltip": "The Huber loss scale parameter. Only used if one of the huber loss modes (huber or smooth l1) is selected with loss_type. default is 1.0"}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("ARGS",)
|
||||
RETURN_NAMES = ("loss_args",)
|
||||
FUNCTION = "create_config"
|
||||
CATEGORY = "FluxTrainer"
|
||||
|
||||
def create_config(self, **kwargs):
|
||||
return (kwargs,)
|
||||
|
||||
class OptimizerConfigProdigy:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -246,7 +314,7 @@ class OptimizerConfigProdigy:
|
||||
"lr_warmup_steps": ("INT",{"default": 0, "min": 0, "tooltip": "learning rate warmup steps"}),
|
||||
"lr_scheduler_num_cycles": ("INT",{"default": 1, "min": 1, "tooltip": "learning rate scheduler num cycles"}),
|
||||
"lr_scheduler_power": ("FLOAT",{"default": 1.0, "min": 0.0, "tooltip": "learning rate scheduler power"}),
|
||||
"weight_decay": ("FLOAT",{"default": 0.0, "tooltip": "weight decay (L2 penalty)"}),
|
||||
"weight_decay": ("FLOAT",{"default": 0.0, "step": 0.0001, "tooltip": "weight decay (L2 penalty)"}),
|
||||
"decouple": ("BOOLEAN",{"default": True, "tooltip": "use AdamW style weight decay"}),
|
||||
"use_bias_correction": ("BOOLEAN",{"default": False, "tooltip": "turn on Adam's bias correction"}),
|
||||
"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"}),
|
||||
@@ -261,7 +329,7 @@ class OptimizerConfigProdigy:
|
||||
|
||||
def create_config(self, weight_decay, decouple, min_snr_gamma, use_bias_correction, extra_optimizer_args, **kwargs):
|
||||
kwargs["optimizer_type"] = "prodigy"
|
||||
extra_args = [arg.strip() for arg in extra_optimizer_args.strip().split(',') if arg.strip()]
|
||||
extra_args = [arg.strip() for arg in extra_optimizer_args.strip().split('|') if arg.strip()]
|
||||
node_args = [
|
||||
f"weight_decay={weight_decay}",
|
||||
f"decouple={decouple}",
|
||||
@@ -272,6 +340,92 @@ class OptimizerConfigProdigy:
|
||||
|
||||
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
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"lr": ("FLOAT",{"default": 1.0, "min": 0.0, "step": 1e-7, "tooltip": "Learning rate adjustment parameter. Increases or decreases the Prodigy learning rate."}),
|
||||
"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_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": "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"}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("ARGS",)
|
||||
RETURN_NAMES = ("optimizer_settings",)
|
||||
FUNCTION = "create_config"
|
||||
CATEGORY = "FluxTrainer"
|
||||
|
||||
def create_config(self, min_snr_gamma, use_bias_correction, extra_optimizer_args, **kwargs):
|
||||
kwargs["optimizer_type"] = "ProdigyPlusScheduleFree"
|
||||
kwargs["lr_scheduler"] = "constant"
|
||||
extra_args = [arg.strip() for arg in extra_optimizer_args.strip().split('|') if arg.strip()]
|
||||
node_args = [
|
||||
f"use_bias_correction={use_bias_correction}",
|
||||
]
|
||||
kwargs["optimizer_args"] = node_args + extra_args
|
||||
kwargs["min_snr_gamma"] = min_snr_gamma if min_snr_gamma != 0.0 else None
|
||||
|
||||
return (kwargs,)
|
||||
|
||||
class InitFluxLoRATraining:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -281,18 +435,14 @@ 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": 256, "step": 1, "tooltip": "network dim"}),
|
||||
"network_alpha": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 256.0, "step": 0.01, "tooltip": "network alpha"}),
|
||||
"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"}),
|
||||
#"unet_lr": ("FLOAT", {"default": 1e-4, "min": 0.0, "max": 10.0, "step": 0.00001, "tooltip": "unet learning rate"}),
|
||||
#"max_train_epochs": ("INT", {"default": 4, "min": 1, "max": 1000, "step": 1, "tooltip": "max number of training epochs"}),
|
||||
"max_train_steps": ("INT", {"default": 1500, "min": 1, "max": 100000, "step": 1, "tooltip": "max number of training steps"}),
|
||||
#"text_encoder_lr": ("FLOAT", {"default": 0, "min": 0.0, "max": 10.0, "step": 0.00001, "tooltip": "text encoder learning rate"}),
|
||||
"apply_t5_attn_mask": ("BOOLEAN", {"default": True, "tooltip": "apply t5 attention mask"}),
|
||||
#"t5xxl_max_token_length": ("INT", {"default": 512, "min": 64, "max": 4096, "step": 8, "tooltip": "dev uses 512, schnell 256"}),
|
||||
"cache_latents": (["disk", "memory", "disabled"], {"tooltip": "caches text encoder outputs"}),
|
||||
"cache_text_encoder_outputs": (["disk", "memory", "disabled"], {"tooltip": "caches text encoder outputs"}),
|
||||
"split_mode": ("BOOLEAN", {"default": False, "tooltip": "[EXPERIMENTAL] use split mode for Flux model, network arg `train_blocks=single` is required"}),
|
||||
"blocks_to_swap": ("INT", {"default": 0, "tooltip": "Previously known as split_mode, number of blocks to swap to save memory, default to enable is 18"}),
|
||||
"weighting_scheme": (["logit_normal", "sigma_sqrt", "mode", "cosmap", "none"],),
|
||||
"logit_mean": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "mean to use when using the logit_normal weighting scheme"}),
|
||||
"logit_std": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01,"tooltip": "std to use when using the logit_normal weighting scheme"}),
|
||||
@@ -305,15 +455,23 @@ class InitFluxLoRATraining:
|
||||
"highvram": ("BOOLEAN", {"default": False, "tooltip": "memory mode"}),
|
||||
"fp8_base": ("BOOLEAN", {"default": True, "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"], {"default": "bf16", "tooltip": "the dtype to save checkpoints as"}),
|
||||
"save_dtype": (["fp32", "fp16", "bf16", "fp8_e4m3fn", "fp8_e5m2"], {"default": "bf16", "tooltip": "the dtype to save checkpoints as"}),
|
||||
"attention_mode": (["sdpa", "xformers", "disabled"], {"default": "sdpa", "tooltip": "memory efficient attention mode"}),
|
||||
"sample_prompts": ("STRING", {"multiline": True, "default": "illustration of a kitten | photograph of a turtle", "tooltip": "validation sample prompts, for multiple prompts, separate by `|`"}),
|
||||
},
|
||||
"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"}),
|
||||
"train_clip_l": (['disabled', 'use_gradient_dtype', 'use_fp8'], {"default": 'disabled', "tooltip": "also train the clip_l text encoder using specified dtype"}),
|
||||
"text_encoder_lr": ("FLOAT", {"default": 0, "min": 0.0, "max": 10.0, "step": 0.00001, "tooltip": "text encoder learning rate"}),
|
||||
"train_text_encoder": (['disabled', 'clip_l', 'clip_l_fp8', 'clip_l+T5', 'clip_l+T5_fp8'], {"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"}),
|
||||
"T5_lr": ("FLOAT", {"default": 0, "min": 0.0, "max": 10.0, "step": 0.000001, "tooltip": "text encoder learning rate"}),
|
||||
"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"
|
||||
},
|
||||
}
|
||||
|
||||
@@ -323,7 +481,8 @@ class InitFluxLoRATraining:
|
||||
CATEGORY = "FluxTrainer"
|
||||
|
||||
def init_training(self, flux_models, dataset, optimizer_settings, sample_prompts, output_name, attention_mode,
|
||||
gradient_dtype, save_dtype, split_mode, additional_args=None, resume_args=None, train_clip_l='disabled', **kwargs,):
|
||||
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, network_config=None, **kwargs):
|
||||
mm.soft_empty_cache()
|
||||
|
||||
output_dir = os.path.abspath(kwargs.get("output_dir"))
|
||||
@@ -339,11 +498,13 @@ 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:
|
||||
args, _ = parser.parse_known_args(args=[additional_args])
|
||||
print(f"additional_args: {additional_args}")
|
||||
args, _ = parser.parse_known_args(args=shlex.split(additional_args))
|
||||
else:
|
||||
args, _ = parser.parse_known_args()
|
||||
#print(args)
|
||||
|
||||
if kwargs.get("cache_latents") == "memory":
|
||||
kwargs["cache_latents"] = True
|
||||
@@ -387,16 +548,16 @@ class InitFluxLoRATraining:
|
||||
"persistent_data_loader_workers": False,
|
||||
"max_data_loader_n_workers": 0,
|
||||
"seed": 42,
|
||||
"gradient_checkpointing": True,
|
||||
"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",
|
||||
"text_encoder_lr": 0,
|
||||
"t5xxl_max_token_length": 512,
|
||||
"alpha_mask": dataset["alpha_mask"],
|
||||
"network_train_unet_only": True if train_clip_l == 'disabled' else False,
|
||||
"fp8_base_unet": True if train_clip_l=='use_gradient_dtype' else False,
|
||||
"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},
|
||||
@@ -410,31 +571,70 @@ class InitFluxLoRATraining:
|
||||
}
|
||||
config_dict.update(gradient_dtype_settings.get(gradient_dtype, {}))
|
||||
|
||||
split_mode_settings = {
|
||||
True: {"split_mode": True, "network_args": ["train_blocks=single"]},
|
||||
False: {"split_mode": False, "network_args": ["train_blocks=all"]}
|
||||
}
|
||||
config_dict.update(split_mode_settings.get(split_mode, {}))
|
||||
if train_text_encoder != 'disabled':
|
||||
if T5_lr != "NaN":
|
||||
config_dict["text_encoder_lr"] = clip_l_lr
|
||||
if T5_lr != "NaN":
|
||||
config_dict["text_encoder_lr"] = [clip_l_lr, T5_lr]
|
||||
|
||||
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 flux_models["lora_path"]:
|
||||
config_dict["network_weights"] = flux_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)
|
||||
|
||||
#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:
|
||||
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 = FluxNetworkTrainer()
|
||||
training_loop = network_trainer.init_train(args)
|
||||
|
||||
epochs_count = network_trainer.num_train_epochs
|
||||
|
||||
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)
|
||||
|
||||
trainer = {
|
||||
"network_trainer": network_trainer,
|
||||
"training_loop": training_loop,
|
||||
@@ -453,7 +653,7 @@ class InitFluxTraining:
|
||||
"learning_rate": ("FLOAT", {"default": 4e-6, "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"}),
|
||||
"apply_t5_attn_mask": ("BOOLEAN", {"default": True, "tooltip": "apply t5 attention mask"}),
|
||||
"t5xxl_max_token_length": ("INT", {"default": 512, "min": 64, "max": 4096, "step": 8, "tooltip": "dev uses 512, schnell 256"}),
|
||||
"t5xxl_max_token_length": ("INT", {"default": 512, "min": 64, "max": 4096, "step": 8, "tooltip": "dev and LibreFlux uses 512, schnell 256"}),
|
||||
"cache_latents": (["disk", "memory", "disabled"], {"tooltip": "caches text encoder outputs"}),
|
||||
"cache_text_encoder_outputs": (["disk", "memory", "disabled"], {"tooltip": "caches text encoder outputs"}),
|
||||
"weighting_scheme": (["logit_normal", "sigma_sqrt", "mode", "cosmap", "none"],),
|
||||
@@ -466,8 +666,7 @@ class InitFluxTraining:
|
||||
"model_prediction_type": (["raw", "additive", "sigma_scaled"], {"tooltip": "How to interpret and process the model prediction: raw (use as is), additive (add to noisy input), sigma_scaled (apply sigma scaling)"}),
|
||||
"cpu_offload_checkpointing": ("BOOLEAN", {"default": True, "tooltip": "offload the gradient checkpointing to CPU. This reduces VRAM usage for about 2GB"}),
|
||||
"optimizer_fusing": (['fused_backward_pass', 'blockwise_fused_optimizers'], {"tooltip": "reduces memory use"}),
|
||||
"single_blocks_to_swap": ("INT", {"default": 0, "min": 0, "max": 100, "step": 1, "tooltip": "number of single blocks to swap. The default is 0. This option must be combined with blockwise_fused_optimizers"}),
|
||||
"double_blocks_to_swap": ("INT", {"default": 6, "min": 0, "max": 100, "step": 1, "tooltip": "number of double blocks to swap. This option must be combined with blockwise_fused_optimizers"}),
|
||||
"blocks_to_swap": ("INT", {"default": 0, "min": 0, "max": 100, "step": 1, "tooltip": "Sets the number of blocks (~640MB) to swap during the forward and backward passes, increasing this number lowers the overall VRAM used during training at the expense of training speed (s/it)."}),
|
||||
"guidance_scale": ("FLOAT", {"default": 1.0, "min": 1.0, "max": 32.0, "step": 0.01, "tooltip": "guidance scale"}),
|
||||
"discrete_flow_shift": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001, "tooltip": "for the Euler Discrete Scheduler, default is 3.0"}),
|
||||
"highvram": ("BOOLEAN", {"default": False, "tooltip": "memory mode"}),
|
||||
@@ -500,11 +699,15 @@ class InitFluxTraining:
|
||||
if free <= required_free_space:
|
||||
raise ValueError(f"Most likely insufficient disk space to complete training. Required: {required_free_space/2**30}GB. Available: {free/2**30}GB")
|
||||
|
||||
dataset_toml = toml.dumps(json.loads(dataset))
|
||||
dataset_config = dataset["datasets"]
|
||||
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:
|
||||
args, _ = parser.parse_known_args(args=[additional_args])
|
||||
print(f"additional_args: {additional_args}")
|
||||
args, _ = parser.parse_known_args(args=shlex.split(additional_args))
|
||||
else:
|
||||
args, _ = parser.parse_known_args()
|
||||
|
||||
@@ -554,6 +757,7 @@ class InitFluxTraining:
|
||||
"dataset_config": dataset_toml,
|
||||
"output_name": f"{output_name}_{save_dtype}",
|
||||
"mem_eff_save": True,
|
||||
"disable_mmap_load_safetensors": True,
|
||||
|
||||
}
|
||||
optimizer_fusing_settings = {
|
||||
@@ -697,13 +901,17 @@ class FluxTrainLoop:
|
||||
initial_global_step = network_trainer.global_step
|
||||
|
||||
target_global_step = network_trainer.global_step + steps
|
||||
pbar = comfy.utils.ProgressBar(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,
|
||||
)
|
||||
pbar.update(steps_done)
|
||||
#pbar.update(steps_done)
|
||||
|
||||
# Also break if the global steps have reached the max train steps
|
||||
if network_trainer.global_step >= network_trainer.args.max_train_steps:
|
||||
@@ -715,6 +923,80 @@ class FluxTrainLoop:
|
||||
}
|
||||
return (trainer, network_trainer.global_step)
|
||||
|
||||
class FluxTrainAndValidateLoop:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {"required": {
|
||||
"network_trainer": ("NETWORKTRAINER",),
|
||||
"validate_at_steps": ("INT", {"default": 250, "min": 1, "max": 10000, "step": 1, "tooltip": "the step point in training to validate/save"}),
|
||||
"save_at_steps": ("INT", {"default": 250, "min": 1, "max": 10000, "step": 1, "tooltip": "the step point in training to validate/save"}),
|
||||
},
|
||||
"optional": {
|
||||
"validation_settings": ("VALSETTINGS",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("NETWORKTRAINER", "INT",)
|
||||
RETURN_NAMES = ("network_trainer", "steps",)
|
||||
FUNCTION = "train"
|
||||
CATEGORY = "FluxTrainer"
|
||||
|
||||
def train(self, network_trainer, validate_at_steps, save_at_steps, validation_settings=None):
|
||||
with torch.inference_mode(False):
|
||||
training_loop = network_trainer["training_loop"]
|
||||
network_trainer = network_trainer["network_trainer"]
|
||||
|
||||
target_global_step = network_trainer.args.max_train_steps
|
||||
comfy_pbar = comfy.utils.ProgressBar(target_global_step)
|
||||
network_trainer.comfy_pbar = comfy_pbar
|
||||
|
||||
network_trainer.optimizer_train_fn()
|
||||
|
||||
while network_trainer.global_step < target_global_step:
|
||||
next_validate_step = ((network_trainer.global_step // validate_at_steps) + 1) * validate_at_steps
|
||||
next_save_step = ((network_trainer.global_step // save_at_steps) + 1) * save_at_steps
|
||||
|
||||
steps_done = training_loop(
|
||||
break_at_steps=min(next_validate_step, next_save_step),
|
||||
epoch=network_trainer.current_epoch.value,
|
||||
)
|
||||
|
||||
# Check if we need to validate
|
||||
if network_trainer.global_step % validate_at_steps == 0:
|
||||
self.validate(network_trainer, validation_settings)
|
||||
|
||||
# Check if we need to save
|
||||
if network_trainer.global_step % save_at_steps == 0:
|
||||
self.save(network_trainer)
|
||||
|
||||
# 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)
|
||||
|
||||
def validate(self, network_trainer, validation_settings=None):
|
||||
params = (
|
||||
network_trainer.current_epoch.value,
|
||||
network_trainer.global_step,
|
||||
validation_settings
|
||||
)
|
||||
network_trainer.optimizer_eval_fn()
|
||||
image_tensors = network_trainer.sample_images(*params)
|
||||
network_trainer.optimizer_train_fn()
|
||||
print("Validating at step:", network_trainer.global_step)
|
||||
|
||||
def save(self, network_trainer):
|
||||
ckpt_name = train_util.get_step_ckpt_name(network_trainer.args, "." + network_trainer.args.save_model_as, network_trainer.global_step)
|
||||
network_trainer.optimizer_eval_fn()
|
||||
network_trainer.save_model(ckpt_name, network_trainer.accelerator.unwrap_model(network_trainer.network), network_trainer.global_step, network_trainer.current_epoch.value + 1)
|
||||
network_trainer.optimizer_train_fn()
|
||||
print("Saving at step:", network_trainer.global_step)
|
||||
|
||||
class FluxTrainSave:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -749,7 +1031,9 @@ class FluxTrainSave:
|
||||
|
||||
lora_path = os.path.join(trainer.args.output_dir, ckpt_name)
|
||||
if copy_to_comfy_lora_folder:
|
||||
shutil.copy(lora_path, os.path.join(folder_paths.models_dir, "loras", "flux_trainer", ckpt_name))
|
||||
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)
|
||||
@@ -775,6 +1059,8 @@ class FluxTrainSaveModel:
|
||||
trainer = network_trainer["network_trainer"]
|
||||
global_step = trainer.global_step
|
||||
|
||||
trainer.optimizer_eval_fn()
|
||||
|
||||
ckpt_name = train_util.get_step_ckpt_name(trainer.args, "." + trainer.args.save_model_as, global_step)
|
||||
flux_train_utils.save_flux_model_on_epoch_end_or_stepwise(
|
||||
trainer.args,
|
||||
@@ -809,6 +1095,7 @@ class FluxTrainEnd:
|
||||
RETURN_NAMES = ("lora_name", "metadata", "lora_path",)
|
||||
FUNCTION = "endtrain"
|
||||
CATEGORY = "FluxTrainer"
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def endtrain(self, network_trainer, save_state):
|
||||
with torch.inference_mode(False):
|
||||
@@ -821,6 +1108,7 @@ class FluxTrainEnd:
|
||||
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)
|
||||
@@ -863,6 +1151,60 @@ class FluxTrainResume:
|
||||
|
||||
return (resume_args, )
|
||||
|
||||
class FluxTrainBlockSelect:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"include": ("STRING", {"default": "lora_unet_single_blocks_20_linear2", "multiline": True, "tooltip": "blocks to include in the LoRA network, to select multiple blocks either input them as "}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("ARGS", )
|
||||
RETURN_NAMES = ("block_args", )
|
||||
FUNCTION = "block_select"
|
||||
CATEGORY = "FluxTrainer"
|
||||
|
||||
def block_select(self, include):
|
||||
import re
|
||||
|
||||
# Split the input string by commas to handle multiple ranges/blocks
|
||||
elements = include.split(',')
|
||||
|
||||
# Initialize a list to collect block names
|
||||
blocks = []
|
||||
|
||||
# Pattern to find ranges like (10-20)
|
||||
pattern = re.compile(r'\((\d+)-(\d+)\)')
|
||||
|
||||
# Extract the prefix and suffix from the first element
|
||||
prefix_suffix_pattern = re.compile(r'(.*)_blocks_(.*)')
|
||||
|
||||
for element in elements:
|
||||
element = element.strip()
|
||||
match = prefix_suffix_pattern.match(element)
|
||||
if match:
|
||||
prefix = match.group(1) + "_blocks_"
|
||||
suffix = match.group(2)
|
||||
matches = pattern.findall(suffix)
|
||||
if matches:
|
||||
for start, end in matches:
|
||||
# Generate block names for the range and add them to the list
|
||||
blocks.extend([f"{prefix}{i}{suffix.replace(f'({start}-{end})', '', 1)}" for i in range(int(start), int(end) + 1)])
|
||||
else:
|
||||
# If no range is found, add the block name directly
|
||||
blocks.append(element)
|
||||
else:
|
||||
blocks.append(element)
|
||||
|
||||
# Construct the `include` string
|
||||
include_string = ','.join(blocks)
|
||||
|
||||
block_args = {
|
||||
"include": f"only_if_contains={include_string}",
|
||||
}
|
||||
|
||||
return (block_args, )
|
||||
|
||||
class FluxTrainValidationSettings:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -911,22 +1253,13 @@ class FluxTrainValidate:
|
||||
network_trainer = network_trainer["network_trainer"]
|
||||
|
||||
params = (
|
||||
network_trainer.accelerator,
|
||||
network_trainer.args,
|
||||
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
|
||||
)
|
||||
|
||||
split_mode = getattr(network_trainer.args, 'split_mode', False)
|
||||
if split_mode:
|
||||
image_tensors = network_trainer.sample_images_split_mode(*params)
|
||||
else:
|
||||
image_tensors = flux_train_utils.sample_images(*params)
|
||||
network_trainer.optimizer_eval_fn()
|
||||
with torch.inference_mode(False):
|
||||
image_tensors = network_trainer.sample_images(*params)
|
||||
|
||||
trainer = {
|
||||
"network_trainer": network_trainer,
|
||||
@@ -1007,6 +1340,7 @@ class FluxKohyaInferenceSampler:
|
||||
"guidance_scale": ("FLOAT", {"default": 3.5, "min": 1.0, "max": 32.0, "step": 0.05, "tooltip": "guidance scale"}),
|
||||
"seed": ("INT", {"default": 42,"min": 0, "max": 0xffffffffffffffff, "step": 1}),
|
||||
"use_fp8": ("BOOLEAN", {"default": True, "tooltip": "use fp8 weights"}),
|
||||
"apply_t5_attn_mask": ("BOOLEAN", {"default": True, "tooltip": "use t5 attention mask"}),
|
||||
"prompt": ("STRING", {"multiline": True, "default": "illustration of a kitten", "tooltip": "prompt"}),
|
||||
|
||||
},
|
||||
@@ -1017,7 +1351,7 @@ class FluxKohyaInferenceSampler:
|
||||
FUNCTION = "sample"
|
||||
CATEGORY = "FluxTrainer"
|
||||
|
||||
def sample(self, flux_models, lora_name, steps, width, height, guidance_scale, seed, prompt, use_fp8, lora_method):
|
||||
def sample(self, flux_models, lora_name, steps, width, height, guidance_scale, seed, prompt, use_fp8, lora_method, apply_t5_attn_mask):
|
||||
|
||||
from .library import flux_utils as flux_utils
|
||||
from .library import strategy_flux as strategy_flux
|
||||
@@ -1030,7 +1364,7 @@ class FluxKohyaInferenceSampler:
|
||||
import gc
|
||||
|
||||
device = "cuda"
|
||||
apply_t5_attn_mask = True
|
||||
|
||||
|
||||
if use_fp8:
|
||||
accelerator = accelerate.Accelerator(mixed_precision="bf16")
|
||||
@@ -1075,8 +1409,7 @@ class FluxKohyaInferenceSampler:
|
||||
# AE
|
||||
ae = flux_utils.load_ae("dev", ae, ae_dtype, loading_device)
|
||||
ae.eval()
|
||||
#if is_fp8(ae_dtype):
|
||||
# ae = accelerator.prepare(ae)
|
||||
|
||||
|
||||
# LoRA
|
||||
lora_models: List[lora_flux.LoRANetwork] = []
|
||||
@@ -1118,7 +1451,7 @@ class FluxKohyaInferenceSampler:
|
||||
clip_l.to(ae_dtype)
|
||||
t5xxl.to(ae_dtype)
|
||||
with accelerator.autocast():
|
||||
_, t5_out, txt_ids, t5_attn_mask = encoding_strategy.encode_tokens(
|
||||
l_pooled, t5_out, txt_ids, t5_attn_mask = encoding_strategy.encode_tokens(
|
||||
tokenize_strategy, [clip_l, t5xxl], tokens_and_masks, apply_t5_attn_mask
|
||||
)
|
||||
else:
|
||||
@@ -1224,6 +1557,7 @@ class FluxKohyaInferenceSampler:
|
||||
flux_dtype: torch.dtype,
|
||||
):
|
||||
timesteps = get_schedule(num_steps, img.shape[1], shift=not is_schnell)
|
||||
print(timesteps)
|
||||
|
||||
# denoise initial noise
|
||||
if accelerator:
|
||||
@@ -1232,9 +1566,11 @@ class FluxKohyaInferenceSampler:
|
||||
model, img, img_ids, t5_out, txt_ids, l_pooled, timesteps=timesteps, guidance=guidance, t5_attn_mask=t5_attn_mask
|
||||
)
|
||||
else:
|
||||
with torch.autocast(device_type=device.type, dtype=flux_dtype), torch.no_grad():
|
||||
x = denoise(
|
||||
model, img, img_ids, t5_out, txt_ids, l_pooled, timesteps=timesteps, guidance=guidance, t5_attn_mask=t5_attn_mask
|
||||
with torch.autocast(device_type=device.type, dtype=flux_dtype):
|
||||
l_pooled, _, _, _ = encoding_strategy.encode_tokens(tokenize_strategy, [clip_l, None], tokens_and_masks)
|
||||
with torch.autocast(device_type=device.type, dtype=flux_dtype):
|
||||
_, t5_out, txt_ids, t5_attn_mask = encoding_strategy.encode_tokens(
|
||||
tokenize_strategy, [None, t5xxl], tokens_and_masks, apply_t5_attn_mask
|
||||
)
|
||||
|
||||
return x
|
||||
@@ -1373,7 +1709,7 @@ class ExtractFluxLoRA:
|
||||
"finetuned_model": (folder_paths.get_filename_list("unet"), ),
|
||||
"output_path": ("STRING", {"default": f"{str(os.path.join(folder_paths.models_dir, 'loras', 'Flux'))}"}),
|
||||
"dim": ("INT", {"default": 4, "min": 2, "max": 1024, "step": 2, "tooltip": "LoRA rank"}),
|
||||
"save_dtype": (["fp32", "fp16", "bf16", "fp8_e4m3fn"], {"default": "bf16", "tooltip": "the dtype to save the LoRA as"}),
|
||||
"save_dtype": (["fp32", "fp16", "bf16", "fp8_e4m3fn", "fp8_e5m2"], {"default": "bf16", "tooltip": "the dtype to save the LoRA as"}),
|
||||
"load_device": (["cpu", "cuda"], {"default": "cuda", "tooltip": "the device to load the model to"}),
|
||||
"store_device": (["cpu", "cuda"], {"default": "cpu", "tooltip": "the device to store the LoRA as"}),
|
||||
"clamp_quantile": ("FLOAT", {"default": 0.99, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "clamp quantile"}),
|
||||
@@ -1425,7 +1761,13 @@ NODE_CLASS_MAPPINGS = {
|
||||
"FluxTrainSaveModel": FluxTrainSaveModel,
|
||||
"ExtractFluxLoRA": ExtractFluxLoRA,
|
||||
"OptimizerConfigProdigy": OptimizerConfigProdigy,
|
||||
"FluxTrainResume": FluxTrainResume
|
||||
"FluxTrainResume": FluxTrainResume,
|
||||
"FluxTrainBlockSelect": FluxTrainBlockSelect,
|
||||
"TrainDatasetRegularization": TrainDatasetRegularization,
|
||||
"FluxTrainAndValidateLoop": FluxTrainAndValidateLoop,
|
||||
"OptimizerConfigProdigyPlusScheduleFree": OptimizerConfigProdigyPlusScheduleFree,
|
||||
"FluxTrainerLossConfig": FluxTrainerLossConfig,
|
||||
"TrainNetworkConfig": TrainNetworkConfig,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"InitFluxLoRATraining": "Init Flux LoRA Training",
|
||||
@@ -1446,5 +1788,11 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"FluxTrainSaveModel": "Flux Train Save Model",
|
||||
"ExtractFluxLoRA": "Extract Flux LoRA",
|
||||
"OptimizerConfigProdigy": "Optimizer Config Prodigy",
|
||||
"FluxTrainResume": "Flux Train Resume"
|
||||
"FluxTrainResume": "Flux Train Resume",
|
||||
"FluxTrainBlockSelect": "Flux Train Block Select",
|
||||
"TrainDatasetRegularization": "Train Dataset Regularization",
|
||||
"FluxTrainAndValidateLoop": "Flux Train And Validate Loop",
|
||||
"OptimizerConfigProdigyPlusScheduleFree": "Optimizer Config ProdigyPlusScheduleFree",
|
||||
"FluxTrainerLossConfig": "Flux Trainer Loss Config",
|
||||
"TrainNetworkConfig": "Train Network Config",
|
||||
}
|
||||
|
||||
+467
@@ -0,0 +1,467 @@
|
||||
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 .sd3_train_network import Sd3NetworkTrainer
|
||||
from .library import sd3_train_utils as sd3_train_utils
|
||||
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 SD3ModelSelect:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"transformer": (folder_paths.get_filename_list("checkpoints"), ),
|
||||
"clip_l": (folder_paths.get_filename_list("clip"), ),
|
||||
"clip_g": (folder_paths.get_filename_list("clip"), ),
|
||||
"t5": (folder_paths.get_filename_list("clip"), ),
|
||||
},
|
||||
"optional": {
|
||||
"lora_path": ("STRING",{"multiline": True, "forceInput": True, "default": "", "tooltip": "pre-trained LoRA path to load (network_weights)"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("TRAIN_SD3_MODELS",)
|
||||
RETURN_NAMES = ("sd3_models",)
|
||||
FUNCTION = "loadmodel"
|
||||
CATEGORY = "FluxTrainer/SD3"
|
||||
|
||||
def loadmodel(self, transformer, clip_l, clip_g, t5, lora_path=""):
|
||||
|
||||
transformer_path = folder_paths.get_full_path("checkpoints", transformer)
|
||||
clip_l_path = folder_paths.get_full_path("clip", clip_l)
|
||||
clip_g_path = folder_paths.get_full_path("clip", clip_g)
|
||||
t5_path = folder_paths.get_full_path("clip", t5)
|
||||
|
||||
sd3_models = {
|
||||
"transformer": transformer_path,
|
||||
"clip_l": clip_l_path,
|
||||
"clip_g": clip_g_path,
|
||||
"t5": t5_path,
|
||||
"lora_path": lora_path
|
||||
}
|
||||
|
||||
return (sd3_models,)
|
||||
|
||||
class InitSD3LoRATraining:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"sd3_models": ("TRAIN_SD3_MODELS",),
|
||||
"dataset": ("JSON",),
|
||||
"optimizer_settings": ("ARGS",),
|
||||
"output_name": ("STRING", {"default": "sd35_lora", "multiline": False}),
|
||||
"output_dir": ("STRING", {"default": "sd35_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": 2048, "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-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"}),
|
||||
"cache_latents": (["disk", "memory", "disabled"], {"tooltip": "caches text encoder outputs"}),
|
||||
"cache_text_encoder_outputs": (["disk", "memory", "disabled"], {"tooltip": "caches text encoder outputs"}),
|
||||
"training_shift ": ("FLOAT", {"default": 3.0, "min": 0.0, "max": 10.0, "step": 0.0001, "tooltip": "shift value for the training distribution of timesteps"}),
|
||||
"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 SD3.5L and 22 for SD3.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": "bf16", "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', 'clip_l_fp8', 'clip_l+T5', 'clip_l+T5_fp8'], {"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"}),
|
||||
"T5_lr": ("FLOAT", {"default": 0, "min": 0.0, "max": 10.0, "step": 0.000001, "tooltip": "text encoder learning rate"}),
|
||||
"sample_prompts": ("STRING", {"multiline": True, "default": "illustration of a kitten | photograph of a turtle", "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"}),
|
||||
},
|
||||
"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/SD3"
|
||||
|
||||
def init_training(self, sd3_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, clip_g_lr=0, T5_lr=0, loss_args=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()
|
||||
sd3_train_utils.add_sd3_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:
|
||||
prompts = sample_prompts.split('|')
|
||||
else:
|
||||
prompts = [sample_prompts]
|
||||
|
||||
config_dict = {
|
||||
"sample_prompts": prompts,
|
||||
"save_precision": save_dtype,
|
||||
"mixed_precision": "bf16",
|
||||
"num_cpu_threads_per_process": 1,
|
||||
"pretrained_model_name_or_path": sd3_models["transformer"],
|
||||
"clip_l": sd3_models["clip_l"],
|
||||
"clip_g": sd3_models["clip_g"],
|
||||
"t5xxl": sd3_models["t5"],
|
||||
"save_model_as": "safetensors",
|
||||
"persistent_data_loader_workers": False,
|
||||
"max_data_loader_n_workers": 0,
|
||||
"seed": 42,
|
||||
"network_module": ".networks.lora_sd3",
|
||||
"dataset_config": dataset_toml,
|
||||
"output_name": f"{output_name}_rank{kwargs.get('network_dim')}_{save_dtype}",
|
||||
"loss_type": "l2",
|
||||
"t5xxl_max_token_length": 512,
|
||||
"alpha_mask": dataset["alpha_mask"],
|
||||
"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,
|
||||
}
|
||||
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, 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":
|
||||
config_dict["gradient_checkpointing"] = True
|
||||
config_dict["cpu_offload_checkpointing"] = True
|
||||
else:
|
||||
config_dict["gradient_checkpointing"] = True
|
||||
|
||||
if sd3_models["lora_path"]:
|
||||
config_dict["network_weights"] = sd3_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 = Sd3NetworkTrainer()
|
||||
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 SD3TrainLoop:
|
||||
@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"
|
||||
|
||||
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 SD3TrainLoRASave:
|
||||
@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"
|
||||
|
||||
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 SD3TrainEnd:
|
||||
@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"
|
||||
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 SD3TrainValidationSettings:
|
||||
@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": 4, "min": 1.0, "max": 32.0, "step": 0.05, "tooltip": "guidance scale"}),
|
||||
"seed": ("INT", {"default": 42,"min": 0, "max": 0xffffffffffffffff, "step": 1}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("VALSETTINGS", )
|
||||
RETURN_NAMES = ("validation_settings", )
|
||||
FUNCTION = "set"
|
||||
CATEGORY = "FluxTrainer"
|
||||
|
||||
def set(self, **kwargs):
|
||||
validation_settings = kwargs
|
||||
print(validation_settings)
|
||||
|
||||
return (validation_settings,)
|
||||
|
||||
class SD3TrainValidate:
|
||||
@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"
|
||||
|
||||
def validate(self, network_trainer, validation_settings=None):
|
||||
training_loop = network_trainer["training_loop"]
|
||||
network_trainer = network_trainer["network_trainer"]
|
||||
|
||||
params = (
|
||||
network_trainer.current_epoch.value,
|
||||
network_trainer.global_step,
|
||||
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 = {
|
||||
"SD3ModelSelect": SD3ModelSelect,
|
||||
"InitSD3LoRATraining": InitSD3LoRATraining,
|
||||
"SD3TrainValidationSettings": SD3TrainValidationSettings,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"SD3ModelSelect": "SD3 Model Select",
|
||||
"InitSD3LoRATraining": "Init SD3 LoRA Training",
|
||||
"SD3TrainValidationSettings": "SD3 Train Validation Settings",
|
||||
}
|
||||
+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
-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.0"
|
||||
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"]
|
||||
|
||||
|
||||
+4
-1
@@ -5,7 +5,7 @@ diffusers>=0.25.0
|
||||
ftfy>=6.1.1
|
||||
opencv-python>=4.7.0.68
|
||||
einops>=0.7.0
|
||||
bitsandbytes>=0.43.3
|
||||
bitsandbytes>=0.44.0
|
||||
prodigyopt>=1.0
|
||||
lion-pytorch>=0.0.6
|
||||
safetensors>=0.4.4
|
||||
@@ -20,3 +20,6 @@ came_pytorch
|
||||
matplotlib
|
||||
# for T5XXL tokenizer (SD3/FLUX)
|
||||
sentencepiece>=0.2.0
|
||||
protobuf
|
||||
schedulefree>=1.2.7
|
||||
prodigy-plus-schedule-free>=1.9.0
|
||||
@@ -0,0 +1,486 @@
|
||||
import argparse
|
||||
import math
|
||||
|
||||
from typing import Any, Optional
|
||||
|
||||
import torch
|
||||
from accelerate import Accelerator
|
||||
from .library import sd3_models, strategy_sd3, utils
|
||||
from .library.device_utils import init_ipex, clean_memory_on_device
|
||||
|
||||
init_ipex()
|
||||
|
||||
from .library import flux_models, flux_utils, sd3_train_utils, sd3_utils, strategy_base, strategy_sd3, train_util
|
||||
from . import train_network
|
||||
from .library.utils import setup_logging
|
||||
|
||||
setup_logging()
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class Sd3NetworkTrainer(train_network.NetworkTrainer):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.sample_prompts_te_outputs = None
|
||||
|
||||
def assert_extra_args(self, args, train_dataset_group: train_util.DatasetGroup):
|
||||
# super().assert_extra_args(args, train_dataset_group)
|
||||
# sdxl_train_util.verify_sdxl_training_args(args)
|
||||
|
||||
if args.fp8_base_unet:
|
||||
args.fp8_base = True # if fp8_base_unet is enabled, fp8_base is also enabled for SD3
|
||||
|
||||
if args.cache_text_encoder_outputs_to_disk and not args.cache_text_encoder_outputs:
|
||||
logger.warning(
|
||||
"cache_text_encoder_outputs_to_disk is enabled, so cache_text_encoder_outputs is also enabled / cache_text_encoder_outputs_to_diskが有効になっているため、cache_text_encoder_outputsも有効になります"
|
||||
)
|
||||
args.cache_text_encoder_outputs = True
|
||||
|
||||
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は使えません"
|
||||
|
||||
# prepare CLIP-L/CLIP-G/T5XXL training flags
|
||||
self.train_clip = not args.network_train_unet_only
|
||||
self.train_t5xxl = False # default is False even if args.network_train_unet_only is False
|
||||
|
||||
if args.max_token_length is not None:
|
||||
logger.warning("max_token_length is not used in Flux training / max_token_lengthはFluxのトレーニングでは使用されません")
|
||||
|
||||
assert (
|
||||
args.blocks_to_swap is None or args.blocks_to_swap == 0
|
||||
) or not args.cpu_offload_checkpointing, "blocks_to_swap is not supported with cpu_offload_checkpointing / blocks_to_swapはcpu_offload_checkpointingと併用できません"
|
||||
|
||||
train_dataset_group.verify_bucket_reso_steps(32) # TODO check this
|
||||
|
||||
# enumerate resolutions from dataset for positional embeddings
|
||||
self.resolutions = train_dataset_group.get_resolutions()
|
||||
|
||||
def load_target_model(self, args, weight_dtype, accelerator):
|
||||
# currently offload to cpu for some models
|
||||
|
||||
# if the file is fp8 and we are using fp8_base, we can load it as is (fp8)
|
||||
loading_dtype = None if args.fp8_base else weight_dtype
|
||||
|
||||
# if we load to cpu, flux.to(fp8) takes a long time, so we should load to gpu in future
|
||||
state_dict = utils.load_safetensors(
|
||||
args.pretrained_model_name_or_path, "cpu", disable_mmap=args.disable_mmap_load_safetensors, dtype=loading_dtype
|
||||
)
|
||||
mmdit = sd3_utils.load_mmdit(state_dict, loading_dtype, "cpu")
|
||||
self.model_type = mmdit.model_type
|
||||
mmdit.set_pos_emb_random_crop_rate(args.pos_emb_random_crop_rate)
|
||||
|
||||
# set resolutions for positional embeddings
|
||||
if args.enable_scaled_pos_embed:
|
||||
latent_sizes = [round(math.sqrt(res[0] * res[1])) // 8 for res in self.resolutions] # 8 is stride for latent
|
||||
latent_sizes = list(set(latent_sizes)) # remove duplicates
|
||||
logger.info(f"Prepare scaled positional embeddings for resolutions: {self.resolutions}, sizes: {latent_sizes}")
|
||||
mmdit.enable_scaled_pos_embed(True, latent_sizes)
|
||||
|
||||
if args.fp8_base:
|
||||
# check dtype of model
|
||||
if mmdit.dtype == torch.float8_e4m3fnuz or mmdit.dtype == torch.float8_e5m2 or mmdit.dtype == torch.float8_e5m2fnuz:
|
||||
raise ValueError(f"Unsupported fp8 model dtype: {mmdit.dtype}")
|
||||
elif mmdit.dtype == torch.float8_e4m3fn:
|
||||
logger.info("Loaded fp8 SD3 model")
|
||||
else:
|
||||
logger.info(
|
||||
"Cast SD3 model to fp8. This may take a while. You can reduce the time by using fp8 checkpoint."
|
||||
)
|
||||
mmdit.to(torch.float8_e4m3fn)
|
||||
self.is_swapping_blocks = args.blocks_to_swap is not None and args.blocks_to_swap > 0
|
||||
if self.is_swapping_blocks:
|
||||
# Swap blocks between CPU and GPU to reduce memory usage, in forward and backward passes.
|
||||
logger.info(f"enable block swap: blocks_to_swap={args.blocks_to_swap}")
|
||||
mmdit.enable_block_swap(args.blocks_to_swap, accelerator.device)
|
||||
|
||||
clip_l = sd3_utils.load_clip_l(
|
||||
args.clip_l, weight_dtype, "cpu", disable_mmap=args.disable_mmap_load_safetensors, state_dict=state_dict
|
||||
)
|
||||
clip_l.eval()
|
||||
clip_g = sd3_utils.load_clip_g(
|
||||
args.clip_g, weight_dtype, "cpu", disable_mmap=args.disable_mmap_load_safetensors, state_dict=state_dict
|
||||
)
|
||||
clip_g.eval()
|
||||
|
||||
# if the file is fp8 and we are using fp8_base (not unet), we can load it as is (fp8)
|
||||
if args.fp8_base and not args.fp8_base_unet:
|
||||
loading_dtype = None # as is
|
||||
else:
|
||||
loading_dtype = weight_dtype
|
||||
|
||||
# loading t5xxl to cpu takes a long time, so we should load to gpu in future
|
||||
t5xxl = sd3_utils.load_t5xxl(
|
||||
args.t5xxl, loading_dtype, "cpu", disable_mmap=args.disable_mmap_load_safetensors, state_dict=state_dict
|
||||
)
|
||||
t5xxl.eval()
|
||||
if args.fp8_base and not args.fp8_base_unet:
|
||||
# check dtype of model
|
||||
if t5xxl.dtype == torch.float8_e4m3fnuz or t5xxl.dtype == torch.float8_e5m2 or t5xxl.dtype == torch.float8_e5m2fnuz:
|
||||
raise ValueError(f"Unsupported fp8 model dtype: {t5xxl.dtype}")
|
||||
elif t5xxl.dtype == torch.float8_e4m3fn:
|
||||
logger.info("Loaded fp8 T5XXL model")
|
||||
|
||||
vae = sd3_utils.load_vae(
|
||||
args.vae, weight_dtype, "cpu", disable_mmap=args.disable_mmap_load_safetensors, state_dict=state_dict
|
||||
)
|
||||
|
||||
return mmdit.model_type, [clip_l, clip_g, t5xxl], vae, mmdit
|
||||
|
||||
def get_tokenize_strategy(self, args):
|
||||
logger.info(f"t5xxl_max_token_length: {args.t5xxl_max_token_length}")
|
||||
return strategy_sd3.Sd3TokenizeStrategy(args.t5xxl_max_token_length, args.tokenizer_cache_dir)
|
||||
|
||||
def get_tokenizers(self, tokenize_strategy: strategy_sd3.Sd3TokenizeStrategy):
|
||||
return [tokenize_strategy.clip_l, tokenize_strategy.clip_g, tokenize_strategy.t5xxl]
|
||||
|
||||
def get_latents_caching_strategy(self, args):
|
||||
latents_caching_strategy = strategy_sd3.Sd3LatentsCachingStrategy(
|
||||
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_sd3.Sd3TextEncodingStrategy(
|
||||
args.apply_lg_attn_mask,
|
||||
args.apply_t5_attn_mask,
|
||||
args.clip_l_dropout_rate,
|
||||
args.clip_g_dropout_rate,
|
||||
args.t5_dropout_rate,
|
||||
)
|
||||
|
||||
def post_process_network(self, args, accelerator, network, text_encoders, unet):
|
||||
# check t5xxl is trained or not
|
||||
self.train_t5xxl = network.train_t5xxl
|
||||
|
||||
if self.train_t5xxl and args.cache_text_encoder_outputs:
|
||||
raise ValueError(
|
||||
"T5XXL is trained, so cache_text_encoder_outputs cannot be used / T5XXL学習時はcache_text_encoder_outputsは使用できません"
|
||||
)
|
||||
|
||||
def get_models_for_text_encoding(self, args, accelerator, text_encoders):
|
||||
if args.cache_text_encoder_outputs:
|
||||
if self.train_clip and not self.train_t5xxl:
|
||||
return text_encoders[0:2] + [None] # only CLIP-L/CLIP-G is needed for encoding because T5XXL is cached
|
||||
else:
|
||||
return None # no text encoders are needed for encoding because both are cached
|
||||
else:
|
||||
return text_encoders # CLIP-L, CLIP-G and T5XXL are needed for encoding
|
||||
|
||||
def get_text_encoders_train_flags(self, args, text_encoders):
|
||||
return [self.train_clip, self.train_clip, self.train_t5xxl]
|
||||
|
||||
def get_text_encoder_outputs_caching_strategy(self, args):
|
||||
if args.cache_text_encoder_outputs:
|
||||
# if the text encoders is trained, we need tokenization, so is_partial is True
|
||||
return strategy_sd3.Sd3TextEncoderOutputsCachingStrategy(
|
||||
args.cache_text_encoder_outputs_to_disk,
|
||||
args.text_encoder_batch_size,
|
||||
args.skip_cache_check,
|
||||
is_partial=self.train_clip or self.train_t5xxl,
|
||||
apply_lg_attn_mask=args.apply_lg_attn_mask,
|
||||
apply_t5_attn_mask=args.apply_t5_attn_mask,
|
||||
)
|
||||
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
|
||||
logger.info("move text encoders to gpu")
|
||||
text_encoders[0].to(accelerator.device, dtype=weight_dtype) # always not fp8
|
||||
text_encoders[1].to(accelerator.device, dtype=weight_dtype) # always not fp8
|
||||
text_encoders[2].to(accelerator.device) # may be fp8
|
||||
|
||||
if text_encoders[2].dtype == torch.float8_e4m3fn:
|
||||
# if we load fp8 weights, the model is already fp8, so we use it as is
|
||||
self.prepare_text_encoder_fp8(2, text_encoders[2], text_encoders[2].dtype, weight_dtype)
|
||||
else:
|
||||
# otherwise, we need to convert it to target dtype
|
||||
text_encoders[2].to(weight_dtype)
|
||||
|
||||
with accelerator.autocast():
|
||||
dataset.new_cache_text_encoder_outputs(text_encoders, accelerator)
|
||||
|
||||
# cache sample prompts
|
||||
if args.sample_prompts is not None:
|
||||
logger.info(f"cache Text Encoder outputs for sample prompt: {args.sample_prompts}")
|
||||
|
||||
tokenize_strategy: strategy_sd3.Sd3TokenizeStrategy = strategy_base.TokenizeStrategy.get_strategy()
|
||||
text_encoding_strategy: strategy_sd3.Sd3TextEncodingStrategy = strategy_base.TextEncodingStrategy.get_strategy()
|
||||
|
||||
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 .library.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)
|
||||
|
||||
sample_prompts_te_outputs = {} # key: prompt, value: text encoder outputs
|
||||
with accelerator.autocast(), torch.no_grad():
|
||||
for prompt_dict in prompts:
|
||||
for p in [prompt_dict.get("prompt", ""), prompt_dict.get("negative_prompt", "")]:
|
||||
if p not in sample_prompts_te_outputs:
|
||||
logger.info(f"cache Text Encoder outputs for prompt: {p}")
|
||||
tokens_and_masks = tokenize_strategy.tokenize(p)
|
||||
sample_prompts_te_outputs[p] = text_encoding_strategy.encode_tokens(
|
||||
tokenize_strategy,
|
||||
text_encoders,
|
||||
tokens_and_masks,
|
||||
args.apply_lg_attn_mask,
|
||||
args.apply_t5_attn_mask,
|
||||
)
|
||||
self.sample_prompts_te_outputs = sample_prompts_te_outputs
|
||||
|
||||
accelerator.wait_for_everyone()
|
||||
|
||||
# move back to cpu
|
||||
if not self.is_train_text_encoder(args):
|
||||
logger.info("move CLIP-L back to cpu")
|
||||
text_encoders[0].to("cpu")
|
||||
logger.info("move CLIP-G back to cpu")
|
||||
text_encoders[1].to("cpu")
|
||||
logger.info("move t5XXL back to cpu")
|
||||
text_encoders[2].to("cpu")
|
||||
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)
|
||||
text_encoders[2].to(accelerator.device)
|
||||
|
||||
# def call_unet(self, args, accelerator, unet, noisy_latents, timesteps, text_conds, batch, weight_dtype):
|
||||
# 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)
|
||||
|
||||
# noise_pred = unet(noisy_latents, timesteps, text_embedding, vector_embedding)
|
||||
# return noise_pred
|
||||
|
||||
def sample_images(self, epoch, global_step, validation_settings):
|
||||
text_encoders = self.get_models_for_text_encoding(self.args, self.accelerator, self.text_encoder)
|
||||
image_tensors = sd3_train_utils.sample_images(
|
||||
self.accelerator, self.args, epoch, global_step, self.unet, self.vae, text_encoders, self.sample_prompts_te_outputs, validation_settings
|
||||
)
|
||||
|
||||
return image_tensors.permute(0, 2, 3, 1)
|
||||
|
||||
def get_noise_scheduler(self, args: argparse.Namespace, device: torch.device) -> Any:
|
||||
# this scheduler is not used in training, but used to get num_train_timesteps etc.
|
||||
noise_scheduler = sd3_train_utils.FlowMatchEulerDiscreteScheduler(num_train_timesteps=1000, shift=args.training_shift)
|
||||
return noise_scheduler
|
||||
|
||||
def encode_images_to_latents(self, args, accelerator, vae, images):
|
||||
return vae.encode(images)
|
||||
|
||||
def shift_scale_latents(self, args, latents):
|
||||
return sd3_models.SDVAE.process_in(latents)
|
||||
|
||||
def get_noise_pred_and_target(
|
||||
self,
|
||||
args,
|
||||
accelerator,
|
||||
noise_scheduler,
|
||||
latents,
|
||||
batch,
|
||||
text_encoder_conds,
|
||||
unet: flux_models.Flux,
|
||||
network,
|
||||
weight_dtype,
|
||||
train_unet,
|
||||
):
|
||||
# Sample noise that we'll add to the latents
|
||||
noise = torch.randn_like(latents)
|
||||
|
||||
# get noisy model input and timesteps
|
||||
noisy_model_input, timesteps, sigmas = sd3_train_utils.get_noisy_model_input_and_timesteps(
|
||||
args, latents, noise, accelerator.device, weight_dtype
|
||||
)
|
||||
|
||||
# ensure the hidden state will require grad
|
||||
if args.gradient_checkpointing:
|
||||
noisy_model_input.requires_grad_(True)
|
||||
for t in text_encoder_conds:
|
||||
if t is not None and t.dtype.is_floating_point:
|
||||
t.requires_grad_(True)
|
||||
|
||||
# Predict the noise residual
|
||||
lg_out, t5_out, lg_pooled, l_attn_mask, g_attn_mask, t5_attn_mask = text_encoder_conds
|
||||
text_encoding_strategy = strategy_base.TextEncodingStrategy.get_strategy()
|
||||
context, lg_pooled = text_encoding_strategy.concat_encodings(lg_out, t5_out, lg_pooled)
|
||||
if not args.apply_lg_attn_mask:
|
||||
l_attn_mask = None
|
||||
g_attn_mask = None
|
||||
if not args.apply_t5_attn_mask:
|
||||
t5_attn_mask = None
|
||||
|
||||
# call model
|
||||
with accelerator.autocast():
|
||||
# TODO support attention mask
|
||||
model_pred = unet(noisy_model_input, timesteps, context=context, y=lg_pooled)
|
||||
|
||||
# Follow: Section 5 of https://arxiv.org/abs/2206.00364.
|
||||
# Preconditioning of the model outputs.
|
||||
model_pred = model_pred * (-sigmas) + noisy_model_input
|
||||
|
||||
# these weighting schemes use a uniform timestep sampling
|
||||
# and instead post-weight the loss
|
||||
weighting = sd3_train_utils.compute_loss_weighting_for_sd3(weighting_scheme=args.weighting_scheme, sigmas=sigmas)
|
||||
|
||||
# flow matching loss
|
||||
target = latents
|
||||
|
||||
# differential output preservation
|
||||
if "custom_attributes" in batch:
|
||||
diff_output_pr_indices = []
|
||||
for i, custom_attributes in enumerate(batch["custom_attributes"]):
|
||||
if "diff_output_preservation" in custom_attributes and custom_attributes["diff_output_preservation"]:
|
||||
diff_output_pr_indices.append(i)
|
||||
|
||||
if len(diff_output_pr_indices) > 0:
|
||||
network.set_multiplier(0.0)
|
||||
with torch.no_grad(), accelerator.autocast():
|
||||
model_pred_prior = unet(
|
||||
noisy_model_input[diff_output_pr_indices],
|
||||
timesteps[diff_output_pr_indices],
|
||||
context=context[diff_output_pr_indices],
|
||||
y=lg_pooled[diff_output_pr_indices],
|
||||
)
|
||||
network.set_multiplier(1.0) # may be overwritten by "network_multipliers" in the next step
|
||||
|
||||
model_pred_prior = model_pred_prior * (-sigmas[diff_output_pr_indices]) + noisy_model_input[diff_output_pr_indices]
|
||||
|
||||
# weighting for differential output preservation is not needed because it is already applied
|
||||
|
||||
target[diff_output_pr_indices] = model_pred_prior.to(target.dtype)
|
||||
|
||||
return model_pred, target, timesteps, weighting
|
||||
|
||||
def post_process_loss(self, loss, args, timesteps, noise_scheduler):
|
||||
return loss
|
||||
|
||||
def get_sai_model_spec(self, args):
|
||||
return train_util.get_sai_model_spec(None, args, False, True, False, sd3=self.model_type)
|
||||
|
||||
def update_metadata(self, metadata, args):
|
||||
metadata["ss_apply_lg_attn_mask"] = args.apply_lg_attn_mask
|
||||
metadata["ss_apply_t5_attn_mask"] = args.apply_t5_attn_mask
|
||||
metadata["ss_weighting_scheme"] = args.weighting_scheme
|
||||
metadata["ss_logit_mean"] = args.logit_mean
|
||||
metadata["ss_logit_std"] = args.logit_std
|
||||
metadata["ss_mode_scale"] = args.mode_scale
|
||||
|
||||
def is_text_encoder_not_needed_for_training(self, args):
|
||||
return args.cache_text_encoder_outputs and not self.is_train_text_encoder(args)
|
||||
|
||||
def prepare_text_encoder_grad_ckpt_workaround(self, index, text_encoder):
|
||||
if index == 0 or index == 1: # CLIP-L/CLIP-G
|
||||
return super().prepare_text_encoder_grad_ckpt_workaround(index, text_encoder)
|
||||
else: # T5XXL
|
||||
text_encoder.encoder.embed_tokens.requires_grad_(True)
|
||||
|
||||
def prepare_text_encoder_fp8(self, index, text_encoder, te_weight_dtype, weight_dtype):
|
||||
if index == 0 or index == 1: # CLIP-L/CLIP-G
|
||||
clip_type = "CLIP-L" if index == 0 else "CLIP-G"
|
||||
logger.info(f"prepare CLIP-{clip_type} for fp8: set to {te_weight_dtype}, set embeddings to {weight_dtype}")
|
||||
text_encoder.to(te_weight_dtype) # fp8
|
||||
text_encoder.text_model.embeddings.to(dtype=weight_dtype)
|
||||
else: # T5XXL
|
||||
|
||||
def prepare_fp8(text_encoder, target_dtype):
|
||||
def forward_hook(module):
|
||||
def forward(hidden_states):
|
||||
hidden_gelu = module.act(module.wi_0(hidden_states))
|
||||
hidden_linear = module.wi_1(hidden_states)
|
||||
hidden_states = hidden_gelu * hidden_linear
|
||||
hidden_states = module.dropout(hidden_states)
|
||||
|
||||
hidden_states = module.wo(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
return forward
|
||||
|
||||
for module in text_encoder.modules():
|
||||
if module.__class__.__name__ in ["T5LayerNorm", "Embedding"]:
|
||||
# print("set", module.__class__.__name__, "to", target_dtype)
|
||||
module.to(target_dtype)
|
||||
if module.__class__.__name__ in ["T5DenseGatedActDense"]:
|
||||
# print("set", module.__class__.__name__, "hooks")
|
||||
module.forward = forward_hook(module)
|
||||
|
||||
if flux_utils.get_t5xxl_actual_dtype(text_encoder) == torch.float8_e4m3fn and text_encoder.dtype == weight_dtype:
|
||||
logger.info(f"T5XXL already prepared for fp8")
|
||||
else:
|
||||
logger.info(f"prepare T5XXL for fp8: set to {te_weight_dtype}, set embeddings to {weight_dtype}, add hooks")
|
||||
text_encoder.to(te_weight_dtype) # fp8
|
||||
prepare_fp8(text_encoder, weight_dtype)
|
||||
|
||||
def on_step_start(self, args, accelerator, network, text_encoders, unet, batch, weight_dtype):
|
||||
# drop cached text encoder outputs
|
||||
text_encoder_outputs_list = batch.get("text_encoder_outputs_list", None)
|
||||
if text_encoder_outputs_list is not None:
|
||||
text_encodoing_strategy: strategy_sd3.Sd3TextEncodingStrategy = strategy_base.TextEncodingStrategy.get_strategy()
|
||||
text_encoder_outputs_list = text_encodoing_strategy.drop_cached_text_encoder_outputs(*text_encoder_outputs_list)
|
||||
batch["text_encoder_outputs_list"] = text_encoder_outputs_list
|
||||
|
||||
def prepare_unet_with_accelerator(
|
||||
self, args: argparse.Namespace, accelerator: Accelerator, unet: torch.nn.Module
|
||||
) -> torch.nn.Module:
|
||||
if not self.is_swapping_blocks:
|
||||
return super().prepare_unet_with_accelerator(args, accelerator, unet)
|
||||
|
||||
# if we doesn't swap blocks, we can move the model to device
|
||||
mmdit: sd3_models.MMDiT = unet
|
||||
mmdit = accelerator.prepare(mmdit, device_placement=[not self.is_swapping_blocks])
|
||||
accelerator.unwrap_model(mmdit).move_to_device_except_swap_blocks(accelerator.device) # reduce peak memory usage
|
||||
accelerator.unwrap_model(mmdit).prepare_block_swap_before_forward()
|
||||
|
||||
return mmdit
|
||||
|
||||
|
||||
def setup_parser() -> argparse.ArgumentParser:
|
||||
parser = train_network.setup_parser()
|
||||
train_util.add_dit_training_arguments(parser)
|
||||
sd3_train_utils.add_sd3_training_arguments(parser)
|
||||
return parser
|
||||
@@ -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)
|
||||
+1
-1
@@ -154,7 +154,7 @@ def train(args):
|
||||
vae.requires_grad_(False)
|
||||
vae.eval()
|
||||
|
||||
train_dataset_group.new_cache_latents(vae, accelerator.is_main_process)
|
||||
train_dataset_group.new_cache_latents(vae, accelerator)
|
||||
|
||||
vae.to("cpu")
|
||||
clean_memory_on_device(accelerator.device)
|
||||
|
||||
+194
-78
@@ -19,6 +19,7 @@ from .library.device_utils import init_ipex, clean_memory_on_device
|
||||
init_ipex()
|
||||
|
||||
from accelerate.utils import set_seed
|
||||
from accelerate import Accelerator
|
||||
from diffusers import DDPMScheduler
|
||||
from .library import deepspeed_utils, model_util, strategy_base, strategy_sd
|
||||
|
||||
@@ -61,6 +62,7 @@ class NetworkTrainer:
|
||||
avr_loss,
|
||||
lr_scheduler,
|
||||
lr_descriptions,
|
||||
optimizer=None,
|
||||
keys_scaled=None,
|
||||
mean_norm=None,
|
||||
maximum_norm=None,
|
||||
@@ -93,11 +95,35 @@ class NetworkTrainer:
|
||||
logs[f"lr/d*lr/{lr_desc}"] = (
|
||||
lr_scheduler.optimizers[-1].param_groups[i]["d"] * lr_scheduler.optimizers[-1].param_groups[i]["lr"]
|
||||
)
|
||||
if (
|
||||
args.optimizer_type.lower().endswith("ProdigyPlusScheduleFree".lower()) and optimizer is not None
|
||||
): # tracking d*lr value of unet.
|
||||
logs["lr/d*lr"] = (
|
||||
optimizer.param_groups[0]["d"] * optimizer.param_groups[0]["lr"]
|
||||
)
|
||||
else:
|
||||
idx = 0
|
||||
if not args.network_train_unet_only:
|
||||
logs["lr/textencoder"] = float(lrs[0])
|
||||
idx = 1
|
||||
|
||||
for i in range(idx, len(lrs)):
|
||||
logs[f"lr/group{i}"] = float(lrs[i])
|
||||
if args.optimizer_type.lower().startswith("DAdapt".lower()) or args.optimizer_type.lower() == "Prodigy".lower():
|
||||
logs[f"lr/d*lr/group{i}"] = (
|
||||
lr_scheduler.optimizers[-1].param_groups[i]["d"] * lr_scheduler.optimizers[-1].param_groups[i]["lr"]
|
||||
)
|
||||
if (
|
||||
args.optimizer_type.lower().endswith("ProdigyPlusScheduleFree".lower()) and optimizer is not None
|
||||
):
|
||||
logs[f"lr/d*lr/group{i}"] = (
|
||||
optimizer.param_groups[i]["d"] * optimizer.param_groups[i]["lr"]
|
||||
)
|
||||
|
||||
return logs
|
||||
|
||||
def assert_extra_args(self, args, train_dataset_group):
|
||||
pass
|
||||
train_dataset_group.verify_bucket_reso_steps(64)
|
||||
|
||||
def load_target_model(self, args, weight_dtype, accelerator):
|
||||
text_encoder, vae, unet, _ = train_util.load_target_model(args, weight_dtype, accelerator)
|
||||
@@ -117,7 +143,7 @@ class NetworkTrainer:
|
||||
|
||||
def get_latents_caching_strategy(self, args):
|
||||
latents_caching_strategy = strategy_sd.SdSdxlLatentsCachingStrategy(
|
||||
True, args.cache_latents_to_disk, args.vae_batch_size, False
|
||||
True, args.cache_latents_to_disk, args.vae_batch_size, args.skip_cache_check
|
||||
)
|
||||
return latents_caching_strategy
|
||||
|
||||
@@ -130,6 +156,7 @@ class NetworkTrainer:
|
||||
def get_models_for_text_encoding(self, args, accelerator, text_encoders):
|
||||
"""
|
||||
Returns a list of models that will be used for text encoding. SDXL uses wrapped and unwrapped models.
|
||||
FLUX.1 and SD3 may cache some outputs of the text encoder, so return the models that will be used for encoding (not cached).
|
||||
"""
|
||||
return text_encoders
|
||||
|
||||
@@ -144,7 +171,7 @@ class NetworkTrainer:
|
||||
for t_enc in text_encoders:
|
||||
t_enc.to(accelerator.device, dtype=weight_dtype)
|
||||
|
||||
def call_unet(self, args, accelerator, unet, noisy_latents, timesteps, text_conds, batch, weight_dtype):
|
||||
def call_unet(self, args, accelerator, unet, noisy_latents, timesteps, text_conds, batch, weight_dtype, **kwargs):
|
||||
noise_pred = unet(noisy_latents, timesteps, text_conds[0]).sample
|
||||
return noise_pred
|
||||
|
||||
@@ -158,6 +185,9 @@ class NetworkTrainer:
|
||||
|
||||
# region SD/SDXL
|
||||
|
||||
def post_process_network(self, args, accelerator, network, text_encoders, unet):
|
||||
pass
|
||||
|
||||
def get_noise_scheduler(self, args: argparse.Namespace, device: torch.device) -> Any:
|
||||
noise_scheduler = DDPMScheduler(
|
||||
beta_start=0.00085, beta_end=0.012, beta_schedule="scaled_linear", num_train_timesteps=1000, clip_sample=False
|
||||
@@ -188,7 +218,7 @@ class NetworkTrainer:
|
||||
):
|
||||
# Sample noise, sample a random timestep for each image, and add noise to the latents,
|
||||
# with noise offset and/or multires noise if specified
|
||||
noise, noisy_latents, timesteps, huber_c = train_util.get_noise_noisy_latents_and_timesteps(args, noise_scheduler, latents)
|
||||
noise, noisy_latents, timesteps = train_util.get_noise_noisy_latents_and_timesteps(args, noise_scheduler, latents)
|
||||
|
||||
# ensure the hidden state will require grad
|
||||
if args.gradient_checkpointing:
|
||||
@@ -216,7 +246,31 @@ class NetworkTrainer:
|
||||
else:
|
||||
target = noise
|
||||
|
||||
return noise_pred, target, timesteps, huber_c, None
|
||||
# differential output preservation
|
||||
if "custom_attributes" in batch:
|
||||
diff_output_pr_indices = []
|
||||
for i, custom_attributes in enumerate(batch["custom_attributes"]):
|
||||
if "diff_output_preservation" in custom_attributes and custom_attributes["diff_output_preservation"]:
|
||||
diff_output_pr_indices.append(i)
|
||||
|
||||
if len(diff_output_pr_indices) > 0:
|
||||
network.set_multiplier(0.0)
|
||||
with torch.no_grad(), accelerator.autocast():
|
||||
noise_pred_prior = self.call_unet(
|
||||
args,
|
||||
accelerator,
|
||||
unet,
|
||||
noisy_latents,
|
||||
timesteps,
|
||||
text_encoder_conds,
|
||||
batch,
|
||||
weight_dtype,
|
||||
indices=diff_output_pr_indices,
|
||||
)
|
||||
network.set_multiplier(1.0) # may be overwritten by "network_multipliers" in the next step
|
||||
target[diff_output_pr_indices] = noise_pred_prior.to(target.dtype)
|
||||
|
||||
return noise_pred, target, timesteps, None
|
||||
|
||||
def post_process_loss(self, loss, args, timesteps, noise_scheduler):
|
||||
if args.min_snr_gamma:
|
||||
@@ -226,16 +280,31 @@ class NetworkTrainer:
|
||||
if args.v_pred_like_loss:
|
||||
loss = add_v_prediction_like_loss(loss, timesteps, noise_scheduler, args.v_pred_like_loss)
|
||||
if args.debiased_estimation_loss:
|
||||
loss = apply_debiased_estimation(loss, timesteps, noise_scheduler)
|
||||
loss = apply_debiased_estimation(loss, timesteps, noise_scheduler, args.v_parameterization)
|
||||
return loss
|
||||
|
||||
def get_sai_model_spec(self, args):
|
||||
return train_util.get_sai_model_spec(None, args, self.is_sdxl, True, False)
|
||||
|
||||
def update_metadata(self, metadata, args):
|
||||
pass
|
||||
|
||||
def is_text_encoder_not_needed_for_training(self, args):
|
||||
return False # use for sample images
|
||||
|
||||
def update_metadata(self, metadata, args):
|
||||
def prepare_text_encoder_grad_ckpt_workaround(self, index, text_encoder):
|
||||
# set top parameter requires_grad = True for gradient checkpointing works
|
||||
text_encoder.text_model.embeddings.requires_grad_(True)
|
||||
|
||||
def prepare_text_encoder_fp8(self, index, text_encoder, te_weight_dtype, weight_dtype):
|
||||
text_encoder.text_model.embeddings.to(dtype=weight_dtype)
|
||||
|
||||
def prepare_unet_with_accelerator(
|
||||
self, args: argparse.Namespace, accelerator: Accelerator, unet: torch.nn.Module
|
||||
) -> torch.nn.Module:
|
||||
return accelerator.prepare(unet)
|
||||
|
||||
def on_step_start(self, args, accelerator, network, text_encoders, unet, batch, weight_dtype):
|
||||
pass
|
||||
|
||||
# endregion
|
||||
@@ -318,7 +387,7 @@ class NetworkTrainer:
|
||||
collator = train_util.collator_class(current_epoch, current_step, ds_for_collator)
|
||||
|
||||
if args.debug_dataset:
|
||||
train_dataset_group.set_current_strategies()
|
||||
train_dataset_group.set_current_strategies() # dasaset needs to know the strategies explicitly
|
||||
train_util.debug_dataset(train_dataset_group)
|
||||
return
|
||||
if len(train_dataset_group) == 0:
|
||||
@@ -332,7 +401,7 @@ class NetworkTrainer:
|
||||
train_dataset_group.is_latent_cacheable()
|
||||
), "when caching latents, either color_aug or random_crop cannot be used / latentをキャッシュするときはcolor_augとrandom_cropは使えません"
|
||||
|
||||
self.assert_extra_args(args, train_dataset_group)
|
||||
self.assert_extra_args(args, train_dataset_group) # may change some args
|
||||
|
||||
# prepare accelerator
|
||||
logger.info("preparing accelerator")
|
||||
@@ -381,7 +450,7 @@ class NetworkTrainer:
|
||||
vae.requires_grad_(False)
|
||||
vae.eval()
|
||||
|
||||
train_dataset_group.new_cache_latents(vae, True)
|
||||
train_dataset_group.new_cache_latents(vae, accelerator)
|
||||
|
||||
vae.to("cpu")
|
||||
clean_memory_on_device(accelerator.device)
|
||||
@@ -434,17 +503,24 @@ class NetworkTrainer:
|
||||
)
|
||||
args.scale_weight_norms = False
|
||||
|
||||
self.post_process_network(args, accelerator, network, text_encoders, unet)
|
||||
|
||||
# apply network to unet and text_encoder
|
||||
train_unet = not args.network_train_text_encoder_only
|
||||
train_text_encoder = self.is_train_text_encoder(args)
|
||||
network.apply_to(text_encoder, unet, train_text_encoder, train_unet)
|
||||
|
||||
if args.network_weights is not None:
|
||||
# FIXME consider alpha of weights
|
||||
# FIXME consider alpha of weights: this assumes that the alpha is not changed
|
||||
info = network.load_weights(args.network_weights)
|
||||
accelerator.print(f"load network weights from {args.network_weights}: {info}")
|
||||
|
||||
if args.gradient_checkpointing:
|
||||
if args.cpu_offload_checkpointing:
|
||||
unet.enable_gradient_checkpointing(cpu_offload=True)
|
||||
else:
|
||||
unet.enable_gradient_checkpointing()
|
||||
|
||||
for t_enc, flag in zip(text_encoders, self.get_text_encoders_train_flags(args, text_encoders)):
|
||||
if flag:
|
||||
if t_enc.supports_gradient_checkpointing:
|
||||
@@ -455,9 +531,21 @@ class NetworkTrainer:
|
||||
# Prepare classes necessary for learning
|
||||
accelerator.print("prepare optimizer, data loader etc.")
|
||||
|
||||
# Ensure backward compatibility
|
||||
# make backward compatibility for text_encoder_lr
|
||||
support_multiple_lrs = hasattr(network, "prepare_optimizer_params_with_multiple_te_lrs")
|
||||
if support_multiple_lrs:
|
||||
text_encoder_lr = args.text_encoder_lr
|
||||
else:
|
||||
# toml backward compatibility
|
||||
if args.text_encoder_lr is None or isinstance(args.text_encoder_lr, float) or isinstance(args.text_encoder_lr, int):
|
||||
text_encoder_lr = args.text_encoder_lr
|
||||
else:
|
||||
text_encoder_lr = None if len(args.text_encoder_lr) == 0 else args.text_encoder_lr[0]
|
||||
try:
|
||||
results = network.prepare_optimizer_params(args.text_encoder_lr, args.unet_lr, args.learning_rate)
|
||||
if support_multiple_lrs:
|
||||
results = network.prepare_optimizer_params_with_multiple_te_lrs(text_encoder_lr, args.unet_lr, args.learning_rate)
|
||||
else:
|
||||
results = network.prepare_optimizer_params(text_encoder_lr, args.unet_lr, args.learning_rate)
|
||||
if type(results) is tuple:
|
||||
trainable_params = results[0]
|
||||
lr_descriptions = results[1]
|
||||
@@ -465,11 +553,7 @@ class NetworkTrainer:
|
||||
trainable_params = results
|
||||
lr_descriptions = None
|
||||
except TypeError as e:
|
||||
# logger.warning(f"{e}")
|
||||
# accelerator.print(
|
||||
# "Deprecated: use prepare_optimizer_params(text_encoder_lr, unet_lr, learning_rate) instead of prepare_optimizer_params(text_encoder_lr, unet_lr)"
|
||||
# )
|
||||
trainable_params = network.prepare_optimizer_params(args.text_encoder_lr, args.unet_lr)
|
||||
trainable_params = network.prepare_optimizer_params(text_encoder_lr, args.unet_lr)
|
||||
lr_descriptions = None
|
||||
|
||||
# if len(trainable_params) == 0:
|
||||
@@ -483,6 +567,7 @@ class NetworkTrainer:
|
||||
# accelerator.print(f"trainable_params: {k} = {v}")
|
||||
|
||||
optimizer_name, optimizer_args, optimizer = train_util.get_optimizer(args, trainable_params)
|
||||
self.optimizer_train_fn, self.optimizer_eval_fn = train_util.get_optimizer_train_eval_fn(optimizer, args)
|
||||
|
||||
# prepare dataloader
|
||||
# strategies are set here because they cannot be referenced in another process. Copy them with the dataset
|
||||
@@ -501,14 +586,14 @@ class NetworkTrainer:
|
||||
persistent_workers=args.persistent_data_loader_workers,
|
||||
)
|
||||
|
||||
# # Calculate the number of learning steps
|
||||
# if args.max_train_epochs is not None:
|
||||
# args.max_train_steps = args.max_train_epochs * math.ceil(
|
||||
# len(train_dataloader) / accelerator.num_processes / args.gradient_accumulation_steps
|
||||
# )
|
||||
# accelerator.print(
|
||||
# f"override steps. steps for {args.max_train_epochs} epochs is {args.max_train_steps}"
|
||||
# )
|
||||
# 学習ステップ数を計算する
|
||||
if args.max_train_epochs is not None:
|
||||
args.max_train_steps = args.max_train_epochs * math.ceil(
|
||||
len(train_dataloader) / accelerator.num_processes / args.gradient_accumulation_steps
|
||||
)
|
||||
accelerator.print(
|
||||
f"override steps. steps for {args.max_train_epochs} epochs is / 指定エポックまでのステップ数: {args.max_train_steps}"
|
||||
)
|
||||
|
||||
# Send learning steps to the dataset side as well
|
||||
train_dataset_group.set_max_train_steps(args.max_train_steps)
|
||||
@@ -538,30 +623,33 @@ class NetworkTrainer:
|
||||
args.mixed_precision != "no"
|
||||
), "fp8_base requires mixed precision='fp16' or 'bf16'"
|
||||
accelerator.print("enable fp8 training for U-Net.")
|
||||
unet_weight_dtype = torch.float8_e4m3fn
|
||||
unet_weight_dtype = torch.float8_e4m3fn if args.fp8_dtype == "e4m3" else torch.float8_e5m2
|
||||
accelerator.print(f"unet_weight_dtype: {unet_weight_dtype}")
|
||||
|
||||
if not args.fp8_base_unet and not args.network_train_unet_only:
|
||||
accelerator.print("enable fp8 training for Text Encoder.")
|
||||
te_weight_dtype = weight_dtype if args.fp8_base_unet else torch.float8_e4m3fn
|
||||
te_weight_dtype = torch.float8_e4m3fn if args.fp8_dtype == "e4m3" else torch.float8_e5m2
|
||||
|
||||
# unet.to(accelerator.device) # this makes faster `to(dtype)` below, but consumes 23 GB VRAM
|
||||
# unet.to(dtype=unet_weight_dtype) # without moving to gpu, this takes a lot of time and main memory
|
||||
|
||||
unet.to(accelerator.device, dtype=unet_weight_dtype) # this seems to be safer than above
|
||||
# logger.info(f"set U-Net weight dtype to {unet_weight_dtype}, device to {accelerator.device}")
|
||||
# unet.to(accelerator.device, dtype=unet_weight_dtype) # this seems to be safer than above
|
||||
logger.info(f"set U-Net weight dtype to {unet_weight_dtype}")
|
||||
unet.to(dtype=unet_weight_dtype) # do not move to device because unet is not prepared by accelerator
|
||||
|
||||
unet.requires_grad_(False)
|
||||
unet.to(dtype=unet_weight_dtype)
|
||||
for t_enc in text_encoders:
|
||||
for i, t_enc in enumerate(text_encoders):
|
||||
t_enc.requires_grad_(False)
|
||||
|
||||
# in case of cpu, dtype is already set to fp32 because cpu does not support fp8/fp16/bf16
|
||||
if t_enc.device.type != "cpu":
|
||||
t_enc.to(dtype=te_weight_dtype)
|
||||
if hasattr(t_enc, "text_model") and hasattr(t_enc.text_model, "embeddings"):
|
||||
|
||||
# nn.Embedding not support FP8
|
||||
t_enc.text_model.embeddings.to(dtype=(weight_dtype if te_weight_dtype != weight_dtype else te_weight_dtype))
|
||||
elif hasattr(t_enc, "encoder") and hasattr(t_enc.encoder, "embeddings"):
|
||||
t_enc.encoder.embeddings.to(dtype=(weight_dtype if te_weight_dtype != weight_dtype else te_weight_dtype))
|
||||
if te_weight_dtype != weight_dtype:
|
||||
self.prepare_text_encoder_fp8(i, t_enc, te_weight_dtype, weight_dtype)
|
||||
|
||||
# acceleratorがなんかよろしくやってくれるらしい / accelerator will do something good
|
||||
if args.deepspeed:
|
||||
@@ -579,7 +667,8 @@ class NetworkTrainer:
|
||||
training_model = ds_model
|
||||
else:
|
||||
if train_unet:
|
||||
unet = accelerator.prepare(unet)
|
||||
# default implementation is: unet = accelerator.prepare(unet)
|
||||
unet = self.prepare_unet_with_accelerator(args, accelerator, unet) # accelerator does some magic here
|
||||
else:
|
||||
unet.to(accelerator.device, dtype=unet_weight_dtype) # move to device because unet is not prepared by accelerator
|
||||
if train_text_encoder:
|
||||
@@ -602,12 +691,12 @@ class NetworkTrainer:
|
||||
if args.gradient_checkpointing:
|
||||
# according to TI example in Diffusers, train is required
|
||||
unet.train()
|
||||
for t_enc, frag in zip(text_encoders, self.get_text_encoders_train_flags(args, text_encoders)):
|
||||
for i, (t_enc, frag) in enumerate(zip(text_encoders, self.get_text_encoders_train_flags(args, text_encoders))):
|
||||
t_enc.train()
|
||||
|
||||
# set top parameter requires_grad = True for gradient checkpointing works
|
||||
if frag:
|
||||
t_enc.text_model.embeddings.requires_grad_(True)
|
||||
self.prepare_text_encoder_grad_ckpt_workaround(i, t_enc)
|
||||
|
||||
else:
|
||||
unet.eval()
|
||||
@@ -706,7 +795,7 @@ class NetworkTrainer:
|
||||
"ss_training_started_at": training_started_at, # unix timestamp
|
||||
"ss_output_name": args.output_name,
|
||||
"ss_learning_rate": args.learning_rate,
|
||||
"ss_text_encoder_lr": args.text_encoder_lr,
|
||||
"ss_text_encoder_lr": text_encoder_lr,
|
||||
"ss_unet_lr": args.unet_lr,
|
||||
"ss_num_train_images": train_dataset_group.num_train_images,
|
||||
"ss_num_reg_images": train_dataset_group.num_reg_images,
|
||||
@@ -752,9 +841,10 @@ class NetworkTrainer:
|
||||
"ss_ip_noise_gamma_random_strength": args.ip_noise_gamma_random_strength,
|
||||
"ss_loss_type": args.loss_type,
|
||||
"ss_huber_schedule": args.huber_schedule,
|
||||
"ss_huber_scale": args.huber_scale,
|
||||
"ss_huber_c": args.huber_c,
|
||||
"ss_fp8_base": args.fp8_base,
|
||||
"ss_fp8_base_unet": args.fp8_base_unet,
|
||||
"ss_fp8_base": bool(args.fp8_base),
|
||||
"ss_fp8_base_unet": bool(args.fp8_base_unet),
|
||||
}
|
||||
|
||||
self.update_metadata(metadata, args) # architecture specific metadata
|
||||
@@ -985,9 +1075,9 @@ class NetworkTrainer:
|
||||
|
||||
# callback for step start
|
||||
if hasattr(accelerator.unwrap_model(network), "on_step_start"):
|
||||
on_step_start = accelerator.unwrap_model(network).on_step_start
|
||||
on_step_start_for_network = accelerator.unwrap_model(network).on_step_start
|
||||
else:
|
||||
on_step_start = lambda *args, **kwargs: None
|
||||
on_step_start_for_network = lambda *args, **kwargs: None
|
||||
|
||||
# function for saving/removing
|
||||
def save_model(ckpt_name, unwrapped_nw, steps, epoch_no, force_sync_upload=False):
|
||||
@@ -1026,14 +1116,18 @@ class NetworkTrainer:
|
||||
self.global_step = 0
|
||||
# training loop
|
||||
if initial_step > 0: # only if skip_until_initial_step is specified
|
||||
self.global_step = initial_step
|
||||
logger.info(f"skipping epoch {epoch_to_start} because initial_step (multiplied) is {initial_step}")
|
||||
initial_step -= epoch_to_start * len(train_dataloader)
|
||||
for skip_epoch in range(epoch_to_start): # skip epochs
|
||||
logger.info(f"skipping epoch {skip_epoch+1} because initial_step (multiplied) is {initial_step}")
|
||||
initial_step -= len(train_dataloader)
|
||||
|
||||
# log device and dtype for each model
|
||||
logger.info(f"unet dtype: {unet_weight_dtype}, device: {unet.device}")
|
||||
for t_enc in text_encoders:
|
||||
logger.info(f"text_encoder dtype: {t_enc.dtype}, device: {t_enc.device}")
|
||||
for i, t_enc in enumerate(text_encoders):
|
||||
params_itr = t_enc.parameters()
|
||||
params_itr.__next__() # skip the first parameter
|
||||
params_itr.__next__() # skip the second parameter. because CLIP first two parameters are embeddings
|
||||
param_3rd = params_itr.__next__()
|
||||
logger.info(f"text_encoder [{i}] dtype: {param_3rd.dtype}, device: {t_enc.device}")
|
||||
|
||||
clean_memory_on_device(accelerator.device)
|
||||
|
||||
@@ -1054,10 +1148,13 @@ class NetworkTrainer:
|
||||
self.lr_scheduler = lr_scheduler
|
||||
self.save_model = save_model
|
||||
self.remove_model = remove_model
|
||||
self.comfy_pbar = None
|
||||
|
||||
progress_bar = tqdm(range(args.max_train_steps - initial_step), smoothing=0, disable=False, desc="steps")
|
||||
|
||||
def training_loop(break_at_steps, epoch):
|
||||
steps_done = 0
|
||||
|
||||
#accelerator.print(f"\nepoch {epoch+1}/{num_train_epochs}")
|
||||
progress_bar.set_description(f"Epoch {epoch + 1}/{num_train_epochs} - steps")
|
||||
|
||||
@@ -1069,14 +1166,20 @@ class NetworkTrainer:
|
||||
|
||||
skipped_dataloader = None
|
||||
if self.initial_step > 0:
|
||||
skipped_dataloader = accelerator.skip_first_batches(train_dataloader, initial_step)
|
||||
initial_step = 0
|
||||
skipped_dataloader = accelerator.skip_first_batches(train_dataloader, self.initial_step - 1)
|
||||
self.initial_step = 1
|
||||
|
||||
for step, batch in enumerate(skipped_dataloader or train_dataloader):
|
||||
current_step.value = self.global_step
|
||||
if self.initial_step > 0:
|
||||
self.initial_step -= 1
|
||||
continue
|
||||
|
||||
with accelerator.accumulate(training_model):
|
||||
on_step_start(text_encoder, unet)
|
||||
on_step_start_for_network(text_encoder, unet)
|
||||
|
||||
# temporary, for batch processing
|
||||
self.on_step_start(args, accelerator, network, text_encoders, unet, batch, weight_dtype)
|
||||
|
||||
if "latents" in batch and batch["latents"] is not None:
|
||||
latents = batch["latents"].to(accelerator.device).to(dtype=weight_dtype)
|
||||
@@ -1104,26 +1207,22 @@ class NetworkTrainer:
|
||||
# print(f"set multiplier: {multipliers}")
|
||||
accelerator.unwrap_model(network).set_multiplier(multipliers)
|
||||
|
||||
text_encoder_conds = []
|
||||
text_encoder_outputs_list = batch.get("text_encoder_outputs_list", None)
|
||||
if text_encoder_outputs_list is not None:
|
||||
text_encoder_conds = text_encoder_outputs_list # List of text encoder outputs
|
||||
if (
|
||||
text_encoder_conds is None
|
||||
or len(text_encoder_conds) == 0
|
||||
or text_encoder_conds[0] is None
|
||||
or train_text_encoder
|
||||
):
|
||||
|
||||
if len(text_encoder_conds) == 0 or text_encoder_conds[0] is None or train_text_encoder:
|
||||
# TODO this does not work if 'some text_encoders are trained' and 'some are not and not cached'
|
||||
with torch.set_grad_enabled(train_text_encoder), accelerator.autocast():
|
||||
# Get the text embedding for conditioning
|
||||
if args.weighted_captions:
|
||||
# SD only
|
||||
encoded_text_encoder_conds = get_weighted_text_embeddings(
|
||||
tokenizers[0],
|
||||
text_encoder,
|
||||
batch["captions"],
|
||||
accelerator.device,
|
||||
args.max_token_length // 75 if args.max_token_length else 1,
|
||||
clip_skip=args.clip_skip,
|
||||
input_ids_list, weights_list = tokenize_strategy.tokenize_with_weights(batch["captions"])
|
||||
encoded_text_encoder_conds = text_encoding_strategy.encode_tokens_with_weights(
|
||||
tokenize_strategy,
|
||||
self.get_models_for_text_encoding(args, accelerator, text_encoders),
|
||||
input_ids_list,
|
||||
weights_list,
|
||||
)
|
||||
else:
|
||||
input_ids = [ids.to(accelerator.device) for ids in batch["input_ids_list"]]
|
||||
@@ -1135,13 +1234,17 @@ class NetworkTrainer:
|
||||
if args.full_fp16:
|
||||
encoded_text_encoder_conds = [c.to(weight_dtype) for c in encoded_text_encoder_conds]
|
||||
|
||||
# if text_encoder_conds is not cached, use encoded_text_encoder_conds
|
||||
if len(text_encoder_conds) == 0:
|
||||
text_encoder_conds = encoded_text_encoder_conds
|
||||
else:
|
||||
# if encoded_text_encoder_conds is not None, update cached text_encoder_conds
|
||||
for i in range(len(encoded_text_encoder_conds)):
|
||||
if encoded_text_encoder_conds[i] is not None:
|
||||
text_encoder_conds[i] = encoded_text_encoder_conds[i]
|
||||
|
||||
# sample noise, call unet, get target
|
||||
noise_pred, target, timesteps, huber_c, weighting = self.get_noise_pred_and_target(
|
||||
noise_pred, target, timesteps, weighting = self.get_noise_pred_and_target(
|
||||
args,
|
||||
accelerator,
|
||||
noise_scheduler,
|
||||
@@ -1154,9 +1257,8 @@ class NetworkTrainer:
|
||||
train_unet,
|
||||
)
|
||||
|
||||
loss = train_util.conditional_loss(
|
||||
noise_pred.float(), target.float(), reduction="none", loss_type=args.loss_type, huber_c=huber_c
|
||||
)
|
||||
huber_c = train_util.get_huber_threshold_if_needed(args, timesteps, noise_scheduler)
|
||||
loss = train_util.conditional_loss(noise_pred.float(), target.float(), args.loss_type, "none", huber_c)
|
||||
if weighting is not None:
|
||||
loss = loss * weighting
|
||||
if args.masked_loss or ("alpha_masks" in batch and batch["alpha_masks"] is not None):
|
||||
@@ -1204,17 +1306,18 @@ class NetworkTrainer:
|
||||
if args.scale_weight_norms:
|
||||
progress_bar.set_postfix(**{**max_mean_logs, **logs})
|
||||
|
||||
if args.logging_dir is not None:
|
||||
if len(accelerator.trackers) > 0:
|
||||
logs = self.generate_step_logs(
|
||||
args, current_loss, avr_loss, lr_scheduler, lr_descriptions, keys_scaled, mean_norm, maximum_norm
|
||||
args, current_loss, avr_loss, lr_scheduler, lr_descriptions, optimizer, keys_scaled, mean_norm, maximum_norm
|
||||
)
|
||||
accelerator.log(logs, step=self.global_step)
|
||||
|
||||
if self.global_step >= break_at_steps:
|
||||
break
|
||||
steps_done += 1
|
||||
self.comfy_pbar.update(1)
|
||||
|
||||
if args.logging_dir is not None:
|
||||
if len(accelerator.trackers) > 0:
|
||||
logs = {"loss/epoch": self.loss_recorder.moving_average}
|
||||
accelerator.log(logs, step=epoch + 1)
|
||||
|
||||
@@ -1251,8 +1354,15 @@ def setup_parser() -> argparse.ArgumentParser:
|
||||
deepspeed_utils.add_deepspeed_arguments(parser)
|
||||
train_util.add_optimizer_arguments(parser)
|
||||
config_util.add_config_arguments(parser)
|
||||
train_util.add_dit_training_arguments(parser)
|
||||
custom_train_functions.add_custom_train_arguments(parser)
|
||||
|
||||
parser.add_argument(
|
||||
"--cpu_offload_checkpointing",
|
||||
action="store_true",
|
||||
help="[EXPERIMENTAL] enable offloading of tensors to CPU during checkpointing for U-Net or DiT, if supported"
|
||||
" / 勾配チェックポイント時にテンソルをCPUにオフロードする(U-NetまたはDiTのみ、サポートされている場合)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--no_metadata", action="store_true", help="do not save metadata in output model / メタデータを出力先モデルに保存しない"
|
||||
)
|
||||
@@ -1265,7 +1375,19 @@ def setup_parser() -> argparse.ArgumentParser:
|
||||
)
|
||||
|
||||
parser.add_argument("--unet_lr", type=float, default=None, help="learning rate for U-Net / U-Netの学習率")
|
||||
parser.add_argument("--text_encoder_lr", type=float, default=None, help="learning rate for Text Encoder / Text Encoderの学習率")
|
||||
parser.add_argument(
|
||||
"--text_encoder_lr",
|
||||
type=float,
|
||||
default=None,
|
||||
nargs="*",
|
||||
help="learning rate for Text Encoder, can be multiple / Text Encoderの学習率、複数指定可能",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--fp8_base_unet",
|
||||
action="store_true",
|
||||
help="use fp8 for U-Net (or DiT), Text Encoder is fp16 or bf16"
|
||||
" / U-Net(またはDiT)にfp8を使用する。Text Encoderはfp16またはbf16",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--network_weights", type=str, default=None, help="pretrained weights for network / 学習するネットワークの初期重み"
|
||||
@@ -1361,12 +1483,6 @@ def setup_parser() -> argparse.ArgumentParser:
|
||||
help="initial step number including all epochs, 0 means first step (same as not specifying). overwrites initial_epoch."
|
||||
+ " / 初期ステップ数、全エポックを含むステップ数、0で最初のステップ(未指定時と同じ)。initial_epochを上書きする",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--fp8_base_unet",
|
||||
action="store_true",
|
||||
help="use fp8 for U-Net (or DiT), Text Encoder is fp16 or bf16"
|
||||
" / U-Net(またはDiT)にfp8を使用する。Text Encoderはfp16またはbf16",
|
||||
)
|
||||
# parser.add_argument("--loraplus_lr_ratio", default=None, type=float, help="LoRA+ learning rate ratio")
|
||||
# parser.add_argument("--loraplus_unet_lr_ratio", default=None, type=float, help="LoRA+ UNet learning rate ratio")
|
||||
# parser.add_argument("--loraplus_text_encoder_lr_ratio", default=None, type=float, help="LoRA+ text encoder learning rate ratio")
|
||||
|
||||
Reference in New Issue
Block a user